252 lines
8.2 KiB
Python
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)
|