from io import BytesIO import json from pathlib import Path import subprocess import tempfile import unittest from unittest.mock import patch from PIL import Image from starlette.testclient import TestClient from ocr_compare import server 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_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()