431 lines
16 KiB
Python
431 lines
16 KiB
Python
"""Tests for which completion an outcome report attaches to.
|
|
|
|
`POST /outcome` is the only ground truth the router gets, so attributing one
|
|
to the wrong conversation is worse than losing it: a model gets penalized for
|
|
work it never did. The guard is a short window plus a refusal -- if more than
|
|
one conversation was served inside it, the report is refused rather than
|
|
guessed at.
|
|
|
|
The window was not a window. `observed_at` is written by
|
|
`datetime.now(timezone.utc).isoformat()`, which separates date from time with
|
|
'T', while `datetime('now', ...)` returns a space. Compared as strings, 'T'
|
|
sorts after ' ', so the time of day never participated once the dates matched
|
|
and a 120-second window admitted everything served that day. Measured on the
|
|
live DB: 14 rows across 2 sessions where the correct comparison matched 0.
|
|
|
|
These tests write timestamps through the same call the dispatcher uses, so
|
|
they stay honest if that format ever changes.
|
|
"""
|
|
|
|
import sqlite3
|
|
from datetime import datetime, timedelta, timezone
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
from starlette.testclient import TestClient
|
|
|
|
import dispatcher
|
|
from dispatcher import (
|
|
AMBIGUOUS,
|
|
SEED_CATEGORY,
|
|
_find_outcome_row,
|
|
_most_recent_if_unambiguous,
|
|
)
|
|
from metrics import quota_burn
|
|
|
|
ROOT = Path(__file__).resolve().parent.parent
|
|
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
|
|
|
|
# Short, so "earlier today" is unambiguously outside it. The old row below sits
|
|
# at the first instant of the current UTC day -- the earliest time that still
|
|
# shares today's date, which is what the string comparison needed to go wrong.
|
|
WINDOW_SECONDS = 5
|
|
|
|
|
|
def _now() -> datetime:
|
|
return datetime.now(timezone.utc)
|
|
|
|
|
|
def _start_of_today() -> datetime:
|
|
return _now().replace(hour=0, minute=0, second=0, microsecond=1)
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def short_window(monkeypatch):
|
|
monkeypatch.setattr(
|
|
dispatcher.cfg.verification,
|
|
"outcome_attribution_window_seconds",
|
|
WINDOW_SECONDS,
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def db(tmp_path):
|
|
conn = sqlite3.connect(tmp_path / "test.db")
|
|
conn.row_factory = sqlite3.Row
|
|
conn.executescript(SCHEMA_SQL)
|
|
yield conn
|
|
conn.close()
|
|
|
|
|
|
def _observe(conn, *, when: datetime, session_key: str, request_id: str,
|
|
category: str = "coding_general", session_dir: object = None):
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO energy_observations (
|
|
model_id, provider, request_id, session_key, task_category,
|
|
session_dir, observed_at
|
|
) VALUES ('m', 'neuralwatt', ?, ?, ?, ?, ?)
|
|
""",
|
|
# Written exactly the way log_observation writes it.
|
|
(request_id, session_key, category, session_dir, when.isoformat()),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def test_a_session_from_earlier_today_is_outside_the_window(db):
|
|
"""The regression: same calendar date is not the same as inside 5 seconds."""
|
|
_observe(db, when=_start_of_today(), session_key="morning", request_id="old")
|
|
_observe(db, when=_now(), session_key="now", request_id="new")
|
|
|
|
row = _most_recent_if_unambiguous(db)
|
|
|
|
assert row is not AMBIGUOUS, "an 8-hours-stale session must not create ambiguity"
|
|
assert row["request_id"] == "new"
|
|
|
|
|
|
def test_an_old_session_alone_is_not_recent_enough_to_attribute(db):
|
|
"""With nothing inside the window there is nothing to attach to -- 404, not a guess."""
|
|
_observe(db, when=_start_of_today(), session_key="morning", request_id="old")
|
|
|
|
assert _most_recent_if_unambiguous(db) is None
|
|
|
|
|
|
def test_two_live_conversations_are_refused(db):
|
|
"""The guard the window exists to serve, still working."""
|
|
_observe(db, when=_now(), session_key="alice", request_id="a")
|
|
_observe(db, when=_now(), session_key="bob", request_id="b")
|
|
|
|
assert _most_recent_if_unambiguous(db) is AMBIGUOUS
|
|
|
|
|
|
def test_one_conversation_across_several_turns_is_not_ambiguous(db):
|
|
"""Several completions, one session -- the ordinary case must still attribute."""
|
|
for i in range(3):
|
|
_observe(db, when=_now(), session_key="alice", request_id=f"a{i}")
|
|
|
|
row = _most_recent_if_unambiguous(db)
|
|
|
|
assert row is not AMBIGUOUS
|
|
assert row["request_id"] == "a2"
|
|
|
|
|
|
def test_the_reference_sweep_is_not_a_conversation(db):
|
|
"""seed_energy's traffic must never make a real report look ambiguous."""
|
|
_observe(db, when=_now(), session_key=None, request_id="seed",
|
|
category=SEED_CATEGORY)
|
|
_observe(db, when=_now(), session_key="alice", request_id="a")
|
|
|
|
row = _most_recent_if_unambiguous(db)
|
|
|
|
assert row is not AMBIGUOUS
|
|
assert row["request_id"] == "a"
|
|
|
|
|
|
def test_quota_burn_counts_only_the_last_thirty_days(db, tmp_path, monkeypatch):
|
|
"""Same comparison, same fix -- an old row must not inflate the burn figure.
|
|
|
|
The row sits at the START of the day 30 days ago: the cutoff falls on that
|
|
same date but later in it, which is exactly the case a string comparison
|
|
gets wrong. A row from 45 days ago would be excluded either way and would
|
|
pin nothing.
|
|
"""
|
|
monkeypatch.setattr(dispatcher.cfg.objective, "plan_kwh_per_period", 6.25)
|
|
just_outside = (_now() - timedelta(days=30)).replace(
|
|
hour=0, minute=0, second=0, microsecond=1
|
|
)
|
|
db.execute(
|
|
"""
|
|
INSERT INTO energy_observations (model_id, provider, energy_kwh, observed_at)
|
|
VALUES ('m', 'neuralwatt', 1.0, ?)
|
|
""",
|
|
(just_outside.isoformat(),),
|
|
)
|
|
db.execute(
|
|
"""
|
|
INSERT INTO energy_observations (model_id, provider, energy_kwh, observed_at)
|
|
VALUES ('m', 'neuralwatt', 0.25, ?)
|
|
""",
|
|
(_now().isoformat(),),
|
|
)
|
|
db.commit()
|
|
|
|
assert quota_burn(db, dispatcher.cfg)["metered_kwh_30d"] == pytest.approx(0.25)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Degraded classifications must not train proficiency.
|
|
# ---------------------------------------------------------------------------
|
|
#
|
|
# A fallback-classified request is recorded with task_category=general_chat,
|
|
# and general_chat is a fully scored category AND the configured
|
|
# fallback_category. Without this guard, outcome reports on mislabelled outage
|
|
# traffic fold into proficiency(model, general_chat) via feedback.py and drag
|
|
# real scores toward whatever was flowing while the classifier was down.
|
|
|
|
|
|
def _decision_with_source(conn, request_id: str, source: str) -> None:
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO route_decisions
|
|
(kind, task_category, task_tier, required_context_tokens,
|
|
selected_model, selected_provider, classification_source,
|
|
request_id, observed_at)
|
|
VALUES ('chat','general_chat',2,100,'m','neuralwatt',?,?,?)
|
|
""",
|
|
(source, request_id, _now().isoformat()),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
@pytest.mark.parametrize("source", ["fallback", "session_stale", "session_history"])
|
|
def test_degraded_source_outcome_is_not_attributable(db, source):
|
|
"""A guessed or borrowed category must not move proficiency."""
|
|
_decision_with_source(db, "rid-degraded", source)
|
|
assert dispatcher._outcome_is_attributable(db, "rid-degraded") is False
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"source", ["classifier", "cached", "override", "classifier_cloud"]
|
|
)
|
|
def test_trusted_source_outcome_is_attributable(db, source):
|
|
"""A measured category still trains proficiency, cloud included.
|
|
|
|
classifier_cloud is a real classification by a real model — degrading its
|
|
attribution too would throw away good ground truth.
|
|
"""
|
|
_decision_with_source(db, "rid-trusted", source)
|
|
assert dispatcher._outcome_is_attributable(db, "rid-trusted") is True
|
|
|
|
|
|
def test_unknown_decision_row_fails_open(db):
|
|
"""No matching decision -> attributable.
|
|
|
|
Withholding ground truth needs positive knowledge the category was a
|
|
guess. Over-applying an exclusion already starved feedback once (the
|
|
client_capped lesson), so the default direction is to keep the signal.
|
|
|
|
NOTE: this test passes on unmodified main by design — it pins the
|
|
fail-open rule against a future over-correction, not a fixed defect.
|
|
"""
|
|
assert dispatcher._outcome_is_attributable(db, "rid-does-not-exist") is True
|
|
assert dispatcher._outcome_is_attributable(db, None) is True
|
|
|
|
|
|
def test_degraded_sources_match_the_cascade_steps():
|
|
"""The exclusion set and the cascade's degraded steps cannot drift apart.
|
|
|
|
If a new degraded source is added to Classification without adding it
|
|
here, its outcomes silently start training proficiency again.
|
|
"""
|
|
literal_sources = set(
|
|
dispatcher.Classification.model_fields["source"].annotation.__args__
|
|
)
|
|
trusted = {"classifier", "override", "cached", "classifier_cloud"}
|
|
assert dispatcher.DEGRADED_CLASSIFICATION_SOURCES == literal_sources - trusted
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Conversation-id attribution (Contract rule 4)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def _conversation_key(conversation_id: str) -> str:
|
|
return "c:" + conversation_id
|
|
|
|
|
|
def _observe_local(conn, *, when: datetime, request_id: str,
|
|
session_key: object = None, session_dir: object = None,
|
|
call_type: str = "coding_general"):
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO local_energy_observations (
|
|
model_id, call_type, request_id, session_dir, session_key, observed_at
|
|
) VALUES ('m', ?, ?, ?, ?, ?)
|
|
""",
|
|
(call_type, request_id, session_dir, session_key, when.isoformat()),
|
|
)
|
|
conn.commit()
|
|
|
|
|
|
def _open_db(path) -> sqlite3.Connection:
|
|
conn = sqlite3.connect(path)
|
|
conn.row_factory = sqlite3.Row
|
|
return conn
|
|
|
|
|
|
def test_conversation_id_resolves_exact_conversation_over_shared_source_dir(db):
|
|
"""Two conversations under one cwd: naming one picks it exactly.
|
|
|
|
Both rows sit inside the window and share a session_dir, so the old path
|
|
(no request_id, no conversation_id) can only refuse -- two live
|
|
conversations is exactly the ambiguous case. conversation_id cuts through.
|
|
"""
|
|
now = _now()
|
|
_observe(db, when=now, session_key=_conversation_key("older"),
|
|
request_id="old", session_dir="/shared")
|
|
_observe(db, when=now, session_key=_conversation_key("newer"),
|
|
request_id="new", session_dir="/shared")
|
|
|
|
# Both live: the fallback can't pick a winner -- the 409 case.
|
|
assert _most_recent_if_unambiguous(db) is AMBIGUOUS
|
|
|
|
row = _find_outcome_row(db, None, None, conversation_id="older")
|
|
|
|
assert row is not None and row is not AMBIGUOUS
|
|
assert row["request_id"] == "old"
|
|
|
|
|
|
def test_conversation_id_resolves_even_when_request_id_omitted(db):
|
|
"""A lone conversation is picked by conversation_id alone (no request_id)."""
|
|
now = _now()
|
|
_observe(db, when=now, session_key=_conversation_key("conv"), request_id="rid")
|
|
|
|
row = _find_outcome_row(db, None, None, conversation_id="conv")
|
|
|
|
assert row is not None and row is not AMBIGUOUS
|
|
assert row["request_id"] == "rid"
|
|
|
|
|
|
def test_request_id_wins_over_conversation_id(db):
|
|
"""Contract rule 4: request_id resolves first, even when it disagrees."""
|
|
now = _now()
|
|
_observe(db, when=now, session_key=_conversation_key("by-conv"), request_id="want")
|
|
_observe(db, when=now, session_key=_conversation_key("other"), request_id="other")
|
|
|
|
row = _find_outcome_row(db, "want", None, conversation_id="other")
|
|
|
|
assert row["request_id"] == "want"
|
|
|
|
|
|
def test_local_dispatch_row_attributed_by_conversation_id(db):
|
|
"""A local-dispatch answer reports in by conversation_id too."""
|
|
now = _now()
|
|
_observe_local(db, when=now, session_key=_conversation_key("conv"),
|
|
request_id="local-rid")
|
|
|
|
row = _find_outcome_row(db, None, None, conversation_id="conv")
|
|
|
|
assert row is not None and row is not AMBIGUOUS
|
|
assert row["provider"] == "ollama-local"
|
|
assert row["request_id"] == "local-rid"
|
|
|
|
|
|
def test_no_header_local_row_resolves_through_session_dir(db):
|
|
"""A local row with NULL session_key (no header) still resolves.
|
|
|
|
COALESCE(session_key, session_dir) keeps pre-header local rows joinable by
|
|
their working directory, exactly as the old session_dir-only expression did.
|
|
"""
|
|
now = _now()
|
|
_observe_local(db, when=now, session_key=None, session_dir="/work",
|
|
request_id="rid")
|
|
|
|
row = _most_recent_if_unambiguous(db)
|
|
|
|
assert row is not None and row is not AMBIGUOUS
|
|
assert row["request_id"] == "rid"
|
|
|
|
|
|
@pytest.fixture
|
|
def client(tmp_path, monkeypatch):
|
|
"""A TestClient wired to a temp DB seeded from schema.sql."""
|
|
conn = _open_db(tmp_path / "test.db")
|
|
conn.executescript(SCHEMA_SQL)
|
|
conn.close()
|
|
monkeypatch.setattr(dispatcher.cfg.database, "path", str(tmp_path / "test.db"))
|
|
monkeypatch.setenv("NEURALWATT_API_KEY", "test-key")
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "test-key")
|
|
with TestClient(dispatcher.app) as c:
|
|
yield c
|
|
|
|
|
|
def test_unknown_conversation_id_404_and_writes_nothing(client, tmp_path):
|
|
"""An unmatched conversation_id is a 404, and records no outcome."""
|
|
resp = client.post("/outcome", json={"ok": True, "conversation_id": "no-such"})
|
|
|
|
assert resp.status_code == 404
|
|
with _open_db(tmp_path / "test.db") as conn:
|
|
assert conn.execute("SELECT COUNT(*) FROM verifications").fetchone()[0] == 0
|
|
|
|
|
|
def test_invalid_conversation_id_is_422(client):
|
|
"""A present-but-invalid conversation_id is a client bug: 422."""
|
|
resp = client.post(
|
|
"/outcome", json={"ok": True, "conversation_id": "has a space"}
|
|
)
|
|
assert resp.status_code == 422
|
|
|
|
|
|
def test_h3_local_table_predates_session_key_migrates_on_first_request(
|
|
tmp_path, monkeypatch
|
|
):
|
|
"""A DB whose local_energy_observations lacks session_key still works.
|
|
|
|
SQLite cannot drop a column, so the table is recreated WITHOUT session_key
|
|
to simulate one that predates the column, then _ensure_tables migrates it.
|
|
Outcome attribution must not 500 against the pre-migration schema.
|
|
"""
|
|
db_path = tmp_path / "test.db"
|
|
conn = _open_db(db_path)
|
|
conn.executescript(SCHEMA_SQL)
|
|
conn.execute("DROP TABLE local_energy_observations")
|
|
conn.execute(
|
|
"""
|
|
CREATE TABLE local_energy_observations (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
model_id TEXT NOT NULL,
|
|
call_type TEXT NOT NULL,
|
|
request_id TEXT,
|
|
session_dir TEXT,
|
|
avg_power_watts REAL,
|
|
duration_seconds REAL,
|
|
energy_kwh REAL,
|
|
cost_usd REAL,
|
|
carbon_g_co2eq REAL,
|
|
meter TEXT,
|
|
observed_at TEXT NOT NULL
|
|
)
|
|
"""
|
|
)
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO energy_observations (
|
|
model_id, provider, request_id, session_key, task_category, observed_at
|
|
) VALUES ('m', 'neuralwatt', 'cloud-rid', 'c:conv', 'coding_general', ?)
|
|
""",
|
|
(_now().isoformat(),),
|
|
)
|
|
conn.commit()
|
|
conn.close()
|
|
|
|
# The import-time _ensure_tables already ran against the real path; point
|
|
# it at OUR temp DB so the migration covers this pre-column table.
|
|
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
|
|
dispatcher._ensure_tables()
|
|
|
|
# Idempotent: running the migration again must be a no-op.
|
|
dispatcher._ensure_tables()
|
|
|
|
monkeypatch.setenv("NEURALWATT_API_KEY", "test-key")
|
|
monkeypatch.setenv("OPENROUTER_API_KEY", "test-key")
|
|
with TestClient(dispatcher.app) as client:
|
|
seeded = client.post(
|
|
"/outcome", json={"ok": True, "request_id": "cloud-rid"}
|
|
)
|
|
assert seeded.status_code == 200
|
|
unknown = client.post(
|
|
"/outcome", json={"ok": True, "conversation_id": "unknown"}
|
|
)
|
|
assert unknown.status_code == 404
|