88 lines
2.9 KiB
Python
88 lines
2.9 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Gap 11: A/B reference analyze endpoint (POST /api/v1/audio/reference-analyze).
|
|
|
|
- path mode: local path (media/browse server mode) -> LUFS/dBTP/centroid + url
|
|
- file_id mode: uploaded/processed file
|
|
"""
|
|
import os
|
|
import tempfile
|
|
|
|
import numpy as np
|
|
import soundfile as sf
|
|
from fastapi.testclient import TestClient
|
|
|
|
from app.main import app
|
|
|
|
client = TestClient(app)
|
|
|
|
|
|
def _write_ref_wav(path, dbfs=-20.0, freq=1000.0, sec=2.0, sr=44100):
|
|
t = np.arange(int(sr * sec)) / sr
|
|
x = (10 ** (dbfs / 20.0)) * np.sin(2 * np.pi * freq * t)
|
|
sf.write(path, np.stack([x, x]).T, sr, subtype="FLOAT")
|
|
|
|
|
|
def test_reference_analyze_path():
|
|
tmp = os.path.join(tempfile.gettempdir(), "ref_ab_test.wav")
|
|
_write_ref_wav(tmp)
|
|
r = client.post("/api/v1/audio/reference-analyze", json={"path": tmp})
|
|
assert r.status_code == 200, r.text
|
|
d = r.json()
|
|
assert d["filename"].endswith(".wav")
|
|
assert d["sample_rate"] == 44100
|
|
assert d["channels"] == 2
|
|
assert abs(d["duration"] - 2.0) < 0.05, d
|
|
# 1 kHz sine @ -20 dBFS -> ~ -23 LUFS (cung ly luan test_loudness)
|
|
assert -25.0 < d["lufs"] < -21.0, d
|
|
assert abs(d["true_peak_db"] - (-20.0)) < 0.3, d
|
|
# spectral centroid cua sine 1 kHz ~ 1000 Hz
|
|
assert 850 < d["centroid_hz"] < 1150, d
|
|
assert d["url"].startswith("/api/v1/media/file")
|
|
|
|
|
|
def test_reference_analyze_file_id():
|
|
# upload trigger celery analyze_audio_task.delay — stub để test không cần Redis
|
|
import types
|
|
|
|
import app.tasks.worker as worker
|
|
_orig = worker.analyze_audio_task
|
|
worker.analyze_audio_task = types.SimpleNamespace(delay=lambda file_id: types.SimpleNamespace(id="stub"))
|
|
try:
|
|
tmp = os.path.join(tempfile.gettempdir(), "ref_ab_upload.wav")
|
|
_write_ref_wav(tmp, dbfs=-12.0, freq=440.0)
|
|
with open(tmp, "rb") as f:
|
|
up = client.post("/api/v1/audio/upload", files={"file": ("ref.wav", f, "audio/wav")})
|
|
finally:
|
|
worker.analyze_audio_task = _orig
|
|
assert up.status_code == 200, up.text
|
|
fid = up.json()["file_id"]
|
|
r = client.post("/api/v1/audio/reference-analyze", json={"file_id": fid})
|
|
assert r.status_code == 200, r.text
|
|
d = r.json()
|
|
assert d["url"].startswith("/api/v1/audio/download/")
|
|
assert -17.5 < d["lufs"] < -13.0, d
|
|
|
|
|
|
def test_reference_analyze_missing():
|
|
r = client.post("/api/v1/audio/reference-analyze", json={})
|
|
assert r.status_code == 400
|
|
|
|
|
|
def test_reference_analyze_bad_path():
|
|
r = client.post("/api/v1/audio/reference-analyze", json={"path": "Z:/no/such/file.wav"})
|
|
assert r.status_code == 404
|
|
|
|
|
|
def test_reference_analyze_bad_ext():
|
|
tmp = os.path.join(tempfile.gettempdir(), "ref_ab_bad.exe")
|
|
with open(tmp, "wb") as f:
|
|
f.write(b"MZ\x90\x00")
|
|
try:
|
|
r = client.post("/api/v1/audio/reference-analyze", json={"path": tmp})
|
|
assert r.status_code == 400, r.text
|
|
finally:
|
|
try:
|
|
os.remove(tmp)
|
|
except OSError:
|
|
pass
|