from io import BytesIO import json from pathlib import Path import subprocess import sys import tempfile from types import SimpleNamespace import unittest from unittest.mock import patch from PIL import Image from starlette.testclient import TestClient from ocr_compare import server from ocr_compare import classic_worker from ocr_compare.fields import assign_fields from ocr_compare.vision_worker import _missing_neighbor_box, _retry_table, extract def photo(width=800, height=500): stream = BytesIO() Image.new("RGB", (width, height), "white").save(stream, "JPEG") return stream.getvalue() class LocalComparisonTests(unittest.TestCase): def test_adjacent_field_code_needs_an_observed_value(self): lines = [ {"text": "D.3", "box": [10, 100, 35, 120]}, {"text": "ZXR 400", "box": [31, 96, 150, 130]}, {"text": "P.3", "box": [10, 200, 36, 220]}, {"text": "BENZIN", "box": [32, 197, 125, 228]}, {"text": "R", "box": [10, 290, 30, 310]}, ] found = {item["code"]: item for item in assign_fields(lines)} self.assertEqual(found["D.3"]["value"], "ZXR 400") self.assertEqual(found["P.3"]["value"], "BENZIN") self.assertNotIn("R", found) self.assertEqual(found["D.3"]["box"], [31, 96, 150, 130]) self.assertIsNone(found["D.3"]["confidence"]) def test_merged_code_is_reviewable_and_value_box_is_not_reused(self): lines = [ {"text": "D3330D", "box": [100, 100, 230, 135]}, {"text": "A", "box": [100, 200, 120, 230]}, {"text": "16", "box": [135, 200, 160, 230]}, {"text": "TESTNUMMER", "box": [170, 200, 320, 230]}, {"text": "C.1.2 Vorname(n)", "box": [10, 300, 240, 330]}, ] found = {item["code"]: item for item in assign_fields(lines)} self.assertEqual(found["D.3"]["value"], "330D") self.assertEqual(found["D.3"]["method"], "merged_code") self.assertEqual(found["16"]["value"], "TESTNUMMER") self.assertNotIn("A", found) self.assertNotIn("C.1.2", found) def test_dotted_type_codes_are_distinct_from_ambiguous_undotted_codes(self): lines = [ {"text": "2.1", "box": [100, 100, 130, 120]}, {"text": "0005", "box": [140, 100, 195, 120]}, {"text": "2.2", "box": [210, 100, 240, 120]}, {"text": "ABC123", "box": [250, 100, 330, 120]}, {"text": "21", "box": [100, 200, 130, 220]}, {"text": "UNSICHER", "box": [140, 200, 230, 220]}, ] found = {item["code"]: item for item in assign_fields(lines)} self.assertEqual(found["2.1"]["value"], "0005") self.assertEqual(found["2.2"]["value"], "ABC123") self.assertNotIn("21", found) def test_undotted_type_codes_need_the_observed_b_table_sequence(self): lines = [{"text": "table", "cells": ["B", "01.02.2003", "21", "0005", "22", "ABC123"], "box": [10, 10, 600, 120]}] found = {item["code"]: item for item in assign_fields(lines, allow_adjacent=False)} self.assertEqual(found["2.1"]["value"], "0005") self.assertEqual(found["2.2"]["value"], "ABC123") self.assertEqual(found["2.1"]["method"], "vl_table_context_code") self.assertNotIn("21", found) def test_holder_values_below_captions_join_only_touching_rows(self): lines = [ {"text": "C.1.1 Name oder Firmenname", "box": [10, 10, 260, 40]}, {"text": "MUSTER", "box": [20, 55, 180, 95]}, {"text": "C.1.2 Vorname(n)", "box": [10, 150, 210, 180]}, {"text": "ANNA", "box": [20, 195, 130, 235]}, {"text": "MARIA", "box": [125, 195, 250, 235]}, {"text": "C.1.3 Anschrift", "box": [10, 280, 170, 310]}, {"text": "TESTSTRASSE", "box": [20, 325, 240, 365]}, {"text": "6", "box": [235, 325, 255, 365]}, {"text": "12345", "box": [20, 370, 100, 410]}, {"text": "TESTSTADT", "box": [105, 370, 250, 410]}, {"text": "X Naechste HU", "box": [10, 480, 180, 520]}, ] found = {item["code"]: item for item in assign_fields(lines)} self.assertEqual(found["C.1.1"]["value"], "MUSTER") self.assertEqual(found["C.1.2"]["value"], "ANNA MARIA") self.assertEqual(found["C.1.3"]["value"], "TESTSTRASSE 6\n12345 TESTSTADT") self.assertEqual(found["C.1.3"]["box"], [20, 325, 255, 410]) def test_vision_caption_and_two_lines_in_one_layout_region(self): lines = [ {"text": "C.I.I Name oder Firmenname", "box": [10, 10, 250, 40]}, {"text": "MUSTER", "box": [20, 55, 130, 95]}, {"text": "C.1.3 Anschrift", "box": [10, 150, 180, 180]}, {"text": "TESTSTRASSE 6", "box": [20, 195, 260, 270]}, {"text": "12345 TESTSTADT", "box": [20, 195, 260, 270]}, {"text": "X Naechste HU", "box": [10, 300, 180, 350]}, ] found = {item["code"]: item for item in assign_fields(lines, allow_adjacent=False)} self.assertEqual(found["C.1.1"]["value"], "MUSTER") self.assertEqual(found["C.1.3"]["value"], "TESTSTRASSE 6\n12345 TESTSTADT") def test_missing_surname_caption_is_only_a_reviewable_nearby_suggestion(self): lines = [ {"text": "MUSTER", "box": [30, 100, 170, 145]}, {"text": "C.1.2 Vorname(n)", "box": [20, 210, 220, 245]}, {"text": "ANNA", "box": [30, 265, 140, 305]}, ] found = {item["code"]: item for item in assign_fields(lines)} self.assertEqual(found["C.1.1"]["method"], "surname_above_forename") self.assertEqual(found["C.1.1"]["value"], "MUSTER") self.assertNotIn("C.1.1", {item["code"] for item in assign_fields(lines, allow_adjacent=False)}) def test_vision_same_region_prefers_next_line_over_a_later_address(self): source = "PaddleOCR-VL-Layoutblock text; keine eigene Zeilenbox" lines = [ {"text": "C.1.2 Vorname(n)", "box": [10, 100, 250, 200], "source": source}, {"text": "ANNA", "box": [10, 100, 250, 200], "source": source}, {"text": "TESTSTRASSE", "box": [10, 220, 300, 270], "source": source}, ] found = {item["code"]: item for item in assign_fields(lines, allow_adjacent=False)} self.assertEqual(found["C.1.2"]["value"], "ANNA") self.assertEqual(found["C.1.2"]["method"], "caption_region_next_line") def test_vision_address_does_not_absorb_the_next_layout_region(self): source = "PaddleOCR-VL-Layoutblock text; keine eigene Zeilenbox" lines = [ {"text": "C.1.3 Anschrift", "box": [10, 100, 180, 140], "source": source}, {"text": "TESTSTRASSE 6 12345 TESTSTADT", "box": [10, 160, 310, 240], "source": source}, {"text": "Nächste HU", "box": [10, 300, 200, 350], "source": source}, ] found = {item["code"]: item for item in assign_fields(lines, allow_adjacent=False)} self.assertEqual(found["C.1.3"]["value"], "TESTSTRASSE 6 12345 TESTSTADT") def test_large_empty_table_is_retried_but_populated_one_is_not(self): block = {"block_label": "table", "block_bbox": [100, 10, 850, 790], "block_content": "
"} self.assertTrue(_retry_table(block, 1000, 800)) self.assertFalse(_retry_table({**block, "block_content": "TEXT " * 30}, 1000, 800)) self.assertFalse(_retry_table({**block, "block_bbox": [100, 10, 500, 400]}, 1000, 800)) def test_missing_neighbor_retry_requires_an_uncovered_column(self): right = {"block_label": "table", "block_bbox": [600, 200, 860, 650]} box = _missing_neighbor_box(right, [right], 1000, 800) self.assertIsNotNone(box) self.assertLess(box[0], 600) covered = {"block_label": "table", "block_bbox": [290, 200, 600, 650]} self.assertIsNone(_missing_neighbor_box(right, [right, covered], 1000, 800)) def test_overlapping_vision_tiles_do_not_erase_identical_field_readings(self): lines = [ {"text": "D.1 | TESTMARKE", "cells": ["D.1", "TESTMARKE"], "box": [100, 0, 500, 500]}, {"text": "D.1 | TESTMARKE", "cells": ["D.1", "TESTMARKE"], "box": [200, 0, 600, 500]}, ] fields = assign_fields(lines, allow_adjacent=False) self.assertEqual([(item["code"], item["value"]) for item in fields], [("D.1", "TESTMARKE")]) def test_wide_vision_timeout_retries_same_image_in_local_tiles(self): payload = {"lines": [{"text": "D.1 | TESTMARKE", "cells": ["D.1", "TESTMARKE"], "box": [10, 10, 100, 100]}], "regions": [], "retry_used": True} timed_out = subprocess.CompletedProcess([], 124, b"") tiled = subprocess.CompletedProcess([], 0, json.dumps(payload).encode()) with patch.dict(server.os.environ, {"OCR_VISION_MODE": "ssh", "OCR_VISION_SSH_TARGET": "test@example.invalid", "OCR_VISION_REMOTE_DIR": "/tmp/test"}): with patch.object(server.subprocess, "run", side_effect=[timed_out, tiled]) as remote: result = server._vision(photo(3600, 1800)) self.assertEqual(result["status"], "ok") self.assertEqual(result["fields"][0]["value"], "TESTMARKE") self.assertEqual(remote.call_count, 2) self.assertIn("--tiles-only", remote.call_args.args[0][-1]) def test_corners_allow_perspective_and_reject_crossing(self): source = photo(1200, 900) upright = {"corners": [[.1, .08], [.9, .13], [.84, .91], [.13, .85]], "rotation": 0} output, (width, height) = server._prepare(source, json.dumps(upright)) self.assertGreater(len(output), 1000) self.assertGreater(width, 700) self.assertGreater(height, 600) crossed = {"corners": [[.1, .1], [.9, .9], [.9, .1], [.1, .9]], "rotation": 0} with self.assertRaisesRegex(server.InputError, "invalid_corners"): server._prepare(source, json.dumps(crossed)) def test_realistic_phone_resolution_is_allowed(self): output, size = server._prepare(photo(4096, 3072), None) self.assertEqual(size, (4000, 3000)) self.assertGreater(len(output), 1000) def test_one_failed_engine_does_not_hide_the_other(self): classic = {"status": "ok", "model": "test classic", "elapsed_ms": 1, "lines": [{"text": "P.3 BENZIN", "box": [1, 1, 150, 30]}], "fields": [], "regions": []} vision = {"status": "error", "error": "vision_failed", "elapsed_ms": 2} with patch.object(server, "_classic", return_value=classic), patch.object(server, "_vision", return_value=vision): with TestClient(server.app) as client: response = client.post("/api/compare", content=photo(), headers={"content-type": "image/jpeg"}) self.assertEqual(response.status_code, 200) self.assertEqual(response.json()["results"]["classic"], classic) self.assertEqual(response.json()["results"]["vision"], vision) def test_fine_text_choice_reaches_classic_engine_only(self): classic = {"status": "error", "error": "classic_failed", "elapsed_ms": 0} vision = {"status": "error", "error": "vision_disabled", "elapsed_ms": 0} with patch.object(server, "_classic", return_value=classic) as run_classic, patch.object( server, "_vision", return_value=vision) as run_vision: with TestClient(server.app) as client: response = client.post("/api/compare", content=photo(), headers={ "content-type": "image/jpeg", "x-ocr-detail": "fine"}) invalid = client.post("/api/compare", content=photo(), headers={ "content-type": "image/jpeg", "x-ocr-detail": "unknown"}) self.assertEqual(response.status_code, 200) self.assertTrue(run_classic.call_args.kwargs["fine_text"]) self.assertEqual(run_vision.call_count, 1) self.assertEqual(invalid.status_code, 422) self.assertEqual(invalid.json()["error"], "invalid_detail") def test_classic_worker_uses_requested_detector_resolution(self): with tempfile.TemporaryDirectory() as directory: root = Path(directory) for name in ("PP-OCRv5_mobile_det", "latin_PP-OCRv5_mobile_rec"): model = root / name model.mkdir() (model / "inference.pdiparams").touch() seen = [] class FakePaddleOCR: def __init__(self, **kwargs): seen.append(kwargs) def predict(self, _image): return [SimpleNamespace(json={"res": {"rec_texts": ["B"], "rec_polys": [[[1, 1], [2, 1], [2, 2], [1, 2]]]}})] for options, expected in (([], 960), (["--fine-text"], 2048)): output = root / "result.json" with patch.object(classic_worker, "MODELS", root), patch.object( sys, "argv", ["classic_worker", str(root / "input.jpg"), str(output), *options]), patch.dict( sys.modules, {"paddleocr": SimpleNamespace(PaddleOCR=FakePaddleOCR)}): classic_worker.main() self.assertEqual(seen[-1]["text_det_limit_side_len"], expected) self.assertEqual(seen[-1]["text_det_limit_type"], "max") self.assertEqual(json.loads(output.read_text())["lines"][0]["text"], "B") def test_worker_launch_errors_are_reported_per_engine(self): with tempfile.TemporaryDirectory() as directory: model_home = Path(directory) for name in ("PP-OCRv5_mobile_det", "latin_PP-OCRv5_mobile_rec"): model = model_home / "official_models" / name model.mkdir(parents=True) (model / "inference.pdiparams").touch() with patch.object(server, "MODEL_HOME", model_home), patch.object( server.subprocess, "run", side_effect=OSError("worker unavailable")): classic = server._classic(photo(), model_home) self.assertEqual(classic["error"], "classic_failed") with patch.dict(server.os.environ, {"OCR_VISION_MODE": "ssh", "OCR_VISION_SSH_TARGET": "test@example.invalid", "OCR_VISION_REMOTE_DIR": "/tmp/test"}): with patch.object(server.subprocess, "run", side_effect=OSError("ssh unavailable")): vision = server._vision(photo()) self.assertEqual(vision["error"], "vision_failed") def test_invalid_input_is_rejected_before_inference(self): with TestClient(server.app) as client: response = client.post("/api/compare", content=b"not image", headers={"content-type": "image/jpeg"}) self.assertEqual(response.status_code, 422) self.assertEqual(response.json()["error"], "invalid_image") def test_vision_table_keeps_all_rows_with_region_evidence(self): payload = {"width": 900, "height": 600, "parsing_res_list": [{ "block_label": "table", "block_bbox": [20, 20, 850, 500], "block_content": "" + "
D.3SAMPLE 123
P.3BENZIN
", }]} result = extract(payload) self.assertEqual(len(result["lines"]), 2) fields = {f["code"]: f for f in assign_fields(result["lines"], allow_adjacent=False)} self.assertEqual(fields["P.3"]["value"], "BENZIN") self.assertEqual(fields["D.3"]["box"], [20, 20, 850, 500]) if __name__ == "__main__": unittest.main()