# Copyright (c) 2026 Tobias Karusseit
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.

cmake_minimum_required(VERSION 3.18)

# Prefer pyproject metadata when driven by scikit-build-core.
if(DEFINED SKBUILD_PROJECT_NAME)
    project(${SKBUILD_PROJECT_NAME} VERSION ${SKBUILD_PROJECT_VERSION} LANGUAGES CXX)
else()
    project(cthreads LANGUAGES CXX)
endif()

set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
set(CMAKE_CXX_EXTENSIONS OFF)
set(CMAKE_POSITION_INDEPENDENT_CODE ON)

# --- Python + pybind11 -------------------------------------------------------
find_package(Python COMPONENTS Interpreter Development.Module REQUIRED)

find_package(pybind11 CONFIG QUIET)
if(NOT pybind11_FOUND)
    include(FetchContent)
    FetchContent_Declare(
        pybind11
        GIT_REPOSITORY https://github.com/pybind/pybind11.git
        GIT_TAG        v2.13.6
    )
    FetchContent_MakeAvailable(pybind11)
endif()

# --- Extension module: import cthreads._ext ----------------------------------
set(_CTHREADS_LINALG_SOURCES
    impl/linalg/array.cpp
    impl/linalg/tiling.cpp
)

pybind11_add_module(_ext MODULE
    bindings/module.cpp
    ${_CTHREADS_LINALG_SOURCES}
)

target_include_directories(_ext PRIVATE
    ${CMAKE_CURRENT_SOURCE_DIR}/headers
)

# Enable AVX2/FMA so Array float/double kernels hit the SIMD path on supported CPUs.
if(MSVC)
    target_compile_options(_ext PRIVATE /arch:AVX2)
else()
    target_compile_options(_ext PRIVATE -mavx2 -mfma)
endif()

if(UNIX AND NOT APPLE)
    target_link_libraries(_ext PRIVATE ${CMAKE_DL_LIBS})
endif()

# --- Standalone linalg micro-bench (dot + matmul, light/medium/heavy) --------
set(_CTHREADS_BENCH_SRC
    "${CMAKE_CURRENT_SOURCE_DIR}/../../../demo/bench/cpp/linalg_bench.cpp"
)
if(EXISTS "${_CTHREADS_BENCH_SRC}")
    add_executable(linalg_bench
        ${_CTHREADS_BENCH_SRC}
        ${_CTHREADS_LINALG_SOURCES}
    )
    target_include_directories(linalg_bench PRIVATE
        ${CMAKE_CURRENT_SOURCE_DIR}/headers
    )
    if(MSVC)
        target_compile_options(linalg_bench PRIVATE /arch:AVX2 /O2)
    else()
        target_compile_options(linalg_bench PRIVATE -mavx2 -mfma -O3)
    endif()
    set_target_properties(linalg_bench PROPERTIES
        RUNTIME_OUTPUT_DIRECTORY "${CMAKE_CURRENT_SOURCE_DIR}/../../../demo/bench/cpp"
        RUNTIME_OUTPUT_DIRECTORY_RELEASE "${CMAKE_CURRENT_SOURCE_DIR}/../../../demo/bench/cpp"
        RUNTIME_OUTPUT_DIRECTORY_DEBUG "${CMAKE_CURRENT_SOURCE_DIR}/../../../demo/bench/cpp"
        RUNTIME_OUTPUT_DIRECTORY_RELWITHDEBINFO "${CMAKE_CURRENT_SOURCE_DIR}/../../../demo/bench/cpp"
    )
    message(STATUS "linalg_bench: ${_CTHREADS_BENCH_SRC}")
endif()

# Python package dir (sibling of this cpp/ tree). Never use cpp/ as output/build.
set(_cthreads_py_out "${CMAKE_CURRENT_SOURCE_DIR}/../python/cthreads")

# Bare cmake (no pip): build straight into the package tree.
if(NOT DEFINED SKBUILD)
    set_target_properties(_ext PROPERTIES
        LIBRARY_OUTPUT_DIRECTORY          "${_cthreads_py_out}"
        RUNTIME_OUTPUT_DIRECTORY          "${_cthreads_py_out}"
        LIBRARY_OUTPUT_DIRECTORY_DEBUG    "${_cthreads_py_out}"
        LIBRARY_OUTPUT_DIRECTORY_RELEASE  "${_cthreads_py_out}"
        LIBRARY_OUTPUT_DIRECTORY_RELWITHDEBINFO "${_cthreads_py_out}"
        RUNTIME_OUTPUT_DIRECTORY_DEBUG    "${_cthreads_py_out}"
        RUNTIME_OUTPUT_DIRECTORY_RELEASE  "${_cthreads_py_out}"
        RUNTIME_OUTPUT_DIRECTORY_RELWITHDEBINFO "${_cthreads_py_out}"
    )
    message(STATUS "Ext out: ${_cthreads_py_out}")
endif()

# pip/scikit-build: keep a copy beside Python sources. Editable redirect loads
# the package from this tree; PathFinder then picks up _ext next to __init__.py.
if(DEFINED SKBUILD)
    add_custom_command(
        TARGET _ext POST_BUILD
        COMMAND ${CMAKE_COMMAND} -E make_directory "${_cthreads_py_out}"
        COMMAND ${CMAKE_COMMAND} -E copy_if_different
            "$<TARGET_FILE:_ext>"
            "${_cthreads_py_out}/$<TARGET_FILE_NAME:_ext>"
        COMMENT "Copy _ext into Python package directory"
    )
endif()

# Wheels: install into platlib/cthreads.
# Editable: install into NULL so we do NOT register a redirect mapping to a
# missing/ephemeral path (which breaks `import cthreads._ext` on Windows).
if(DEFINED SKBUILD_STATE AND SKBUILD_STATE STREQUAL "editable")
    install(TARGETS _ext
        LIBRARY DESTINATION "${SKBUILD_NULL_DIR}"
        RUNTIME DESTINATION "${SKBUILD_NULL_DIR}"
    )
else()
    install(TARGETS _ext
        LIBRARY DESTINATION cthreads
        RUNTIME DESTINATION cthreads
    )
endif()

message(STATUS "Python:  ${Python_EXECUTABLE} (${Python_VERSION})")
if(DEFINED SKBUILD_STATE)
    message(STATUS "SKBUILD_STATE: ${SKBUILD_STATE}")
endif()
