293 lines
16 KiB
Python
293 lines
16 KiB
Python
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": "<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")
|
|
|
|
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()
|