859 lines
30 KiB
Python
859 lines
30 KiB
Python
"""Tests for the multi-provider main loop in poller.py.
|
|
|
|
Covers:
|
|
- two providers both upsert
|
|
- one provider fails with RequestException while the other succeeds (failure isolation)
|
|
- disabled provider is skipped
|
|
- zero-row provider is logged as skip
|
|
- per-provider mark_stale scopes staleness to the queried provider
|
|
- upstream mark_stale (provider=None) retains global behavior
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import sqlite3
|
|
from datetime import datetime, timedelta, timezone
|
|
from pathlib import Path
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
import requests
|
|
|
|
import poller
|
|
from config import RouterConfig, load_config
|
|
|
|
ROOT = Path(__file__).resolve().parent.parent
|
|
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
|
|
REAL_CFG = load_config(str(ROOT / "config" / "config.yaml"))
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Helpers
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _fake_response(json_data: dict):
|
|
fake = MagicMock()
|
|
fake.json.return_value = json_data
|
|
fake.raise_for_status.return_value = None
|
|
fake.status_code = 200
|
|
return fake
|
|
|
|
|
|
def _catalog_item(model_id: str) -> dict:
|
|
return {
|
|
"id": model_id,
|
|
"metadata": {
|
|
"display_name": model_id.replace("-", " ").title(),
|
|
"description": "test model",
|
|
"pricing": {
|
|
"input_per_million": 0.30,
|
|
"output_per_million": 0.60,
|
|
"cached_input_per_million": 0.20,
|
|
"pricing_tbd": False,
|
|
},
|
|
"capabilities": {
|
|
"tools": True,
|
|
"json_mode": True,
|
|
"vision": False,
|
|
"reasoning": False,
|
|
},
|
|
"limits": {
|
|
"max_context_length": 131072,
|
|
"max_output_tokens": 8192,
|
|
},
|
|
"deprecated": False,
|
|
},
|
|
}
|
|
|
|
|
|
def _openrouter_model(
|
|
model_id: str,
|
|
prompt: str = "0.000001",
|
|
context_length: int = 131072,
|
|
):
|
|
return {
|
|
"id": model_id,
|
|
"name": model_id.replace("/", " ").title(),
|
|
"canonical_slug": model_id,
|
|
"pricing": {"prompt": prompt, "completion": "0.000002"},
|
|
"top_provider": {
|
|
"context_length": context_length,
|
|
"max_completion_tokens": 8192,
|
|
},
|
|
"architecture": {"input_modalities": ["text"]},
|
|
"supported_parameters": [],
|
|
"reasoning": {},
|
|
}
|
|
|
|
|
|
def _seed_models(
|
|
conn: sqlite3.Connection,
|
|
models: list[dict],
|
|
provider: str = "neuralwatt",
|
|
last_updated_days_ago: int = 0,
|
|
) -> None:
|
|
ts = (
|
|
datetime.now(timezone.utc) - timedelta(days=last_updated_days_ago)
|
|
).isoformat()
|
|
for m in models:
|
|
deprecated = bool(m.get("deprecated", False))
|
|
availability = "deprecated" if deprecated else m.get("availability", "active")
|
|
conn.execute(
|
|
"""
|
|
INSERT OR REPLACE 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 (?, ?, ?, ?, 0.30, 0.60, 0.20,
|
|
131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
?, ?)
|
|
""",
|
|
(
|
|
m["model_id"],
|
|
provider,
|
|
m.get("base_model_id", m["model_id"]),
|
|
m.get(
|
|
"display_name",
|
|
f"Model {m['model_id'].title().replace('_', ' ')}",
|
|
),
|
|
availability,
|
|
ts,
|
|
),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Fixtures
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@pytest.fixture()
|
|
def tmp_db(tmp_path):
|
|
"""Create an empty temp DB with schema and return a helper tuple.
|
|
|
|
Returns ``(seed_and_check, cfg)`` where ``seed_and_check`` is a callable
|
|
that opens a connection, lets the caller seed it, and returns the
|
|
connection. The connection must be closed by the test. main() opens its
|
|
own connection, so we cannot keep a fixture-level connection open.
|
|
"""
|
|
db_path = str(tmp_path / "test.db")
|
|
# REAL_CFG has require_allowlist=True for openrouter. Force it off so the
|
|
# tmp_db fixture isn't blocked by an empty allowlist table.
|
|
providers = {
|
|
_prov: _prov_cfg.model_copy(update={"require_allowlist": False})
|
|
for _prov, _prov_cfg in REAL_CFG.dispatch_providers.items()
|
|
}
|
|
cfg = REAL_CFG.model_copy(
|
|
update={
|
|
"database": REAL_CFG.database.model_copy(update={"path": db_path}),
|
|
"freshness": REAL_CFG.freshness.model_copy(update={"stale_after_days": 300}),
|
|
"dispatch_providers": providers,
|
|
}
|
|
)
|
|
|
|
def connect():
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
return conn
|
|
|
|
c = connect()
|
|
c.executescript(SCHEMA_SQL)
|
|
c.close()
|
|
|
|
yield connect, cfg
|
|
|
|
|
|
@pytest.fixture()
|
|
def monkeypatch_cfg(monkeypatch):
|
|
"""Return a helper that sets up a minimal dispatch_providers config."""
|
|
def _make(providers: dict) -> RouterConfig:
|
|
return REAL_CFG.model_copy(
|
|
update={
|
|
"database": REAL_CFG.database.model_copy(
|
|
update={"path": str(tmp_path_provider[0] / "test.db")}
|
|
),
|
|
"freshness": REAL_CFG.freshness.model_copy(
|
|
update={"stale_after_days": 300}
|
|
),
|
|
"dispatch_providers": providers,
|
|
}
|
|
)
|
|
|
|
# We'll create tmp paths on demand. Use a list so the closure gets updated.
|
|
tmp_dirs: list[Path] = []
|
|
|
|
def _make_with_dir(providers: dict, tmp_path: Path) -> tuple[RouterConfig, str]:
|
|
db_path = str(tmp_path / "test.db")
|
|
cfg = REAL_CFG.model_copy(
|
|
update={
|
|
"database": REAL_CFG.database.model_copy(update={"path": db_path}),
|
|
"freshness": REAL_CFG.freshness.model_copy(
|
|
update={"stale_after_days": 300}
|
|
),
|
|
"dispatch_providers": providers,
|
|
}
|
|
)
|
|
tmp_dirs.append(tmp_path)
|
|
return cfg, db_path
|
|
|
|
return _make_with_dir, tmp_dirs
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def test_two_providers_both_upsert(tmp_db, monkeypatch):
|
|
"""Two enabled providers are fetched, upserted, and total reported."""
|
|
connect, cfg = tmp_db
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response(
|
|
{
|
|
"data": [
|
|
_catalog_item("nw-model-a"),
|
|
_catalog_item("nw-model-b"),
|
|
]
|
|
}
|
|
),
|
|
)
|
|
|
|
exit_code = poller.main()
|
|
|
|
conn = connect()
|
|
try:
|
|
nw_rows = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
or_rows = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='openrouter'"
|
|
).fetchone()[0]
|
|
assert nw_rows == 2
|
|
assert or_rows == 2
|
|
assert exit_code == 0
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_one_provider_fails_other_succeeds(tmp_db, monkeypatch):
|
|
"""Failure in one provider's fetch does not abort the other."""
|
|
connect, cfg = tmp_db
|
|
|
|
def _failing_get(url, timeout, **kwargs):
|
|
if "/credits" in url:
|
|
# OpenRouter account-balance poll happens before the catalog fetch.
|
|
return _fake_response({"data": {"total_credits": 50.0, "total_usage": 0.0}})
|
|
if "openrouter.ai" in url:
|
|
return _fake_response({"data": [_openrouter_model("openai/gpt-6-astra")]})
|
|
# NeuralWatt catalog fails
|
|
raise requests.RequestException("network error")
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(poller.requests, "get", _failing_get)
|
|
|
|
exit_code = poller.main()
|
|
|
|
conn = connect()
|
|
try:
|
|
nw_rows = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
or_rows = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='openrouter'"
|
|
).fetchone()[0]
|
|
assert nw_rows == 0, "failed provider should have 0 rows"
|
|
assert or_rows == 1, "succeeded provider should have 1 row"
|
|
assert exit_code == 0
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_disabled_provider_skipped(tmp_db, monkeypatch):
|
|
"""Provider with enabled=False is never fetched or upserted."""
|
|
connect, cfg = tmp_db
|
|
# Configure only neuralwatt as disabled
|
|
nw_cfg = cfg.dispatch_providers.get("neuralwatt")
|
|
if nw_cfg is None:
|
|
nw_cfg = REAL_CFG.dispatch_providers["neuralwatt"]
|
|
disabled_cfg = nw_cfg.model_copy(update={"enabled": False})
|
|
|
|
cfg = cfg.model_copy(
|
|
update={
|
|
"dispatch_providers": {"neuralwatt": disabled_cfg}
|
|
}
|
|
)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
# If the provider is skipped, requests.get should never be called
|
|
monkeypatch.setattr(poller.requests, "get", MagicMock(side_effect=RuntimeError("should not be called")))
|
|
|
|
exit_code = poller.main()
|
|
|
|
conn = connect()
|
|
try:
|
|
nw_rows = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
assert nw_rows == 0
|
|
assert exit_code == 0
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_zero_row_provider_logged_as_skip(tmp_path, monkeypatch):
|
|
"""A provider that returns an empty catalog is logged and skipped."""
|
|
db_path = str(tmp_path / "test.db")
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.executescript(SCHEMA_SQL)
|
|
conn.close()
|
|
|
|
cfg = REAL_CFG.model_copy(
|
|
update={
|
|
"database": REAL_CFG.database.model_copy(update={"path": db_path}),
|
|
"freshness": REAL_CFG.freshness.model_copy(
|
|
update={"stale_after_days": 300}
|
|
),
|
|
"dispatch_providers": {"neuralwatt": REAL_CFG.dispatch_providers["neuralwatt"]},
|
|
}
|
|
)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response({"data": []}),
|
|
)
|
|
|
|
import io
|
|
|
|
stderr_capture = io.StringIO()
|
|
monkeypatch.setattr("sys.stderr", stderr_capture)
|
|
|
|
exit_code = poller.main()
|
|
|
|
assert exit_code == 0
|
|
stderr_content = stderr_capture.getvalue()
|
|
assert "0 models" in stderr_content or "skipping" in stderr_content
|
|
|
|
|
|
def test_in_process_tiering_runs_after_upsert(tmp_path, monkeypatch):
|
|
"""Successful poll assigns tiers so rows are routable immediately."""
|
|
db_path = str(tmp_path / "test.db")
|
|
conn_init = sqlite3.connect(db_path)
|
|
conn_init.row_factory = sqlite3.Row
|
|
conn_init.executescript(SCHEMA_SQL)
|
|
conn_init.close()
|
|
|
|
cfg = REAL_CFG.model_copy(
|
|
update={
|
|
"database": REAL_CFG.database.model_copy(update={"path": db_path}),
|
|
"freshness": REAL_CFG.freshness.model_copy(
|
|
update={"stale_after_days": 300}
|
|
),
|
|
"dispatch_providers": {
|
|
"openrouter": REAL_CFG.dispatch_providers["openrouter"].model_copy(
|
|
update={"require_allowlist": False}
|
|
),
|
|
},
|
|
}
|
|
)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response(
|
|
{"data": [_openrouter_model("openai/gpt-6-astra", context_length=200000)]}
|
|
),
|
|
)
|
|
|
|
exit_code = poller.main()
|
|
|
|
assert exit_code == 0
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
try:
|
|
row = conn.execute(
|
|
"SELECT tier FROM models WHERE provider='openrouter' AND model_id='openai/gpt-6-astra'"
|
|
).fetchone()
|
|
assert row is not None
|
|
assert row["tier"] is not None
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_per_provider_mark_stale_scopes_to_provider(tmp_db, monkeypatch):
|
|
"""Staleness update for one provider doesn't affect the other."""
|
|
connect, cfg = tmp_db
|
|
|
|
# Seed stale rows for both providers
|
|
old_ts = (datetime.now(timezone.utc) - timedelta(days=10)).isoformat()
|
|
conn = connect()
|
|
try:
|
|
_seed_models(
|
|
conn,
|
|
[{"model_id": "nw_old"}],
|
|
provider="neuralwatt",
|
|
last_updated_days_ago=10,
|
|
)
|
|
# Insert openrouter row with availability='active' and old timestamp
|
|
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 ('or_old', 'openrouter', 'or_old', 'Old OR',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(old_ts,),
|
|
)
|
|
# Also insert a 'fresh' old-provider row that should NOT become stale
|
|
fresh_ts = datetime.now(timezone.utc).isoformat()
|
|
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 ('nw_fresh', 'neuralwatt', 'nw_fresh', 'Fresh NW',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(fresh_ts,),
|
|
)
|
|
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 ('or_fresh', 'openrouter', 'or_fresh', 'Fresh OR',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(fresh_ts,),
|
|
)
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
# Fresh timestamp for upsert
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response({"data": [_catalog_item("nw_fresh")]}),
|
|
)
|
|
|
|
cfg = cfg.model_copy(
|
|
update={
|
|
"freshness": cfg.freshness.model_copy(update={"stale_after_days": 5})
|
|
}
|
|
)
|
|
|
|
poller.main()
|
|
|
|
conn = connect()
|
|
try:
|
|
# nw_old should have been marked stale (was active, not fetched this run)
|
|
stale_nw = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
# or_old should NOT be stale — mark_stale was only called for neuralwatt provider
|
|
stale_or = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='openrouter'"
|
|
).fetchone()[0]
|
|
assert stale_nw == 1, f"Expected 1 stale neuralwatt row, got {stale_nw}"
|
|
# openrouter's stale rows remain because mark_stale scoped to neuralwatt
|
|
assert stale_or == 1, f"Expected OR stale rows unchanged, got {stale_or} (or_old was already stale from old_ts)"
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_unknown_provider_key_skipped(tmp_path, monkeypatch):
|
|
"""An unknown provider key in dispatch_providers is logged and skipped."""
|
|
db_path = str(tmp_path / "test.db")
|
|
conn_init = sqlite3.connect(db_path)
|
|
conn_init.executescript(SCHEMA_SQL)
|
|
conn_init.close()
|
|
|
|
# Use a known entry as a template for the unknown provider shape.
|
|
template = REAL_CFG.dispatch_providers.get(
|
|
"neuralwatt", REAL_CFG.dispatch_providers["openrouter"]
|
|
)
|
|
unknown_cfg = template.model_copy(update={"enabled": True})
|
|
|
|
cfg_with_unknown = REAL_CFG.model_copy(
|
|
update={
|
|
"database": REAL_CFG.database.model_copy(update={"path": db_path}),
|
|
"freshness": REAL_CFG.freshness.model_copy(
|
|
update={"stale_after_days": 300}
|
|
),
|
|
"dispatch_providers": {
|
|
"neuralwatt": REAL_CFG.dispatch_providers["neuralwatt"],
|
|
"unknown_provider": unknown_cfg,
|
|
},
|
|
}
|
|
)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg_with_unknown)
|
|
# Unknown provider should be skipped before any fetch.
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
MagicMock(
|
|
side_effect=lambda url, timeout: (
|
|
_fake_response({"data": [_catalog_item("model_a")]})
|
|
)
|
|
),
|
|
)
|
|
|
|
exit_code = poller.main()
|
|
|
|
assert exit_code == 0
|
|
# unknown_provider produced no rows
|
|
unknown_rows = sqlite3.connect(db_path).execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='unknown_provider'"
|
|
).fetchone()[0]
|
|
assert unknown_rows == 0
|
|
|
|
|
|
def test_mark_stale_global_behavior_unaffected():
|
|
"""mark_stale() with no provider arg retains the old global behavior."""
|
|
import tempfile
|
|
|
|
with tempfile.NamedTemporaryFile(suffix=".db") as f:
|
|
conn = sqlite3.connect(f.name)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.executescript(SCHEMA_SQL)
|
|
|
|
old_ts = (datetime.now(timezone.utc) - timedelta(days=10)).isoformat()
|
|
# Insert two rows from different providers, both active and old
|
|
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 ('model_a', 'nw', 'model_a', 'A',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(old_ts,),
|
|
)
|
|
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 ('model_b', 'openrouter', 'model_b', 'B',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(old_ts,),
|
|
)
|
|
conn.commit()
|
|
|
|
# Call mark_stale without provider (global behavior)
|
|
mark_stale_fn = poller.mark_stale
|
|
mark_stale_fn(conn, REAL_CFG)
|
|
conn.commit()
|
|
|
|
stale_total = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale'"
|
|
).fetchone()[0]
|
|
|
|
assert stale_total == 2, f"All rows should be stale globally, got {stale_total}"
|
|
|
|
conn.close()
|
|
|
|
|
|
def test_per_provider_mark_stale_scoped(tmp_path):
|
|
"""mark_stale(provider=...) only marks rows for that provider."""
|
|
db_path = str(tmp_path / "test.db")
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.executescript(SCHEMA_SQL)
|
|
|
|
old_ts = (datetime.now(timezone.utc) - timedelta(days=10)).isoformat()
|
|
|
|
# Seed one active, old row for each provider
|
|
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 ('nw_old', 'neuralwatt', 'nw_old', 'A',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(old_ts,),
|
|
)
|
|
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 ('or_old', 'openrouter', 'or_old', 'B',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(old_ts,),
|
|
)
|
|
conn.commit()
|
|
|
|
# mark_stale scoped to neuralwatt
|
|
poller.mark_stale(conn, REAL_CFG, provider="neuralwatt")
|
|
|
|
stale_nw = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
stale_or = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='openrouter'"
|
|
).fetchone()[0]
|
|
|
|
assert stale_nw == 1, "neuralwatt row should be stale"
|
|
assert stale_or == 0, "openrouter row should NOT be stale when scoped to neuralwatt"
|
|
|
|
conn.close()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Balance polling tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
OPENROUTER_BALANCE_URL = REAL_CFG.dispatch_providers["openrouter"].balance_url
|
|
|
|
|
|
def _url_keyed_get(url: str, timeout: object, **kwargs: object) -> MagicMock:
|
|
"""Dispatch mocked HTTP responses by URL for balance + catalog tests."""
|
|
if "/credits" in url:
|
|
return _fake_response(
|
|
{"data": {"total_credits": 50.0, "total_usage": 7.0}}
|
|
)
|
|
if "openrouter.ai" in url:
|
|
return _fake_response({"data": [_openrouter_model("openai/gpt-6-astra")]})
|
|
if "neuralwatt" in url:
|
|
return _fake_response({"data": [_catalog_item("nw-model-a")]})
|
|
raise RuntimeError(f"unexpected URL in test mock: {url}")
|
|
|
|
|
|
def test_balance_poll_success_records_row(tmp_db, monkeypatch):
|
|
"""An OpenRouter balance poll writes one provider_balance_observations row."""
|
|
connect, cfg = tmp_db
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(poller.requests, "get", _url_keyed_get)
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-test")
|
|
|
|
exit_code = poller.main()
|
|
assert exit_code == 0
|
|
|
|
conn = connect()
|
|
try:
|
|
row = conn.execute(
|
|
"SELECT provider, balance_usd FROM provider_balance_observations "
|
|
"WHERE provider = ?",
|
|
("openrouter",),
|
|
).fetchone()
|
|
assert row is not None, "expected a balance row for openrouter"
|
|
assert row["provider"] == "openrouter"
|
|
assert row["balance_usd"] == 43.0
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_balance_parser_registry_matches_config_registry():
|
|
"""poller._BALANCE_PARSERS keys equal config.PROVIDERS_WITH_BALANCE_PARSERS."""
|
|
from config import PROVIDERS_WITH_BALANCE_PARSERS
|
|
|
|
assert set(poller._BALANCE_PARSERS) == PROVIDERS_WITH_BALANCE_PARSERS
|
|
|
|
|
|
def test_balance_poll_failure_is_isolated(tmp_db, monkeypatch):
|
|
"""A balance-poll exception does not abort the rest of the poll run."""
|
|
connect, cfg = tmp_db
|
|
|
|
def _get(url, timeout, **kwargs):
|
|
if "/credits" in url:
|
|
raise requests.ConnectionError("balance endpoint unreachable")
|
|
if "openrouter.ai" in url:
|
|
return _fake_response({"data": [_openrouter_model("openai/gpt-6-astra")]})
|
|
if "neuralwatt" in url:
|
|
return _fake_response({"data": [_catalog_item("nw-model-a")]})
|
|
raise RuntimeError(f"unexpected URL: {url}")
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(poller.requests, "get", _get)
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-test")
|
|
|
|
import io
|
|
|
|
stderr_capture = io.StringIO()
|
|
monkeypatch.setattr("sys.stderr", stderr_capture)
|
|
|
|
exit_code = poller.main()
|
|
assert exit_code == 0
|
|
assert "[openrouter] balance poll FAILED" in stderr_capture.getvalue()
|
|
|
|
conn = connect()
|
|
try:
|
|
# NeuralWatt catalog still upserted
|
|
assert conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0] == 1
|
|
# No balance row because the poll failed
|
|
assert conn.execute(
|
|
"SELECT COUNT(*) FROM provider_balance_observations"
|
|
).fetchone()[0] == 0
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_balance_poll_skipped_when_api_key_missing(tmp_db, monkeypatch):
|
|
"""Without OPENROUTER_API_KEY the balance poll is skipped but the run continues."""
|
|
connect, cfg = tmp_db
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(poller.requests, "get", _url_keyed_get)
|
|
monkeypatch.delenv("OPENROUTER_API_KEY", raising=False)
|
|
|
|
import io
|
|
|
|
stderr_capture = io.StringIO()
|
|
monkeypatch.setattr("sys.stderr", stderr_capture)
|
|
|
|
exit_code = poller.main()
|
|
assert exit_code == 0
|
|
assert "[openrouter] balance poll skipped: OPENROUTER_API_KEY not set" in stderr_capture.getvalue()
|
|
|
|
conn = connect()
|
|
try:
|
|
assert conn.execute(
|
|
"SELECT COUNT(*) FROM provider_balance_observations"
|
|
).fetchone()[0] == 0
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def test_balance_parser_drift_caught_not_fatal(tmp_db, monkeypatch):
|
|
"""A malformed credits response is logged and the run continues."""
|
|
connect, cfg = tmp_db
|
|
|
|
def _get(url, timeout, **kwargs):
|
|
if "/credits" in url:
|
|
return _fake_response({"data": {}})
|
|
if "openrouter.ai" in url:
|
|
return _fake_response({"data": [_openrouter_model("openai/gpt-6-astra")]})
|
|
if "neuralwatt" in url:
|
|
return _fake_response({"data": [_catalog_item("nw-model-a")]})
|
|
raise RuntimeError(f"unexpected URL: {url}")
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(poller.requests, "get", _get)
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-test")
|
|
|
|
import io
|
|
|
|
stderr_capture = io.StringIO()
|
|
monkeypatch.setattr("sys.stderr", stderr_capture)
|
|
|
|
exit_code = poller.main()
|
|
assert exit_code == 0
|
|
assert "[openrouter] balance poll FAILED" in stderr_capture.getvalue()
|
|
|
|
|
|
def test_disabled_provider_with_balance_url_gets_no_http_call(tmp_db, monkeypatch):
|
|
"""An enabled=False provider with balance_url configured is skipped entirely."""
|
|
connect, cfg = tmp_db
|
|
|
|
or_cfg = REAL_CFG.dispatch_providers["openrouter"]
|
|
disabled_or = or_cfg.model_copy(update={"enabled": False})
|
|
cfg = cfg.model_copy(
|
|
update={
|
|
"dispatch_providers": {
|
|
"neuralwatt": REAL_CFG.dispatch_providers["neuralwatt"],
|
|
"openrouter": disabled_or,
|
|
}
|
|
}
|
|
)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
call_log: list[str] = []
|
|
|
|
def _get(url, timeout, **kwargs):
|
|
call_log.append(url)
|
|
if "neuralwatt" in url:
|
|
return _fake_response({"data": [_catalog_item("nw-model-a")]})
|
|
if "openrouter.ai" in url:
|
|
return _fake_response({"data": [_openrouter_model("openai/gpt-6-astra")]})
|
|
raise RuntimeError(f"unexpected URL: {url}")
|
|
|
|
monkeypatch.setattr(poller.requests, "get", _get)
|
|
|
|
exit_code = poller.main()
|
|
assert exit_code == 0
|
|
assert not any("/credits" in u for u in call_log)
|
|
assert not any("openrouter.ai" in u for u in call_log)
|