"""Tests for the SDXL invisible watermark detection service.""" import base64 import io import numpy as np import pytest from fastapi.testclient import TestClient from PIL import Image from app import _compute_byte_entropy, app client = TestClient(app) # ── Fixture helpers ──────────────────────────────────────────────────────── def _make_clean_image( width: int = 256, height: int = 256, mode: str = "RGB", ) -> bytes: """Create a clean noisy image with no watermark. Uses random pixel noise so that watermark decoder output has high entropy (appears random), ensuring no false-positive detection. Args: width: Image width in pixels. height: Image height in pixels. mode: PIL image mode (RGB, RGBA, L). Returns: Raw PNG bytes. """ rng = np.random.RandomState(42) if mode == "L": arr = rng.randint(0, 256, (height, width), dtype=np.uint8) elif mode == "RGBA": arr = rng.randint(0, 256, (height, width, 4), dtype=np.uint8) else: arr = rng.randint(0, 256, (height, width, 3), dtype=np.uint8) img = Image.fromarray(arr, mode=mode) buf = io.BytesIO() img.save(buf, format="PNG") return buf.getvalue() def _make_watermarked_image() -> bytes: """Create a 512x512 image with an SDXL-style watermark embedded. Uses WatermarkEncoder to embed the SDV2 pattern (b"SDV2" + 13 null bytes = 17 bytes = 136 bits) via the dwtDct method. Returns: Raw PNG bytes of the watermarked image. """ from imwatermark import WatermarkEncoder # Create a 512x512 test image (larger for reliable watermark encoding) img = Image.new("RGB", (512, 512), color=(100, 150, 200)) rgb_array = np.array(img) # Convert RGB to BGR for OpenCV/imwatermark import cv2 bgr_array = cv2.cvtColor(rgb_array, cv2.COLOR_RGB2BGR) # Encode SDV2 watermark: 4 bytes "SDV2" + 13 null bytes = 17 bytes watermark_payload = b"SDV2" + b"\x00" * 13 encoder = WatermarkEncoder() encoder.set_watermark("bytes", watermark_payload) watermarked_bgr = encoder.encode(bgr_array, "dwtDct") # Convert back to RGB PIL image watermarked_rgb = cv2.cvtColor(watermarked_bgr, cv2.COLOR_BGR2RGB) watermarked_img = Image.fromarray(watermarked_rgb) buf = io.BytesIO() watermarked_img.save(buf, format="PNG") return buf.getvalue() def _b64(raw: bytes) -> str: """Encode raw bytes as base64 string.""" return base64.b64encode(raw).decode("utf-8") # ── 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"] == "sdxl_watermark_detector" def test_health_returns_device(self): """Health endpoint includes device=cpu.""" resp = client.get("/health") data = resp.json() assert data["device"] == "cpu" def test_health_returns_model_loaded(self): """Health endpoint includes model_loaded=True.""" resp = client.get("/health") data = resp.json() assert data["model_loaded"] is True # ── 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"] == "sdxl_watermark_detector" assert data["device"] == "cpu" assert data["model_loaded"] is True 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 ───────────────────────────────────────── class TestPredictCleanImage: """Tests for POST /predict with a clean (no watermark) image.""" def test_clean_image_returns_neutral(self): """Clean image with no watermark returns probability=0.5.""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 assert data["class"] == "real" # ── Predict endpoint: watermarked image ─────────────────────────────────── class TestPredictWatermarkedImage: """Tests for POST /predict with an SDXL-watermarked image.""" def test_sdxl_watermark_detected(self): """Image with SDXL watermark returns high probability.""" watermarked = _make_watermarked_image() resp = client.post( "/predict", json={"image_data": _b64(watermarked)}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] > 0.5 assert data["class"] == "fake" # ── Predict endpoint: image format handling ─────────────────────────────── class TestPredictImageFormats: """Tests for various image formats and sizes.""" def test_grayscale_image_handled(self): """Grayscale image is converted and processed without error.""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image(mode="L"))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 def test_rgba_image_handled(self): """RGBA image is converted and processed without error.""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image(mode="RGBA"))}, ) assert resp.status_code == 200 data = resp.json() assert 0.0 <= data["probability"] <= 1.0 def test_small_image_handled(self): """Small 32x32 image is processed without error.""" resp = client.post( "/predict", json={ "image_data": _b64(_make_clean_image(width=32, height=32)), }, ) 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.""" def test_invalid_base64_returns_error(self): """Invalid base64 string returns 400.""" resp = client.post( "/predict", json={"image_data": "!!!not-valid-base64!!!"}, ) assert resp.status_code == 400 def test_non_image_bytes_returns_neutral(self): """Non-image 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 # ── Predict endpoint: response schema ───────────────────────────────────── class TestPredictResponseSchema: """Tests for the standard response format.""" def test_response_has_all_required_fields(self): """Response contains model, probability, prediction, class, inference_time.""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image())}, ) data = resp.json() required = { "model", "probability", "prediction", "class", "inference_time", } assert required.issubset(data.keys()) def test_model_name_in_response(self): """Response model field matches sdxl_watermark_detector.""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image())}, ) data = resp.json() assert data["model"] == "sdxl_watermark_detector" def test_probability_is_float_in_range(self): """Probability is a float in [0, 1].""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image())}, ) data = resp.json() assert isinstance(data["probability"], float) assert 0.0 <= data["probability"] <= 1.0 def test_inference_time_is_positive_float(self): """Inference time is a positive float.""" resp = client.post( "/predict", json={"image_data": _b64(_make_clean_image())}, ) data = resp.json() assert isinstance(data["inference_time"], float) assert data["inference_time"] > 0 def test_threshold_controls_prediction(self): """Prediction uses > threshold (0.5 is neutral, not fake).""" resp = client.post( "/predict", json={ "image_data": _b64(_make_clean_image()), "threshold": 0.5, }, ) data = resp.json() # probability=0.5, threshold=0.5 -> 0.5 > 0.5 is False -> real assert data["prediction"] == 0 assert data["class"] == "real" # ── Entropy helper ───────────────────────────────────────────────────────── class TestComputeByteEntropy: """Tests for the _compute_byte_entropy helper function.""" def test_empty_bytes_returns_zero(self): """Empty input returns 0.0 entropy.""" assert _compute_byte_entropy(b"") == 0.0 def test_single_byte_returns_zero(self): """All identical bytes have zero entropy.""" assert _compute_byte_entropy(b"\x00" * 100) == 0.0 def test_two_equally_distributed_returns_one(self): """Two equally distributed byte values have entropy=1.0.""" data = b"\x00\x01" * 50 entropy = _compute_byte_entropy(data) assert abs(entropy - 1.0) < 0.01 def test_high_entropy_random_bytes(self): """Random-like bytes have high entropy (close to 8.0).""" # 256 distinct byte values each appearing once data = bytes(range(256)) entropy = _compute_byte_entropy(data) assert entropy == 8.0 def test_low_entropy_structured_bytes(self): """Structured bytes have low entropy.""" # "SDV2" repeated = limited unique bytes data = b"SDV2" * 10 entropy = _compute_byte_entropy(data) assert entropy < 3.0 def test_entropy_is_non_negative(self): """Entropy is always non-negative.""" assert _compute_byte_entropy(b"\xff") >= 0.0 assert _compute_byte_entropy(b"hello world") >= 0.0