from __future__ import annotations

import base64
import json
import urllib.error
import urllib.request
from dataclasses import dataclass


class OllamaError(RuntimeError):
    pass


@dataclass(frozen=True)
class OllamaResult:
    text: str
    eval_count: int
    total_duration_ns: int


class OllamaClient:
    def __init__(self, base_url: str, model: str, timeout: int) -> None:
        self.base_url = base_url.rstrip("/")
        self.model = model
        self.timeout = timeout

    def _request(self, path: str, body: dict | None = None) -> dict:
        data = None if body is None else json.dumps(body).encode("utf-8")
        request = urllib.request.Request(
            f"{self.base_url}{path}",
            data=data,
            headers={"Content-Type": "application/json"},
            method="GET" if body is None else "POST",
        )
        try:
            with urllib.request.urlopen(request, timeout=self.timeout) as response:
                payload = response.read()
        except (urllib.error.URLError, TimeoutError) as exc:
            raise OllamaError("Ollama indisponivel ou excedeu o tempo limite") from exc
        try:
            decoded = json.loads(payload)
        except json.JSONDecodeError as exc:
            raise OllamaError("Ollama devolveu JSON invalido") from exc
        if not isinstance(decoded, dict):
            raise OllamaError("Ollama devolveu uma resposta inesperada")
        return decoded

    def health(self) -> dict:
        version = self._request("/api/version")
        tags = self._request("/api/tags")
        models = [item.get("name", "") for item in tags.get("models", []) if isinstance(item, dict)]
        loaded = []
        try:
            running = self._request("/api/ps")
            loaded = [item for item in running.get("models", []) if isinstance(item, dict)]
        except OllamaError:
            # Versoes antigas do Ollama podem nao disponibilizar /api/ps. Isso nao
            # impede inferencia; apenas deixa o estado da GPU como desconhecido.
            loaded = []
        active = next(
            (item for item in loaded if item.get("name", "") == self.model),
            None,
        )
        return {
            "version": version.get("version", ""),
            "model_available": self.model in models,
            "models_count": len(models),
            "gpu": {
                "model_loaded": active is not None,
                "vram_bytes": int(active.get("size_vram", 0) or 0) if active else 0,
            },
        }

    def load(self) -> None:
        """Carrega o modelo sem inferencia e mantem-no residente no Ollama."""
        self._request(
            "/api/generate",
            {"model": self.model, "prompt": "", "stream": False, "keep_alive": -1},
        )

    def transcribe(self, image: bytes, prompt: str, document_type: str = "") -> OllamaResult:
        model_name = self.model.lower()
        if model_name.startswith("glm-ocr"):
            options = {
                "temperature": 0,
                "top_p": 0.00001,
                "top_k": 1,
                "num_predict": 2400,
                "repeat_penalty": 1.1,
                # O pacote Ollama glm-ocr pode repetir blocos Markdown vazios ate
                # num_predict depois de ja ter devolvido a transcricao correcta.
                "stop": ["```markdown", "```\n```"],
            }
        else:
            options = {
                "temperature": 0,
                "num_predict": 1400,
                "repeat_penalty": 1.15,
                "repeat_last_n": 256,
            }
            if document_type.strip().lower() in {"cc_frente_identificacao", "rc_validade_focada"}:
                # Esta tarefa devolve no maximo duas linhas. Um limite curto evita os
                # loops de repeticao observados nas fotografias de cartoes sem alterar
                # as opcoes das transcricoes gerais do modelo congelado.
                options = {
                    "temperature": 0,
                    "num_predict": 160,
                    "repeat_penalty": 1.05,
                    "repeat_last_n": 64,
                }
        encoded_image = base64.b64encode(image).decode("ascii")
        payload = self._request(
            "/api/generate",
            {
                "model": self.model,
                "prompt": prompt,
                "images": [encoded_image],
                "stream": False,
                "think": False,
                "keep_alive": -1,
                "options": options,
            },
        )
        text = str(payload.get("response", "")).strip()
        if model_name.startswith("glm-ocr"):
            for marker in ("\n```markdown", "\n```\n```"):
                if marker in text:
                    text = text.split(marker, 1)[0].rstrip()
        if text == "[NAO_LEGIVEL]":
            text = ""
        return OllamaResult(
            text=text,
            eval_count=int(payload.get("eval_count", 0) or 0),
            total_duration_ns=int(payload.get("total_duration", 0) or 0),
        )
