"""
High-Performance Python Bridge to Native C++ SIMD Simulation Core.
Allows Python to orchestrate high-level game rules while delegating the
CPU-intensive Hot Path (100,000+ entities displacement, flat spatial hash, AOI queries)
to compiled C++ AVX2 native code.
"""

import os
import sys
import ctypes
from typing import List, Tuple, Optional


def resolve_sim_core_dll_path(custom_path: Optional[str] = None) -> str:
    """
    Resolves the binary path for freeexile_sim_core dynamic library.
    Checks environment overrides, unified build directories in server_cpp,
    and legacy fallback paths in server/engine_native.
    """
    if custom_path:
        return custom_path

    env_path = os.environ.get("FREEEXILE_SIM_CORE_DLL")
    if env_path:
        return env_path

    repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))

    if sys.platform == "win32":
        lib_names = ["freeexile_sim_core.dll"]
    elif sys.platform == "darwin":
        lib_names = ["libfreeexile_sim_core.dylib", "freeexile_sim_core.dylib"]
    else:
        lib_names = ["libfreeexile_sim_core.so", "freeexile_sim_core.so"]

    search_dirs = [
        os.path.join(repo_root, "server_cpp", "build_ninja"),
        os.path.join(repo_root, "server_cpp", "build", "Release"),
        os.path.join(repo_root, "server_cpp", "build", "Debug"),
        os.path.join(repo_root, "server_cpp", "build"),
        os.path.join(repo_root, "server_cpp", "bin"),
        os.path.join(repo_root, "server_cpp"),
        os.path.join(repo_root, "server", "engine_native"),
    ]

    for directory in search_dirs:
        for name in lib_names:
            candidate = os.path.join(directory, name)
            if os.path.exists(candidate):
                return candidate

    return os.path.join(repo_root, "server_cpp", "build_ninja", lib_names[0])


class FatalNativeEngineInitError(RuntimeError):
    """Raised when native C++ engine binary fails to load and strict mode is active."""
    pass


class NativeEngineBridge:
    def __init__(self, dll_path: Optional[str] = None, strict_mode: Optional[bool] = None):
        self.dll_path = resolve_sim_core_dll_path(dll_path)
        if strict_mode is not None:
            self.strict_mode = strict_mode
        else:
            self.strict_mode = os.environ.get("FREEEXILE_STRICT_NATIVE", "").strip().lower() in {"1", "true", "yes", "on"}
        self.is_loaded = False
        self._lib: Optional[ctypes.CDLL] = None
        self._load_library()

    def _load_library(self) -> None:
        if os.path.exists(self.dll_path):
            try:
                self._lib = ctypes.CDLL(self.dll_path)
                self._setup_function_signatures()
                self.is_loaded = True
            except Exception as e:
                msg = f"[NativeEngineBridge] Warning: Failed to load {self.dll_path}: {e}"
                print(msg)
                self.is_loaded = False
                if self.strict_mode:
                    raise FatalNativeEngineInitError(msg) from e
            else:
                pass
        else:
            msg = f"[NativeEngineBridge] Notice: Native DLL not found at {self.dll_path}. Running in fallback mode."
            print(msg)
            self.is_loaded = False
            if self.strict_mode:
                raise FatalNativeEngineInitError(
                    f"Fatal: Native C++ engine DLL not found at '{self.dll_path}' while FREEEXILE_STRICT_NATIVE is active. "
                    "Halting to prevent 50x performance degradation."
                )

    def _setup_function_signatures(self) -> None:
        if not self._lib:
            return

        # engine_init(float cell_size)
        self._lib.engine_init.argtypes = [ctypes.c_float]
        self._lib.engine_init.restype = None

        # engine_add_entity(int32_t entity_id, float x, float y, float z, float move_speed, float collision_radius, float hp) -> int32_t
        self._lib.engine_add_entity.argtypes = [
            ctypes.c_int32,
            ctypes.c_float,
            ctypes.c_float,
            ctypes.c_float,
            ctypes.c_float,
            ctypes.c_float,
            ctypes.c_float,
        ]
        self._lib.engine_add_entity.restype = ctypes.c_int32

        # engine_set_velocity(int32_t entity_idx, float vx, float vy)
        self._lib.engine_set_velocity.argtypes = [ctypes.c_int32, ctypes.c_float, ctypes.c_float]
        self._lib.engine_set_velocity.restype = None

        # engine_set_evasion(int32_t entity_idx, bool is_evading)
        self._lib.engine_set_evasion.argtypes = [ctypes.c_int32, ctypes.c_bool]
        self._lib.engine_set_evasion.restype = None

        # engine_step_tick(float dt)
        self._lib.engine_step_tick.argtypes = [ctypes.c_float]
        self._lib.engine_step_tick.restype = None

        # engine_query_aoi(float center_x, float center_y, int32_t radius_cells, int32_t* out_entity_ids, int32_t max_results) -> int32_t
        self._lib.engine_query_aoi.argtypes = [
            ctypes.c_float,
            ctypes.c_float,
            ctypes.c_int32,
            ctypes.POINTER(ctypes.c_int32),
            ctypes.c_int32,
        ]
        self._lib.engine_query_aoi.restype = ctypes.c_int32

        # engine_get_entity_count() -> int32_t
        self._lib.engine_get_entity_count.argtypes = []
        self._lib.engine_get_entity_count.restype = ctypes.c_int32

        # engine_get_entity_pos(int32_t idx, float* out_x, float* out_y, float* out_z)
        self._lib.engine_get_entity_pos.argtypes = [
            ctypes.c_int32,
            ctypes.POINTER(ctypes.c_float),
            ctypes.POINTER(ctypes.c_float),
            ctypes.POINTER(ctypes.c_float),
        ]
        self._lib.engine_get_entity_pos.restype = None

    def init(self, cell_size: float = 64.0) -> None:
        if self.is_loaded and self._lib:
            self._lib.engine_init(ctypes.c_float(cell_size))

    def add_entity(
        self,
        entity_id: int,
        x: float,
        y: float,
        z: float = 0.0,
        move_speed: float = 6.0,
        collision_radius: float = 0.5,
        hp: float = 1000.0,
    ) -> int:
        if self.is_loaded and self._lib:
            return self._lib.engine_add_entity(
                ctypes.c_int32(entity_id),
                ctypes.c_float(x),
                ctypes.c_float(y),
                ctypes.c_float(z),
                ctypes.c_float(move_speed),
                ctypes.c_float(collision_radius),
                ctypes.c_float(hp),
            )
        return -1

    def set_velocity(self, entity_idx: int, vx: float, vy: float) -> None:
        if self.is_loaded and self._lib:
            self._lib.engine_set_velocity(
                ctypes.c_int32(entity_idx),
                ctypes.c_float(vx),
                ctypes.c_float(vy),
            )

    def set_evasion(self, entity_idx: int, is_evading: bool) -> None:
        if self.is_loaded and self._lib:
            self._lib.engine_set_evasion(ctypes.c_int32(entity_idx), ctypes.c_bool(is_evading))

    def step_tick(self, dt: float) -> None:
        if self.is_loaded and self._lib:
            self._lib.engine_step_tick(ctypes.c_float(dt))

    def query_aoi(self, center_x: float, center_y: float, radius_cells: int = 1, max_results: int = 512) -> List[int]:
        if not (self.is_loaded and self._lib):
            return []

        out_buffer = (ctypes.c_int32 * max_results)()
        count = self._lib.engine_query_aoi(
            ctypes.c_float(center_x),
            ctypes.c_float(center_y),
            ctypes.c_int32(radius_cells),
            out_buffer,
            ctypes.c_int32(max_results),
        )
        return [out_buffer[i] for i in range(count)]

    def get_entity_pos(self, entity_idx: int) -> Tuple[float, float, float]:
        if not (self.is_loaded and self._lib):
            return (0.0, 0.0, 0.0)

        out_x = ctypes.c_float()
        out_y = ctypes.c_float()
        out_z = ctypes.c_float()
        self._lib.engine_get_entity_pos(
            ctypes.c_int32(entity_idx),
            ctypes.byref(out_x),
            ctypes.byref(out_y),
            ctypes.byref(out_z),
        )
        return (out_x.value, out_y.value, out_z.value)

    def get_entity_count(self) -> int:
        if self.is_loaded and self._lib:
            return self._lib.engine_get_entity_count()
        return 0
