From 751391e6a229f1b5791eef4da3e2e8ba567896f9 Mon Sep 17 00:00:00 2001 From: Justin Visser Date: Mon, 10 Aug 2026 12:05:59 +0200 Subject: [PATCH] feat: add recommendation service adapters --- backend/app/adapters/anthropic/llm.py | 339 ++++++++++++++++++++++++ backend/app/adapters/spotify/client.py | 48 +++- backend/app/adapters/spotify/errors.py | 4 +- backend/app/adapters/spotify/mapping.py | 60 ++++- backend/app/config.py | 2 + backend/app/domain/models.py | 18 ++ backend/app/observability/timing.py | 98 +++++++ backend/app/ports/protocols.py | 82 ++++++ backend/tests/test_spotify_client.py | 51 +++- 9 files changed, 676 insertions(+), 26 deletions(-) create mode 100644 backend/app/adapters/anthropic/llm.py create mode 100644 backend/app/observability/timing.py create mode 100644 backend/app/ports/protocols.py diff --git a/backend/app/adapters/anthropic/llm.py b/backend/app/adapters/anthropic/llm.py new file mode 100644 index 0000000..98f5c5e --- /dev/null +++ b/backend/app/adapters/anthropic/llm.py @@ -0,0 +1,339 @@ +"""Anthropic implementation of structured intent and streamed reranking.""" + +import json +import re +from collections.abc import AsyncIterator +from typing import Literal, cast + +from anthropic import AsyncAnthropic +from anthropic.lib._parse._transform import transform_schema +from anthropic.types import Message, TextBlockParam, Usage +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from app.config import Settings +from app.domain.models import ( + ConversationTurn, + Familiarity, + Intent, + PreviousRecommendation, + RerankSelection, + Track, + TrackCandidate, +) +from app.observability.timing import record_llm_tokens +from app.ports.protocols import RecommenderOutputError +from app.prompts import INTENT_SYSTEM_PROMPT, RERANK_SYSTEM_PROMPT + +Effort = Literal["low", "medium", "high", "xhigh", "max"] +_RECOMMENDATION_ARRAY = re.compile(r'"recommendations"\s*:\s*\[') + + +class CandidateOutput(BaseModel): + """One bounded candidate in the intent response.""" + + model_config = ConfigDict(extra="forbid") + + title: str = Field(min_length=1, max_length=200) + artist: str = Field(min_length=1, max_length=200) + + +class IntentOutput(BaseModel): + """Bounded structured output for the intent call.""" + + model_config = ConfigDict(extra="forbid") + + mood: list[str] = Field(default_factory=list, max_length=6) + activity: str | None = Field(default=None, max_length=100) + era: list[str] = Field(default_factory=list, max_length=5) + languages: list[str] = Field(default_factory=list, max_length=8) + genres: list[str] = Field(default_factory=list, max_length=8) + familiarity: Familiarity + is_refinement: bool + intent_summary: str = Field(min_length=1, max_length=300, pattern=r"^[^\r\n]+$") + candidates: list[CandidateOutput] = Field(min_length=30, max_length=40) + + +class RerankSelectionOutput(BaseModel): + """One validated object extracted from the rerank stream.""" + + model_config = ConfigDict(extra="forbid") + + track_id: str = Field(min_length=1, max_length=200) + justification: str = Field(min_length=1, max_length=300, pattern=r"^[^\r\n]+$") + + +class RerankOutput(BaseModel): + """Complete bounded rerank output used for final validation.""" + + model_config = ConfigDict(extra="forbid") + + recommendations: list[RerankSelectionOutput] = Field(min_length=1, max_length=50) + + +class AnthropicRecommender: + """Use two Anthropic calls for intent generation and grounded ranking.""" + + def __init__(self, client: AsyncAnthropic, settings: Settings) -> None: + """Bind the shared asynchronous client and immutable settings.""" + self.client = client + self.settings = settings + + async def create_intent( + self, + query: str, + history: tuple[ConversationTurn, ...], + previous_recommendations: tuple[PreviousRecommendation, ...], + taste_summary: str, + candidate_count: int, + ) -> Intent: + """Interpret a request through Anthropic structured output.""" + message = await self.client.messages.parse( + model=self.settings.llm_model, + max_tokens=self.settings.intent_max_tokens, + output_config={"effort": _parse_effort(self.settings.intent_effort)}, + output_format=IntentOutput, + system=[_system_block(INTENT_SYSTEM_PROMPT)], + messages=[ + { + "role": "user", + "content": _render_intent_input( + query, + history, + previous_recommendations, + taste_summary, + candidate_count, + ), + } + ], + ) + _validate_stop_reason(message) + _record_usage(message.usage) + parsed = message.parsed_output + if parsed is None: + raise RecommenderOutputError("Intent response contained no structured output") + validated = IntentOutput.model_validate(parsed.model_dump()) + if len(validated.candidates) != candidate_count: + raise RecommenderOutputError("Intent response returned the wrong candidate count") + return _to_intent(validated) + + async def stream_rerank( + self, + intent: Intent, + grounded_tracks: tuple[Track, ...], + taste_summary: str, + history: tuple[ConversationTurn, ...], + selection_count: int, + correction: str | None = None, + ) -> AsyncIterator[RerankSelection]: + """Yield each complete valid selection while the JSON is streaming.""" + schema = transform_schema(RerankOutput.model_json_schema()) + parser = _RecommendationObjectParser() + async with self.client.messages.stream( + model=self.settings.llm_model, + max_tokens=self.settings.rerank_max_tokens, + output_config={ + "effort": _parse_effort(self.settings.rerank_effort), + "format": {"type": "json_schema", "schema": schema}, + }, + system=[_system_block(RERANK_SYSTEM_PROMPT)], + messages=[ + { + "role": "user", + "content": _render_rerank_input( + intent, + grounded_tracks, + taste_summary, + history, + selection_count, + correction, + ), + } + ], + ) as stream: + async for text_delta in stream.text_stream: + for selection in parser.feed(text_delta): + yield RerankSelection( + track_id=selection.track_id, + justification=selection.justification, + ) + final_message = await stream.get_final_message() + + _validate_stop_reason(final_message) + _record_usage(final_message.usage) + try: + validated = RerankOutput.model_validate_json(parser.complete_text) + except ValidationError as error: + raise RecommenderOutputError("Rerank response failed final validation") from error + if len(validated.recommendations) > selection_count: + raise RecommenderOutputError("Rerank response returned too many selections") + + +class _RecommendationObjectParser: + def __init__(self) -> None: + self.complete_text = "" + self._scan_index = 0 + self._object_start: int | None = None + self._object_depth = 0 + self._is_in_string = False + self._is_escaped = False + self._has_found_array = False + + def feed(self, text_delta: str) -> list[RerankSelectionOutput]: + self.complete_text += text_delta + if not self._has_found_array: + match = _RECOMMENDATION_ARRAY.search(self.complete_text) + if match is None: + return [] + self._has_found_array = True + self._scan_index = match.end() + + selections: list[RerankSelectionOutput] = [] + while self._scan_index < len(self.complete_text): + character = self.complete_text[self._scan_index] + completed = self._scan_character(character) + self._scan_index += 1 + if completed is not None: + selections.append(completed) + return selections + + def _scan_character(self, character: str) -> RerankSelectionOutput | None: + if self._object_start is None: + if character == "{": + self._object_start = self._scan_index + self._object_depth = 1 + return None + + if self._is_in_string: + if self._is_escaped: + self._is_escaped = False + elif character == "\\": + self._is_escaped = True + elif character == '"': + self._is_in_string = False + return None + + if character == '"': + self._is_in_string = True + elif character == "{": + self._object_depth += 1 + elif character == "}": + self._object_depth -= 1 + if self._object_depth == 0: + return self._finish_object() + return None + + def _finish_object(self) -> RerankSelectionOutput: + assert self._object_start is not None + object_text = self.complete_text[self._object_start : self._scan_index + 1] + self._object_start = None + try: + return RerankSelectionOutput.model_validate_json(object_text) + except ValidationError as error: + raise RecommenderOutputError("Rerank item failed validation") from error + + +def _to_intent(output: IntentOutput) -> Intent: + return Intent( + mood=tuple(output.mood), + activity=output.activity, + era=tuple(output.era), + languages=tuple(output.languages), + genres=tuple(output.genres), + familiarity=output.familiarity, + is_refinement=output.is_refinement, + intent_summary=output.intent_summary, + candidates=tuple( + TrackCandidate(title=candidate.title, artist=candidate.artist) + for candidate in output.candidates + ), + ) + + +def _render_intent_input( + query: str, + history: tuple[ConversationTurn, ...], + previous_recommendations: tuple[PreviousRecommendation, ...], + taste_summary: str, + candidate_count: int, +) -> str: + payload = { + "query": query, + "history": [turn.__dict__ for turn in history], + "prior_recommendations": [ + recommendation.__dict__ for recommendation in previous_recommendations + ], + "taste_profile": taste_summary, + "required_candidate_count": candidate_count, + } + return json.dumps(payload, ensure_ascii=True, separators=(",", ":")) + + +def _render_rerank_input( + intent: Intent, + grounded_tracks: tuple[Track, ...], + taste_summary: str, + history: tuple[ConversationTurn, ...], + selection_count: int, + correction: str | None, +) -> str: + payload = { + "intent": { + "mood": intent.mood, + "activity": intent.activity, + "era": intent.era, + "languages": intent.languages, + "genres": intent.genres, + "familiarity": intent.familiarity, + "intent_summary": intent.intent_summary, + }, + "grounded_pool": [ + {"track_id": track.id, "title": track.title, "artists": track.artists} + for track in grounded_tracks + ], + "taste_profile": taste_summary, + "history": [turn.__dict__ for turn in history], + "requested_selection_count": selection_count, + "correction": correction, + } + return json.dumps(payload, ensure_ascii=True, separators=(",", ":")) + + +def _system_block(prompt: str) -> TextBlockParam: + return { + "type": "text", + "text": prompt, + "cache_control": {"type": "ephemeral"}, + } + + +def _parse_effort(value: str) -> Effort: + allowed = {"low", "medium", "high", "xhigh", "max"} + if value not in allowed: + raise ValueError(f"Unsupported Anthropic effort: {value}") + return cast(Effort, value) + + +def _validate_stop_reason(message: Message) -> None: + stop_reason_messages = { + None: "Anthropic response had no stop reason", + "max_tokens": "Anthropic response reached its token limit", + "stop_sequence": "Anthropic response hit an unexpected stop sequence", + "tool_use": "Anthropic response attempted tool use", + "pause_turn": "Anthropic response paused before completion", + "refusal": "Anthropic response was refused", + "model_context_window_exceeded": "Anthropic context window was exceeded", + } + if message.stop_reason == "end_turn": + return + raise RecommenderOutputError( + stop_reason_messages.get(message.stop_reason, "Anthropic response stopped unexpectedly") + ) + + +def _record_usage(usage: Usage) -> None: + cache_creation_tokens = usage.cache_creation_input_tokens or 0 + cache_read_tokens = usage.cache_read_input_tokens or 0 + record_llm_tokens( + usage.input_tokens + cache_creation_tokens + cache_read_tokens, + usage.output_tokens, + ) diff --git a/backend/app/adapters/spotify/client.py b/backend/app/adapters/spotify/client.py index c1cf9c3..5aee722 100644 --- a/backend/app/adapters/spotify/client.py +++ b/backend/app/adapters/spotify/client.py @@ -12,15 +12,22 @@ from app.adapters.spotify.errors import ( SpotifyUnavailableError, ) from app.adapters.spotify.mapping import ( - CreatedPlaylist, CurrentUser, parse_created_playlist, parse_current_user, + parse_saved_track_page, parse_search_tracks, + parse_top_artists, + parse_track_page, ) from app.adapters.spotify.session import SpotifySession from app.config import Settings -from app.domain.models import Track +from app.domain.models import CreatedPlaylist, Track +from app.observability.timing import increment_spotify_calls +from app.ports.protocols import TimeRange + +SPOTIFY_SEARCH_LIMIT = 10 +SPOTIFY_PAGE_LIMIT = 50 class SpotifyClient: @@ -39,6 +46,8 @@ class SpotifyClient: async def search_tracks(self, query: str, limit: int = 10) -> list[Track]: """Search Spotify tracks and return only valid mapped results.""" + if not 1 <= limit <= SPOTIFY_SEARCH_LIMIT: + raise ValueError("Spotify search limit must be between 1 and 10") response = await self._request( "GET", "/search", @@ -46,6 +55,40 @@ class SpotifyClient: ) return parse_search_tracks(response.json()) + async def fetch_top_artists(self, time_range: TimeRange, limit: int) -> list[str]: + """Fetch the user's top artists for a supported time range.""" + response = await self._request( + "GET", + "/me/top/artists", + params={"time_range": time_range, "limit": limit}, + ) + return parse_top_artists(response.json()) + + async def fetch_top_tracks(self, time_range: TimeRange, limit: int) -> list[Track]: + """Fetch the user's top tracks for a supported time range.""" + response = await self._request( + "GET", + "/me/top/tracks", + params={"time_range": time_range, "limit": limit}, + ) + return parse_track_page(response.json()) + + async def fetch_saved_tracks(self, limit: int) -> list[Track]: + """Fetch a bounded saved-track sample across Spotify pages.""" + tracks: list[Track] = [] + while len(tracks) < limit: + page_limit = min(SPOTIFY_PAGE_LIMIT, limit - len(tracks)) + response = await self._request( + "GET", + "/me/tracks", + params={"limit": page_limit, "offset": len(tracks)}, + ) + page = parse_saved_track_page(response.json()) + tracks.extend(page) + if len(page) < page_limit: + break + return tracks + async def fetch_current_user(self) -> CurrentUser: """Fetch the authenticated Spotify user's stable identity.""" response = await self._request("GET", "/me") @@ -110,6 +153,7 @@ class SpotifyClient: params: dict[str, str | int] | None, json: dict[str, object] | None, ) -> httpx2.Response: + increment_spotify_calls() return await self.http.request( method, f"{self.settings.spotify_api_base_url.rstrip('/')}{path}", diff --git a/backend/app/adapters/spotify/errors.py b/backend/app/adapters/spotify/errors.py index e45ad14..9665573 100644 --- a/backend/app/adapters/spotify/errors.py +++ b/backend/app/adapters/spotify/errors.py @@ -1,5 +1,7 @@ """Typed failures raised by the Spotify adapter.""" +from app.ports.protocols import CatalogQuotaExhaustedError + class SpotifyError(Exception): """Base class for Spotify adapter failures.""" @@ -9,7 +11,7 @@ class SpotifyAuthenticationError(SpotifyError): """Spotify rejected authentication or token refresh.""" -class SpotifyRateLimitedError(SpotifyError): +class SpotifyRateLimitedError(SpotifyError, CatalogQuotaExhaustedError): """Spotify rate limited a request that could not be retried.""" def __init__(self, retry_after_seconds: float | None, reason: str | None = None) -> None: diff --git a/backend/app/adapters/spotify/mapping.py b/backend/app/adapters/spotify/mapping.py index 44793dc..0ce2b0b 100644 --- a/backend/app/adapters/spotify/mapping.py +++ b/backend/app/adapters/spotify/mapping.py @@ -4,7 +4,8 @@ from collections.abc import Mapping from dataclasses import dataclass from typing import cast -from app.domain.models import Track +from app.domain.matching import track_key +from app.domain.models import CreatedPlaylist, Track @dataclass(frozen=True) @@ -15,14 +16,6 @@ class CurrentUser: display_name: str -@dataclass(frozen=True) -class CreatedPlaylist: - """The application-owned result of creating a Spotify playlist.""" - - id: str - url: str - - def parse_search_tracks(payload: object) -> list[Track]: """Map valid Spotify search items and discard malformed entries.""" root = _as_mapping(payload) @@ -32,17 +25,54 @@ def parse_search_tracks(payload: object) -> list[Track]: return [] parsed_tracks: list[Track] = [] + seen_ids: set[str] = set() + seen_keys: set[str] = set() for item in items: parsed_track = _parse_track(item) - if parsed_track is not None: - parsed_tracks.append(parsed_track) + if parsed_track is None: + continue + normalized_key = track_key(parsed_track) + if parsed_track.id in seen_ids or normalized_key in seen_keys: + continue + seen_ids.add(parsed_track.id) + seen_keys.add(normalized_key) + parsed_tracks.append(parsed_track) return parsed_tracks +def parse_top_artists(payload: object) -> list[str]: + """Map a top-artists page to valid artist names.""" + root = _as_mapping(payload) + items = root.get("items") if root is not None else None + if not isinstance(items, list): + return [] + return [ + name for item in items if (name := _required_string(_as_mapping(item), "name")) is not None + ] + + +def parse_track_page(payload: object) -> list[Track]: + """Map a direct Spotify track page to domain tracks.""" + root = _as_mapping(payload) + items = root.get("items") if root is not None else None + return _parse_track_items(items) + + +def parse_saved_track_page(payload: object) -> list[Track]: + """Map a saved-track wrapper page to domain tracks.""" + root = _as_mapping(payload) + items = root.get("items") if root is not None else None + if not isinstance(items, list): + return [] + return _parse_track_items( + [wrapper.get("track") for item in items if (wrapper := _as_mapping(item)) is not None] + ) + + def parse_current_user(payload: object) -> CurrentUser: """Map a Spotify current-user response into stable identity fields.""" root = _as_mapping(payload) - account_id = _required_string(root, "account_id") + account_id = _required_string(root, "id") or _required_string(root, "account_id") display_name = _required_string(root, "display_name") if account_id is None or display_name is None: raise ValueError("Spotify returned an invalid current-user response") @@ -90,6 +120,12 @@ def _parse_track(payload: object) -> Track | None: ) +def _parse_track_items(payload: object) -> list[Track]: + if not isinstance(payload, list): + return [] + return [track for item in payload if (track := _parse_track(item)) is not None] + + def _parse_artists(payload: object) -> tuple[str, ...] | None: if not isinstance(payload, list) or not payload: return None diff --git a/backend/app/config.py b/backend/app/config.py index 62d1620..40f7923 100644 --- a/backend/app/config.py +++ b/backend/app/config.py @@ -41,6 +41,8 @@ class Settings(BaseSettings): llm_model: str = "claude-sonnet-5" intent_effort: str = "low" rerank_effort: str = "medium" + intent_max_tokens: int = 8192 + rerank_max_tokens: int = 4096 # Pipeline shape. candidate_count is the main call-1 latency lever and # the hallucination budget: at "new to you" familiarity a large share of diff --git a/backend/app/domain/models.py b/backend/app/domain/models.py index 4ef2612..6fbdd78 100644 --- a/backend/app/domain/models.py +++ b/backend/app/domain/models.py @@ -25,6 +25,24 @@ class TrackCandidate: artist: str +@dataclass(frozen=True) +class ConversationTurn: + """One bounded user or assistant message supplied by the client.""" + + role: str + content: str + + +@dataclass(frozen=True) +class PreviousRecommendation: + """One earlier recommendation available to a refinement request.""" + + rank: int + track_id: str + title: str + artists: tuple[str, ...] + + class Familiarity(StrEnum): """How strongly a request should favor known or unknown music.""" diff --git a/backend/app/observability/timing.py b/backend/app/observability/timing.py new file mode 100644 index 0000000..315761c --- /dev/null +++ b/backend/app/observability/timing.py @@ -0,0 +1,98 @@ +"""Request timing context and the counters emitted with completion logs.""" + +import time +import uuid +from contextvars import ContextVar +from dataclasses import dataclass +from typing import cast + +import structlog +from starlette.types import ASGIApp, Receive, Scope, Send + + +@dataclass +class RequestCounters: + """Mutable counters scoped to one ASGI request context.""" + + spotify_calls: int = 0 + cache_hits: int = 0 + llm_input_tokens: int = 0 + llm_output_tokens: int = 0 + + +_COUNTERS: ContextVar[RequestCounters | None] = ContextVar("request_counters", default=None) +_REQUEST_ID: ContextVar[str | None] = ContextVar("request_id", default=None) + + +def increment_spotify_calls() -> None: + """Count one Spotify HTTP request when a request context is active.""" + counters = _COUNTERS.get() + if counters is not None: + counters.spotify_calls += 1 + + +def increment_cache_hits() -> None: + """Count one in-process cache hit when a request context is active.""" + counters = _COUNTERS.get() + if counters is not None: + counters.cache_hits += 1 + + +def record_llm_tokens(input_tokens: int, output_tokens: int) -> None: + """Accumulate model token usage when the provider reports it.""" + counters = _COUNTERS.get() + if counters is not None: + counters.llm_input_tokens += input_tokens + counters.llm_output_tokens += output_tokens + + +def current_request_id() -> str | None: + """Return the active request identifier when one exists.""" + return _REQUEST_ID.get() + + +class RequestTimingMiddleware: + """Log duration and counters after the complete response body is sent.""" + + def __init__(self, app: ASGIApp) -> None: + """Wrap an ASGI application.""" + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + """Install request context and measure an HTTP exchange.""" + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + request_id = uuid.uuid4().hex + state = scope.setdefault("state", {}) + cast(dict[str, object], state)["request_id"] = request_id + started_at = time.monotonic() + counters = RequestCounters() + counter_token = _COUNTERS.set(counters) + request_token = _REQUEST_ID.set(request_id) + try: + await self.app(scope, receive, send) + finally: + self._log_completion(scope, request_id, started_at, counters) + _COUNTERS.reset(counter_token) + _REQUEST_ID.reset(request_token) + + @staticmethod + def _log_completion( + scope: Scope, + request_id: str, + started_at: float, + counters: RequestCounters, + ) -> None: + structlog.get_logger().info( + "request_complete", + request_id=request_id, + method=scope.get("method"), + path=scope.get("path"), + duration_ms=round((time.monotonic() - started_at) * 1000), + spotify_calls=counters.spotify_calls, + cache_hits=counters.cache_hits, + llm_input_tokens=counters.llm_input_tokens, + llm_output_tokens=counters.llm_output_tokens, + ) diff --git a/backend/app/ports/protocols.py b/backend/app/ports/protocols.py new file mode 100644 index 0000000..dc40d3d --- /dev/null +++ b/backend/app/ports/protocols.py @@ -0,0 +1,82 @@ +"""Structural ports implemented by external service adapters.""" + +from collections.abc import AsyncIterator +from typing import Literal, Protocol + +from app.domain.models import ( + ConversationTurn, + CreatedPlaylist, + Intent, + PreviousRecommendation, + RerankSelection, + Track, +) + +TimeRange = Literal["short_term", "long_term"] + + +class CatalogQuotaExhaustedError(Exception): + """A catalog quota stopped further resolution attempts.""" + + +class RecommenderOutputError(Exception): + """The recommender returned unusable structured output.""" + + +class MusicCatalog(Protocol): + """Read the restricted Spotify surface used by discovery.""" + + async def search_tracks(self, query: str, limit: int = 10) -> list[Track]: + """Search tracks by fielded or bare-text query.""" + ... + + async def fetch_top_artists(self, time_range: TimeRange, limit: int) -> list[str]: + """Fetch bounded top artist names for one time range.""" + ... + + async def fetch_top_tracks(self, time_range: TimeRange, limit: int) -> list[Track]: + """Fetch bounded top tracks for one time range.""" + ... + + async def fetch_saved_tracks(self, limit: int) -> list[Track]: + """Fetch a bounded sample of saved tracks.""" + ... + + +class Recommender(Protocol): + """Interpret discovery intent and stream a grounded reranking.""" + + async def create_intent( + self, + query: str, + history: tuple[ConversationTurn, ...], + previous_recommendations: tuple[PreviousRecommendation, ...], + taste_summary: str, + candidate_count: int, + ) -> Intent: + """Interpret a query and propose a bounded candidate set.""" + ... + + def stream_rerank( + self, + intent: Intent, + grounded_tracks: tuple[Track, ...], + taste_summary: str, + history: tuple[ConversationTurn, ...], + selection_count: int, + correction: str | None = None, + ) -> AsyncIterator[RerankSelection]: + """Stream validated selections from the grounded pool.""" + ... + + +class PlaylistWriter(Protocol): + """Write a Spotify playlist without retrying ambiguous mutations.""" + + async def create_playlist(self, name: str, description: str) -> CreatedPlaylist: + """Create a private playlist.""" + ... + + async def add_tracks_to_playlist(self, playlist_id: str, track_uris: list[str]) -> None: + """Add ordered Spotify track URIs to a playlist.""" + ... diff --git a/backend/tests/test_spotify_client.py b/backend/tests/test_spotify_client.py index 14b86a0..4840673 100644 --- a/backend/tests/test_spotify_client.py +++ b/backend/tests/test_spotify_client.py @@ -228,6 +228,31 @@ def test_search_maps_valid_fields_and_drops_malformed_item() -> None: asyncio.run(run()) +def test_taste_endpoints_map_supported_response_shapes() -> None: + async def run() -> None: + async def handler(request: httpx2.Request) -> httpx2.Response: + if request.url.path == "/v1/me/top/artists": + assert request.url.params["time_range"] == "short_term" + return httpx2.Response(200, json={"items": [{"name": "Top Artist"}]}) + if request.url.path == "/v1/me/top/tracks": + return httpx2.Response(200, json={"items": [_track_payload()]}) + assert request.url.path == "/v1/me/tracks" + assert request.url.params["offset"] == "0" + return httpx2.Response(200, json={"items": [{"track": _track_payload()}]}) + + async with httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http: + client = _client(http) + artists = await client.fetch_top_artists("short_term", 50) + top_tracks = await client.fetch_top_tracks("long_term", 50) + saved_tracks = await client.fetch_saved_tracks(100) + + assert artists == ["Top Artist"] + assert [track.id for track in top_tracks] == ["track-1"] + assert [track.id for track in saved_tracks] == ["track-1"] + + asyncio.run(run()) + + async def _search_with_handler(handler: TransportHandler) -> list[Track]: async with httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http: return list(await _client(http).search_tracks("mapped")) @@ -251,18 +276,22 @@ def _search_payload() -> dict[str, object]: "tracks": { "total": 0, "items": [ - { - "id": "track-1", - "uri": "spotify:track:1", - "name": "Mapped song", - "artists": [{"name": "First artist"}, {"name": "Second artist"}], - "album": { - "name": "Mapped album", - "images": [{"url": "https://images.example/cover.jpg"}], - }, - "external_urls": {"spotify": "https://open.spotify.com/track/track-1"}, - }, + _track_payload(), {"id": "missing-required-fields"}, ], } } + + +def _track_payload() -> dict[str, object]: + return { + "id": "track-1", + "uri": "spotify:track:1", + "name": "Mapped song", + "artists": [{"name": "First artist"}, {"name": "Second artist"}], + "album": { + "name": "Mapped album", + "images": [{"url": "https://images.example/cover.jpg"}], + }, + "external_urls": {"spotify": "https://open.spotify.com/track/track-1"}, + }