cmake_minimum_required(VERSION 3.27)

project(elfes_native LANGUAGES CXX)

option(BUILD_CUDA "Build ELFES CUDA providers" ON)
set(ELFES_NEIGHBOR_LIST_UPSTREAM_COMMIT "fecce19e8e757a6dfac2e41fb0ca5d5bed47fb00")

find_package(Python REQUIRED COMPONENTS Interpreter Development.Module NumPy)
find_package(pybind11 CONFIG REQUIRED)

execute_process(
    COMMAND
        "${Python_EXECUTABLE}" -c
        "import torch; from torch.utils.cpp_extension import include_paths, library_paths; print(torch.__version__); print(torch.version.cuda or ''); print('|'.join(include_paths())); print('|'.join(library_paths())); print(int(torch.compiled_with_cxx11_abi()))"
    RESULT_VARIABLE ELFES_TORCH_INFO_RESULT
    OUTPUT_VARIABLE ELFES_TORCH_INFO
    ERROR_VARIABLE ELFES_TORCH_INFO_ERROR
    OUTPUT_STRIP_TRAILING_WHITESPACE
)
if(NOT ELFES_TORCH_INFO_RESULT EQUAL 0)
    message(FATAL_ERROR "Failed to inspect the installed PyTorch package:\n${ELFES_TORCH_INFO_ERROR}")
endif()
string(REPLACE "\n" ";" ELFES_TORCH_INFO "${ELFES_TORCH_INFO}")
list(GET ELFES_TORCH_INFO 0 ELFES_TORCH_VERSION)
list(GET ELFES_TORCH_INFO 1 ELFES_TORCH_CUDA_VERSION)
list(GET ELFES_TORCH_INFO 2 ELFES_TORCH_INCLUDE_PATHS)
list(GET ELFES_TORCH_INFO 3 ELFES_TORCH_LIBRARY_PATHS)
list(GET ELFES_TORCH_INFO 4 ELFES_TORCH_CXX11_ABI)
string(REPLACE "|" ";" ELFES_TORCH_INCLUDE_DIRS "${ELFES_TORCH_INCLUDE_PATHS}")
string(REPLACE "|" ";" ELFES_TORCH_LIBRARY_DIRS "${ELFES_TORCH_LIBRARY_PATHS}")
if(ELFES_TORCH_VERSION VERSION_LESS 2.12)
    message(FATAL_ERROR "ELFES requires PyTorch 2.12 or newer, found ${ELFES_TORCH_VERSION}")
endif()

find_library(ELFES_C10_LIBRARY c10 PATHS ${ELFES_TORCH_LIBRARY_DIRS} NO_DEFAULT_PATH REQUIRED)
find_library(ELFES_TORCH_LIBRARY torch PATHS ${ELFES_TORCH_LIBRARY_DIRS} NO_DEFAULT_PATH REQUIRED)
find_library(ELFES_TORCH_CPU_LIBRARY torch_cpu PATHS ${ELFES_TORCH_LIBRARY_DIRS} NO_DEFAULT_PATH REQUIRED)
find_library(ELFES_TORCH_PYTHON_LIBRARY torch_python PATHS ${ELFES_TORCH_LIBRARY_DIRS} NO_DEFAULT_PATH REQUIRED)

if(BUILD_CUDA)
    if(ELFES_TORCH_CUDA_VERSION STREQUAL "")
        message(FATAL_ERROR "BUILD_CUDA=ON requires a CUDA-enabled PyTorch installation")
    endif()
    include(CheckLanguage)
    check_language(CUDA)
    if(NOT CMAKE_CUDA_COMPILER)
        message(FATAL_ERROR "BUILD_CUDA=ON requires a CUDA toolkit containing nvcc")
    endif()
    if(NOT DEFINED CMAKE_CUDA_ARCHITECTURES AND NOT DEFINED ENV{CUDAARCHS})
        set(CMAKE_CUDA_ARCHITECTURES native)
    endif()
    enable_language(CUDA)
    find_package(CUDAToolkit REQUIRED)
    string(REGEX MATCH "^[0-9]+" ELFES_TORCH_CUDA_MAJOR "${ELFES_TORCH_CUDA_VERSION}")
    string(REGEX MATCH "^[0-9]+" ELFES_TOOLKIT_CUDA_MAJOR "${CUDAToolkit_VERSION}")
    if(NOT ELFES_TORCH_CUDA_MAJOR STREQUAL ELFES_TOOLKIT_CUDA_MAJOR)
        message(FATAL_ERROR "PyTorch uses CUDA ${ELFES_TORCH_CUDA_VERSION}, but the local toolkit is ${CUDAToolkit_VERSION}; their major versions must match")
    endif()
    find_library(ELFES_C10_CUDA_LIBRARY c10_cuda PATHS ${ELFES_TORCH_LIBRARY_DIRS} NO_DEFAULT_PATH REQUIRED)
    find_library(ELFES_TORCH_CUDA_LIBRARY torch_cuda PATHS ${ELFES_TORCH_LIBRARY_DIRS} NO_DEFAULT_PATH REQUIRED)
endif()

add_subdirectory(src/third_party/gpu-lite)
add_subdirectory(src/third_party/sphericart)
add_subdirectory(src/third_party/tonari)

if(BUILD_CUDA)
    set(ELFES_BUILD_WITH_CUDA True)
else()
    set(ELFES_BUILD_WITH_CUDA False)
endif()
if(ELFES_TORCH_CUDA_VERSION STREQUAL "")
    set(ELFES_BUILD_TORCH_CUDA_VERSION None)
else()
    set(ELFES_BUILD_TORCH_CUDA_VERSION "\"${ELFES_TORCH_CUDA_VERSION}\"")
endif()
configure_file(
    cmake/_build_info.py.in
    "${CMAKE_CURRENT_BINARY_DIR}/_build_info.py"
    @ONLY
)

pybind11_add_module(sphericart MODULE src/elfes/native/sphericart.cpp)
target_compile_features(sphericart PRIVATE cxx_std_17)
target_link_libraries(sphericart PRIVATE elfes_sphericart_core)

install(TARGETS sphericart LIBRARY DESTINATION elfes/native)
install(TARGETS elfes_sphericart_torch LIBRARY DESTINATION elfes/native)
install(TARGETS elfes_sphericart_torch_cuda_stream LIBRARY DESTINATION elfes/native)
install(FILES "${CMAKE_CURRENT_BINARY_DIR}/_build_info.py" DESTINATION elfes/native)
install(
    FILES src/third_party/sphericart/LICENSE
    DESTINATION elfes/native/licenses
    RENAME sphericart-LICENSE
)
install(
    FILES src/third_party/gpu-lite/LICENSE
    DESTINATION elfes/native/licenses
    RENAME gpu-lite-LICENSE
)
install(
    FILES src/third_party/tonari/LICENSE
    DESTINATION elfes/native/licenses
    RENAME tonari-LICENSE
)
