test: add scenario eval runner, recorded fixtures, and baseline comparison
This commit is contained in:
parent
3dc1af5f0c
commit
0664cc2d27
38 changed files with 6986 additions and 5 deletions
284
eval/live.py
Normal file
284
eval/live.py
Normal file
|
|
@ -0,0 +1,284 @@
|
|||
"""Exercise the live recommendation stream and assert behavior properties."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
import httpx
|
||||
from scenario import Scenario
|
||||
from wire import validate_event
|
||||
|
||||
FACET_TERMS: dict[tuple[str, str], tuple[str, ...]] = {
|
||||
("activity", "coding"): ("coding", "programming"),
|
||||
("activity", "workout"): ("workout", "gym"),
|
||||
("familiarity", "mix"): ("mix", "familiar", "new", "discover"),
|
||||
("familiarity", "new"): ("new", "unfamiliar", "discover", "surprise"),
|
||||
("language", "nl"): ("dutch", "nederlandse"),
|
||||
("era", "1990s"): ("1990", "nineties", "90s"),
|
||||
("mood", "energetic"): ("energetic", "high-energy", "high energy"),
|
||||
("mood", "relaxed"): ("relaxed", "laid-back", "laid back", "couch-friendly"),
|
||||
("mood", "subdued"): ("subdued", "background", "unobtrusive", "not too present"),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LiveLimits:
|
||||
"""Quality and latency thresholds for one live request."""
|
||||
|
||||
minimum_tracks: int
|
||||
maximum_tracks: int
|
||||
minimum_artists: int
|
||||
first_track_budget_ms: int
|
||||
total_budget_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScenarioResult:
|
||||
"""Serializable observations and property failures for one scenario."""
|
||||
|
||||
scenario: Scenario
|
||||
events: list[dict[str, object]] = field(default_factory=list)
|
||||
failures: list[str] = field(default_factory=list)
|
||||
checks: dict[str, bool] = field(default_factory=dict)
|
||||
first_track_ms: int | None = None
|
||||
total_ms: int | None = None
|
||||
|
||||
@property
|
||||
def tracks(self) -> list[dict[str, object]]:
|
||||
"""Return track event payloads in streamed order."""
|
||||
return [event for event in self.events if event.get("type") == "track"]
|
||||
|
||||
@property
|
||||
def passed(self) -> bool:
|
||||
"""Return whether all required properties passed."""
|
||||
return not self.failures
|
||||
|
||||
def as_report(self) -> dict[str, object]:
|
||||
"""Render the stable JSON report shape."""
|
||||
return {
|
||||
"key": self.scenario.key,
|
||||
"query": self.scenario.query,
|
||||
"after": self.scenario.after,
|
||||
"status": "passed" if self.passed else "failed",
|
||||
"first_track_ms": self.first_track_ms,
|
||||
"total_ms": self.total_ms,
|
||||
"checks": self.checks,
|
||||
"failures": self.failures,
|
||||
"events": self.events,
|
||||
"tracks": self.tracks,
|
||||
}
|
||||
|
||||
|
||||
async def run_live_scenarios(
|
||||
scenarios: tuple[Scenario, ...],
|
||||
base_url: str,
|
||||
report_root: Path,
|
||||
limits: LiveLimits,
|
||||
) -> tuple[ScenarioResult, ...]:
|
||||
"""Run all scenarios sequentially and write one report for each."""
|
||||
report_root.mkdir(parents=True, exist_ok=True)
|
||||
results_by_key: dict[str, ScenarioResult] = {}
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=None) as client:
|
||||
for scenario in scenarios:
|
||||
parent = results_by_key.get(scenario.after) if scenario.after is not None else None
|
||||
result = await run_live_scenario(client, scenario, limits, parent)
|
||||
_write_report(report_root / f"{scenario.key}.json", result.as_report())
|
||||
results_by_key[scenario.key] = result
|
||||
return tuple(results_by_key[scenario.key] for scenario in scenarios)
|
||||
|
||||
|
||||
async def run_live_scenario(
|
||||
client: httpx.AsyncClient,
|
||||
scenario: Scenario,
|
||||
limits: LiveLimits,
|
||||
parent: ScenarioResult | None = None,
|
||||
) -> ScenarioResult:
|
||||
"""Run one strict NDJSON request and evaluate its properties."""
|
||||
result = ScenarioResult(scenario)
|
||||
if scenario.after is not None and (parent is None or not parent.tracks):
|
||||
result.failures.append("parent scenario did not produce usable recommendations")
|
||||
return result
|
||||
payload = build_request(scenario, parent)
|
||||
started_at = time.monotonic()
|
||||
try:
|
||||
async with asyncio.timeout(limits.total_budget_ms / 1000):
|
||||
async with client.stream("POST", "/api/recommendations", json=payload) as response:
|
||||
if response.status_code != 200:
|
||||
body = (await response.aread()).decode("utf-8", errors="replace")
|
||||
result.failures.append(f"HTTP {response.status_code}: {body}")
|
||||
else:
|
||||
await _consume_response(response, result, started_at)
|
||||
except TimeoutError:
|
||||
result.failures.append("stream exceeded the configured total latency budget")
|
||||
except (httpx.HTTPError, ValueError) as error:
|
||||
result.failures.append(f"stream failed: {error}")
|
||||
result.total_ms = _elapsed_ms(started_at)
|
||||
_evaluate_result(result, limits)
|
||||
return result
|
||||
|
||||
|
||||
def build_request(scenario: Scenario, parent: ScenarioResult | None) -> dict[str, object]:
|
||||
"""Build the frozen request shape, including client-owned turn context."""
|
||||
payload: dict[str, object] = {"schema_version": 1, "query": scenario.query}
|
||||
if parent is None:
|
||||
return payload
|
||||
prior = []
|
||||
labels = []
|
||||
for event in parent.tracks:
|
||||
track = event["track"]
|
||||
if not isinstance(track, dict):
|
||||
continue
|
||||
rank = event["rank"]
|
||||
artists = track["artists"]
|
||||
prior_artists = list(artists)[:10] if isinstance(artists, list) else []
|
||||
prior.append(
|
||||
{
|
||||
"rank": rank,
|
||||
"track_id": track["id"],
|
||||
"title": track["title"],
|
||||
"artists": prior_artists,
|
||||
}
|
||||
)
|
||||
labels.append(
|
||||
f"{rank}. {track['title']} by {', '.join(str(name) for name in prior_artists)}"
|
||||
)
|
||||
payload["history"] = [
|
||||
{"role": "user", "content": parent.scenario.query},
|
||||
{"role": "assistant", "content": "\n".join(labels)[:2000]},
|
||||
]
|
||||
payload["prior_recommendations"] = prior
|
||||
return payload
|
||||
|
||||
|
||||
def _evaluate_result(result: ScenarioResult, limits: LiveLimits) -> None:
|
||||
event_types = [str(event.get("type")) for event in result.events]
|
||||
terminal_type = event_types[-1] if event_types else None
|
||||
order_is_valid = (
|
||||
bool(event_types)
|
||||
and event_types[0] == "metadata"
|
||||
and terminal_type in {"done", "error"}
|
||||
and all(event_type in {"track", "warning"} for event_type in event_types[1:-1])
|
||||
and event_types.count("metadata") == 1
|
||||
)
|
||||
_record_check(result, "event_order", order_is_valid, "event order is invalid")
|
||||
if terminal_type == "error":
|
||||
terminal = result.events[-1]
|
||||
result.failures.append(f"terminal error: {terminal.get('code', 'unknown')}")
|
||||
|
||||
tracks = result.tracks
|
||||
if terminal_type == "done":
|
||||
done_count = result.events[-1].get("track_count")
|
||||
_record_check(
|
||||
result,
|
||||
"done_track_count",
|
||||
done_count == len(tracks),
|
||||
"done event track count did not match streamed tracks",
|
||||
)
|
||||
track_ids = [_track_field(event, "id") for event in tracks]
|
||||
_record_check(
|
||||
result,
|
||||
"unique_track_ids",
|
||||
len(track_ids) == len(set(track_ids)),
|
||||
"duplicate track ids were returned",
|
||||
)
|
||||
justifications = [event.get("justification") for event in tracks]
|
||||
_record_check(
|
||||
result,
|
||||
"non_empty_justifications",
|
||||
all(isinstance(value, str) and bool(value.strip()) for value in justifications),
|
||||
"a track justification was empty",
|
||||
)
|
||||
_record_check(
|
||||
result,
|
||||
"track_count",
|
||||
limits.minimum_tracks <= len(tracks) <= limits.maximum_tracks,
|
||||
f"track count {len(tracks)} is outside configured bounds",
|
||||
)
|
||||
artists = {str(artist).casefold() for event in tracks for artist in _track_artists(event)}
|
||||
_record_check(
|
||||
result,
|
||||
"distinct_artists",
|
||||
len(artists) >= limits.minimum_artists,
|
||||
f"distinct artist count {len(artists)} is below configured minimum",
|
||||
)
|
||||
_record_check(
|
||||
result,
|
||||
"first_track_latency",
|
||||
result.first_track_ms is not None and result.first_track_ms <= limits.first_track_budget_ms,
|
||||
"time to first track exceeded the configured budget",
|
||||
)
|
||||
_record_check(
|
||||
result,
|
||||
"total_latency",
|
||||
result.total_ms is not None and result.total_ms <= limits.total_budget_ms,
|
||||
"total latency exceeded the configured budget",
|
||||
)
|
||||
_evaluate_facets(result)
|
||||
|
||||
|
||||
def _evaluate_facets(result: ScenarioResult) -> None:
|
||||
metadata = next(
|
||||
(event for event in result.events if event.get("type") == "metadata"),
|
||||
None,
|
||||
)
|
||||
summary = str(metadata.get("intent_summary", "")).casefold() if metadata else ""
|
||||
for name, value in result.scenario.facets.items():
|
||||
if not isinstance(value, str):
|
||||
continue
|
||||
terms = FACET_TERMS.get((name, value), (value.casefold(),))
|
||||
_record_check(
|
||||
result,
|
||||
f"facet_{name}",
|
||||
any(term in summary for term in terms),
|
||||
f"intent summary did not express {name}={value}",
|
||||
)
|
||||
|
||||
|
||||
async def _consume_response(
|
||||
response: httpx.Response,
|
||||
result: ScenarioResult,
|
||||
started_at: float,
|
||||
) -> None:
|
||||
content_type = response.headers.get("content-type", "").split(";", 1)[0]
|
||||
if content_type != "application/x-ndjson":
|
||||
result.failures.append(f"unexpected content type: {content_type or 'missing'}")
|
||||
return
|
||||
async for line in response.aiter_lines():
|
||||
if not line:
|
||||
result.failures.append("NDJSON stream contained an empty line")
|
||||
continue
|
||||
event = validate_event(line)
|
||||
event_payload = event.model_dump(mode="json")
|
||||
result.events.append(event_payload)
|
||||
if event_payload["type"] == "track" and result.first_track_ms is None:
|
||||
result.first_track_ms = _elapsed_ms(started_at)
|
||||
|
||||
|
||||
def _record_check(result: ScenarioResult, name: str, passed: bool, failure: str) -> None:
|
||||
result.checks[name] = passed
|
||||
if not passed:
|
||||
result.failures.append(failure)
|
||||
|
||||
|
||||
def _track_field(event: dict[str, object], key: str) -> str:
|
||||
track = event.get("track")
|
||||
return str(track.get(key, "")) if isinstance(track, dict) else ""
|
||||
|
||||
|
||||
def _track_artists(event: dict[str, object]) -> list[object]:
|
||||
track = event.get("track")
|
||||
artists = track.get("artists") if isinstance(track, dict) else None
|
||||
return artists if isinstance(artists, list) else []
|
||||
|
||||
|
||||
def _elapsed_ms(started_at: float) -> int:
|
||||
return round((time.monotonic() - started_at) * 1000)
|
||||
|
||||
|
||||
def _write_report(path: Path, report: dict[str, object]) -> None:
|
||||
path.write_text(
|
||||
json.dumps(report, ensure_ascii=True, indent=2, sort_keys=True) + "\n",
|
||||
encoding="ascii",
|
||||
)
|
||||
Loading…
Add table
Add a link
Reference in a new issue