drum-sample-extractor / scripts /test_upload_error_visibility_and_fallback.py
ChatGPT
fix: make upload and fallback robust
5a90820
Raw
History Blame Contribute Delete
2.76 kB
import io
import json
import sys
import tempfile
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
import numpy as np
import soundfile as sf
from fastapi.testclient import TestClient
import pipeline_runner as runner
from app import app
def make_wav_buffer() -> io.BytesIO:
sr = 22050
y = np.zeros(sr, dtype=np.float32)
for pos in [0, sr // 4, sr // 2, 3 * sr // 4]:
n = min(1600, len(y) - pos)
t = np.arange(n) / sr
y[pos:pos + n] += 0.7 * np.exp(-t * 42) * np.sin(2 * np.pi * 90 * t)
buf = io.BytesIO()
sf.write(buf, y, sr, format="WAV")
buf.seek(0)
return buf
def test_config_exposes_runtime_diagnostics() -> None:
client = TestClient(app)
payload = client.get("/api/config").json()
assert "runtime" in payload
assert payload["runtime"]["backends"]["none"]["available"] is True
assert payload["defaults"]["separation_backend"] == "spleeter"
def test_invalid_upload_extension_is_actionable() -> None:
client = TestClient(app)
response = client.post(
"/api/jobs",
files={"file": ("not-audio.txt", b"hello", "text/plain")},
data={"params": json.dumps({})},
)
assert response.status_code == 400
assert "Unsupported audio file extension" in response.json()["detail"]
def test_spleeter_falls_back_to_full_mix_without_demucs() -> None:
original_spleeter = runner._extract_spleeter_separation
original_demucs = runner._extract_demucs_separation
runner._extract_spleeter_separation = lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("synthetic spleeter failure"))
runner._extract_demucs_separation = lambda *args, **kwargs: (_ for _ in ()).throw(RuntimeError("demucs should not be used for Spleeter fallback"))
try:
with tempfile.TemporaryDirectory() as tmp:
wav_path = Path(tmp) / "input.wav"
wav_path.write_bytes(make_wav_buffer().getvalue())
params = runner.PipelineParams(separation_backend="spleeter", stem="drums", use_disk_cache=False, allow_backend_fallback=True)
audio, sr, detail, context = runner._load_or_extract_separation(wav_path, params)
assert sr == 44100
assert audio.size > 0
assert context is None
assert "fallback full mix" in detail
assert "Spleeter failed" in detail
finally:
runner._extract_spleeter_separation = original_spleeter
runner._extract_demucs_separation = original_demucs
if __name__ == "__main__":
test_config_exposes_runtime_diagnostics()
test_invalid_upload_extension_is_actionable()
test_spleeter_falls_back_to_full_mix_without_demucs()
print("upload/error/fallback contract passed")