from __future__ import annotations

import json
import re
import time
from typing import Any

import httpx

from .config import Settings
from . import langfuse_client as lf
from .langfuse_client import get_current_trace_id


class BaseProvider:
    name = "heuristic"

    async def complete(self, system_prompt: str, user_prompt: str) -> str:
        return ""


class AnthropicProvider(BaseProvider):
    name = "anthropic"

    def __init__(self, api_key: str, model: str) -> None:
        self.api_key = api_key
        self.model = model

    async def complete(self, system_prompt: str, user_prompt: str) -> str:
        headers = {
            "x-api-key": self.api_key,
            "anthropic-version": "2023-06-01",
            "content-type": "application/json",
        }
        payload = {
            "model": self.model,
            "max_tokens": 1400,
            "temperature": 0.2,
            "system": system_prompt,
            "messages": [
                {
                    "role": "user",
                    "content": user_prompt,
                }
            ],
        }
        t0 = time.monotonic()
        async with httpx.AsyncClient(timeout=45) as client:
            response = await client.post("https://api.anthropic.com/v1/messages", headers=headers, json=payload)
            response.raise_for_status()
            data = response.json()
        latency_ms = (time.monotonic() - t0) * 1000
        chunks = data.get("content", [])
        texts = [chunk.get("text", "") for chunk in chunks if chunk.get("type") == "text"]
        result = "\n".join(filter(None, texts)).strip()
        # Langfuse: capturar tokens nativos de Anthropic
        usage = data.get("usage", {})
        lf.record_generation(
            trace_name="anthropic.complete",
            model=self.model,
            provider_name=self.name,
            system_prompt=system_prompt,
            user_prompt=user_prompt,
            output=result,
            latency_ms=latency_ms,
            input_tokens=usage.get("input_tokens"),
            output_tokens=usage.get("output_tokens"),
            trace_id=get_current_trace_id(),
        )
        return result


class OllamaProvider(BaseProvider):
    name = "ollama"

    def __init__(self, base_url: str, model: str) -> None:
        self.base_url = base_url.rstrip("/")
        self.model = model

    async def complete(self, system_prompt: str, user_prompt: str) -> str:
        payload = {
            "model": self.model,
            "stream": False,
            "messages": [
                {"role": "system", "content": system_prompt},
                {"role": "user", "content": user_prompt},
            ],
        }
        t0 = time.monotonic()
        async with httpx.AsyncClient(timeout=60) as client:
            response = await client.post(f"{self.base_url}/api/chat", json=payload)
            response.raise_for_status()
            data = response.json()
        latency_ms = (time.monotonic() - t0) * 1000
        result = (data.get("message") or {}).get("content", "").strip()
        # Langfuse: Ollama devuelve prompt_eval_count y eval_count en la respuesta
        lf.record_generation(
            trace_name="ollama.complete",
            model=self.model,
            provider_name=self.name,
            system_prompt=system_prompt,
            user_prompt=user_prompt,
            output=result,
            latency_ms=latency_ms,
            input_tokens=data.get("prompt_eval_count"),
            output_tokens=data.get("eval_count"),
            trace_id=get_current_trace_id(),
        )
        return result


class HeuristicProvider(BaseProvider):
    name = "heuristic"


_JSON_BLOCK_RE = re.compile(r"```(?:json)?\s*(\{.*?\})\s*```", re.DOTALL)


def extract_json_document(raw: str) -> dict[str, Any] | None:
    if not raw:
        return None
    candidate = raw.strip()
    fenced = _JSON_BLOCK_RE.search(candidate)
    if fenced:
        candidate = fenced.group(1).strip()
    start = candidate.find("{")
    end = candidate.rfind("}")
    if start != -1 and end != -1:
        candidate = candidate[start : end + 1]
    try:
        return json.loads(candidate)
    except json.JSONDecodeError:
        return None


def build_provider(settings: Settings) -> BaseProvider:
    if settings.provider_name == "anthropic" and settings.anthropic_api_key:
        return AnthropicProvider(settings.anthropic_api_key, settings.anthropic_model)
    if settings.provider_name == "ollama":
        return OllamaProvider(settings.ollama_base_url, settings.ollama_model)
    return HeuristicProvider()
