"""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", )