#!/usr/bin/env sh
set -eu

use_openblas="${FASTPLS_USE_OPENBLAS:-auto}"
openblas_root="${OPENBLAS_ROOT:-}"
use_accelerate="${FASTPLS_USE_ACCELERATE:-auto}"
use_cuda="${FASTPLS_USE_CUDA:-auto}"
require_cuda="${FASTPLS_REQUIRE_CUDA:-${PACKAGE_REQUIRE_CUDA:-0}}"
diagnostic_cuda="${FASTPLS_CUDA_DIAGNOSTIC_ONLY:-0}"
cuda_root="${CUDA_ROOT:-${CUDA_HOME:-${CUDA_PATH:-}}}"
use_metal="${FASTPLS_USE_METAL:-auto}"
system_name="$(uname -s)"

pkg_cxxflags=""
pkg_cppflags=""
pkg_libs='$(LAPACK_LIBS) $(BLAS_LIBS)'
cuda_objects=""
metal_objects=""
metal_mm_rule=""

openblas_cppflags=""
openblas_libs=""
openblas_origin=""

find_openblas_from_root() {
  [ -n "${openblas_root}" ] || return 1
  header="$(find -H "${openblas_root}/include" -name openblas_config.h -print 2>/dev/null | head -n 1 || true)"
  library="$(find -H "${openblas_root}" \( -name 'libopenblas.so' -o -name 'libopenblas.dylib' -o -name 'libopenblas.a' \) -print 2>/dev/null | head -n 1 || true)"
  [ -n "${header}" ] && [ -n "${library}" ] || return 1
  openblas_cppflags="-I$(dirname "${header}")"
  openblas_library_dir="$(dirname "${library}")"
  openblas_libs="-L${openblas_library_dir} -Wl,-rpath,${openblas_library_dir} -lopenblas"
  openblas_origin="${openblas_root}"
}

find_openblas() {
  if find_openblas_from_root; then
    return 0
  fi
  if command -v pkg-config >/dev/null 2>&1 && pkg-config --exists openblas; then
    openblas_cppflags="$(pkg-config --cflags openblas)"
    openblas_libs="$(pkg-config --libs openblas)"
    openblas_origin="pkg-config"
    return 0
  fi
  if [ "${system_name}" = "Darwin" ] && [ -d "/opt/homebrew/opt/openblas" ]; then
    openblas_root="/opt/homebrew/opt/openblas"
    find_openblas_from_root
    return
  fi
  return 1
}

if [ "${use_openblas}" = "true" ] || [ "${use_openblas}" = "TRUE" ]; then
  use_openblas="1"
elif [ "${use_openblas}" = "false" ] || [ "${use_openblas}" = "FALSE" ]; then
  use_openblas="0"
fi

if [ "${use_openblas}" = "auto" ]; then
  if [ "${system_name}" != "Darwin" ] && find_openblas; then
    use_openblas="1"
  else
    use_openblas="0"
  fi
elif [ "${use_openblas}" = "1" ]; then
  if ! find_openblas; then
    echo "ERROR: FASTPLS_USE_OPENBLAS=1 but OpenBLAS was not found." >&2
    echo "Install the OpenBLAS development package and pkg-config, or set OPENBLAS_ROOT." >&2
    exit 1
  fi
elif [ "${use_openblas}" != "0" ]; then
  echo "ERROR: FASTPLS_USE_OPENBLAS must be auto, 0, or 1." >&2
  exit 1
fi

if [ "${system_name}" = "Darwin" ] &&
   [ "${use_openblas}" != "1" ] &&
   [ "${use_accelerate}" != "0" ] &&
   [ "${use_accelerate}" != "false" ] &&
  [ "${use_accelerate}" != "FALSE" ]; then
  echo "fastPLS configure: using Apple Accelerate for CPU BLAS/LAPACK"
  pkg_cppflags="-DFASTPLS_USE_ACCELERATE"
  pkg_libs="-framework Accelerate"
elif [ "${use_openblas}" = "1" ]; then
  echo "fastPLS configure: using OpenBLAS (${openblas_origin}); thread count is controlled by OpenBLAS"
  pkg_cppflags="-DFASTPLS_USE_OPENBLAS ${openblas_cppflags}"
  if [ "${system_name}" = "Darwin" ]; then
    openblas_dl_lib=""
  else
    openblas_dl_lib="-ldl"
  fi
  pkg_libs="${openblas_libs} ${openblas_dl_lib}"
elif [ "${system_name}" != "Darwin" ]; then
  echo "fastPLS configure: OpenBLAS not found; using the BLAS/LAPACK supplied by R"
fi

if [ "${system_name}" != "Darwin" ] && [ "${use_openblas}" != "1" ]; then
  pkg_libs="${pkg_libs} -ldl"
fi

normalize_switch() {
  case "$1" in
    true|TRUE|yes|YES) printf "%s" "1" ;;
    false|FALSE|no|NO) printf "%s" "0" ;;
    *) printf "%s" "$1" ;;
  esac
}

use_cuda="$(normalize_switch "${use_cuda}")"
require_cuda="$(normalize_switch "${require_cuda}")"
diagnostic_cuda="$(normalize_switch "${diagnostic_cuda}")"

if [ "${require_cuda}" = "1" ]; then
  use_cuda="1"
fi
if [ "${diagnostic_cuda}" = "1" ] && [ "${require_cuda}" = "1" ]; then
  echo "ERROR: strict CUDA and diagnostic-only modes cannot be combined." >&2
  exit 1
fi

find_cuda_root() {
  if [ -n "${cuda_root}" ]; then
    printf "%s" "${cuda_root}"
    return
  fi

  for d in /usr/local/cuda /opt/cuda /usr/local/cuda-*; do
    if [ -d "$d" ]; then
      printf "%s" "$d"
      return
    fi
  done
}

find_cuda_include_dir() {
  root="$1"
  for directory in "${root}/include" "${root}"/targets/*/include; do
    if [ -f "${directory}/cuda_runtime.h" ] &&
       [ -f "${directory}/cublas_v2.h" ] &&
       [ -f "${directory}/cusolverDn.h" ] &&
       [ -f "${directory}/curand.h" ]; then
      printf "%s" "${directory}"
      return 0
    fi
  done
  return 1
}

has_cuda_library() {
  directory="$1"
  library="$2"
  for candidate in \
    "${directory}/lib${library}.so" \
    "${directory}/lib${library}.so."* \
    "${directory}/lib${library}.a"; do
    [ -f "${candidate}" ] && return 0
  done
  return 1
}

find_cuda_library_dir() {
  root="$1"
  for directory in \
    "${root}/lib" \
    "${root}/lib64" \
    "${root}"/targets/*/lib \
    "${root}"/targets/*/lib64; do
    [ -d "${directory}" ] || continue
    if has_cuda_library "${directory}" cudart &&
       has_cuda_library "${directory}" cublas &&
       has_cuda_library "${directory}" cusolver &&
       has_cuda_library "${directory}" curand; then
      printf "%s" "${directory}"
      return 0
    fi
  done
  return 1
}

find_cuda_host_cxx() {
  if [ -n "${FASTPLS_CUDA_HOST_CXX:-}" ]; then
    if [ ! -x "${FASTPLS_CUDA_HOST_CXX}" ]; then
      echo "ERROR: FASTPLS_CUDA_HOST_CXX is not executable." >&2
      return 1
    fi
    printf "%s" "${FASTPLS_CUDA_HOST_CXX}"
    return 0
  fi
  for compiler in /usr/bin/g++ /usr/bin/clang++ /usr/bin/c++; do
    if [ -x "${compiler}" ]; then
      printf "%s" "${compiler}"
      return 0
    fi
  done
  compiler="$(command -v c++ 2>/dev/null || true)"
  case "${compiler}" in
    ""|*conda*|*CUDA*|*cuda*) return 1 ;;
    *) printf "%s" "${compiler}" ;;
  esac
}

probe_cuda_toolkit() {
  probe_directory="${TMPDIR:-/tmp}/fastpls-cuda-probe.$$"
  probe_source="${probe_directory}/probe.cu"
  probe_binary="${probe_directory}/probe"
  probe_log="${probe_directory}/probe.log"
  mkdir -p "${probe_directory}"
  cat > "${probe_source}" <<'EOF'
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <curand.h>

int main() {
  cublasHandle_t cublas = nullptr;
  cusolverDnHandle_t cusolver = nullptr;
  curandGenerator_t curand = nullptr;
  (void) cudaFree(nullptr);
  (void) cublasCreate(&cublas);
  (void) cusolverDnCreate(&cusolver);
  (void) curandCreateGenerator(&curand, CURAND_RNG_PSEUDO_DEFAULT);
  return 0;
}
EOF
  if PATH="$(dirname "${cuda_host_cxx}"):/usr/bin:/bin:${PATH}" \
      "${cuda_nvcc}" -ccbin "${cuda_host_cxx}" -std=c++17 \
      -I"${cuda_include_dir}" "${probe_source}" \
      -L"${cuda_lib_dir}" -lcudart -lcublas -lcusolver -lcurand \
      -Xlinker=-rpath -Xlinker="${cuda_lib_dir}" \
      -o "${probe_binary}" >"${probe_log}" 2>&1; then
    rm -rf "${probe_directory}"
    return 0
  fi
  echo "fastPLS configure: CUDA compile/link probe failed:" >&2
  sed -n '1,120p' "${probe_log}" >&2
  rm -rf "${probe_directory}"
  return 1
}

cuda_root="$(find_cuda_root || true)"
cuda_lib_dir=""
cuda_include_dir=""
cuda_nvcc=""
cuda_host_cxx=""
cuda_probe_reason="CUDA Toolkit was not found"
have_cuda=0

if [ "${use_cuda}" != "0" ] && [ "${diagnostic_cuda}" != "1" ] &&
   [ -n "${cuda_root}" ]; then
  cuda_include_dir="$(find_cuda_include_dir "${cuda_root}" || true)"
  cuda_lib_dir="$(find_cuda_library_dir "${cuda_root}" || true)"
  cuda_nvcc="${cuda_root}/bin/nvcc"
  cuda_host_cxx="$(find_cuda_host_cxx || true)"
  if [ ! -x "${cuda_nvcc}" ]; then
    cuda_probe_reason="nvcc was not found under the selected CUDA Toolkit"
  elif [ -z "${cuda_include_dir}" ]; then
    cuda_probe_reason="required CUDA headers were not found"
  elif [ -z "${cuda_lib_dir}" ]; then
    cuda_probe_reason="required CUDA libraries were not found together"
  elif [ -z "${cuda_host_cxx}" ]; then
    cuda_probe_reason="a system CUDA host C++ compiler was not found"
  elif probe_cuda_toolkit; then
    have_cuda=1
    cuda_probe_reason="CUDA compile/link probe succeeded"
  else
    cuda_probe_reason="CUDA compile/link probe failed"
  fi
fi

if [ "${require_cuda}" = "1" ] && [ "${have_cuda}" != "1" ]; then
  echo "ERROR: FASTPLS_REQUIRE_CUDA=1 but CUDA support cannot be built: ${cuda_probe_reason}." >&2
  echo "Set CUDA_ROOT or CUDA_HOME to the CUDA Toolkit directory, for example:" >&2
  echo "  /usr/local/cuda" >&2
  exit 1
fi

if [ "${use_cuda}" = "1" ] && [ "${have_cuda}" != "1" ]; then
  echo "fastPLS configure: FASTPLS_USE_CUDA=1 was set but the CUDA Toolkit was not found or is incomplete; building CPU-only package"
  echo "fastPLS configure: set FASTPLS_REQUIRE_CUDA=1 to make missing CUDA a configuration error"
fi

if [ "${have_cuda}" = "1" ]; then
  echo "fastPLS configure: CUDA Toolkit found at ${cuda_root}"
  echo "fastPLS configure: CUDA include directory ${cuda_include_dir}"
  echo "fastPLS configure: CUDA library directory ${cuda_lib_dir}"
  echo "fastPLS configure: CUDA host compiler ${cuda_host_cxx}"
  cuda_flags="-DFASTPLS_HAS_CUDA -DFASTPLS_HAS_CUDA_KERNELS -I${cuda_include_dir}"
  cuda_objects=" core_cuda_backend.o"
  if [ -n "${pkg_cppflags}" ]; then
    pkg_cppflags="${pkg_cppflags} ${cuda_flags}"
  else
    pkg_cppflags="${cuda_flags}"
  fi
  pkg_cxxflags="${cuda_flags}"
  pkg_libs="-L${cuda_lib_dir} -Wl,-rpath,${cuda_lib_dir} -lcudart -lcublas -lcusolver -lcurand ${pkg_libs}"
elif [ "${diagnostic_cuda}" = "1" ]; then
  echo "fastPLS configure: building in explicit CUDA diagnostic-only mode"
  pkg_cppflags="${pkg_cppflags} -DFASTPLS_CUDA_DIAGNOSTIC_ONLY"
  pkg_cxxflags="${pkg_cxxflags} -DFASTPLS_CUDA_DIAGNOSTIC_ONLY"
else
  echo "fastPLS configure: ${cuda_probe_reason}; building without CUDA"
fi

have_metal=0
if [ "${system_name}" = "Darwin" ]; then
  system_metal_framework="/System/Library/Frameworks/Metal.framework"
  metal_sdk=""
  if [ ! -d "${system_metal_framework}" ]; then
    metal_sdk="$(xcrun --sdk macosx --show-sdk-path 2>/dev/null || true)"
  fi
  if [ "${use_metal}" = "0" ] || [ "${use_metal}" = "false" ] || [ "${use_metal}" = "FALSE" ]; then
    have_metal=0
  elif [ "${use_metal}" = "1" ]; then
    have_metal=1
  elif [ "${use_metal}" = "auto" ] &&
       { [ -d "${system_metal_framework}" ] ||
         { [ -n "${metal_sdk}" ] && [ -d "${metal_sdk}/System/Library/Frameworks/Metal.framework" ]; }; }; then
    have_metal=1
  fi
fi

if [ "${have_metal}" = "1" ]; then
  if [ -n "${metal_sdk}" ]; then
    echo "fastPLS configure: Apple Metal backend enabled (SDK: ${metal_sdk})"
  else
    echo "fastPLS configure: Apple Metal backend enabled using system frameworks"
  fi
  cp src/core_metal_backend.mm.in src/core_metal_backend.mm
  metal_objects=" core_metal_backend.o"
  metal_flags="-DFASTPLS_HAS_METAL"
  if [ -n "${pkg_cppflags}" ]; then
    pkg_cppflags="${pkg_cppflags} ${metal_flags}"
  else
    pkg_cppflags="${metal_flags}"
  fi
  pkg_cxxflags="${pkg_cxxflags} ${metal_flags}"
  pkg_libs="${pkg_libs} -framework Foundation -framework Metal -framework MetalPerformanceShaders"
  metal_mm_rule='
%.o: %.mm
	$(CXX17) -std=gnu++17 $(ALL_CPPFLAGS) $(PKG_CPPFLAGS) $(CXX17FLAGS) $(CXX17PICFLAGS) $(PKG_CXXFLAGS) -x objective-c++ -fobjc-arc -c $< -o $@'
elif [ "${system_name}" = "Darwin" ]; then
  if [ "${use_metal}" = "0" ] || [ "${use_metal}" = "false" ] || [ "${use_metal}" = "FALSE" ]; then
    echo "fastPLS configure: Apple Metal backend disabled by FASTPLS_USE_METAL=${use_metal}"
  else
    echo "fastPLS configure: Apple Metal frameworks not found; building without Metal"
  fi
fi

sed \
  -e "s|@FASTPLS_PKG_CXXFLAGS@|${pkg_cxxflags}|g" \
  -e "s|@FASTPLS_PKG_CPPFLAGS@|${pkg_cppflags}|g" \
  -e "s|@FASTPLS_METAL_OBJECTS@|${metal_objects}|g" \
  -e "s|@FASTPLS_CUDA_OBJECTS@|${cuda_objects}|g" \
  -e "s|@FASTPLS_PKG_LIBS@|${pkg_libs}|g" \
  src/Makevars.in > src/Makevars

if [ -n "${metal_mm_rule}" ]; then
  printf "%s\n" "${metal_mm_rule}" >> src/Makevars
fi

if [ -n "${cuda_objects}" ]; then
  cat >> src/Makevars <<EOF

CUDA_NVCC = ${cuda_nvcc}
CUDA_HOST_CXX = ${cuda_host_cxx}
CUDA_NVCCFLAGS = -std=c++17 --extended-lambda --expt-relaxed-constexpr -ccbin \$(CUDA_HOST_CXX) -I../inst/include/ ${cuda_flags} -Xcompiler -fPIC

%.o: %.cu
	PATH=/usr/bin:/bin:\$\$PATH \$(CUDA_NVCC) \$(CUDA_NVCCFLAGS) -c \$< -o \$@
EOF
fi
