"""SDXL invisible watermark detection service. Detects Stable Diffusion XL invisible watermarks embedded in images using the invisible-watermark library. SDXL embeds a specific 136-bit pattern (starting with "SDV2") via DWT-DCT encoding during generation. Returns a probability score based on whether the decoded bytes match known SDXL watermark patterns or exhibit non-random structure. """ import base64 import io import logging import math import os import sys import time from collections import Counter from typing import Any, Dict import cv2 import numpy as np import uvicorn from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from PIL import Image from pydantic import BaseModel # ── Logging ──────────────────────────────────────────────────────────────── logging.basicConfig( level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", handlers=[logging.StreamHandler(sys.stdout)], ) logger = logging.getLogger(__name__) # ── Config ───────────────────────────────────────────────────────────────── MODEL_NAME = "sdxl_watermark_detector" MODEL_PORT = int(os.environ.get("MODEL_PORT", "9002")) MAX_IMAGE_DIMENSION = 4096 _PRODUCTION = os.environ.get("PRODUCTION", "false").lower() == "true" # SDXL embeds 136 bits (17 bytes). The first 4 bytes are "SDV2". _WATERMARK_BITS = 136 _SDV2_PREFIX = b"SDV2" # ── Pydantic models ─────────────────────────────────────────────────────── class ImageInput(BaseModel): """Request body for /predict.""" image_data: str threshold: float = 0.5 # ── Detection helpers ────────────────────────────────────────────────────── def _compute_byte_entropy(data: bytes) -> float: """Compute Shannon entropy of a byte sequence. Args: data: Raw bytes to analyze. Returns: Entropy value in bits per byte (0.0 to 8.0). Returns 0.0 for empty input. """ if not data: return 0.0 length = len(data) counts = Counter(data) entropy = 0.0 for count in counts.values(): probability = count / length if probability > 0: entropy -= probability * math.log2(probability) return entropy def _hamming_distance(a: bytes, b: bytes) -> int: """Count the number of differing bits between two byte sequences. Args: a: First byte sequence. b: Second byte sequence (same length as a). Returns: Number of differing bits. """ dist = 0 for x, y in zip(a, b): dist += bin(x ^ y).count("1") return dist def _decode_watermark(bgr_array: np.ndarray) -> float: """Attempt to decode SDXL watermark from a BGR numpy array. Tries two decoding methods (dwtDct, dwtDctSvd) and scores the result based on pattern matching and entropy analysis. The SDV2 prefix check uses Hamming distance to tolerate minor bit errors introduced by lossy compression or PNG round-trips. The entropy threshold is strict (< 2.0) with a minimum unique byte count to avoid false positives on JPEG DCT noise. Args: bgr_array: OpenCV BGR image as numpy array. Returns: Probability score: 0.70 for exact SDV2 match, 0.65 for fuzzy SDV2 match, 0.5 for random/no signal. Scores are capped low to prevent flipping borderline model verdicts via PROVENANCE_WEIGHT boost. """ from imwatermark import WatermarkDecoder methods = ["dwtDct", "dwtDctSvd"] for method in methods: try: decoder = WatermarkDecoder("bytes", _WATERMARK_BITS) watermark_bytes = decoder.decode(bgr_array, method) if watermark_bytes is None or len(watermark_bytes) == 0: continue logger.debug( "Decoded %d bytes via %s: %s", len(watermark_bytes), method, watermark_bytes[:8].hex(), ) # Exact SDV2 prefix match if watermark_bytes[:4] == _SDV2_PREFIX: logger.info( "SDXL watermark detected via %s: SDV2 prefix matched", method, ) return 0.70 # Fuzzy SDV2 match: tolerate up to 2 bit errors in the # first 4 bytes (32 bits) to handle minor round-trip noise. # Higher tolerance causes false positives on real images. prefix_hamming = _hamming_distance(watermark_bytes[:4], _SDV2_PREFIX) if prefix_hamming <= 2: logger.info( "SDXL watermark detected via %s: fuzzy SDV2 match " "(hamming=%d)", method, prefix_hamming, ) return 0.65 except Exception as exc: logger.debug("Watermark decode failed with %s: %s", method, exc) continue return 0.5 def _detect_watermark(image_bytes: bytes) -> float: """Run watermark detection on raw image bytes. Decodes image bytes with PIL, converts to RGB numpy array, and attempts watermark decoding. Args: image_bytes: Raw bytes of the image file. Returns: Float probability in [0, 1]. 0.5 means neutral (no signal). """ try: image = Image.open(io.BytesIO(image_bytes)) if image.width > MAX_IMAGE_DIMENSION or image.height > MAX_IMAGE_DIMENSION: logger.warning( "Image too large: %dx%d, max %d", image.width, image.height, MAX_IMAGE_DIMENSION, ) return 0.5 image = image.convert("RGB") # invisible-watermark library operates on RGB arrays (not BGR). # The encoder/decoder use channel-independent DWT transforms. rgb_array = np.array(image) return _decode_watermark(rgb_array) except Exception as exc: logger.warning("Watermark detection error: %s", exc) return 0.5 # ── FastAPI app ──────────────────────────────────────────────────────────── app = FastAPI( title="SDXL Invisible Watermark Detector Service", description=( "Detects Stable Diffusion XL invisible watermarks embedded " "in images using DWT-DCT steganographic decoding." ), version="1.0.0", docs_url=None if _PRODUCTION else "/docs", redoc_url=None if _PRODUCTION else "/redoc", openapi_url=None if _PRODUCTION else "/openapi.json", ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) @app.get("/") async def root(): """Root endpoint with service information.""" if _PRODUCTION: return {"status": "ok"} return { "model_name": MODEL_NAME, "description": ( "SDXL invisible watermark detector using " "DWT-DCT steganographic decoding" ), "device": "cpu", "model_loaded": True, } @app.get("/health") async def health(): """Health check endpoint.""" if _PRODUCTION: return {"status": "healthy"} return { "status": "healthy", "model_name": MODEL_NAME, "device": "cpu", "model_loaded": True, } @app.post("/predict") async def predict(payload: ImageInput) -> Dict[str, Any]: """Detect SDXL invisible watermark in a base64-encoded image. Args: payload: Base64-encoded image data and optional threshold. Returns: Dict with model name, probability, prediction, class, and inference time. """ if not payload.image_data: raise HTTPException( status_code=400, detail="No image data provided.", ) start = time.time() try: image_bytes = base64.b64decode(payload.image_data) except Exception as exc: raise HTTPException( status_code=400, detail=f"Invalid base64 data: {exc}", ) if not image_bytes: raise HTTPException( status_code=400, detail="Empty image data after base64 decode.", ) probability = _detect_watermark(image_bytes) prediction = 1 if probability > payload.threshold else 0 class_label = "fake" if prediction == 1 else "real" inference_time = time.time() - start logger.info( "Watermark check: %s (prob=%.4f, %.3fs)", class_label, probability, inference_time, ) return { "model": MODEL_NAME, "probability": float(probability), "prediction": int(prediction), "class": class_label, "inference_time": float(inference_time), } if __name__ == "__main__": logger.info( "Starting SDXL watermark detector service on port %d", MODEL_PORT, ) uvicorn.run(app, host="0.0.0.0", port=MODEL_PORT)