deepsafe's picture
sync from GitHub (0154d02)
4b0b144 verified
Raw History Blame Contribute Delete
8.82 kB
"""Tests for ShiftySpeech Audio Deepfake Detection API.
Tests cover:
- Health endpoint
- Predict endpoint with real audio
- Predict endpoint with fake audio
- Audio preprocessing (pad/trim/tile)
- Error handling for invalid input
- Response schema validation
"""
import base64
import io
import os
import sys
import warnings
import numpy as np
import pytest
warnings.filterwarnings("ignore", category=DeprecationWarning)
# Monkey-patch omegaconf before importing api module
import omegaconf._utils as _omegaconf_utils
if not hasattr(_omegaconf_utils, "is_primitive_type"):
_omegaconf_utils.is_primitive_type = lambda t: t in (int, float, bool, str, bytes)
# Add model code path for local testing
SERVICE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
MODEL_CODE_PATH = os.path.join(
SERVICE_DIR, "synthetic_speech_detection", "SSL_Anti-spoofing"
)
if MODEL_CODE_PATH not in sys.path:
sys.path.insert(0, MODEL_CODE_PATH)
# Patch the api module constants for local testing
import api
api.MODEL_CODE_PATH = MODEL_CODE_PATH
api.WEIGHTS_PATH = os.path.join(SERVICE_DIR, "weights", "hfg_aug_1_2.pt")
api.XLSR_DIR = os.path.join(SERVICE_DIR, "models")
from fastapi.testclient import TestClient
client = TestClient(api.app)
# Dataset paths
DATASET_DIR = os.path.join(
os.path.dirname(SERVICE_DIR),
os.pardir,
os.pardir,
"dataset",
"audio",
)
DATASET_DIR = os.path.normpath(DATASET_DIR)
REAL_DIR = os.path.join(DATASET_DIR, "real")
FAKE_DIR = os.path.join(DATASET_DIR, "fake")
# Check if model weights are available for integration tests
WEIGHTS_AVAILABLE = os.path.exists(api.WEIGHTS_PATH) and os.path.exists(
os.path.join(SERVICE_DIR, "models", "xlsr2_300m.pt")
)
DATASET_AVAILABLE = os.path.isdir(REAL_DIR) and os.path.isdir(FAKE_DIR)
def _make_wav_bytes(duration_s: float = 1.0, sr: int = 16000) -> bytes:
"""Generate a simple sine wave WAV file as bytes."""
import soundfile as sf
t = np.linspace(0, duration_s, int(sr * duration_s), endpoint=False)
audio = 0.5 * np.sin(2 * np.pi * 440 * t).astype(np.float32)
buf = io.BytesIO()
sf.write(buf, audio, sr, format="WAV")
buf.seek(0)
return buf.read()
def _encode_file(path: str) -> str:
"""Read a file and return base64 encoded string."""
with open(path, "rb") as f:
return base64.b64encode(f.read()).decode("utf-8")
class TestHealthEndpoint:
"""Tests for the /health endpoint."""
def test_health_returns_200(self):
response = client.get("/health")
assert response.status_code == 200
def test_health_contains_model_name(self):
response = client.get("/health")
data = response.json()
assert data["model"] == "shiftyspeech"
def test_health_contains_device(self):
response = client.get("/health")
data = response.json()
assert data["device"] == "cpu"
def test_health_contains_status(self):
response = client.get("/health")
data = response.json()
assert data["status"] in ("healthy", "degraded")
class TestPreprocessAudio:
"""Tests for audio preprocessing logic."""
def test_preprocess_short_audio_tiles(self):
"""Short audio should be tiled to TARGET_SAMPLES."""
wav_bytes = _make_wav_bytes(duration_s=0.5, sr=16000)
tensor = api.preprocess_audio(wav_bytes)
assert tensor.shape == (1, api.TARGET_SAMPLES)
def test_preprocess_long_audio_trims(self):
"""Long audio should be trimmed to TARGET_SAMPLES."""
wav_bytes = _make_wav_bytes(duration_s=10.0, sr=16000)
tensor = api.preprocess_audio(wav_bytes)
assert tensor.shape == (1, api.TARGET_SAMPLES)
def test_preprocess_exact_length(self):
"""Audio at exact TARGET_SAMPLES should pass through."""
duration = api.TARGET_SAMPLES / api.SAMPLE_RATE
wav_bytes = _make_wav_bytes(duration_s=duration, sr=16000)
tensor = api.preprocess_audio(wav_bytes)
assert tensor.shape == (1, api.TARGET_SAMPLES)
def test_preprocess_resamples_from_8khz(self):
"""Audio at 8kHz should be resampled to 16kHz."""
wav_bytes = _make_wav_bytes(duration_s=1.0, sr=8000)
tensor = api.preprocess_audio(wav_bytes)
assert tensor.shape == (1, api.TARGET_SAMPLES)
def test_preprocess_invalid_input_raises(self):
"""Invalid audio bytes should raise ValueError."""
with pytest.raises(ValueError):
api.preprocess_audio(b"not audio data")
@pytest.mark.skipif(
not WEIGHTS_AVAILABLE,
reason="Model weights not available locally",
)
class TestPredictEndpoint:
"""Integration tests for the /predict endpoint (requires weights)."""
def test_predict_returns_200(self):
wav_bytes = _make_wav_bytes(duration_s=2.0)
b64 = base64.b64encode(wav_bytes).decode("utf-8")
response = client.post("/predict", json={"audio_data": b64})
assert response.status_code == 200
def test_predict_response_schema(self):
wav_bytes = _make_wav_bytes(duration_s=2.0)
b64 = base64.b64encode(wav_bytes).decode("utf-8")
response = client.post("/predict", json={"audio_data": b64})
data = response.json()
assert "model" in data
assert "probability" in data
assert "prediction" in data
assert "class" in data
assert "inference_time" in data
assert data["model"] == "shiftyspeech"
def test_predict_probability_in_range(self):
wav_bytes = _make_wav_bytes(duration_s=2.0)
b64 = base64.b64encode(wav_bytes).decode("utf-8")
response = client.post("/predict", json={"audio_data": b64})
data = response.json()
assert 0.0 <= data["probability"] <= 1.0
def test_predict_class_matches_prediction(self):
wav_bytes = _make_wav_bytes(duration_s=2.0)
b64 = base64.b64encode(wav_bytes).decode("utf-8")
response = client.post("/predict", json={"audio_data": b64})
data = response.json()
if data["prediction"] == 1:
assert data["class"] == "fake"
else:
assert data["class"] == "real"
def test_predict_custom_threshold(self):
wav_bytes = _make_wav_bytes(duration_s=2.0)
b64 = base64.b64encode(wav_bytes).decode("utf-8")
response = client.post(
"/predict",
json={"audio_data": b64, "threshold": 0.99},
)
data = response.json()
assert response.status_code == 200
# With threshold=0.99, only very high prob_fake => fake
if data["probability"] < 0.99:
assert data["prediction"] == 0
assert data["class"] == "real"
def test_predict_inference_time_positive(self):
wav_bytes = _make_wav_bytes(duration_s=2.0)
b64 = base64.b64encode(wav_bytes).decode("utf-8")
response = client.post("/predict", json={"audio_data": b64})
data = response.json()
assert data["inference_time"] > 0
@pytest.mark.skipif(
not WEIGHTS_AVAILABLE or not DATASET_AVAILABLE,
reason="Model weights or dataset not available",
)
class TestRealDataset:
"""Integration tests using actual dataset files."""
def test_predict_real_audio(self):
"""Test prediction on a real audio file."""
path = os.path.join(REAL_DIR, "real_0.wav")
if not os.path.exists(path):
pytest.skip("real_0.wav not found")
b64 = _encode_file(path)
response = client.post("/predict", json={"audio_data": b64})
assert response.status_code == 200
data = response.json()
assert 0.0 <= data["probability"] <= 1.0
def test_predict_fake_audio(self):
"""Test prediction on a fake audio file."""
path = os.path.join(FAKE_DIR, "fake_1.wav")
if not os.path.exists(path):
pytest.skip("fake_1.wav not found")
b64 = _encode_file(path)
response = client.post("/predict", json={"audio_data": b64})
assert response.status_code == 200
data = response.json()
assert 0.0 <= data["probability"] <= 1.0
class TestPredictValidation:
"""Tests for input validation on /predict."""
def test_predict_missing_audio_data(self):
response = client.post("/predict", json={})
assert response.status_code == 422
def test_predict_invalid_base64(self):
response = client.post("/predict", json={"audio_data": "not-valid-base64!!!"})
# Should return 500 (decode error) or 422
assert response.status_code in (400, 422, 500)
def test_predict_threshold_out_of_range(self):
response = client.post(
"/predict",
json={"audio_data": "dGVzdA==", "threshold": 1.5},
)
assert response.status_code == 422