trackembeddingapi / tests /test_api.py
kxmWebwe's picture
fix
bff0fe1
Raw
History Blame Contribute Delete
4.96 kB
from __future__ import annotations
import numpy as np
from fastapi.testclient import TestClient
from openmusic_analysis.api import create_app
from openmusic_analysis.settings import 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"}
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