"""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