Spaces:
Running
Running
Download tests/test_remix.py from Rthur2003/crowncode-backend: direct link, hf CLI and curl.
- Browser
- Download file 12 kB
-
https://huggingface.co/spaces/Rthur2003/crowncode-backend/resolve/main/tests/test_remix.py
- Command line
-
hf download hf://spaces/Rthur2003/crowncode-backend/tests/test_remix.py
-
curl -L -o test_remix.py https://huggingface.co/spaces/Rthur2003/crowncode-backend/resolve/main/tests/test_remix.py
12 kB
| """ | |
| Tests for the Creator Studio AI Remix engine (/api/remix/blend). | |
| Uses real synthesized audio — a pure tone for basic plumbing/error-path | |
| tests, and a click track (periodic impulses) for tempo-detection tests, | |
| since librosa.beat.beat_track needs actual rhythmic structure to lock onto. | |
| """ | |
| import io | |
| import json | |
| import numpy as np | |
| import pytest | |
| import soundfile as sf | |
| from fastapi.testclient import TestClient | |
| from app.routes.remix import _remix_rate_store | |
| def _reset_remix_rate_limiter(): | |
| """The rate limiter's store is module-level and shared across every | |
| test in this file (6 requests/60s per IP, and TestClient always uses | |
| the same fake client IP) — without resetting it, tests that pass in | |
| isolation start failing with 429 once enough tests run before them in | |
| the same process.""" | |
| _remix_rate_store.clear() | |
| yield | |
| _remix_rate_store.clear() | |
| def _tone_wav(freq: float = 440.0, duration: float = 3.0, sr: int = 22050) -> io.BytesIO: | |
| t = np.linspace(0, duration, int(sr * duration)) | |
| y = (0.3 * np.sin(2 * np.pi * freq * t)).astype(np.float32) | |
| buf = io.BytesIO() | |
| sf.write(buf, y, sr, format="WAV") | |
| buf.seek(0) | |
| return buf | |
| def _click_track_wav(bpm: float, duration: float = 6.0, sr: int = 22050) -> io.BytesIO: | |
| """A percussive click every beat at the given BPM — enough rhythmic | |
| structure for librosa's onset-based beat tracker to detect real tempo.""" | |
| n_samples = int(sr * duration) | |
| y = np.zeros(n_samples, dtype=np.float32) | |
| beat_interval_sec = 60.0 / bpm | |
| click_len = int(sr * 0.02) | |
| decay = np.exp(-np.linspace(0, 12, click_len)).astype(np.float32) | |
| t = 0.0 | |
| while t < duration: | |
| start = int(t * sr) | |
| end = min(start + click_len, n_samples) | |
| y[start:end] += decay[: end - start] | |
| t += beat_interval_sec | |
| buf = io.BytesIO() | |
| sf.write(buf, y, sr, format="WAV") | |
| buf.seek(0) | |
| return buf | |
| def _valid_options(**overrides) -> str: | |
| base = {"matchTempo": True, "matchKey": False, "crossfadeSeconds": 2.0, "crossfadeCurve": "equal_power"} | |
| base.update(overrides) | |
| return json.dumps(base) | |
| def test_remix_blend_rejects_non_audio_track_a(client: TestClient) -> None: | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options()}, | |
| files={ | |
| "track_a": ("a.txt", io.BytesIO(b"not audio"), "text/plain"), | |
| "track_b": ("b.wav", _tone_wav(440), "audio/wav"), | |
| }, | |
| ) | |
| assert response.status_code == 400 | |
| assert response.json()["detail"]["code"] == "invalid_file_type" | |
| def test_remix_blend_rejects_non_audio_track_b(client: TestClient) -> None: | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options()}, | |
| files={ | |
| "track_a": ("a.wav", _tone_wav(440), "audio/wav"), | |
| "track_b": ("b.bin", io.BytesIO(b"not audio"), "application/octet-stream"), | |
| }, | |
| ) | |
| assert response.status_code == 400 | |
| assert response.json()["detail"]["code"] == "invalid_file_type" | |
| def test_remix_blend_rejects_invalid_options_json(client: TestClient) -> None: | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": "{not json"}, | |
| files={ | |
| "track_a": ("a.wav", _tone_wav(440), "audio/wav"), | |
| "track_b": ("b.wav", _tone_wav(220), "audio/wav"), | |
| }, | |
| ) | |
| assert response.status_code == 422 | |
| def test_remix_blend_rejects_silent_track(client: TestClient) -> None: | |
| silent = io.BytesIO() | |
| sf.write(silent, np.zeros(22050, dtype=np.float32), 22050, format="WAV") | |
| silent.seek(0) | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options()}, | |
| files={ | |
| "track_a": ("a.wav", silent, "audio/wav"), | |
| "track_b": ("b.wav", _tone_wav(440), "audio/wav"), | |
| }, | |
| ) | |
| assert response.status_code == 400 | |
| assert response.json()["detail"]["code"] == "validation_error" | |
| def test_remix_blend_returns_wav_with_analysis_header(client: TestClient) -> None: | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options(matchTempo=False)}, | |
| files={ | |
| "track_a": ("a.wav", _tone_wav(440, duration=2.0), "audio/wav"), | |
| "track_b": ("b.wav", _tone_wav(330, duration=2.0), "audio/wav"), | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| assert response.headers["content-type"] == "audio/wav" | |
| assert "X-Remix-Analysis" in response.headers | |
| analysis = json.loads(response.headers["X-Remix-Analysis"]) | |
| assert "trackABpm" in analysis or "track_a_bpm" in analysis # camelCase alias in JSON | |
| assert response.content[:4] == b"RIFF" | |
| def test_remix_blend_output_is_real_blended_audio_not_passthrough(client: TestClient) -> None: | |
| """The output WAV must actually contain a crossfaded blend — verify via | |
| duration: two 2s tracks with a 1s crossfade should yield ~3s of audio, | |
| not 2s (passthrough) or 4s (naive concatenation).""" | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options(matchTempo=False, crossfadeSeconds=1.0)}, | |
| files={ | |
| "track_a": ("a.wav", _tone_wav(440, duration=2.0), "audio/wav"), | |
| "track_b": ("b.wav", _tone_wav(330, duration=2.0), "audio/wav"), | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| y, sr = sf.read(io.BytesIO(response.content)) | |
| duration = len(y) / sr | |
| assert 2.8 < duration < 3.2 | |
| def test_remix_blend_matches_tempo_with_click_tracks(client: TestClient) -> None: | |
| """Two click tracks at different real BPMs: with matchTempo on, the | |
| engine must detect both tempos and report a non-trivial stretch rate.""" | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options(matchTempo=True, crossfadeSeconds=1.0)}, | |
| files={ | |
| "track_a": ("a.wav", _click_track_wav(120.0), "audio/wav"), | |
| "track_b": ("b.wav", _click_track_wav(140.0), "audio/wav"), | |
| }, | |
| ) | |
| assert response.status_code == 200 | |
| analysis_raw = response.headers["X-Remix-Analysis"] | |
| analysis = json.loads(analysis_raw) | |
| bpm_a = analysis.get("track_a_bpm", analysis.get("trackABpm")) | |
| bpm_b = analysis.get("track_b_bpm", analysis.get("trackBBpm")) | |
| stretch = analysis.get("applied_stretch_rate", analysis.get("appliedStretchRate")) | |
| assert bpm_a > 1 # real beat detected, not the "beatless" 0 fallback | |
| assert bpm_b > 1 | |
| assert stretch != 1.0 # tempos differed, so a real stretch must have been applied | |
| def test_remix_blend_rate_limit(client: TestClient) -> None: | |
| last_status = None | |
| for _ in range(7): | |
| response = client.post( | |
| "/api/remix/blend", | |
| data={"options": _valid_options(matchTempo=False)}, | |
| files={ | |
| "track_a": ("a.wav", _tone_wav(440, duration=0.5), "audio/wav"), | |
| "track_b": ("b.wav", _tone_wav(330, duration=0.5), "audio/wav"), | |
| }, | |
| ) | |
| last_status = response.status_code | |
| assert last_status == 429 | |
| # ── Multitrack mixer tests ─────────────────────────────────────────── | |
| def test_multitrack_rejects_too_few_tracks(client: TestClient) -> None: | |
| response = client.post( | |
| "/api/remix/multitrack", | |
| files={"tracks": ("a.wav", _tone_wav(440, duration=1.0), "audio/wav")}, | |
| ) | |
| assert response.status_code == 400 | |
| assert response.json()["detail"]["code"] == "invalid_track_count" | |
| def test_multitrack_rejects_too_many_tracks(client: TestClient) -> None: | |
| files = [ | |
| ("tracks", (f"t{i}.wav", _tone_wav(220 + i * 20, duration=0.5), "audio/wav")) | |
| for i in range(7) | |
| ] | |
| response = client.post("/api/remix/multitrack", files=files) | |
| assert response.status_code == 400 | |
| assert response.json()["detail"]["code"] == "invalid_track_count" | |
| def test_multitrack_rejects_non_audio_track(client: TestClient) -> None: | |
| files = [ | |
| ("tracks", ("a.wav", _tone_wav(440, duration=1.0), "audio/wav")), | |
| ("tracks", ("b.txt", io.BytesIO(b"not audio"), "text/plain")), | |
| ] | |
| response = client.post("/api/remix/multitrack", files=files) | |
| assert response.status_code == 400 | |
| assert response.json()["detail"]["code"] == "invalid_file_type" | |
| def test_multitrack_rejects_invalid_options_json(client: TestClient) -> None: | |
| files = [ | |
| ("tracks", ("a.wav", _tone_wav(440, duration=1.0), "audio/wav")), | |
| ("tracks", ("b.wav", _tone_wav(330, duration=1.0), "audio/wav")), | |
| ] | |
| response = client.post( | |
| "/api/remix/multitrack", | |
| data={"options": "{not json"}, | |
| files=files, | |
| ) | |
| assert response.status_code == 422 | |
| def test_multitrack_mixes_three_tracks_and_reports_channels(client: TestClient) -> None: | |
| files = [ | |
| ("tracks", ("a.wav", _tone_wav(220, duration=2.0), "audio/wav")), | |
| ("tracks", ("b.wav", _tone_wav(330, duration=2.0), "audio/wav")), | |
| ("tracks", ("c.wav", _tone_wav(440, duration=2.0), "audio/wav")), | |
| ] | |
| options = json.dumps({ | |
| "channels": [ | |
| {"gainDb": 0.0, "pan": -1.0, "muted": False}, | |
| {"gainDb": -6.0, "pan": 0.0, "muted": False}, | |
| {"gainDb": 0.0, "pan": 1.0, "muted": False}, | |
| ], | |
| "normalizeOutput": True, | |
| }) | |
| response = client.post( | |
| "/api/remix/multitrack", | |
| data={"options": options}, | |
| files=files, | |
| ) | |
| assert response.status_code == 200 | |
| assert response.headers["content-type"] == "audio/wav" | |
| assert "X-Multitrack-Analysis" in response.headers | |
| analysis = json.loads(response.headers["X-Multitrack-Analysis"]) | |
| channels = analysis.get("channels", []) | |
| assert len(channels) == 3 | |
| assert all(not c["muted"] for c in channels) | |
| # Output must be real stereo audio at the expected duration. | |
| y, sr = sf.read(io.BytesIO(response.content)) | |
| assert y.ndim == 2 and y.shape[1] == 2 | |
| assert abs(len(y) / sr - 2.0) < 0.1 | |
| def test_multitrack_muted_track_excluded_from_mix(client: TestClient) -> None: | |
| """A muted track's audio must not appear in the panned-left channel.""" | |
| files = [ | |
| ("tracks", ("a.wav", _tone_wav(220, duration=1.5), "audio/wav")), | |
| ("tracks", ("b.wav", _tone_wav(330, duration=1.5), "audio/wav")), | |
| ] | |
| options = json.dumps({ | |
| "channels": [ | |
| {"gainDb": 0.0, "pan": 0.0, "muted": True}, | |
| {"gainDb": 0.0, "pan": 0.0, "muted": False}, | |
| ], | |
| }) | |
| response = client.post( | |
| "/api/remix/multitrack", | |
| data={"options": options}, | |
| files=files, | |
| ) | |
| assert response.status_code == 200 | |
| analysis = json.loads(response.headers["X-Multitrack-Analysis"]) | |
| channels = analysis["channels"] | |
| assert channels[0]["muted"] is True | |
| assert channels[0]["peakLevel"] == 0.0 | |
| assert channels[1]["muted"] is False | |
| assert channels[1]["peakLevel"] > 0.0 | |
| def test_multitrack_prevents_clipping_when_normalized(client: TestClient) -> None: | |
| """Six loud, centered tracks summed together would clip well past | |
| 1.0 without normalization — verify the output never does.""" | |
| files = [ | |
| ("tracks", (f"t{i}.wav", _tone_wav(200 + i * 50, duration=1.0), "audio/wav")) | |
| for i in range(6) | |
| ] | |
| options = json.dumps({ | |
| "channels": [{"gainDb": 6.0, "pan": 0.0, "muted": False} for _ in range(6)], | |
| "normalizeOutput": True, | |
| }) | |
| response = client.post( | |
| "/api/remix/multitrack", | |
| data={"options": options}, | |
| files=files, | |
| ) | |
| assert response.status_code == 200 | |
| y, sr = sf.read(io.BytesIO(response.content)) | |
| assert float(np.max(np.abs(y))) <= 1.0 | |
| analysis = json.loads(response.headers["X-Multitrack-Analysis"]) | |
| assert analysis["clippingPrevented"] is True | |
| assert analysis["outputPeakLevel"] <= 1.0 | |