Files
6krrt/tests/test_multi_provider_poller.py

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)