377 lines
13 KiB
Python
377 lines
13 KiB
Python
"""Offline tests for poller.freshness: row-count floor, staleness, recovery.
|
|
|
|
Verifies the five freshness behaviors:
|
|
(a) zero rows raises CatalogTooSmall (exit 1), nothing stale
|
|
(b) below-half warns and proceeds, mark_stale called
|
|
(c) absent-row becomes stale (mark_stale flags missing)
|
|
(d) previously-stale row recovers when re-fetched
|
|
(e) RequestException marks nothing, no upserts
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import sqlite3
|
|
import warnings
|
|
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):
|
|
"""Return an object that mimics requests.Response for poller.fetch_neuralwatt."""
|
|
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 a plausible NeuralWatt catalog item 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 _seed_models(
|
|
conn: sqlite3.Connection,
|
|
models: list[dict],
|
|
last_updated_days_ago: int = 0,
|
|
) -> None:
|
|
"""Insert or replace model rows into the DB."""
|
|
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 (?, 'neuralwatt', ?, ?, 0.30, 0.60, 0.20,
|
|
131072, 192500, 8192,
|
|
?, ?, ?, ?,
|
|
?, 'standard', 'default', 'full', 'public', ?, ?,
|
|
?, ?)
|
|
""",
|
|
(
|
|
m["model_id"],
|
|
m.get("base_model_id", m["model_id"]),
|
|
m.get(
|
|
"display_name",
|
|
f"Model {m['model_id'].title().replace('_', ' ')}",
|
|
),
|
|
int(m.get("supports_tools", True)),
|
|
int(m.get("supports_json_mode", True)),
|
|
int(m.get("supports_vision", False)),
|
|
int(m.get("supports_reasoning", False)),
|
|
int(m.get("reasoning_default_enabled", False)),
|
|
int(m.get("pricing_tbd", False)),
|
|
deprecated,
|
|
availability,
|
|
ts,
|
|
),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Fixtures
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@pytest.fixture()
|
|
def tmp_db(tmp_path, monkeypatch):
|
|
"""Create an empty temp DB with schema and a config referencing it.
|
|
|
|
Yields ``(conn, cfg)``. Caller can seed the DB; fixture closes conn at
|
|
end.
|
|
|
|
The fixture overrides ``dispatch_providers`` to contain only ``neuralwatt``
|
|
(the provider under test) so the multi-provider main loop doesn't also
|
|
try to fetch from ``openrouter`` and other providers.
|
|
"""
|
|
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": {"neuralwatt": REAL_CFG.dispatch_providers["neuralwatt"]},
|
|
}
|
|
)
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.executescript(SCHEMA_SQL)
|
|
yield conn, cfg
|
|
conn.close()
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Tests
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def test_zero_rows_exits_nonzero_and_marks_nothing_stale(tmp_db, monkeypatch):
|
|
"""(a) Zero-row fetch no longer aborts the process.
|
|
|
|
In the multi-provider main, CatalogTooSmall is caught and the provider
|
|
is skipped — the run completes with 0 rows upserted and no rows marked
|
|
stale.
|
|
"""
|
|
conn, cfg = tmp_db
|
|
_seed_models(conn, [
|
|
{"model_id": "model_alpha"},
|
|
{"model_id": "model_beta"},
|
|
], last_updated_days_ago=10)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(poller.requests, "get", lambda url, timeout: _fake_response({"data": []}))
|
|
|
|
exit_code = poller.main()
|
|
|
|
total = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
stale = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale'"
|
|
).fetchone()[0]
|
|
|
|
assert total == 2
|
|
assert stale == 0
|
|
|
|
|
|
def test_below_half_warns_but_proceeds(tmp_db, monkeypatch):
|
|
"""(b) Fetched count < DB count // 2 -> warning + proceed.
|
|
|
|
Seeds 10 models; fetch returns 4 (4 < 10//2=5 -> triggers warn).
|
|
mark_stale must be called.
|
|
"""
|
|
conn, cfg = tmp_db
|
|
_seed_models(conn, [
|
|
{"model_id": f"model_{i}"} for i in range(10)
|
|
], last_updated_days_ago=10)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response(
|
|
{"data": [_catalog_item(f"model_{i}") for i in range(4)]}
|
|
),
|
|
)
|
|
|
|
with warnings.catch_warnings(record=True) as caught:
|
|
warnings.simplefilter("always")
|
|
exit_code = poller.main()
|
|
|
|
assert exit_code == 0
|
|
assert any("may be truncated" in str(w.message) for w in caught)
|
|
|
|
|
|
def test_sanity_floor_counts_only_configured_provider(tmp_db, monkeypatch):
|
|
conn, cfg = tmp_db
|
|
_seed_models(conn, [
|
|
{"model_id": f"model_{i}"} for i in range(10)
|
|
], last_updated_days_ago=10)
|
|
other_ts = (datetime.now(timezone.utc) - timedelta(days=10)).isoformat()
|
|
for i in range(10):
|
|
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 (?, 'other-provider', ?, ?, 0.30, 0.60, 0.20,
|
|
131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'active', ?)
|
|
""",
|
|
(f"other_{i}", f"other_{i}", f"Other {i}", other_ts),
|
|
)
|
|
conn.commit()
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response(
|
|
{"data": [_catalog_item(f"model_{i}") for i in range(2)]}
|
|
),
|
|
)
|
|
|
|
with warnings.catch_warnings(record=True) as caught:
|
|
warnings.simplefilter("always")
|
|
exit_code = poller.main()
|
|
|
|
assert exit_code == 0
|
|
assert any(
|
|
"[neuralwatt] fetched only 2 models (current DB has 10)" in str(w.message)
|
|
for w in caught
|
|
), caught
|
|
|
|
|
|
def test_absent_row_becomes_stale(tmp_db, monkeypatch):
|
|
"""(c) Row present in DB but absent from fetch becomes stale.
|
|
|
|
Seeds model_a + model_b (active) with a last_updated age above the
|
|
freshness threshold. Fetch returns only model_a; model_b becomes stale,
|
|
model_a stays active.
|
|
"""
|
|
conn, cfg = tmp_db
|
|
monkeypatch.setattr(cfg.freshness, "stale_after_days", 5)
|
|
_seed_models(conn, [
|
|
{"model_id": "model_a"},
|
|
{"model_id": "model_b"},
|
|
], last_updated_days_ago=10)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response({"data": [_catalog_item("model_a")]}),
|
|
)
|
|
|
|
poller.main()
|
|
|
|
total = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
stale = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
|
|
assert total == 2, f"total: {total}"
|
|
assert stale == 1, f"stale: {stale}"
|
|
neuralwatt_active = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='active' AND provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
assert neuralwatt_active == 1, f"neuralwatt active: {neuralwatt_active}"
|
|
|
|
|
|
def test_previously_stale_row_recovers(tmp_db, monkeypatch):
|
|
"""(d) Previously-stale row returns to active when fetched again.
|
|
|
|
The DB row was staled earlier and not touched since; upsert resets its
|
|
availability to active.
|
|
"""
|
|
conn, cfg = tmp_db
|
|
ts = (datetime.now(timezone.utc) - timedelta(days=180)).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 ('recover_a', 'neuralwatt', 'recover_a', 'Recover A',
|
|
0.30, 0.60, 0.20, 131072, 192500, 8192,
|
|
1, 1, 0, 0,
|
|
0, 'standard', 'default', 'full', 'public', 0, 0,
|
|
'stale', ?)
|
|
""",
|
|
(ts,),
|
|
)
|
|
conn.commit()
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: _fake_response({"data": [_catalog_item("recover_a")]}),
|
|
)
|
|
|
|
poller.main()
|
|
|
|
row = conn.execute(
|
|
"SELECT availability FROM models WHERE model_id='recover_a'"
|
|
).fetchone()
|
|
assert row is not None
|
|
assert row["availability"] == "active"
|
|
|
|
|
|
def test_request_exception_skip_provider(tmp_db, monkeypatch):
|
|
"""(e) RequestException during fetch -> skip provider, DB untouched.
|
|
|
|
With multi-provider main, a fetch error skips the provider rather than
|
|
aborting the whole run. Existing rows stay exactly as they were.
|
|
"""
|
|
conn, cfg = tmp_db
|
|
_seed_models(conn, [
|
|
{"model_id": "stable_a"},
|
|
{"model_id": "stable_b"},
|
|
], last_updated_days_ago=14)
|
|
|
|
monkeypatch.setattr(poller, "load_config", lambda p: cfg)
|
|
monkeypatch.setattr(
|
|
poller.requests,
|
|
"get",
|
|
lambda url, timeout: (_ for _ in ()).throw(
|
|
requests.RequestException("offline")
|
|
),
|
|
)
|
|
|
|
exit_code = poller.main()
|
|
|
|
total = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
stale = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='stale' AND provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
|
|
assert exit_code == 0
|
|
assert total == 2
|
|
assert stale == 0
|
|
neuralwatt_active = conn.execute(
|
|
"SELECT COUNT(*) FROM models WHERE availability='active' AND provider='neuralwatt'"
|
|
).fetchone()[0]
|
|
assert neuralwatt_active == 2
|