Kansho/backend/model_catalog.py
2026-08-26 10:42:20 +02:00

252 lines
8.2 KiB
Python

"""Cached model window metadata for journal budget checks.
OpenRouter exposes context_length and top_provider.max_completion_tokens via
GET /api/v1/models. The catalog is cached so journal calls do not refetch on
every generate. Missing metadata fails closed for journal generation.
"""
from __future__ import annotations
import os
import threading
import time
from urllib.parse import urlparse
import httpx
from prompt_budget import (
ERROR_MODEL_METADATA_UNKNOWN,
JOURNAL_TARGET_COMPLETION_TOKENS,
JournalBudgetError,
ModelWindow,
_int_env,
)
DEFAULT_TTL_SECONDS = 3600
FAKE_CONTEXT = 32_768
_lock = threading.Lock()
# url -> (fetched_at, {model_id: ModelWindow fields as dict})
_catalog: dict[str, tuple[float, dict[str, dict]]] = {}
_overrides: dict[str, ModelWindow] = {}
def catalog_ttl() -> int:
raw = (os.environ.get("KANSHO_MODEL_CATALOG_TTL_SECONDS") or "").strip()
if not raw:
return DEFAULT_TTL_SECONDS
try:
value = int(raw)
except ValueError:
return DEFAULT_TTL_SECONDS
return max(60, value)
def reset_catalog() -> None:
with _lock:
_catalog.clear()
_overrides.clear()
def set_metadata_override(model: str, window: ModelWindow) -> None:
with _lock:
_overrides[model] = window
def is_openrouter_url(url: str) -> bool:
host = (urlparse(url or "").hostname or "").lower()
return host.endswith("openrouter.ai")
def models_url(chat_url: str) -> str:
raw = (chat_url or "").rstrip("/")
if raw.endswith("/chat/completions"):
return raw[: -len("/chat/completions")] + "/models"
if is_openrouter_url(raw):
return "https://openrouter.ai/api/v1/models"
return raw + "/models" if raw else "https://openrouter.ai/api/v1/models"
def _from_env(model: str, *, provider: str | None, source: str) -> ModelWindow | None:
context = _int_env("KANSHO_PROVIDER_CONTEXT_LENGTH", 0)
completion = _int_env("KANSHO_PROVIDER_MAX_COMPLETION_TOKENS", 0)
if context <= 0:
return None
return ModelWindow(
model=model,
context_length=context,
max_completion_tokens=completion or JOURNAL_TARGET_COMPLETION_TOKENS,
source=source,
provider=provider,
cached=False,
)
def _parse_model_row(row: dict) -> ModelWindow | None:
model_id = str(row.get("id") or "").strip()
if not model_id:
return None
top = row.get("top_provider") if isinstance(row.get("top_provider"), dict) else {}
context = row.get("context_length") or top.get("context_length")
completion = top.get("max_completion_tokens")
if completion is None:
limits = row.get("per_request_limits") if isinstance(row.get("per_request_limits"), dict) else {}
completion = limits.get("completion_tokens") or limits.get("max_tokens")
try:
context_n = int(context)
except (TypeError, ValueError):
return None
try:
completion_n = int(completion) if completion is not None else None
except (TypeError, ValueError):
completion_n = None
if context_n <= 0:
return None
return ModelWindow(
model=model_id,
context_length=context_n,
max_completion_tokens=completion_n,
source="openrouter_models",
provider="openrouter",
cached=True,
)
def _fetch_catalog(url: str, key: str) -> dict[str, dict]:
headers = {"Content-Type": "application/json"}
if key:
headers["Authorization"] = f"Bearer {key}"
try:
response = httpx.get(url, headers=headers, timeout=15.0)
except httpx.HTTPError as exc:
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={"models_url": url, "reason": "unreachable"},
) from exc
if response.status_code >= 400:
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={"models_url": url, "http_status": response.status_code},
)
try:
payload = response.json()
except ValueError as exc:
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={"models_url": url, "reason": "invalid_json"},
) from exc
rows = payload.get("data") if isinstance(payload, dict) else payload
if not isinstance(rows, list):
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={"models_url": url, "reason": "unexpected_shape"},
)
parsed: dict[str, dict] = {}
for row in rows:
if not isinstance(row, dict):
continue
window = _parse_model_row(row)
if window:
parsed[window.model] = {
"model": window.model,
"context_length": window.context_length,
"max_completion_tokens": window.max_completion_tokens,
"source": window.source,
"provider": window.provider,
}
if not parsed:
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={"models_url": url, "reason": "empty_catalog"},
)
return parsed
def _window_from_cache(fields: dict, *, model: str) -> ModelWindow:
return ModelWindow(
model=model,
context_length=int(fields["context_length"]),
max_completion_tokens=fields.get("max_completion_tokens"),
source=fields.get("source") or "openrouter_models",
provider=fields.get("provider"),
cached=True,
)
def resolve_generate_metadata(config, *, now: float | None = None) -> ModelWindow:
"""Return a window for the generate provider. Fail closed if unknown."""
model = getattr(config, "model", None) or ""
with _lock:
override = _overrides.get(model)
if override:
return override
if getattr(config, "mode", None) == "fake" or model == "fake":
return ModelWindow(
model=model or "fake",
context_length=FAKE_CONTEXT,
max_completion_tokens=JOURNAL_TARGET_COMPLETION_TOKENS,
source="fake",
provider="fake",
cached=False,
)
env_window = _from_env(model, provider=getattr(config, "name", None), source="env")
if getattr(config, "local", False):
if env_window:
return env_window
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={
"model": model,
"provider": getattr(config, "name", None),
"reason": "local_context_unconfigured",
},
)
url = getattr(config, "url", "") or ""
if not is_openrouter_url(url):
if env_window:
return env_window
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={
"model": model,
"provider": getattr(config, "name", None),
"reason": "non_openrouter_unconfigured",
},
)
catalog_key = models_url(url)
clock = time.time() if now is None else now
ttl = catalog_ttl()
with _lock:
cached = _catalog.get(catalog_key)
if cached and clock - cached[0] < ttl and model in cached[1]:
return _window_from_cache(cached[1][model], model=model)
stale = cached
try:
fetched = _fetch_catalog(catalog_key, getattr(config, "key", "") or "")
except JournalBudgetError:
if stale and model in stale[1]:
window = _window_from_cache(stale[1][model], model=model)
window.source = "openrouter_models_stale"
return window
if env_window:
env_window.source = "env_fallback"
return env_window
raise
with _lock:
_catalog[catalog_key] = (clock, fetched)
if model not in fetched:
if env_window:
env_window.source = "env_fallback"
return env_window
raise JournalBudgetError(
ERROR_MODEL_METADATA_UNKNOWN,
status_code=503,
diagnostics={"model": model, "reason": "model_not_in_catalog"},
)
return _window_from_cache(fetched[model], model=model)