"""Tests for the AudioSeal watermark detection service.""" import base64 import struct from unittest.mock import MagicMock, patch import pytest from fastapi.testclient import TestClient from app import _map_confidence_to_probability, app client = TestClient(app) # ── Fixture helpers ──────────────────────────────────────────────────────── def _make_wav( sample_rate: int = 16000, duration_s: float = 1.0, channels: int = 1, ) -> bytes: """Create a valid PCM WAV file with RIFF header. Generates silence (all zeros) as audio data. Args: sample_rate: Sample rate in Hz. duration_s: Duration in seconds. channels: Number of audio channels (1=mono, 2=stereo). Returns: Raw WAV file bytes. """ bits_per_sample = 16 num_samples = int(sample_rate * duration_s) byte_rate = sample_rate * channels * bits_per_sample // 8 block_align = channels * bits_per_sample // 8 data_size = num_samples * channels * bits_per_sample // 8 header = struct.pack( "<4sI4s4sIHHIIHH4sI", b"RIFF", 36 + data_size, b"WAVE", b"fmt ", 16, 1, # PCM format channels, sample_rate, byte_rate, block_align, bits_per_sample, b"data", data_size, ) return header + b"\x00" * data_size def _make_noise_wav( sample_rate: int = 16000, duration_s: float = 1.0, channels: int = 1, ) -> bytes: """Create a valid PCM WAV with white noise. Args: sample_rate: Sample rate in Hz. duration_s: Duration in seconds. channels: Number of audio channels. Returns: Raw WAV file bytes. """ import numpy as np bits_per_sample = 16 num_samples = int(sample_rate * duration_s) byte_rate = sample_rate * channels * bits_per_sample // 8 block_align = channels * bits_per_sample // 8 data_size = num_samples * channels * bits_per_sample // 8 header = struct.pack( "<4sI4s4sIHHIIHH4sI", b"RIFF", 36 + data_size, b"WAVE", b"fmt ", 16, 1, # PCM format channels, sample_rate, byte_rate, block_align, bits_per_sample, b"data", data_size, ) rng = np.random.RandomState(42) samples = rng.randint(-32768, 32767, size=num_samples * channels, dtype=np.int16) return header + samples.tobytes() def _b64(raw: bytes) -> str: """Encode raw bytes as base64 string.""" return base64.b64encode(raw).decode("utf-8") # ── Mock detector ───────────────────────────────────────────────────────── # All /predict tests mock the AudioSeal model so tests run without GPU # or the audioseal package installed. The mock returns (confidence, message). def _mock_detect_watermark(waveform, sample_rate=16000): """Mock detector that returns low confidence (no watermark).""" return (0.1, "no watermark") def _patch_detector(): """Create a mock detector with detect_watermark method.""" mock = MagicMock() mock.detect_watermark = MagicMock(side_effect=_mock_detect_watermark) return mock # ── Health endpoint ──────────────────────────────────────────────────────── class TestHealthEndpoint: """Tests for GET /health.""" def test_health_returns_healthy(self): """Health endpoint returns healthy status.""" resp = client.get("/health") assert resp.status_code == 200 data = resp.json() assert data["status"] == "healthy" def test_health_returns_model_name(self): """Health endpoint includes correct model_name.""" resp = client.get("/health") data = resp.json() assert data["model_name"] == "audioseal_detector" def test_health_returns_device(self): """Health endpoint includes device field.""" resp = client.get("/health") data = resp.json() assert "device" in data def test_health_returns_model_loaded(self): """Health endpoint includes model_loaded field.""" resp = client.get("/health") data = resp.json() assert "model_loaded" in data # ── Root endpoint ────────────────────────────────────────────────────────── class TestRootEndpoint: """Tests for GET /.""" def test_root_returns_service_info(self): """Root endpoint returns model info.""" resp = client.get("/") assert resp.status_code == 200 data = resp.json() assert data["model_name"] == "audioseal_detector" def test_root_returns_description(self): """Root endpoint includes a description.""" resp = client.get("/") data = resp.json() assert "description" in data assert len(data["description"]) > 0 # ── Predict endpoint: silent audio ────────────────────────────────────── class TestPredictSilentAudio: """Tests for POST /predict with silent (no watermark) WAV audio.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_silent_audio_returns_neutral(self, mock_load, mock_det): """Silent WAV with no watermark returns probability=0.5.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 assert data["class"] == "real" # ── Predict endpoint: white noise ──────────────────────────────────────── class TestPredictWhiteNoise: """Tests for POST /predict with white noise audio.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_white_noise_returns_neutral(self, mock_load, mock_det): """White noise WAV returns probability=0.5 (no watermark).""" resp = client.post( "/predict", json={"audio_data": _b64(_make_noise_wav())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 assert data["class"] == "real" # ── Predict endpoint: sample rate handling ─────────────────────────────── class TestPredictSampleRates: """Tests for different audio sample rates.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_16khz_handled(self, mock_load, mock_det): """16 kHz audio is processed without error.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav(sample_rate=16000))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_44100hz_handled(self, mock_load, mock_det): """44.1 kHz audio is resampled and processed without error.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav(sample_rate=44100))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_48khz_handled(self, mock_load, mock_det): """48 kHz audio is resampled and processed without error.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav(sample_rate=48000))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 # ── Predict endpoint: stereo audio ────────────────────────────────────── class TestPredictStereoAudio: """Tests for stereo audio handling.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_stereo_audio_handled(self, mock_load, mock_det): """Stereo WAV is averaged to mono and processed without error.""" resp = client.post( "/predict", json={ "audio_data": _b64(_make_wav(channels=2, duration_s=1.0)), }, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 # ── Predict endpoint: short audio ─────────────────────────────────────── class TestPredictShortAudio: """Tests for short duration audio.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_short_audio_handled(self, mock_load, mock_det): """Short 0.5s WAV is processed without error.""" resp = client.post( "/predict", json={ "audio_data": _b64(_make_wav(duration_s=0.5)), }, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 # ── Predict endpoint: error handling ───────────────────────────────────── class TestPredictErrorHandling: """Tests for error cases.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_invalid_base64_returns_error(self, mock_load, mock_det): """Invalid base64 string returns 400.""" resp = client.post( "/predict", json={"audio_data": "!!!not-valid-base64!!!"}, ) assert resp.status_code == 400 @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_non_audio_bytes_returns_neutral(self, mock_load, mock_det): """Non-audio binary data returns probability=0.5, not crash.""" random_bytes = base64.b64encode(b"\x00\x01\x02\x03" * 20).decode() resp = client.post( "/predict", json={"audio_data": random_bytes}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 # ── Predict endpoint: response format ──────────────────────────────────── class TestPredictResponseSchema: """Tests for the standard response format.""" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_response_has_all_required_fields(self, mock_load, mock_det): """Response contains model, probability, prediction, class, inference_time.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav())}, ) data = resp.json() required = { "model", "probability", "prediction", "class", "inference_time", } assert required.issubset(data.keys()) @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_model_name_in_response(self, mock_load, mock_det): """Response model field matches audioseal_detector.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav())}, ) data = resp.json() assert data["model"] == "audioseal_detector" @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_probability_is_float_in_range(self, mock_load, mock_det): """Probability is a float in [0, 1].""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav())}, ) data = resp.json() assert isinstance(data["probability"], float) assert 0.0 <= data["probability"] <= 1.0 @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_inference_time_is_positive_float(self, mock_load, mock_det): """Inference time is a positive float.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav())}, ) data = resp.json() assert isinstance(data["inference_time"], float) assert data["inference_time"] > 0 @patch("app.detector", new_callable=_patch_detector) @patch("app.ensure_model_loaded") def test_threshold_controls_prediction(self, mock_load, mock_det): """Prediction uses > threshold (0.5 is neutral, not fake).""" resp = client.post( "/predict", json={"audio_data": _b64(_make_wav()), "threshold": 0.5}, ) data = resp.json() # probability=0.5 (from mock), threshold=0.5 -> not > -> real assert data["prediction"] == 0 assert data["class"] == "real" # ── _map_confidence_to_probability helper ──────────────────────────────── class TestMapConfidenceToProbability: """Tests for the _map_confidence_to_probability helper function.""" def test_high_confidence_returns_high(self): """High confidence (0.9) maps to high probability.""" result = _map_confidence_to_probability(0.9) assert result == 0.5 + 0.9 * 0.5 # 0.95 assert result > 0.9 def test_medium_confidence_returns_medium(self): """Medium confidence (0.5) maps to medium-high probability.""" result = _map_confidence_to_probability(0.5) assert result == 0.5 + 0.5 * 0.5 # 0.75 def test_below_threshold_returns_neutral(self): """Confidence below 0.3 threshold returns 0.5 (neutral).""" assert _map_confidence_to_probability(0.1) == 0.5 assert _map_confidence_to_probability(0.0) == 0.5 assert _map_confidence_to_probability(0.29) == 0.5 def test_boundary_at_threshold(self): """Confidence at exactly 0.3 maps above neutral.""" result = _map_confidence_to_probability(0.3) assert result == 0.5 + 0.3 * 0.5 # 0.65 assert result > 0.5 def test_just_below_threshold(self): """Confidence just below 0.3 returns neutral.""" result = _map_confidence_to_probability(0.299) assert result == 0.5 def test_confidence_one_returns_max(self): """Confidence of 1.0 returns maximum probability (1.0).""" result = _map_confidence_to_probability(1.0) assert result == 1.0