"""Tests for the c2pa-checker provenance detection service.""" import base64 import io import struct import pytest from fastapi.testclient import TestClient from app import app client = TestClient(app) # ── Fixture helpers ──────────────────────────────────────────────────────── def _make_minimal_jpeg() -> bytes: """Create a minimal valid JPEG from an 8x8 RGB image.""" from PIL import Image img = Image.new("RGB", (8, 8)) buf = io.BytesIO() img.save(buf, format="JPEG") return buf.getvalue() def _make_minimal_wav() -> bytes: """Create a minimal valid WAV: RIFF/WAVE header + 1s silence 16kHz mono 16-bit.""" sample_rate = 16000 num_samples = sample_rate # 1 second bits_per_sample = 16 num_channels = 1 byte_rate = sample_rate * num_channels * bits_per_sample // 8 block_align = num_channels * bits_per_sample // 8 data_size = num_samples * block_align header = struct.pack( "<4sI4s4sIHHIIHH4sI", b"RIFF", 36 + data_size, b"WAVE", b"fmt ", 16, 1, # PCM num_channels, sample_rate, byte_rate, block_align, bits_per_sample, b"data", data_size, ) return header + b"\x00" * data_size def _make_minimal_mp4() -> bytes: """Create a minimal valid MP4: ftyp box only.""" 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") # ── 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"] == "c2pa_checker" # ── 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 "model_name" in data assert data["model_name"] == "c2pa_checker" # ── Predict endpoint: clean files ────────────────────────────────────────── class TestPredictCleanFiles: """Tests for POST /predict with clean (no provenance) media.""" def test_clean_jpeg_returns_neutral(self): """Clean JPEG with no C2PA/EXIF markers returns probability=0.5.""" resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 assert data["class"] == "real" def test_clean_wav_returns_neutral(self): """Clean WAV with no provenance returns probability=0.5.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_minimal_wav())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 def test_clean_mp4_returns_neutral(self): """Clean MP4 with no provenance returns probability=0.5.""" resp = client.post( "/predict", json={"video_data": _b64(_make_minimal_mp4())}, ) assert resp.status_code == 200 data = resp.json() assert data["probability"] == 0.5 # ── Predict endpoint: payload key acceptance ─────────────────────────────── class TestPredictPayloadKeys: """Tests for accepted payload keys.""" def test_accepts_image_data(self): """image_data key is accepted.""" resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg())}, ) assert resp.status_code == 200 def test_accepts_audio_data(self): """audio_data key is accepted.""" resp = client.post( "/predict", json={"audio_data": _b64(_make_minimal_wav())}, ) assert resp.status_code == 200 def test_accepts_video_data(self): """video_data key is accepted.""" resp = client.post( "/predict", json={"video_data": _b64(_make_minimal_mp4())}, ) assert resp.status_code == 200 def test_no_payload_key_returns_400(self): """Request with no recognized payload key returns 400.""" resp = client.post("/predict", json={"threshold": 0.5}) assert resp.status_code == 400 # ── 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_empty_base64_returns_error(self): """Empty base64 string returns 400.""" resp = client.post( "/predict", json={"image_data": ""}, ) assert resp.status_code == 400 def test_corrupt_file_bytes_returns_neutral(self): """Corrupt file bytes should return probability=0.5, not crash.""" corrupt = base64.b64encode(b"\x00\x01\x02\x03" * 20).decode() resp = client.post( "/predict", json={"image_data": corrupt}, ) 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_minimal_jpeg())}, ) data = resp.json() required = {"model", "probability", "prediction", "class", "inference_time"} assert required.issubset(data.keys()) def test_probability_is_float_in_range(self): """Probability is a float in [0, 1].""" resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg())}, ) data = resp.json() assert isinstance(data["probability"], float) assert 0.0 <= data["probability"] <= 1.0 def test_prediction_matches_threshold_logic(self): """Prediction=1 when probability > threshold, else 0.""" resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg()), "threshold": 0.6}, ) data = resp.json() # probability=0.5, threshold=0.6 -> 0.5 < 0.6 -> prediction=0, class=real assert data["prediction"] == 0 assert data["class"] == "real" def test_prediction_fake_when_above_threshold(self): """When probability >= threshold, prediction=1 and class=fake.""" # With threshold=0.4 and probability=0.5: 0.5 >= 0.4 -> prediction=1 resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg()), "threshold": 0.4}, ) data = resp.json() assert data["prediction"] == 1 assert data["class"] == "fake" def test_model_name_in_response(self): """Response model field matches c2pa_checker.""" resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg())}, ) data = resp.json() assert data["model"] == "c2pa_checker" def test_inference_time_is_positive_float(self): """Inference time is a positive float.""" resp = client.post( "/predict", json={"image_data": _b64(_make_minimal_jpeg())}, ) data = resp.json() assert isinstance(data["inference_time"], float) assert data["inference_time"] > 0