"""Tests for the VideoSeal watermark detection service.""" import base64 import io from unittest.mock import MagicMock, patch import pytest import torch from fastapi.testclient import TestClient from PIL import Image from app import _map_confidence_to_probability, app client = TestClient(app) # ── Fixture helpers ──────────────────────────────────────────────────────── def _make_image(width: int = 256, height: int = 256) -> bytes: """Create a valid PNG image of the given size. Generates a solid red image. Args: width: Image width in pixels. height: Image height in pixels. Returns: Raw PNG file bytes. """ img = Image.new("RGB", (width, height), color=(255, 0, 0)) buf = io.BytesIO() img.save(buf, format="PNG") return buf.getvalue() def _make_mp4() -> bytes: """Create a minimal MP4 file with a valid ftyp box header. Returns: Raw bytes of a minimal MP4 file. """ return b"\x00\x00\x00\x14ftypmp42\x00\x00\x00\x00mp42" 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 VideoSeal model so tests run without GPU # or the videoseal package installed. The mock returns a dict with pvalue. def _mock_detect(tensor): """Mock detector that returns low confidence (no watermark). Returns a dict with pvalue tensor of 0.1 (below threshold). """ return {"pvalue": torch.tensor(0.1)} def _patch_model(): """Create a mock model with detect method.""" mock = MagicMock() mock.detect = MagicMock(side_effect=_mock_detect) 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"] == "videoseal_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"] == "videoseal_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: clean image returns neutral ───────────────────────── class TestPredictCleanImage: """Tests for POST /predict with a clean (no watermark) image.""" @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_clean_image_returns_neutral(self, mock_load, mock_mdl): """Clean image with no watermark returns probability=0.5.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 assert data["class"] == "real" # ── Predict endpoint: different image sizes ─────────────────────────────── class TestPredictImageSizes: """Tests for POST /predict with various image dimensions.""" @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_128x128_image(self, mock_load, mock_mdl): """128x128 image is processed without error.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image(128, 128))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_512x512_image(self, mock_load, mock_mdl): """512x512 image is processed without error.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image(512, 512))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_1920x1080_image(self, mock_load, mock_mdl): """1920x1080 image is processed without error.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image(1920, 1080))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 # ── Predict endpoint: accepts both modalities ───────────────────────────── class TestPredictModalities: """Tests that /predict accepts both image_data and video_data keys.""" @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_accepts_image_data(self, mock_load, mock_mdl): """Endpoint accepts image_data key.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image())}, ) assert resp.status_code == 200 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_accepts_video_data(self, mock_load, mock_mdl): """Endpoint accepts video_data key.""" resp = client.post( "/predict", json={"video_data": _b64(_make_mp4())}, ) assert resp.status_code == 200 data = resp.json() # Minimal MP4 cannot be decoded by OpenCV, so confidence=0.0 # which maps to probability=0.5 (neutral) assert data["probability"] == 0.5 # ── Predict endpoint: error handling ─────────────────────────────────────── class TestPredictErrorHandling: """Tests for error cases.""" @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_invalid_base64_returns_error(self, mock_load, mock_mdl): """Invalid base64 string returns 400.""" resp = client.post( "/predict", json={"image_data": "!!!not-valid-base64!!!"}, ) assert resp.status_code == 400 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_non_media_bytes_returns_neutral(self, mock_load, mock_mdl): """Non-media binary data returns probability=0.5, not crash.""" random_bytes = base64.b64encode(b"\x00\x01\x02\x03" * 20).decode() resp = client.post( "/predict", json={"image_data": random_bytes}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_no_media_returns_error(self, mock_load, mock_mdl): """Request with neither image_data nor video_data returns 400.""" resp = client.post( "/predict", json={"threshold": 0.5}, ) assert resp.status_code == 400 # ── Predict endpoint: response format ────────────────────────────────────── class TestPredictResponseSchema: """Tests for the standard response format.""" @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_response_has_all_required_fields(self, mock_load, mock_mdl): """Response contains model, probability, prediction, class, inference_time.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image())}, ) data = resp.json() required = { "model", "probability", "prediction", "class", "inference_time", } assert required.issubset(data.keys()) @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_model_name_in_response(self, mock_load, mock_mdl): """Response model field matches videoseal_detector.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image())}, ) data = resp.json() assert data["model"] == "videoseal_detector" @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_probability_is_float_in_range(self, mock_load, mock_mdl): """Probability is a float in [0, 1].""" resp = client.post( "/predict", json={"image_data": _b64(_make_image())}, ) data = resp.json() assert isinstance(data["probability"], float) assert 0.0 <= data["probability"] <= 1.0 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_inference_time_is_positive_float(self, mock_load, mock_mdl): """Inference time is a positive float.""" resp = client.post( "/predict", json={"image_data": _b64(_make_image())}, ) data = resp.json() assert isinstance(data["inference_time"], float) assert data["inference_time"] > 0 @patch("app.model", new_callable=_patch_model) @patch("app.ensure_model_loaded") def test_threshold_controls_prediction(self, mock_load, mock_mdl): """Prediction uses > threshold (0.5 is neutral, not fake).""" resp = client.post( "/predict", json={"image_data": _b64(_make_image()), "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