# CMakeLists.txt for the control module

# Find required dependencies
find_package(Eigen3 REQUIRED)
find_package(Python3 COMPONENTS Interpreter Development.Module REQUIRED)
find_package(Python COMPONENTS Interpreter Development.Module QUIET)

find_package(pybind11 REQUIRED)
set(MAP_LOADER_DEPS)
# Keep artifacts in the active build tree so different entrypoints do not
# overwrite each other's outputs in the source tree.
set(CMAKE_LIBRARY_OUTPUT_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR}/lib)
set(CMAKE_ARCHIVE_OUTPUT_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR}/lib)
set(CMAKE_RUNTIME_OUTPUT_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR}/lib)

if(Torch_FOUND)
    list(APPEND MAP_LOADER_DEPS "${TORCH_LIBRARIES}")
endif()

if(Torch_FOUND)
  set(MPLOAD_SHARED_RPATH "/usr/local/lib")
  if(DEFINED STANDARD_BOTS_LIBTORCH_ROOT)
    if(EXISTS "${STANDARD_BOTS_LIBTORCH_ROOT}/lib")
      list(APPEND MPLOAD_SHARED_RPATH "${STANDARD_BOTS_LIBTORCH_ROOT}/lib")
    endif()
  endif()
  if(DEFINED Torch_INSTALL_PREFIX AND Torch_INSTALL_PREFIX)
    if(EXISTS "${Torch_INSTALL_PREFIX}/lib")
      list(APPEND MPLOAD_SHARED_RPATH "${Torch_INSTALL_PREFIX}/lib")
    endif()
  endif()

  if(Python3_FOUND AND Python3_EXECUTABLE)
    execute_process(
      COMMAND ${Python3_EXECUTABLE} -c "import torch, pathlib; print(pathlib.Path(torch.__file__).resolve().parent)"
      OUTPUT_VARIABLE _torch_site_dir
      OUTPUT_STRIP_TRAILING_WHITESPACE
      ERROR_QUIET)
    if(_torch_site_dir)
      if(EXISTS "${_torch_site_dir}/lib")
        list(APPEND MPLOAD_SHARED_RPATH "${_torch_site_dir}/lib")
      endif()
      get_filename_component(_torch_site_parent "${_torch_site_dir}/.." REALPATH)
      if(EXISTS "${_torch_site_parent}/torch.libs")
        list(APPEND MPLOAD_SHARED_RPATH "${_torch_site_parent}/torch.libs")
      endif()
    endif()
  endif()

  file(GLOB _torch_lib_folders
       "${PROJECT_SOURCE_DIR}/venv/lib/python*/site-packages/torch.libs"
       "${PROJECT_SOURCE_DIR}/../../venv/lib/python*/site-packages/torch.libs")
  foreach(_torch_lib_dir IN LISTS _torch_lib_folders)
    if(IS_DIRECTORY "${_torch_lib_dir}")
      list(APPEND MPLOAD_SHARED_RPATH "${_torch_lib_dir}")
    endif()
  endforeach()

  list(REMOVE_DUPLICATES MPLOAD_SHARED_RPATH)
  string(REPLACE ";" ":" MPLOAD_SHARED_RPATH_STRING "${MPLOAD_SHARED_RPATH}")
  message(STATUS "Map loader module RPATH: ${MPLOAD_SHARED_RPATH_STRING}")
endif()

# Torch-free core used by the base_shaper extension.
add_library(base_shaper_core
    include/base_shaper.cpp
)

set_target_properties(base_shaper_core PROPERTIES
    POSITION_INDEPENDENT_CODE ON
)

target_include_directories(base_shaper_core
    PUBLIC
        ${CMAKE_CURRENT_SOURCE_DIR}
        ${EIGEN3_INCLUDE_DIR}
)

target_compile_features(base_shaper_core PRIVATE cxx_std_17)

target_link_libraries(base_shaper_core
    PUBLIC
        pybind11::headers
)

add_library(shaper_interface_dynamics_adapter
    include/shaper_interface_dynamics.cpp
)

set_target_properties(shaper_interface_dynamics_adapter PROPERTIES
    POSITION_INDEPENDENT_CODE ON
)

target_include_directories(shaper_interface_dynamics_adapter
    PUBLIC
        ${CMAKE_CURRENT_SOURCE_DIR}
        ${EIGEN3_INCLUDE_DIR}
)

target_compile_features(
    shaper_interface_dynamics_adapter PRIVATE cxx_std_17
)

target_link_libraries(shaper_interface_dynamics_adapter
    PUBLIC ReforgeShaper::dynamics
)

if(Torch_FOUND)
  # Torch-enabled implementation used only by the legacy shaper extension.
  add_library(shaper_interface_core
      include/shaper_interface.cpp
      include/map_loader.cpp
  )

  set_target_properties(shaper_interface_core PROPERTIES
      POSITION_INDEPENDENT_CODE ON
  )

  target_include_directories(shaper_interface_core
      PUBLIC
          ${CMAKE_CURRENT_SOURCE_DIR}
          ${EIGEN3_INCLUDE_DIR}
          ${TORCH_INCLUDE_DIRS}
      PRIVATE
          ${Python3_INCLUDE_DIRS}
  )

  target_compile_features(shaper_interface_core PRIVATE cxx_std_17)

  target_link_libraries(shaper_interface_core
      PUBLIC
          base_shaper_core
          pybind11::headers
          shaper_interface_dynamics_adapter
          ${MAP_LOADER_DEPS}
  )
endif()

# Create Python module
pybind11_add_module(base_shaper
    include/base_shaper_bindings.cpp
)

target_link_libraries(base_shaper
    PRIVATE
        base_shaper_core
)

# Set include directories for Python module
target_include_directories(base_shaper
    PRIVATE
        ${CMAKE_CURRENT_SOURCE_DIR}
        ${EIGEN3_INCLUDE_DIR}
)

# Set compile options for Python module
target_compile_features(base_shaper PRIVATE cxx_std_17)

if(Torch_FOUND)
  pybind11_add_module(shaper_interface
      include/shaper_interface_binding.cpp
  )

  target_link_libraries(shaper_interface
      PRIVATE
          shaper_interface_core
  )

  target_include_directories(shaper_interface
      PRIVATE
          ${CMAKE_CURRENT_SOURCE_DIR}
          ${EIGEN3_INCLUDE_DIR}
  )

  target_compile_features(shaper_interface PRIVATE cxx_std_17)
  set_target_properties(shaper_interface PROPERTIES
      BUILD_RPATH "${MPLOAD_SHARED_RPATH}"
      INSTALL_RPATH "${MPLOAD_SHARED_RPATH}"
      INSTALL_RPATH_USE_LINK_PATH ON
  )
endif()

if(BUILD_TESTING)
  add_executable(shaper_interface_dynamics_smoke_tests
      tests/test_shaper_interface_dynamics.cpp
  )
  target_compile_features(
      shaper_interface_dynamics_smoke_tests PRIVATE cxx_std_17
  )
  target_compile_definitions(
      shaper_interface_dynamics_smoke_tests
      PRIVATE
          REFORGE_DYNAMICS_FIXTURE_DIR="${CMAKE_CURRENT_LIST_DIR}/../../native/tests/fixtures/urdf"
  )
  target_link_libraries(
      shaper_interface_dynamics_smoke_tests
      PRIVATE shaper_interface_dynamics_adapter
  )
  add_test(
      NAME shaper_interface_dynamics_smoke_tests
      COMMAND shaper_interface_dynamics_smoke_tests
  )
endif()

install(TARGETS base_shaper
    LIBRARY DESTINATION reforge_core/control
    RUNTIME DESTINATION reforge_core/control
)

if(TARGET shaper_interface)
  install(TARGETS shaper_interface
      LIBRARY DESTINATION reforge_core/control
      RUNTIME DESTINATION reforge_core/control
  )
endif()

if(NOT SKBUILD)
  install(FILES
      include/base_shaper.hpp
      include/shaper_interface.cpp
      include/map_loader.cpp
      DESTINATION reforge_core/control/include
  )
endif()
