Spaces:
Running
Running
Download tests/test_api.py from kxmWebwe/trackembeddingapi: direct link, hf CLI and curl.
- Browser
- Download file 16.7 kB
-
https://huggingface.co/spaces/kxmWebwe/trackembeddingapi/resolve/main/tests/test_api.py
- Command line
-
hf download hf://spaces/kxmWebwe/trackembeddingapi/tests/test_api.py
-
curl -L -o test_api.py https://huggingface.co/spaces/kxmWebwe/trackembeddingapi/resolve/main/tests/test_api.py
16.7 kB
| from __future__ import annotations | |
| from pathlib import Path | |
| import numpy as np | |
| from fastapi.testclient import TestClient | |
| from openmusic_analysis.api import create_app | |
| from openmusic_analysis.application import build_service | |
| from openmusic_analysis.audio.decoder import AudioMetadata, DecodedAudio | |
| from openmusic_analysis.domain import LyricsResolution, LyricsSourceError | |
| from openmusic_analysis.settings import ( | |
| AudioEmotionConfig, | |
| LimitConfig, | |
| Settings, | |
| TemporalAudioConfig, | |
| ) | |
| from .conftest import FailingAudioEncoder, make_service | |
| def client_for(service) -> TestClient: | |
| return TestClient(create_app(service=service, settings=Settings())) | |
| def upload(content: bytes = b"valid fake audio"): | |
| return {"audio": ("track.mp3", content, "audio/mpeg")} | |
| def test_valid_audio_only_request(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post("/v1/tracks/analyze", files=upload()) | |
| assert response.status_code == 200 | |
| assert set(response.json()["representations"]) == {"audio.global", "audio.temporal"} | |
| assert response.json()["schema_version"] == "1" | |
| def test_audio_emotion_global_request_uses_schema_v2(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.emotion.global"}, | |
| ) | |
| assert response.status_code == 200 | |
| body = response.json() | |
| assert body["schema_version"] == "2" | |
| result = body["representations"]["audio.emotion.global"] | |
| assert -1 <= result["valence"] <= 1 | |
| assert -1 <= result["arousal"] <= 1 | |
| assert len(result["mood_distribution"]["labels"]) == 56 | |
| assert result["emotion_embedding"] is None | |
| def test_combined_legacy_and_audio_emotion_representations(service_bundle): | |
| service, decoder, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={ | |
| "requested_representations": [ | |
| "audio.global", | |
| "audio.emotion.global", | |
| "audio.emotion.temporal", | |
| ] | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| assert list(response.json()["representations"]) == [ | |
| "audio.global", | |
| "audio.emotion.global", | |
| "audio.emotion.temporal", | |
| ] | |
| assert decoder.calls == 1 | |
| def test_lyrics_emotion_request_returns_typed_unavailable(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={ | |
| "lyrics": "[Verse]\nHello world\n\n[Припев]\nПривет, мир", | |
| "requested_representations": "lyrics.emotion.global", | |
| }, | |
| ) | |
| assert response.status_code == 503 | |
| assert response.json()["error"]["code"] == "LYRICS_EMOTION_MODEL_UNAVAILABLE" | |
| def test_disabled_deployment_audio_emotion_is_typed_unavailable(): | |
| settings = Settings(audio_emotion=AudioEmotionConfig(enabled=False)) | |
| response = client_for(build_service(settings)).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.emotion.global"}, | |
| ) | |
| assert response.status_code == 503 | |
| assert response.json()["error"]["code"] == "EMOTION_MODEL_UNAVAILABLE" | |
| assert "traceback" not in response.text.lower() | |
| def test_status_reports_lazy_emotion_runtime_state_and_real_availability(): | |
| settings = Settings(audio_emotion=AudioEmotionConfig(enabled=True)) | |
| response = TestClient(create_app(service=build_service(settings), settings=settings)).get( | |
| "/v1/status" | |
| ) | |
| assert response.status_code == 200 | |
| body = response.json() | |
| assert body["model_states"] == { | |
| "audio.emotion.global": "not_loaded", | |
| "audio.emotion.temporal": "not_loaded", | |
| } | |
| assert "audio.emotion.global" in body["available_representations"] | |
| assert "lyrics.emotion.global" not in body["available_representations"] | |
| def test_audio_emotion_inference_failure_is_sanitized(service_bundle): | |
| service, _, _ = service_bundle | |
| encoder = service.registry.analyzer("audio.emotion.global").pipeline.encoder | |
| async def fail(_windows): | |
| raise RuntimeError("private emotion runtime detail") | |
| encoder.predict = fail | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.emotion.global"}, | |
| ) | |
| assert response.status_code == 500 | |
| assert response.json()["error"]["code"] == "EMOTION_INFERENCE_FAILED" | |
| assert "private emotion runtime detail" not in response.text | |
| def test_audio_emotion_malformed_output_is_typed(service_bundle): | |
| from openmusic_analysis.analyzers.emotion import RawEmotionPrediction | |
| service, _, _ = service_bundle | |
| encoder = service.registry.analyzer("audio.emotion.global").pipeline.encoder | |
| async def malformed(windows): | |
| return [ | |
| RawEmotionPrediction( | |
| raw_valence=5.0, | |
| raw_arousal=5.0, | |
| mood_probabilities=(0.5,), | |
| ) | |
| for _ in windows | |
| ] | |
| encoder.predict = malformed | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.emotion.global"}, | |
| ) | |
| assert response.status_code == 500 | |
| assert response.json()["error"]["code"] == "EMOTION_OUTPUT_INVALID" | |
| def test_audio_and_multilingual_lyrics(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"lyrics": "[Verse]\nHello world\n\n[Припев]\nПривет, мир!"}, | |
| ) | |
| assert response.status_code == 200 | |
| assert "lyrics.global" in response.json()["representations"] | |
| def test_requested_representations_avoid_unrequested_work(service_bundle): | |
| service, decoder, encoder = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={ | |
| "lyrics": "lyrics that must not be encoded", | |
| "requested_representations": "audio.global", | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| assert list(response.json()["representations"]) == ["audio.global"] | |
| assert decoder.calls == 1 | |
| assert len(encoder.calls) == 1 | |
| def test_invalid_file_returns_structured_decode_error(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", files=upload(b"bad bytes") | |
| ) | |
| assert response.status_code == 422 | |
| assert response.json()["error"]["code"] == "AUDIO_DECODE_FAILED" | |
| def test_missing_audio_is_structured(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post("/v1/tracks/analyze", data={"lyrics": "text"}) | |
| assert response.status_code == 422 | |
| assert response.json()["error"]["code"] == "MISSING_AUDIO" | |
| def test_unsupported_representation(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.emotion"}, | |
| ) | |
| assert response.status_code == 422 | |
| assert response.json()["error"]["code"] == "UNSUPPORTED_REPRESENTATION" | |
| def test_unsupported_extension(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", files={"audio": ("track.txt", b"data", "text/plain")} | |
| ) | |
| assert response.status_code == 415 | |
| assert response.json()["error"]["code"] == "UNSUPPORTED_AUDIO_FORMAT" | |
| def test_explicit_lyrics_representation_requires_lyrics(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "lyrics.global"}, | |
| ) | |
| assert response.status_code == 422 | |
| assert response.json()["error"]["code"] == "LYRICS_REQUIRED" | |
| def test_invalid_empty_lyrics(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", files=upload(), data={"lyrics": " \n "} | |
| ) | |
| assert response.status_code == 422 | |
| assert response.json()["error"]["code"] == "INVALID_LYRICS" | |
| def test_model_failure_does_not_leak_internal_exception(): | |
| service, _, _ = make_service(audio_encoder=FailingAudioEncoder()) | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.global"}, | |
| ) | |
| assert response.status_code == 500 | |
| assert response.json()["error"]["code"] == "MODEL_INFERENCE_FAILED" | |
| assert "internal model detail" not in response.text | |
| def test_temporal_api_response_uses_adjacent_transition_index(): | |
| service, _, _ = make_service(waveform=np.arange(150, dtype=np.float32)) | |
| service.registry.analyzer("audio.temporal").config = TemporalAudioConfig( | |
| sample_rate=10, | |
| window_seconds=1, | |
| hop_seconds=1, | |
| max_segments=15, | |
| minimum_audio_seconds=1, | |
| ) | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(), | |
| data={"requested_representations": "audio.temporal"}, | |
| ) | |
| assert response.status_code == 200 | |
| temporal = response.json()["representations"]["audio.temporal"] | |
| assert len(temporal["segments"]) == 15 | |
| assert temporal["summary"]["number_of_segments"] == 15 | |
| assert 0 <= temporal["summary"]["largest_transition_index"] <= 13 | |
| class ContentDecoder: | |
| canonical_sample_rate = 10 | |
| def __init__(self) -> None: | |
| self.calls = 0 | |
| def decode(self, source_path) -> DecodedAudio: | |
| self.calls += 1 | |
| content = source_path.read_bytes() | |
| if content in {b"target", b"same"}: | |
| waveform = np.linspace(-1.0, 1.0, 95, dtype=np.float32) | |
| elif content == b"reversed": | |
| waveform = np.linspace(1.0, -1.0, 95, dtype=np.float32) | |
| else: | |
| waveform = np.full(95, 3.0, dtype=np.float32) | |
| return DecodedAudio( | |
| waveform=waveform, | |
| sample_rate=self.canonical_sample_rate, | |
| metadata=AudioMetadata( | |
| duration_ms=9500, | |
| source_sample_rate=10, | |
| source_channels=1, | |
| source_format="fake", | |
| ), | |
| ) | |
| class FakeLyricsResolver: | |
| async def resolve(self, source_path, *, provided_lyrics=None, supplied_metadata=None): | |
| content = Path(source_path).read_bytes() | |
| if content == b"same": | |
| return LyricsResolution( | |
| text="Resolved candidate lyrics", | |
| source="lrclib", | |
| errors=[ | |
| LyricsSourceError( | |
| source="embedded", | |
| code="LYRICS_NOT_FOUND", | |
| message="No embedded lyrics.", | |
| ) | |
| ], | |
| ) | |
| return LyricsResolution( | |
| errors=[ | |
| LyricsSourceError( | |
| source="lrclib", | |
| code="LYRICS_NOT_FOUND", | |
| message="No provider match.", | |
| ) | |
| ] | |
| ) | |
| def test_rank_similar_tracks_returns_descending_global_audio_similarity(service_bundle): | |
| service, _, encoder = service_bundle | |
| decoder = ContentDecoder() | |
| service.decoder = decoder | |
| service.lyrics_resolver = FakeLyricsResolver() | |
| response = client_for(service).post( | |
| "/v1/tracks/rank-similar", | |
| files=[ | |
| ("target_audio", ("target.mp3", b"target", "audio/mpeg")), | |
| ("tracks", ("different.mp3", b"different", "audio/mpeg")), | |
| ("tracks", ("same.mp3", b"same", "audio/mpeg")), | |
| ("tracks", ("reversed.mp3", b"reversed", "audio/mpeg")), | |
| ], | |
| data={ | |
| "target_track_id": "target-id", | |
| "target_lyrics": "Lyrics supplied for target", | |
| "track_ids": ["different-id", "same-id", "reversed-id"], | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| body = response.json() | |
| assert body["representation"] == "audio.global" | |
| assert body["target"]["track_id"] == "target-id" | |
| assert body["target"]["filename"] == "target.mp3" | |
| assert body["target"]["lyrics"]["text"] == "Lyrics supplied for target" | |
| assert body["target"]["lyrics"]["source"] == "request" | |
| assert body["tracks"][0]["track_id"] == "same-id" | |
| assert body["tracks"][0]["similarity"] == 1.0 | |
| assert body["tracks"][0]["lyrics"]["text"] == "Resolved candidate lyrics" | |
| assert body["tracks"][0]["lyrics"]["source"] == "lrclib" | |
| assert body["tracks"][0]["lyrics"]["errors"][0]["source"] == "embedded" | |
| assert [item["similarity"] for item in body["tracks"]] == sorted( | |
| [item["similarity"] for item in body["tracks"]], reverse=True | |
| ) | |
| assert decoder.calls == 4 | |
| assert len(encoder.calls) == 4 | |
| def test_analyze_automatically_adds_lyrics_representation(service_bundle): | |
| service, _, _ = service_bundle | |
| service.lyrics_resolver = FakeLyricsResolver() | |
| response = client_for(service).post( | |
| "/v1/tracks/analyze", | |
| files=upload(b"same"), | |
| data={"title": "Candidate", "artist": "Artist"}, | |
| ) | |
| assert response.status_code == 200 | |
| body = response.json() | |
| assert "lyrics.global" in body["representations"] | |
| assert body["lyrics"]["text"] == "Resolved candidate lyrics" | |
| assert body["lyrics"]["source"] == "lrclib" | |
| def test_rank_similar_tracks_validates_track_id_count(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/rank-similar", | |
| files=[ | |
| ("target_audio", ("target.mp3", b"target", "audio/mpeg")), | |
| ("tracks", ("one.mp3", b"one", "audio/mpeg")), | |
| ("tracks", ("two.mp3", b"two", "audio/mpeg")), | |
| ], | |
| data={"track_ids": ["only-one-id"]}, | |
| ) | |
| assert response.status_code == 422 | |
| assert response.json()["error"]["code"] == "INVALID_TRACK_IDS" | |
| def test_rank_similar_tracks_treats_swagger_empty_arrays_as_omitted(service_bundle): | |
| service, _, _ = service_bundle | |
| response = client_for(service).post( | |
| "/v1/tracks/rank-similar", | |
| files=[ | |
| ("target_audio", ("target.mp3", b"target", "audio/mpeg")), | |
| ("tracks", ("one.mp3", b"one", "audio/mpeg")), | |
| ("tracks", ("two.mp3", b"two", "audio/mpeg")), | |
| ], | |
| data={ | |
| "track_ids": "", | |
| "track_lyrics": "", | |
| "track_titles": "", | |
| "track_artists": "", | |
| "track_albums": "", | |
| "track_isrcs": "", | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| assert [item["track_id"] for item in response.json()["tracks"]] == [None, None] | |
| def test_rank_similar_tracks_respects_candidate_limit(service_bundle): | |
| service, _, _ = service_bundle | |
| settings = Settings(limits=LimitConfig(max_similarity_tracks=1)) | |
| response = TestClient(create_app(service=service, settings=settings)).post( | |
| "/v1/tracks/rank-similar", | |
| files=[ | |
| ("target_audio", ("target.mp3", b"target", "audio/mpeg")), | |
| ("tracks", ("one.mp3", b"one", "audio/mpeg")), | |
| ("tracks", ("two.mp3", b"two", "audio/mpeg")), | |
| ], | |
| ) | |
| assert response.status_code == 413 | |
| assert response.json()["error"]["code"] == "TOO_MANY_CANDIDATE_TRACKS" | |
| def test_rank_similar_openapi_declares_tracks_as_multiple_binary_files(service_bundle): | |
| service, _, _ = service_bundle | |
| schema = client_for(service).get("/openapi.json").json() | |
| operation = schema["paths"]["/v1/tracks/rank-similar"]["post"] | |
| multipart = operation["requestBody"]["content"]["multipart/form-data"] | |
| body = multipart["schema"] | |
| if "$ref" in body: | |
| body_name = body["$ref"].rsplit("/", maxsplit=1)[-1] | |
| body = schema["components"]["schemas"][body_name] | |
| tracks = body["properties"]["tracks"] | |
| assert tracks["type"] == "array" | |
| assert tracks["items"] == {"type": "string", "format": "binary"} | |
| assert tracks["minItems"] == 1 | |
| assert multipart["encoding"]["tracks"]["explode"] is True | |
| assert body["properties"]["target_lyrics"]["example"] == "" | |
| assert body["properties"]["track_lyrics"]["example"] == [] | |
| response_example = operation["responses"]["200"]["content"]["application/json"][ | |
| "example" | |
| ] | |
| assert response_example["tracks"][0]["lyrics"]["text"] is None | |
| assert response_example["tracks"][0]["lyrics"]["source"] is None | |