Files
6krrt/src/poller.py
adlee-was-taken ec7851a819 feat: store total_credits_usd and total_usage_usd from balance poller
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.
2026-09-10 22:05:12 -04:00

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())