"""Shared types and helpers for model-fetching scripts."""

from __future__ import annotations

import json
import re
import sys
from dataclasses import asdict, dataclass
from typing import TextIO


@dataclass(frozen=True)
class Model:
    """Normalized model metadata used to update client configs."""

    id: str
    name: str
    context_length: int
    max_completion_tokens: int
    supports_reasoning: bool
    supports_image: bool


# Known brand names that need non-trivial capitalisation.
# Lowercase lookup → correct display form.
_BRANDS: dict[str, str] = {
    "deepseek": "DeepSeek",
    "glm": "GLM",
    "gpt": "GPT",
    "gemma": "Gemma",
    "kimi": "Kimi",
    "devstral": "Devstral",
    "llama": "Llama",
    "mimo": "MiMo",
    "minimax": "MiniMax",
    "qwen": "Qwen",
    "qwen3": "Qwen3",
    "claude": "Claude",
    "gemini": "Gemini",
    "google": "Google",
    "nvidia": "NVIDIA",
    "nemotron": "Nemotron",
}

# Pure-alpha tokens that should be all-caps.
_ACRONYMS: set[str] = {"oss", "ar"}

# Alpha+digit tokens that should be all-caps.
_MIXED_ACRONYMS: set[str] = {"fp8", "int4", "nvfp4"}


def _split_brand_version(token: str) -> list[str]:
    """Split a brand-version compound like 'Qwen3.6' into ['Qwen', '3.6'].

    Only splits when the leading letters match a known brand **and** the
    version contains a decimal point.  This avoids splitting 'Qwen3'
    (no decimal) which should stay together in forms like 'Qwen3 Coder'.
    """
    m = re.fullmatch(r"([a-zA-Z]{2,})(\d+\.\d+)", token)
    if m and m.group(1).lower() in _BRANDS:
        return [m.group(1), m.group(2)]
    return [token]


def _clean_token(tok: str) -> str:
    low = tok.lower()

    if low in _BRANDS:
        return _BRANDS[low]
    if low in _ACRONYMS:
        return tok.upper()
    if low in _MIXED_ACRONYMS:
        return tok.upper()
    if tok.isalpha() and tok.isupper() and len(tok) >= 2:
        return tok
    if re.fullmatch(r"\d+(?:\.\d+)?", tok):
        return tok

    # '26b' → '26B', '128e' → '128E'
    m = re.fullmatch(r"(\d+)([a-zA-Z]+)", tok)
    if m:
        return m.group(1) + m.group(2).upper()

    # 'v4' → 'V4', 'k2.5' → 'K2.5', 'a4b' → 'A4B', 'a35b' → 'A35B'
    m = re.fullmatch(r"([a-zA-Z]+)(\d+(?:\.\d+)?)([a-zA-Z]*)", tok)
    if m:
        prefix = m.group(1).capitalize()
        nums = m.group(2)
        suffix = m.group(3).upper() if m.group(3) else ""
        return prefix + nums + suffix

    if tok.isalpha():
        return tok.capitalize()

    return tok


def clean_name(raw: str) -> str:
    """Normalise a display name from an upstream catalog.

    - Strip redundant vendor prefixes ('DeepSeek: DeepSeek V4 Pro' → 'DeepSeek V4 Pro')
    - Remove parentheses ('(fast)' → 'fast')
    - Replace hyphens with spaces
    - Split brand-version compounds ('Qwen3.6' → 'Qwen 3.6')
    - Normalise capitalisation for known brands and acronyms
    """
    if ":" in raw:
        raw = raw.split(":")[-1].strip()
    raw = re.sub(r"\(([^)]+)\)", r"\1", raw)
    raw = raw.replace("-", " ")

    tokens: list[str] = []
    for tok in raw.split():
        tokens.extend(_split_brand_version(tok))

    return " ".join(_clean_token(t) for t in tokens)


def emit_tsv(models: list[Model], stream: TextIO) -> None:
    """Write models as TSV: id, name, ctx, max_tokens, reasoning, image."""
    for m in models:
        reasoning = "reasoning" if m.supports_reasoning else ""
        image = "image" if m.supports_image else ""
        stream.write(
            f"{m.id}\t{m.name}\t{m.context_length}\t{m.max_completion_tokens}"
            f"\t{reasoning}\t{image}\n"
        )


def emit_json(models: list[Model], stream: TextIO) -> None:
    """Write models as a JSON array (indent=2, trailing newline)."""
    json.dump([asdict(m) for m in models], stream, indent=2)
    stream.write("\n")
