trackembeddingapi / tests /test_api.py
kxmWebwe's picture
emotion
bf521ca
Raw History Blame Contribute Delete
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