WI-2 of quota-multi-provider-redesign. - total_credits_usd REAL + total_usage_usd REAL columns added to provider_balance_observations table in schema.sql + idempotent ALTER. - _BALANCE_PARSERS now returns 3-tuple (balance, total_credits, total_usage). OpenRouter parser extracts pool size from /credits response; non-pool providers return (x, None, None) for backward compat. - record_balance persists all three columns. - Updated test_provider_balance.py column-list pins and round-trip assertion.
756 lines
30 KiB
Python
756 lines
30 KiB
Python
#!/usr/bin/env python3
|
|
"""
|
|
Pricing/catalog poller for the local LLM router.
|
|
|
|
Fetches the NeuralWatt model catalog (unauthenticated public endpoint),
|
|
normalizes it into the `models` table in router.db, and flags rows that have
|
|
gone stale.
|
|
|
|
Run manually:
|
|
python poller.py
|
|
|
|
Run on a schedule (cron example, every 2 hours):
|
|
0 */2 * * * /usr/bin/python3 /path/to/poller.py >> /var/log/router-poller.log 2>&1
|
|
|
|
Does NOT touch energy data — energy is only available per-completion, not from
|
|
a models list, so the dispatcher writes `energy_observations` instead.
|
|
|
|
NeuralWatt is the only provider. The `provider` column and the (model_id,
|
|
provider) primary key are kept so a second provider can be added without a
|
|
migration.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
import sqlite3
|
|
import sys
|
|
import warnings
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timezone
|
|
from typing import Callable, Optional
|
|
|
|
import requests
|
|
|
|
from config import DispatchProvider, LocalDispatchModel, RouterConfig, load_config
|
|
from tier import apply_tiering
|
|
|
|
MODELS_URL = "https://api.neuralwatt.com/v1/models"
|
|
OPENROUTER_MODELS_URL = "https://openrouter.ai/api/v1/models"
|
|
|
|
REQUEST_TIMEOUT = 20 # seconds
|
|
|
|
OPENROUTER_VIRTUAL_ROUTERS = frozenset({
|
|
"openrouter/auto",
|
|
"openrouter/auto-beta",
|
|
"openrouter/free",
|
|
"openrouter/fusion",
|
|
"openrouter/pareto-code",
|
|
"openrouter/bodybuilder",
|
|
})
|
|
|
|
# Serving-class suffixes. NeuralWatt ships one base model as several catalog
|
|
# rows that differ only by these tokens, and they combine freely — hence ids
|
|
# like 'glm-5.2-short-fast-flex'. They are stripped from the end of the id one
|
|
# segment at a time so a base name that merely *contains* a lookalike token is
|
|
# never misread (e.g. 'deepseek-v4-flash' is not a '-fast' row).
|
|
SUFFIX_FLEX = "flex"
|
|
SUFFIX_FAST = "fast"
|
|
SUFFIX_SHORT = "short"
|
|
SERVING_SUFFIXES = frozenset({SUFFIX_FLEX, SUFFIX_FAST, SUFFIX_SHORT})
|
|
|
|
# Provider-specific account-balance parsers. The registry lives inside the
|
|
# poller because poller imports config, and config importing this dict would
|
|
# create a circular dependency. Tests pin this set equal to
|
|
# config.PROVIDERS_WITH_BALANCE_PARSERS.
|
|
#
|
|
# Each parser returns a 3-tuple:
|
|
# (balance_usd, total_credits_usd | None, total_usage_usd | None)
|
|
# Providers without a pool concept (no /credits endpoint) return
|
|
# (balance, None, None) for backward compat.
|
|
_BALANCE_PARSERS: dict[str, Callable[[dict], tuple[float, Optional[float], Optional[float]]]] = {
|
|
"openrouter": lambda payload: (
|
|
payload["data"]["total_credits"] - payload["data"]["total_usage"],
|
|
payload["data"]["total_credits"],
|
|
payload["data"]["total_usage"],
|
|
)
|
|
}
|
|
|
|
|
|
class CatalogTooSmall(requests.RequestException):
|
|
"""Raised when the fetched catalog has zero rows — not a transient error,
|
|
but a broken catalog that should not silence. Subclassing
|
|
``requests.RequestException`` means the existing ``except
|
|
requests.RequestException`` handler in :func:`main` catches it and we
|
|
exit with code 1 without any extra machinery."""
|
|
|
|
|
|
def _auth_header(prov_cfg: DispatchProvider) -> Optional[dict[str, str]]:
|
|
"""Return a Bearer header if the provider's API key env var is set."""
|
|
key = os.environ.get(prov_cfg.api_key_env)
|
|
if not key:
|
|
return None
|
|
return {"Authorization": f"Bearer {key}"}
|
|
|
|
|
|
def parse_serving_class(model_id: str) -> tuple[str, str, str]:
|
|
"""Derive (latency_class, reasoning_mode, context_variant) from a model id.
|
|
|
|
Returns the schema defaults ('standard', 'default', 'full') for a base
|
|
model. Suffixes are matched as whole '-'-delimited segments only.
|
|
"""
|
|
segments = model_id.lower().split("-")
|
|
found = set()
|
|
while len(segments) > 1 and segments[-1] in SERVING_SUFFIXES:
|
|
found.add(segments.pop())
|
|
|
|
return (
|
|
"flex" if SUFFIX_FLEX in found else "standard",
|
|
"reduced" if SUFFIX_FAST in found else "default",
|
|
"short" if SUFFIX_SHORT in found else "full",
|
|
)
|
|
|
|
|
|
def parse_base_model_id(model_id: str) -> str:
|
|
"""Reduce a catalog id to the model family underneath it.
|
|
|
|
``glm-5.2-short-fast-flex`` and ``glm-5.2`` are the same weights served
|
|
differently, and ``deepseek-ai/DeepSeek-V4-Flash`` is the HF-style
|
|
duplicate of ``deepseek-v4-flash``. Proficiency is a property of the
|
|
weights, not of the queue they sit in, so scores are keyed on this and
|
|
every serving variant inherits from its family. Leaderboard priors work
|
|
the same way — no benchmark rates a ``-flex`` row separately.
|
|
|
|
Note this deliberately collapses ``-fast`` too, even though reasoning
|
|
being off does change answer quality. The eval runner scores ``-fast``
|
|
variants separately and overrides the inherited value; the family is the
|
|
fallback, not the final word.
|
|
"""
|
|
namespace_stripped = model_id.rsplit("/", 1)[-1]
|
|
segments = namespace_stripped.lower().split("-")
|
|
while len(segments) > 1 and segments[-1] in SERVING_SUFFIXES:
|
|
segments.pop()
|
|
return "-".join(segments)
|
|
|
|
|
|
def parse_access_level(display_name: Optional[str], description: Optional[str]) -> str:
|
|
"""Derive an access level from the catalog's prose.
|
|
|
|
NeuralWatt exposes no structured gating field — restricted models are only
|
|
marked in free text ("Private preview (grant-gated)", "(Canary)"). Routing
|
|
to one earns a 403 at dispatch, so this is parsed defensively: anything
|
|
that looks gated is treated as gated.
|
|
"""
|
|
blob = f"{display_name or ''} {description or ''}".lower()
|
|
if "grant-gated" in blob or "private preview" in blob:
|
|
return "preview"
|
|
if "canary" in blob:
|
|
return "canary"
|
|
return "public"
|
|
|
|
|
|
@dataclass
|
|
class ModelRow:
|
|
model_id: str
|
|
provider: str
|
|
base_model_id: str
|
|
display_name: Optional[str]
|
|
cost_per_1m_prompt: Optional[float]
|
|
cost_per_1m_completion: Optional[float]
|
|
cost_per_1m_prompt_cached: Optional[float]
|
|
context_window: Optional[int]
|
|
max_output_tokens: Optional[int]
|
|
supports_tools: bool
|
|
supports_json_mode: bool
|
|
supports_vision: bool
|
|
supports_reasoning: bool
|
|
reasoning_default_enabled: bool
|
|
latency_class: str
|
|
reasoning_mode: str
|
|
context_variant: str
|
|
access_level: str
|
|
pricing_tbd: bool
|
|
deprecated: bool
|
|
|
|
def effective_context_window(self, cfg: RouterConfig) -> Optional[int]:
|
|
"""Usable context, after the safety factor and an output reserve.
|
|
|
|
``context.per_model_overrides`` wins where it is set, which is what it
|
|
is for: the global factor is a guess that has to hold for the whole
|
|
catalog, while a row someone has actually measured deserves its own
|
|
number.
|
|
|
|
Compared with ``is not None`` rather than ``or``, so an override of 0
|
|
reserve tokens means zero rather than silently falling through to the
|
|
default. Present-but-falsy is a trap this project's own eval set tests
|
|
models on; the router should not walk into it.
|
|
|
|
**The reserve is capped at a fraction of the usable window, and that
|
|
cap is load-bearing.** ``max_output_tokens`` is read from the
|
|
provider's catalog, and the two providers do not mean the same thing
|
|
by it. NeuralWatt reports a genuine per-request output cap: at most
|
|
0.16 of context across all 19 rows. OpenRouter reports
|
|
``top_provider.max_completion_tokens``, which is a *ceiling on what
|
|
you may ask for* -- 0.8-0.9 of context on a dozen rows, and on a 1M
|
|
window that is nearly the whole thing.
|
|
|
|
Subtracting it wholesale made ``usable`` negative, which ``max(_, 0)``
|
|
turned into a silent 0, which ``routing.py``'s context filter reads as
|
|
"fits nothing". Measured on the live catalog: 12 of 30 active
|
|
OpenRouter rows sat at 0 and had never been selected once across
|
|
23,000+ decisions, including two 1M-context models. Nothing warned,
|
|
because a row that is never eligible never rejects anything either.
|
|
"""
|
|
if not self.context_window:
|
|
return None
|
|
override = cfg.context.per_model_overrides.get(self.model_id)
|
|
|
|
factor = cfg.context.safety_factor
|
|
if override is not None and override.safety_factor is not None:
|
|
factor = override.safety_factor
|
|
|
|
# 11 of 19 catalog rows report no max_output_tokens, so the configured
|
|
# reserve carries most of the catalog.
|
|
if override is not None and override.output_reserve_tokens is not None:
|
|
reserve = override.output_reserve_tokens
|
|
else:
|
|
reserve = self.max_output_tokens or cfg.context.default_output_reserve_tokens
|
|
|
|
window = int(self.context_window * factor)
|
|
# A per-model override is someone's measurement and is trusted as
|
|
# given; only the catalog-derived reserve is capped.
|
|
if override is None or override.output_reserve_tokens is None:
|
|
reserve = min(reserve, int(window * cfg.context.max_output_reserve_fraction))
|
|
|
|
return max(window - reserve, 0)
|
|
|
|
|
|
def fetch_neuralwatt(provider: str) -> list[ModelRow]:
|
|
resp = requests.get(MODELS_URL, timeout=REQUEST_TIMEOUT)
|
|
resp.raise_for_status()
|
|
payload = resp.json()
|
|
|
|
rows = []
|
|
for m in payload.get("data", []):
|
|
model_id = m.get("id")
|
|
meta = m.get("metadata", {}) or {}
|
|
pricing = meta.get("pricing", {}) or {}
|
|
caps = meta.get("capabilities", {}) or {}
|
|
limits = meta.get("limits", {}) or {}
|
|
reasoning = meta.get("reasoning") or {}
|
|
|
|
supports_reasoning = bool(caps.get("reasoning"))
|
|
# capabilities.reasoning only means "the API accepts a reasoning
|
|
# param" and is true for nearly the whole catalog. default_enabled is
|
|
# the discriminating signal; a few models (the kimi-k2.7-code family)
|
|
# expose no reasoning block at all, so fall back to the capability.
|
|
default_enabled = reasoning.get("default_enabled")
|
|
if default_enabled is None:
|
|
default_enabled = supports_reasoning
|
|
|
|
latency_class, reasoning_mode, context_variant = parse_serving_class(model_id)
|
|
|
|
rows.append(
|
|
ModelRow(
|
|
model_id=model_id,
|
|
provider=provider,
|
|
base_model_id=parse_base_model_id(model_id),
|
|
display_name=meta.get("display_name"),
|
|
cost_per_1m_prompt=pricing.get("input_per_million"),
|
|
cost_per_1m_completion=pricing.get("output_per_million"),
|
|
cost_per_1m_prompt_cached=pricing.get("cached_input_per_million"),
|
|
context_window=limits.get("max_context_length")
|
|
or m.get("max_model_len"),
|
|
max_output_tokens=limits.get("max_output_tokens"),
|
|
supports_tools=bool(caps.get("tools")),
|
|
supports_json_mode=bool(caps.get("json_mode")),
|
|
supports_vision=bool(caps.get("vision")),
|
|
supports_reasoning=supports_reasoning,
|
|
reasoning_default_enabled=bool(default_enabled),
|
|
latency_class=latency_class,
|
|
reasoning_mode=reasoning_mode,
|
|
context_variant=context_variant,
|
|
access_level=parse_access_level(
|
|
meta.get("display_name"), meta.get("description")
|
|
),
|
|
pricing_tbd=bool(pricing.get("pricing_tbd")),
|
|
deprecated=bool(meta.get("deprecated")),
|
|
)
|
|
)
|
|
return rows
|
|
|
|
|
|
def _openrouter_price_to_cost_per_1m(value: Optional[str]) -> Optional[float]:
|
|
"""Convert an OpenRouter per-token price string to cost per 1M tokens."""
|
|
if value is None:
|
|
return None
|
|
try:
|
|
per_token = float(value)
|
|
except ValueError:
|
|
return None
|
|
return per_token * 1_000_000
|
|
|
|
|
|
def parse_openrouter_model(raw_model: dict) -> Optional[ModelRow]:
|
|
"""Normalize a single OpenRouter /v1/models entry into a ModelRow.
|
|
|
|
Returns ``None`` for virtual routers (``openrouter/*``) so callers can
|
|
drop them cleanly without losing the rest of the catalog.
|
|
"""
|
|
model_id = raw_model.get("id") or ""
|
|
if model_id in OPENROUTER_VIRTUAL_ROUTERS:
|
|
return None
|
|
|
|
canonical_slug = raw_model.get("canonical_slug") or model_id
|
|
pricing = raw_model.get("pricing", {}) or {}
|
|
top_provider = raw_model.get("top_provider", {}) or {}
|
|
architecture = raw_model.get("architecture", {}) or {}
|
|
reasoning = raw_model.get("reasoning", {}) or {}
|
|
supported_parameters = raw_model.get("supported_parameters") or []
|
|
input_modalities = architecture.get("input_modalities") or []
|
|
|
|
prompt_price = _openrouter_price_to_cost_per_1m(pricing.get("prompt"))
|
|
completion_price = _openrouter_price_to_cost_per_1m(pricing.get("completion"))
|
|
|
|
variant = model_id.split(":")[-1] if ":" in model_id else "standard"
|
|
|
|
return ModelRow(
|
|
model_id=model_id,
|
|
provider="openrouter",
|
|
base_model_id=canonical_slug,
|
|
display_name=raw_model.get("name"),
|
|
cost_per_1m_prompt=prompt_price,
|
|
cost_per_1m_completion=completion_price,
|
|
cost_per_1m_prompt_cached=None,
|
|
context_window=top_provider.get("context_length")
|
|
or raw_model.get("context_length"),
|
|
max_output_tokens=top_provider.get("max_completion_tokens"),
|
|
supports_tools="tools" in supported_parameters,
|
|
supports_json_mode="response_format" in supported_parameters,
|
|
supports_vision="image" in input_modalities,
|
|
supports_reasoning="reasoning" in supported_parameters,
|
|
reasoning_default_enabled=bool(reasoning.get("default_enabled")),
|
|
latency_class="flex" if ":batch" in model_id else "standard",
|
|
reasoning_mode=reasoning.get("mode") or "none",
|
|
context_variant=variant,
|
|
access_level="public",
|
|
pricing_tbd=prompt_price is None,
|
|
deprecated=False,
|
|
)
|
|
|
|
|
|
def fetch_openrouter(provider: str) -> list[ModelRow]:
|
|
"""Fetch the public OpenRouter catalog and normalize it into ModelRows.
|
|
|
|
The OpenRouter `/v1/models` endpoint is unauthenticated and paginated with
|
|
a default page size of 500 (max 1000). A single request covers the whole
|
|
catalog. ``pricing.web_search`` and ``pricing.cache_read`` differentials
|
|
are intentionally omitted from catalog cost estimates; routing uses the
|
|
prompt and completion price only.
|
|
"""
|
|
resp = requests.get(OPENROUTER_MODELS_URL, timeout=REQUEST_TIMEOUT)
|
|
resp.raise_for_status()
|
|
payload = resp.json()
|
|
|
|
rows: list[ModelRow] = []
|
|
for raw_model in payload.get("data", []):
|
|
row = parse_openrouter_model(raw_model)
|
|
if row is not None:
|
|
rows.append(row)
|
|
return rows
|
|
|
|
|
|
def upsert(conn: sqlite3.Connection, rows: list[ModelRow], cfg: RouterConfig) -> None:
|
|
now = datetime.now(timezone.utc).isoformat()
|
|
for r in rows:
|
|
availability = "deprecated" if r.deprecated else "active"
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO models (
|
|
model_id, provider, base_model_id, display_name,
|
|
cost_per_1m_prompt, cost_per_1m_completion, cost_per_1m_prompt_cached,
|
|
context_window, effective_context_window, max_output_tokens,
|
|
supports_tools, supports_json_mode, supports_vision, supports_reasoning,
|
|
reasoning_default_enabled, latency_class, reasoning_mode,
|
|
context_variant, access_level,
|
|
pricing_tbd, deprecated, availability, last_updated
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(model_id, provider) DO UPDATE SET
|
|
base_model_id = excluded.base_model_id,
|
|
display_name = excluded.display_name,
|
|
cost_per_1m_prompt = excluded.cost_per_1m_prompt,
|
|
cost_per_1m_completion = excluded.cost_per_1m_completion,
|
|
cost_per_1m_prompt_cached = excluded.cost_per_1m_prompt_cached,
|
|
context_window = excluded.context_window,
|
|
effective_context_window = excluded.effective_context_window,
|
|
max_output_tokens = excluded.max_output_tokens,
|
|
supports_tools = excluded.supports_tools,
|
|
supports_json_mode = excluded.supports_json_mode,
|
|
supports_vision = excluded.supports_vision,
|
|
supports_reasoning = excluded.supports_reasoning,
|
|
reasoning_default_enabled = excluded.reasoning_default_enabled,
|
|
latency_class = excluded.latency_class,
|
|
reasoning_mode = excluded.reasoning_mode,
|
|
context_variant = excluded.context_variant,
|
|
access_level = excluded.access_level,
|
|
pricing_tbd = excluded.pricing_tbd,
|
|
deprecated = excluded.deprecated,
|
|
availability = excluded.availability,
|
|
last_updated = excluded.last_updated
|
|
""",
|
|
(
|
|
r.model_id,
|
|
r.provider,
|
|
r.base_model_id,
|
|
r.display_name,
|
|
r.cost_per_1m_prompt,
|
|
r.cost_per_1m_completion,
|
|
r.cost_per_1m_prompt_cached,
|
|
r.context_window,
|
|
r.effective_context_window(cfg),
|
|
r.max_output_tokens,
|
|
int(r.supports_tools),
|
|
int(r.supports_json_mode),
|
|
int(r.supports_vision),
|
|
int(r.supports_reasoning),
|
|
int(r.reasoning_default_enabled),
|
|
r.latency_class,
|
|
r.reasoning_mode,
|
|
r.context_variant,
|
|
r.access_level,
|
|
int(r.pricing_tbd),
|
|
int(r.deprecated),
|
|
availability,
|
|
now,
|
|
),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def _ensure_provider_balance_table(conn: sqlite3.Connection) -> None:
|
|
"""Idempotently create the provider_balance_observations table and index.
|
|
|
|
Live router.db files that predate this feature lack the table; the
|
|
CREATE TABLE IF NOT EXISTS / CREATE INDEX IF NOT EXISTS DDL is safe to
|
|
re-run on every poll. The ALTER TABLEs below add columns to router.db
|
|
files that predate the credit-pool columns; the try/except swallows the
|
|
"duplicate column" error so they are safe to re-run on every poll.
|
|
"""
|
|
conn.execute("""
|
|
CREATE TABLE IF NOT EXISTS provider_balance_observations (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
provider TEXT NOT NULL,
|
|
balance_usd REAL NOT NULL,
|
|
total_credits_usd REAL,
|
|
total_usage_usd REAL,
|
|
observed_at TEXT NOT NULL
|
|
)
|
|
""")
|
|
conn.execute(
|
|
"CREATE INDEX IF NOT EXISTS idx_provider_balance "
|
|
"ON provider_balance_observations (provider, observed_at)"
|
|
)
|
|
try:
|
|
conn.execute(
|
|
"ALTER TABLE provider_balance_observations "
|
|
"ADD COLUMN total_credits_usd REAL"
|
|
)
|
|
except sqlite3.OperationalError:
|
|
pass
|
|
try:
|
|
conn.execute(
|
|
"ALTER TABLE provider_balance_observations "
|
|
"ADD COLUMN total_usage_usd REAL"
|
|
)
|
|
except sqlite3.OperationalError:
|
|
pass
|
|
conn.commit()
|
|
|
|
|
|
def record_balance(
|
|
conn: sqlite3.Connection,
|
|
provider: str,
|
|
balance_usd: float,
|
|
total_credits_usd: Optional[float] = None,
|
|
total_usage_usd: Optional[float] = None,
|
|
) -> None:
|
|
"""Write a provider account-balance observation row.
|
|
|
|
total_credits_usd / total_usage_usd describe the prepaid pool and are
|
|
None for providers that report only a remainder.
|
|
"""
|
|
conn.execute(
|
|
"INSERT INTO provider_balance_observations "
|
|
"(provider, balance_usd, total_credits_usd, total_usage_usd, observed_at) "
|
|
"VALUES (?, ?, ?, ?, ?)",
|
|
(
|
|
provider,
|
|
balance_usd,
|
|
total_credits_usd,
|
|
total_usage_usd,
|
|
datetime.now(timezone.utc).isoformat(),
|
|
),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def _ensure_models_eligible_categories(conn: sqlite3.Connection) -> None:
|
|
"""Idempotently add the models.eligible_categories column.
|
|
|
|
Older router.db files may lack the column because it was added after the
|
|
original schema. The try/except swallows the "duplicate column" error so
|
|
this can run on every poll.
|
|
"""
|
|
try:
|
|
conn.execute("ALTER TABLE models ADD COLUMN eligible_categories TEXT")
|
|
except sqlite3.OperationalError:
|
|
pass
|
|
|
|
|
|
def _effective_context_window_for_local(
|
|
entry: LocalDispatchModel, cfg: RouterConfig
|
|
) -> int:
|
|
"""Mirror of ModelRow.effective_context_window for static config entries.
|
|
|
|
per_model_overrides may specify a custom safety factor or reserve, so they
|
|
are honored exactly the way the catalog upsert honors them.
|
|
"""
|
|
override = cfg.context.per_model_overrides.get(entry.model_id)
|
|
|
|
factor = cfg.context.safety_factor
|
|
if override is not None and override.safety_factor is not None:
|
|
factor = override.safety_factor
|
|
|
|
if override is not None and override.output_reserve_tokens is not None:
|
|
reserve = override.output_reserve_tokens
|
|
else:
|
|
reserve = entry.max_output_tokens
|
|
|
|
usable = int(entry.context_window * factor) - reserve
|
|
return max(usable, 0)
|
|
|
|
|
|
def upsert_local_dispatch_models(
|
|
conn: sqlite3.Connection, cfg: RouterConfig
|
|
) -> None:
|
|
"""Seed/update models table rows from cfg.local_dispatch_models."""
|
|
_ensure_models_eligible_categories(conn)
|
|
now = datetime.now(timezone.utc).isoformat()
|
|
for entry in cfg.local_dispatch_models:
|
|
effective = _effective_context_window_for_local(entry, cfg)
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO models (
|
|
model_id, provider, base_model_id, display_name,
|
|
cost_per_1m_prompt, cost_per_1m_completion, cost_per_1m_prompt_cached,
|
|
context_window, effective_context_window, max_output_tokens,
|
|
supports_tools, supports_json_mode, supports_vision, supports_reasoning,
|
|
reasoning_default_enabled, latency_class, reasoning_mode,
|
|
context_variant, access_level,
|
|
eligible_categories,
|
|
pricing_tbd, deprecated, availability, last_updated, tier
|
|
) VALUES (?, ?, ?, ?, NULL, NULL, NULL, ?, ?, ?, 0, 0, 0, 0, 0, 'standard', 'default', 'full', 'public', ?, 0, 0, 'active', ?, ?)
|
|
ON CONFLICT(model_id, provider) DO UPDATE SET
|
|
base_model_id = excluded.base_model_id,
|
|
display_name = excluded.display_name,
|
|
context_window = excluded.context_window,
|
|
effective_context_window = excluded.effective_context_window,
|
|
max_output_tokens = excluded.max_output_tokens,
|
|
supports_tools = excluded.supports_tools,
|
|
supports_json_mode = excluded.supports_json_mode,
|
|
supports_vision = excluded.supports_vision,
|
|
supports_reasoning = excluded.supports_reasoning,
|
|
reasoning_default_enabled = excluded.reasoning_default_enabled,
|
|
latency_class = excluded.latency_class,
|
|
reasoning_mode = excluded.reasoning_mode,
|
|
context_variant = excluded.context_variant,
|
|
access_level = excluded.access_level,
|
|
eligible_categories = excluded.eligible_categories,
|
|
pricing_tbd = excluded.pricing_tbd,
|
|
deprecated = excluded.deprecated,
|
|
availability = excluded.availability,
|
|
last_updated = excluded.last_updated,
|
|
tier = excluded.tier
|
|
""",
|
|
(
|
|
entry.model_id,
|
|
"ollama-local",
|
|
entry.model_id,
|
|
None,
|
|
entry.context_window,
|
|
effective,
|
|
entry.max_output_tokens,
|
|
",".join(entry.eligible_categories),
|
|
now,
|
|
entry.tier,
|
|
),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def mark_stale(
|
|
conn: sqlite3.Connection,
|
|
cfg: RouterConfig,
|
|
*,
|
|
provider: Optional[str] = None,
|
|
) -> None:
|
|
"""Flag rows that weren't touched by this poll run as stale.
|
|
|
|
When *provider* is given, only rows for that provider are considered;
|
|
when absent (the historical default across the rest of this repo), all
|
|
rows are scoped — matching the previous global behaviour.
|
|
|
|
Assumes the caller has already established the fetched catalog was
|
|
plausible (see the row-count floor check before upsert).
|
|
"""
|
|
where = "availability = 'active' AND julianday('now') - julianday(last_updated) > ?"
|
|
params: list[object] = [cfg.freshness.stale_after_days]
|
|
if provider is not None:
|
|
where += " AND provider = ?"
|
|
params.append(provider)
|
|
|
|
conn.execute(f"UPDATE models SET availability = 'stale' WHERE {where}", params)
|
|
conn.commit()
|
|
|
|
|
|
def main() -> int:
|
|
cfg = load_config("config/config.yaml")
|
|
conn = sqlite3.connect(cfg.database.path)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.execute("PRAGMA foreign_keys = ON")
|
|
|
|
upsert_local_dispatch_models(conn, cfg)
|
|
_ensure_provider_balance_table(conn)
|
|
|
|
# Map provider keys (dispatch_providers dict keys) to their fetch functions.
|
|
# Only entries with keys matching the key are dispatched; unknown keys
|
|
# are logged and skipped so a future provider in config doesn't abort
|
|
# the whole run.
|
|
FETCHERS: dict[str, callable] = {
|
|
"neuralwatt": fetch_neuralwatt,
|
|
"openrouter": fetch_openrouter,
|
|
}
|
|
|
|
total_upserted = 0
|
|
|
|
for provider, prov_cfg in cfg.dispatch_providers.items():
|
|
if not prov_cfg.enabled:
|
|
print(f"[{provider}] skipped (disabled)")
|
|
continue
|
|
|
|
if prov_cfg.balance_url:
|
|
headers = _auth_header(prov_cfg)
|
|
if headers is None:
|
|
print(
|
|
f"[{provider}] balance poll skipped: "
|
|
f"{prov_cfg.api_key_env} not set",
|
|
file=sys.stderr,
|
|
)
|
|
else:
|
|
try:
|
|
resp = requests.get(
|
|
prov_cfg.balance_url, headers=headers, timeout=REQUEST_TIMEOUT
|
|
)
|
|
resp.raise_for_status()
|
|
balance, total_credits, total_usage = _BALANCE_PARSERS[provider](
|
|
resp.json()
|
|
)
|
|
record_balance(conn, provider, balance, total_credits, total_usage)
|
|
print(f"[{provider}] balance ${balance:.2f}")
|
|
except (requests.RequestException, KeyError, TypeError, ValueError) as e:
|
|
print(f"[{provider}] balance poll FAILED: {e}", file=sys.stderr)
|
|
|
|
fetcher = FETCHERS.get(provider)
|
|
if fetcher is None:
|
|
print(f"[{provider}] unknown provider — skipping (no fetcher)")
|
|
continue
|
|
|
|
try:
|
|
rows = fetcher(provider)
|
|
except CatalogTooSmall as e:
|
|
print(f"[{provider}] CATALOG_TOO_SMALL: {e}; skipping provider", file=sys.stderr)
|
|
continue
|
|
except requests.RequestException as e:
|
|
print(f"[{provider}] FAILED: {e}", file=sys.stderr)
|
|
continue
|
|
|
|
if prov_cfg.require_allowlist:
|
|
# Query the allowlist for this provider from the database.
|
|
allowlist_ids = {
|
|
row["model_id"]
|
|
for row in conn.execute(
|
|
"SELECT model_id FROM provider_model_allowlist WHERE provider=?",
|
|
(provider,),
|
|
).fetchall()
|
|
}
|
|
|
|
if len(allowlist_ids) == 0:
|
|
# Empty allowlist — warn, deprecate every existing active row,
|
|
# then skip this provider for today.
|
|
print(
|
|
f"[{provider}] allowlist is empty — deprecating all active rows"
|
|
if conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider=? AND deprecated=0",
|
|
(provider,),
|
|
).fetchone()[0]
|
|
else f"[{provider}] allowlist is empty — nothing active to deprecate",
|
|
file=sys.stderr,
|
|
)
|
|
if conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider=? AND deprecated=0",
|
|
(provider,),
|
|
).fetchone()[0]:
|
|
conn.execute(
|
|
"UPDATE models SET deprecated=1, availability='deprecated' "
|
|
"WHERE provider=? AND deprecated=0",
|
|
(provider,),
|
|
)
|
|
conn.commit()
|
|
rows = []
|
|
continue
|
|
|
|
# Deprecate existing active rows that are NOT on the allowlist.
|
|
placeholders = ",".join("?" for _ in allowlist_ids)
|
|
conn.execute(
|
|
f"UPDATE models SET deprecated=1, availability='deprecated' "
|
|
f"WHERE provider=? AND deprecated=0 AND model_id NOT IN ({placeholders});",
|
|
[provider] + sorted(allowlist_ids),
|
|
)
|
|
|
|
# Filter fetched rows to the allowlist.
|
|
original_count = len(rows)
|
|
rows = [r for r in rows if r.model_id in allowlist_ids]
|
|
skipped = original_count - len(rows)
|
|
if skipped:
|
|
print(
|
|
f"[{provider}] allowlist filter: {len(rows)} kept, {skipped} pruned"
|
|
)
|
|
|
|
if len(rows) == 0:
|
|
print(f"[{provider}] fetched 0 models — skipping provider", file=sys.stderr)
|
|
continue
|
|
|
|
current_count = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider=?", (provider,)
|
|
).fetchone()[0]
|
|
if len(rows) < (current_count // 2) and current_count > 0:
|
|
warnings.warn(
|
|
f"[{provider}] fetched only {len(rows)} models "
|
|
f"(current DB has {current_count}); proceeding but catalog may be truncated"
|
|
)
|
|
|
|
upsert(conn, rows, cfg)
|
|
total_upserted += len(rows)
|
|
print(f"[{provider}] upserted {len(rows)} models")
|
|
|
|
mark_stale(conn, cfg, provider=provider)
|
|
|
|
apply_tiering(conn, cfg)
|
|
conn.close()
|
|
print(f"done, {total_upserted} rows upserted total")
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|