Files
zulassung-ocr-compare/ocr_compare/server.py
T

312 lines
14 KiB
Python

"""Local-only, privacy-preserving comparison of classic OCR and PaddleOCR-VL."""
from __future__ import annotations
import asyncio
from io import BytesIO
import json
import os
from pathlib import Path
import shlex
import subprocess
import sys
import tempfile
import time
import warnings
import cv2
import numpy as np
from PIL import Image, ImageOps, UnidentifiedImageError
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import FileResponse, JSONResponse
from starlette.routing import Route
from .fields import FIELD_CODES, assign_fields
HERE = Path(__file__).resolve().parent
PROJECT = HERE.parent
MODEL_HOME = Path(os.environ.get("OCR_MODEL_HOME", PROJECT / ".cache")).expanduser().resolve()
def _remote_command(seconds: int = 180, *, tiles_only: bool = False) -> str:
remote_dir = os.environ["OCR_VISION_REMOTE_DIR"]
remote_python = os.environ.get("OCR_VISION_REMOTE_PYTHON", "./.venv-vision/bin/python")
runner = f'{shlex.quote(remote_python)} -m ocr_compare.vision_worker'
if tiles_only:
runner += ' --tiles-only'
isolation = 'unshare -Urn ' if os.environ.get("OCR_VISION_REMOTE_UNSHARE") == "1" else ""
device = shlex.quote(os.environ.get("OCR_VISION_REMOTE_DEVICE", "gpu:0"))
model_home = os.environ.get("OCR_VISION_REMOTE_MODEL_HOME")
model_env = f'OCR_MODEL_HOME={shlex.quote(model_home)} ' if model_home else ""
return (
f'set -eu; cd {shlex.quote(remote_dir)}; directory=$(mktemp -d); '
'trap \'rm -rf -- "$directory"\' EXIT; '
f'OCR_PRIVATE_DIR="$directory" OCR_VISION_DEVICE={device} {model_env}'
f'timeout -k 5s {seconds}s {isolation}{runner}'
)
MAX_UPLOAD = 16 * 1024 * 1024
MAX_PIXELS = 24_000_000
MAX_SIDE = 10_000
OCR_SIDE = 4000
Image.MAX_IMAGE_PIXELS = MAX_PIXELS
_slot = asyncio.Semaphore(1)
class InputError(ValueError):
pass
def _geometry(image: Image.Image, header: str | None) -> Image.Image:
try:
spec = json.loads(header) if header else {}
rotation = int(spec.get("rotation", 0))
if rotation not in (0, 90, 180, 270):
raise InputError("invalid_rotation")
points = spec.get("corners")
if points is not None:
if (not isinstance(points, list) or len(points) != 4 or
any(not isinstance(p, list) or len(p) != 2 or
any(not isinstance(v, (int, float)) or isinstance(v, bool) or
not 0 <= v <= 1 for v in p) for p in points)):
raise InputError("invalid_corners")
w, h = image.size
source = np.array([[x * (w - 1), y * (h - 1)] for x, y in points], np.float32)
area = cv2.contourArea(source)
edges = [float(np.linalg.norm(source[(i + 1) % 4] - source[i])) for i in range(4)]
turns = []
for i in range(4):
ax, ay = source[(i + 1) % 4] - source[i]
bx, by = source[(i + 2) % 4] - source[(i + 1) % 4]
turns.append(float(ax * by - ay * bx))
if area < .02 * w * h or min(edges) < 30 or min(turns) <= 0:
raise InputError("invalid_corners")
out_w = round(max(edges[0], edges[2]))
out_h = round(max(edges[1], edges[3]))
if out_w < 30 or out_h < 30 or out_w * out_h > MAX_PIXELS:
raise InputError("invalid_corners")
# The four positions are always TL, TR, BR, BL. This avoids the
# previous fixed-template aspect-ratio and rotation ambiguity.
target = np.array([[0, 0], [out_w - 1, 0],
[out_w - 1, out_h - 1], [0, out_h - 1]], np.float32)
matrix = cv2.getPerspectiveTransform(source, target)
warped = cv2.warpPerspective(np.asarray(image), matrix, (out_w, out_h),
flags=cv2.INTER_LINEAR,
borderMode=cv2.BORDER_CONSTANT,
borderValue=(255, 255, 255))
image = Image.fromarray(warped)
if rotation:
image = image.rotate(-rotation, expand=True)
if max(image.size) > OCR_SIDE:
image.thumbnail((OCR_SIDE, OCR_SIDE), Image.Resampling.LANCZOS)
return image
except (ValueError, TypeError, KeyError, json.JSONDecodeError) as exc:
if isinstance(exc, InputError):
raise
raise InputError("invalid_corners") from None
def _prepare(data: bytes, header: str | None) -> tuple[bytes, tuple[int, int]]:
try:
with warnings.catch_warnings():
warnings.simplefilter("error", Image.DecompressionBombWarning)
with Image.open(BytesIO(data)) as source:
if source.format not in ("JPEG", "PNG", "MPO"):
raise InputError("unsupported_encoding")
w, h = source.size
if min(w, h) < 1 or max(w, h) > MAX_SIDE or w * h > MAX_PIXELS:
raise InputError("image_too_large")
source.load()
image = ImageOps.exif_transpose(source).convert("RGB")
image = _geometry(image, header)
output = BytesIO()
image.save(output, format="JPEG", quality=92, subsampling=0)
return output.getvalue(), image.size
except InputError:
raise
except (OSError, ValueError, UnidentifiedImageError, Image.DecompressionBombError,
Image.DecompressionBombWarning):
raise InputError("invalid_image") from None
def _classic(data: bytes, directory: Path, *, fine_text: bool = False) -> dict:
started = time.monotonic()
models = MODEL_HOME / "official_models"
if any(not (models / name / "inference.pdiparams").is_file() for name in
("PP-OCRv5_mobile_det", "latin_PP-OCRv5_mobile_rec")):
return {"status": "error", "error": "classic_models_missing", "elapsed_ms": 0}
image_path = directory / "same-input.jpg"
output_path = directory / "classic.json"
image_path.write_bytes(data)
env = os.environ.copy()
env.update(PADDLE_PDX_CACHE_HOME=str(MODEL_HOME),
HF_HOME=str(MODEL_HOME / "hf"), HF_HUB_OFFLINE="1",
TRANSFORMERS_OFFLINE="1", PADDLE_PDX_DISABLE_MODEL_SOURCE_CHECK="True",
OMP_NUM_THREADS="1", OPENBLAS_NUM_THREADS="1", OMP_THREAD_LIMIT="1")
try:
command = [sys.executable, "-m", "ocr_compare.classic_worker", str(image_path), str(output_path)]
if fine_text:
command.append("--fine-text")
process = subprocess.run(
command,
cwd=PROJECT, env=env, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
timeout=120, check=False,
)
except subprocess.TimeoutExpired:
return {"status": "error", "error": "classic_timeout", "elapsed_ms": round((time.monotonic()-started)*1000)}
except OSError:
return {"status": "error", "error": "classic_failed", "elapsed_ms": round((time.monotonic()-started)*1000)}
elapsed = round((time.monotonic() - started) * 1000)
if process.returncode or not output_path.is_file():
return {"status": "error", "error": "classic_failed", "elapsed_ms": elapsed}
try:
raw = json.loads(output_path.read_text())
lines = [{"text": row["text"], "box": row["bbox"], "source": "OCR-Zeile"}
for row in raw["lines"] if row.get("bbox")]
except (OSError, ValueError, KeyError, TypeError):
return {"status": "error", "error": "classic_invalid", "elapsed_ms": elapsed}
model = "PP-OCRv5 mobile (Latin), lokal" + (" · Feintext" if fine_text else "")
return {"status": "ok", "model": model,
"elapsed_ms": elapsed, "lines": lines, "fields": assign_fields(lines),
"regions": []}
def _vision(data: bytes) -> dict:
started = time.monotonic()
mode = os.environ.get("OCR_VISION_MODE", "local").lower()
if mode == "off":
return {"status": "error", "error": "vision_disabled", "elapsed_ms": 0}
if mode not in ("local", "ssh"):
return {"status": "error", "error": "vision_config_error", "elapsed_ms": 0}
if mode == "local" and any(not path.is_file() for path in (
MODEL_HOME / "official_models" / "PP-DocLayoutV3" / "inference.pdiparams",
MODEL_HOME / "official_models" / "PaddleOCR-VL-1.6" / "model.safetensors")):
return {"status": "error", "error": "vision_models_missing", "elapsed_ms": 0}
if mode == "ssh" and (not os.environ.get("OCR_VISION_SSH_TARGET") or
not os.environ.get("OCR_VISION_REMOTE_DIR")):
return {"status": "error", "error": "vision_config_error", "elapsed_ms": 0}
with Image.open(BytesIO(data)) as page:
wide = page.width > 2500 and page.width / page.height > 1.65
limit = 90 if wide else 180
def run_worker(seconds: int, *, tiles_only: bool = False):
if mode == "ssh":
command = ["ssh", "-T", "-o", "BatchMode=yes", "-o", "StrictHostKeyChecking=yes",
"-o", "ConnectTimeout=10"]
identity = os.environ.get("OCR_VISION_SSH_IDENTITY")
if identity:
command += ["-i", str(Path(identity).expanduser()), "-o", "IdentitiesOnly=yes"]
command += [os.environ["OCR_VISION_SSH_TARGET"],
_remote_command(seconds, tiles_only=tiles_only)]
env = None
else:
command = [os.environ.get("OCR_VISION_PYTHON", sys.executable),
"-m", "ocr_compare.vision_worker"]
if tiles_only:
command.append("--tiles-only")
env = os.environ.copy()
env.update(OCR_MODEL_HOME=str(MODEL_HOME),
OCR_VISION_DEVICE=os.environ.get("OCR_VISION_DEVICE", "cpu"))
try:
with tempfile.TemporaryDirectory(prefix="ocr-vision-") as private:
if env is not None:
env["OCR_PRIVATE_DIR"] = private
return subprocess.run(command, input=data, stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL, timeout=seconds + 15,
check=False, cwd=PROJECT if mode == "local" else None, env=env)
except subprocess.TimeoutExpired:
return None
except OSError:
return subprocess.CompletedProcess(command, 127, b"")
process = run_worker(limit)
if wide and (process is None or process.returncode == 124):
# Wide documents can form one prohibitively expensive layout block.
# The three generic views cover the same uploaded page without a
# document-specific field template.
process = run_worker(90, tiles_only=True)
elapsed = round((time.monotonic() - started) * 1000)
if process is None or process.returncode == 124:
return {"status": "error", "error": "vision_timeout", "elapsed_ms": elapsed}
if process.returncode or len(process.stdout) > 2_000_000:
return {"status": "error", "error": "vision_failed", "elapsed_ms": elapsed}
try:
raw = json.loads(process.stdout)
if raw.get("error"):
code = "vision_models_missing" if raw["error"] == "models_missing" else "vision_failed"
return {"status": "error", "error": code, "elapsed_ms": elapsed}
lines = raw["lines"]
regions = raw["regions"]
if not isinstance(lines, list) or not isinstance(regions, list):
raise ValueError("invalid collections")
except (ValueError, TypeError, KeyError):
return {"status": "error", "error": "vision_invalid", "elapsed_ms": elapsed}
return {"status": "ok", "model": raw.get("model", "PaddleOCR-VL"),
"elapsed_ms": elapsed, "lines": lines,
"fields": assign_fields(lines, allow_adjacent=False),
"regions": regions, "layout_note": raw.get("layout_note", ""),
"retry_used": bool(raw.get("retry_used", False))}
async def index(_: Request):
return FileResponse(HERE / "index.html", media_type="text/html; charset=utf-8")
async def script(_: Request):
return FileResponse(HERE / "app.js", media_type="text/javascript; charset=utf-8")
async def style(_: Request):
return FileResponse(HERE / "style.css", media_type="text/css; charset=utf-8")
async def health(_: Request):
return JSONResponse({"ready": True, "scope": "localhost only",
"vision_mode": os.environ.get("OCR_VISION_MODE", "local")})
async def compare(request: Request):
if _slot.locked():
return JSONResponse({"error": "busy"}, status_code=429)
if request.headers.get("content-type", "").split(";", 1)[0] not in ("image/jpeg", "image/png"):
return JSONResponse({"error": "unsupported_media_type"}, status_code=415)
detail = request.headers.get("x-ocr-detail", "standard")
if detail not in ("standard", "fine"):
return JSONResponse({"error": "invalid_detail"}, status_code=422)
async with _slot:
data = bytearray()
try:
async with asyncio.timeout(30):
async for chunk in request.stream():
data.extend(chunk)
if len(data) > MAX_UPLOAD:
return JSONResponse({"error": "upload_too_large"}, status_code=413)
except TimeoutError:
return JSONResponse({"error": "upload_timeout"}, status_code=408)
try:
image, (width, height) = await asyncio.to_thread(
_prepare, bytes(data), request.headers.get("x-ocr-geometry"))
except InputError as exc:
return JSONResponse({"error": str(exc)}, status_code=422)
with tempfile.TemporaryDirectory(prefix="ocr-compare-") as path:
directory = Path(path)
classic, vision = await asyncio.gather(
asyncio.to_thread(_classic, image, directory, fine_text=detail == "fine"),
asyncio.to_thread(_vision, image),
)
# Exact JPEG bytes sent to both pipelines are returned for evidence
# highlighting. They never enter a report, log, or external OCR API.
import base64
return JSONResponse({"image": "data:image/jpeg;base64," + base64.b64encode(image).decode(),
"image_size": {"width": width, "height": height},
"field_codes": FIELD_CODES,
"results": {"classic": classic, "vision": vision}})
app = Starlette(routes=[Route("/", index), Route("/app.js", script),
Route("/style.css", style), Route("/api/health", health),
Route("/api/compare", compare, methods=["POST"])])