247 lines
13 KiB
Python
247 lines
13 KiB
Python
|
|
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": "<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_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()
|