#!/usr/bin/env python3
"""Fail if public provider/model count claims drift from providers.toml."""

from __future__ import annotations

import argparse
import json
import re
import sys
from pathlib import Path

SCRIPT_DIR = Path(__file__).resolve().parent
if str(SCRIPT_DIR) not in sys.path:
    sys.path.insert(0, str(SCRIPT_DIR))

from catalog_counts import catalog_counts  # noqa: E402, I001


REQUIRED_SURFACES = (
    ("README.md", "{providers} LLM providers"),
    ("README.md", "{enabled_chat_models} enabled chat routes"),
    ("README.md", "{cataloged_chat_models} cataloged"),
    ("README.es.md", "{providers} proveedores de LLM"),
    ("README.es.md", "{enabled_chat_models} rutas de chat"),
    ("README.es.md", "{cataloged_chat_models} modelos de chat"),
    ("pyproject.toml", "{providers} LLM providers"),
    ("pyproject.toml", "{enabled_chat_models} enabled chat routes"),
    ("pyproject.toml", "{cataloged_chat_models} cataloged"),
    ("server.json", "{providers} LLM providers"),
    ("docs/index.html", "{providers} cataloged LLM providers"),
    ("docs/index.html", "{enabled_chat_models} enabled chat routes"),
    ("docs/index.html", "{cataloged_chat_models} cataloged"),
    (
        "docs/free-alternative-to-openrouter.html",
        "{providers} cataloged provider groups",
    ),
    ("docs/free-claude-api.html", "over {providers} cataloged chat providers"),
    ("docs/free-llm-api-providers-list.html", "These {providers} providers"),
    ("docs/run-coding-agents-on-free-models.html", "{providers} cataloged provider groups"),
    ("docs/run-opencode-on-free-models.html", "{providers} cataloged provider groups"),
    ("assets/demo.svg", "{providers} cataloged providers"),
    ("assets/demo.svg", "{enabled_chat_models} enabled chat routes"),
    ("assets/social-preview.svg", "{providers} cataloged"),
    ("assets/social-preview.svg", "{enabled_chat_models} enabled chat routes"),
    ("assets/social-preview.svg", "{cataloged_chat_models} cataloged chat models"),
    ("assets/tokenmax-results.svg", "{providers} cataloged providers"),
    ("assets/tokenmax-results.svg", "{enabled_chat_models} enabled chat routes"),
    ("assets/tokenmax-results.svg", "{cataloged_chat_models} cataloged chat models"),
    ("plugins/llm-freellmpool/README.md", "{providers} cataloged providers"),
    ("plugins/llm-freellmpool/pyproject.toml", "{providers} cataloged providers"),
)

EXTERNAL_CONTEXT = (
    "freellmapi",
    "litellm",
    "openrouter free models",
    "openrouter's",
)

PROVIDER_ONLY_CONTEXT = (
    "keyless",
    "no API key",
    "need no API",
    "providers ok",
)


def _format(template: str, counts) -> str:
    return template.format(
        providers=counts.providers,
        live_bucket=counts.live_bucket,
        cataloged_bucket=counts.cataloged_bucket,
        enabled_chat_models=counts.enabled_chat_models,
        cataloged_chat_models=counts.cataloged_chat_models,
    )


def _read(root: Path, rel: str) -> str:
    return (root / rel).read_text(encoding="utf-8")


def _check_required_surfaces(root: Path, counts) -> list[str]:
    errors: list[str] = []
    for rel, template in REQUIRED_SURFACES:
        expected = _format(template, counts)
        if expected not in _read(root, rel):
            errors.append(f"{rel}: missing expected count phrase {expected!r}")
    return errors


def _check_reference_table(root: Path, counts) -> list[str]:
    rel = "docs/free-llm-api-providers-list.html"
    html = _read(root, rel)
    rows = {
        match.group("id"): int(match.group("count"))
        for match in re.finditer(
            r'<tr data-provider="(?P<id>[^"]+)">.*?<td class=num>(?P<count>\d+)</td>',
            html,
        )
    }
    expected = {provider.id: provider.enabled_models for provider in counts.by_provider}
    errors = []
    if set(rows) != set(expected):
        missing = ", ".join(sorted(set(expected) - set(rows))) or "none"
        extra = ", ".join(sorted(set(rows) - set(expected))) or "none"
        errors.append(f"{rel}: provider table ids drifted (missing: {missing}; extra: {extra})")
        return errors
    for provider_id, expected_count in expected.items():
        if rows[provider_id] != expected_count:
            errors.append(
                f"{rel}: {provider_id} count is {rows[provider_id]}, expected {expected_count}"
            )
    return errors


def _check_public_drift(root: Path, counts) -> list[str]:
    errors: list[str] = []
    docs = [
        root / "README.md",
        root / "README.es.md",
        root / "FAQ.md",
        *sorted((root / "assets").glob("*.svg")),
        root / "plugins" / "llm-freellmpool" / "README.md",
        root / "plugins" / "llm-freellmpool" / "pyproject.toml",
        *sorted((root / "docs").glob("**/*")),
    ]
    provider_claim = re.compile(
        r"\b(?P<n>\d+)\s+(?:(?:cataloged|pooled|free|chat|LLM)\s+){0,4}providers\b",
        flags=re.IGNORECASE,
    )
    enabled_route_claim = re.compile(r"\b(?P<n>\d+)\s+enabled\s+chat\s+routes\b")
    cataloged_claim = re.compile(r"\b(?P<n>\d+)\s+cataloged\s+(?:chat\s+)?models\b")
    for path in docs:
        if (
            path.name in {"POLISH_PLAN.md", "POLISH_REPORT.md"}
            or path.name.startswith("MODEL_ACTIVITY_AUDIT_")
            or not path.is_file()
            or path.suffix not in {".html", ".md", ".svg", ".toml"}
        ):
            continue
        rel = path.relative_to(root)
        content = path.read_text(encoding="utf-8")
        checks = (
            (
                provider_claim,
                counts.providers,
                "provider count drift",
                EXTERNAL_CONTEXT + PROVIDER_ONLY_CONTEXT,
            ),
            (
                enabled_route_claim,
                counts.enabled_chat_models,
                "enabled route bucket drift",
                EXTERNAL_CONTEXT,
            ),
            (
                cataloged_claim,
                counts.cataloged_chat_models,
                "cataloged model bucket drift",
                EXTERNAL_CONTEXT,
            ),
        )
        for pattern, expected, label, exemptions in checks:
            for match in pattern.finditer(content):
                context_starts = (
                    content.rfind(". ", 0, match.start()),
                    content.rfind("\n\n", 0, match.start()),
                    content.rfind("<tr", 0, match.start()),
                )
                context_start = max(max(context_starts) + 1, match.start() - 200, 0)
                context_ends = [
                    end + 1
                    for token in (".", "!", "?", "\n")
                    if (end := content.find(token, match.end())) >= 0
                ]
                context_end = min(context_ends, default=match.end())
                claim_context = content[context_start:context_end].lower()
                first_party_provider_claim = (
                    label == "provider count drift"
                    and re.match(
                        r"\s+freellmpool\s+can\s+pool\b",
                        content[match.end() :],
                        flags=re.IGNORECASE,
                    )
                    is not None
                )
                if any(token in claim_context for token in exemptions) and not (
                    first_party_provider_claim
                ):
                    continue
                if int(match.group("n")) == expected:
                    continue
                lineno = content.count("\n", 0, match.start()) + 1
                claim = re.sub(r"\s+", " ", match.group(0)).strip()
                errors.append(f"{rel}:{lineno}: {label}: {claim}")
    return errors


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--root", type=Path, default=Path(__file__).resolve().parent.parent)
    parser.add_argument("--json", action="store_true", help="print derived counts as JSON")
    args = parser.parse_args(argv)

    root = args.root.resolve()
    counts = catalog_counts(root)
    if args.json:
        print(
            json.dumps(
                {
                    "providers": counts.providers,
                    "enabled_chat_models": counts.enabled_chat_models,
                    "cataloged_chat_models": counts.cataloged_chat_models,
                    "live_bucket": counts.live_bucket,
                    "cataloged_bucket": counts.cataloged_bucket,
                    "enabled_all_models": counts.enabled_all_models,
                    "cataloged_all_models": counts.cataloged_all_models,
                    "by_provider": {
                        provider.id: {
                            "enabled_models": provider.enabled_models,
                            "cataloged_models": provider.cataloged_models,
                        }
                        for provider in counts.by_provider
                    },
                },
                sort_keys=True,
            )
        )
        return 0

    errors = [
        *_check_required_surfaces(root, counts),
        *_check_reference_table(root, counts),
        *_check_public_drift(root, counts),
    ]
    if errors:
        print("Count drift detected:", file=sys.stderr)
        for error in errors:
            print(f"  - {error}", file=sys.stderr)
        return 1
    print(
        "Count check passed: "
        f"{counts.providers} providers, "
        f"{counts.enabled_chat_models} enabled chat routes ({counts.live_bucket}), "
        f"{counts.cataloged_chat_models} cataloged chat models ({counts.cataloged_bucket})."
    )
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
