from __future__ import annotations

import json
import logging
import os
from typing import Any

import json_repair

from app.core.ai.domain.models import (
    AiCapability,
    AiModelSpec,
    AiProviderType,
    AiResponse,
    AiStructuredRequest,
    AiTextRequest,
    AiVisionRequest,
)
from app.core.ai.domain.ports import AiGatewayPort
from app.core.ai.infrastructure.adapters.antigravity_adapter import (
    AntigravitySdkAdapter,
)
from app.core.ai.infrastructure.adapters.lmstudio_adapter import LmStudioAdapter
from app.core.ai.infrastructure.adapters.openrouter_adapter import (
    OpenRouterAdapter,
)
from app.core.ai.infrastructure.circuit_breaker import ProviderCircuitBreaker
from app.core.ai.infrastructure.registry import CanonicalModelRegistry
from app.core.settings import Settings, get_settings

logger = logging.getLogger("dscons.ai.gateway")


class EnterpriseAiGateway(AiGatewayPort):
    """Hexagonal Architecture Gateway Orchestrator with Plug-and-Play Resilience."""

    def __init__(self, settings: Settings | None = None) -> None:
        self._settings = settings or get_settings()
        self._circuit_breaker = ProviderCircuitBreaker()
        self._antigravity_adapter = AntigravitySdkAdapter(
            api_key=getattr(self._settings, "gemini_api_key", "")
            or os.environ.get("GEMINI_API_KEY", "")
        )
        self._openrouter_adapter = OpenRouterAdapter(
            api_key=getattr(self._settings, "openrouter_api_key", "")
            or os.environ.get("OPENROUTER_API_KEY", ""),
            base_url=getattr(
                self._settings,
                "openrouter_base_url",
                "https://openrouter.ai/api/v1",
            ),
        )
        self._lmstudio_adapter = LmStudioAdapter(
            base_url=getattr(
                self._settings, "lmstudio_base_url", "http://127.0.0.1:1234/v1"
            ),
            api_key=getattr(self._settings, "lmstudio_api_key", "lm-studio"),
        )

    def _get_adapter(self, provider: AiProviderType) -> Any:
        """Resolves the appropriate adapter instance."""
        if provider in (AiProviderType.ANTIGRAVITY, AiProviderType.GEMINI_DIRECT):
            return self._antigravity_adapter
        elif provider == AiProviderType.LMSTUDIO:
            return self._lmstudio_adapter
        return self._openrouter_adapter

    async def generate_text(self, request: AiTextRequest) -> AiResponse:
        """Executes text completion with automatic failover."""
        target_model = request.canonical_model_id or getattr(
            self._settings, "llm_model_name", "openrouter/gemini-2.5-flash"
        )
        if (
            request.preferred_provider == AiProviderType.ANTIGRAVITY
            and "gemini" in target_model.lower()
        ):
            target_model = "openrouter/gemini-2.5-flash"
        elif request.preferred_provider == AiProviderType.LMSTUDIO:
            target_model = "lfm-2.5-2.6b"
        elif (
            request.preferred_provider == AiProviderType.OPENROUTER
            and not target_model.startswith("openrouter/")
        ):
            if f"openrouter/{target_model}" in CanonicalModelRegistry._MODELS:
                target_model = f"openrouter/{target_model}"

        spec = CanonicalModelRegistry.resolve_spec(target_model)

        attempts = [spec]
        fb_spec = CanonicalModelRegistry.get_fallback_spec(spec)
        while fb_spec and fb_spec not in attempts:
            attempts.append(fb_spec)
            fb_spec = CanonicalModelRegistry.get_fallback_spec(fb_spec)

        last_error = None
        for current_spec in attempts:
            provider = current_spec.provider_type
            if not self._circuit_breaker.can_attempt(provider):
                logger.warning(
                    "[AI_GATEWAY] Circuit is OPEN for %s, skipping to fallback.",
                    provider.value,
                )
                continue

            adapter = self._get_adapter(provider)
            try:
                response = await adapter.generate_text(request, current_spec)
                self._circuit_breaker.record_success(provider)
                if current_spec != spec:
                    response.is_fallback = True
                return response
            except Exception as e:
                last_error = e
                self._circuit_breaker.record_failure(provider, e)
                logger.warning(
                    "[AI_GATEWAY] Text generation with %s (%s) failed: %s. Initiating resilient failover...",
                    current_spec.canonical_id,
                    provider.value,
                    e,
                )

        if last_error:
            raise last_error
        raise RuntimeError("All configured AI providers failed text generation.")

    async def analyze_vision(self, request: AiVisionRequest) -> AiResponse:
        """Executes multimodal vision analysis prioritizing native Antigravity SDK."""
        spec = CanonicalModelRegistry.resolve_spec(
            request.canonical_model_id or "openrouter/gemini-2.5-flash"
        )

        attempts = [spec]
        fb_spec = CanonicalModelRegistry.get_fallback_spec(spec)
        while fb_spec and fb_spec not in attempts:
            if AiCapability.MULTIMODAL_VISION in fb_spec.capabilities:
                attempts.append(fb_spec)
            fb_spec = CanonicalModelRegistry.get_fallback_spec(fb_spec)

        last_error = None
        for current_spec in attempts:
            provider = current_spec.provider_type
            if not self._circuit_breaker.can_attempt(provider):
                logger.warning(
                    "[AI_GATEWAY] Circuit is OPEN for %s, bypassing vision call.",
                    provider.value,
                )
                continue

            adapter = self._get_adapter(provider)
            try:
                response = await adapter.analyze_vision(request, current_spec)
                self._circuit_breaker.record_success(provider)
                if current_spec != spec:
                    response.is_fallback = True
                return response
            except Exception as e:
                last_error = e
                self._circuit_breaker.record_failure(provider, e)
                logger.warning(
                    "[AI_GATEWAY] Vision analysis with %s (%s) failed: %s. Attempting fallback...",
                    current_spec.canonical_id,
                    provider.value,
                    e,
                )

        if last_error:
            raise last_error
        raise RuntimeError("All vision-capable AI providers failed.")

    async def generate_structured(
        self, request: AiStructuredRequest
    ) -> dict[str, Any]:
        """Executes structured JSON inference with robust schema guarantees."""
        schema_hint = ""
        if request.schema_dict:
            schema_hint = f"\nCẤU TRÚC JSON BẮT BUỘC:\n{json.dumps(request.schema_dict, ensure_ascii=False, indent=2)}\nChỉ trả về JSON thuần túy."

        text_req = AiTextRequest(
            prompt=f"{request.prompt}\n{schema_hint}",
            system_prompt=request.system_prompt,
            canonical_model_id=request.canonical_model_id,
            preferred_provider=request.preferred_provider,
            agent_code=request.agent_code,
            temperature=request.temperature,
        )

        resp = await self.generate_text(text_req)
        cleaned = resp.content.strip()
        cleaned = cleaned.removeprefix("```json").removeprefix("```").removesuffix("```").strip()

        parsed = json_repair.repair_json(cleaned, return_objects=True)
        if isinstance(parsed, dict):
            return parsed

        try:
            val = json.loads(cleaned)
            if isinstance(val, dict):
                return val
        except Exception:
            pass

        return {"raw_content": resp.content}

    def generate_text_sync(
        self, request: AiTextRequest, timeout: float = 60.0
    ) -> AiResponse:
        """Synchronous wrapper for generate_text with thread safety across event loops."""
        import asyncio
        import concurrent.futures

        async def _run() -> AiResponse:
            return await self.generate_text(request)

        try:
            loop = asyncio.get_running_loop()
        except RuntimeError:
            loop = None

        if loop and loop.is_running():
            with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
                return pool.submit(asyncio.run, _run()).result(timeout=timeout)
        else:
            return asyncio.run(_run())

    def analyze_vision_sync(
        self, request: AiVisionRequest, timeout: float = 60.0
    ) -> AiResponse:
        """Synchronous wrapper for analyze_vision with thread safety across event loops."""
        import asyncio
        import concurrent.futures

        async def _run() -> AiResponse:
            return await self.analyze_vision(request)

        try:
            loop = asyncio.get_running_loop()
        except RuntimeError:
            loop = None

        if loop and loop.is_running():
            with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
                return pool.submit(asyncio.run, _run()).result(timeout=timeout)
        else:
            return asyncio.run(_run())


_GLOBAL_AI_GATEWAY: EnterpriseAiGateway | None = None


def get_ai_gateway() -> EnterpriseAiGateway:
    """Returns singleton Enterprise AI Gateway instance."""
    global _GLOBAL_AI_GATEWAY
    if _GLOBAL_AI_GATEWAY is None:
        _GLOBAL_AI_GATEWAY = EnterpriseAiGateway()
    return _GLOBAL_AI_GATEWAY
