""" 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 @pytest.fixture(autouse=True) 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