Files

293 lines
16 KiB
Python
Raw Permalink Normal View History

2026-10-07 21:00:21 +02:00
from io import BytesIO
import json
from pathlib import Path
import subprocess
import sys
2026-10-07 21:00:21 +02:00
import tempfile
from types import SimpleNamespace
2026-10-07 21:00:21 +02:00
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
2026-10-07 21:00:21 +02:00
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": "<table></table>"}
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")
2026-10-07 21:00:21 +02:00
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": "<table><tr><td>D.3</td><td>SAMPLE 123</td></tr>" +
"<tr><td>P.3</td><td>BENZIN</td></tr></table>",
}]}
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()