import json from pathlib import Path import tempfile import unittest from unittest.mock import patch import cv2 import numpy as np from ocr_compare.synthetic import TARGET_CODES, _render, generate from ocr_compare.benchmark_synthetic import evaluate class SyntheticDatasetTests(unittest.TestCase): def test_document_level_split_and_exact_labels(self): with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "dataset" manifest = generate(output, count=6, seed=31415) self.assertEqual(manifest["source"], "fully synthetic") self.assertEqual(manifest["train_documents"], 5) self.assertEqual(manifest["validation_documents"], 1) rows = [json.loads(line) for line in (output / "ground_truth.jsonl").read_text().splitlines()] self.assertEqual(len(rows), 6) self.assertEqual({row["id"] for row in rows if row["split"] == "val"}, {"synthetic-00005"}) self.assertEqual(len((output / "det_train.txt").read_text().splitlines()), 5) self.assertEqual(len((output / "det_val.txt").read_text().splitlines()), 1) crop_count = 0 for row in rows: self.assertEqual(set(row["fields"]), set(TARGET_CODES)) page = cv2.imread(str(output / row["image"])) self.assertIsNotNone(page) self.assertEqual(page.shape[:2], (row["height"], row["width"])) for item in row["lines"]: self.assertLessEqual(len(item["text"]), 25) self.assertTrue((output / item["crop"]).is_file()) self.assertTrue(all(0 <= x < row["width"] and 0 <= y < row["height"] for x, y in item["polygon"])) crop_count += 1 self.assertEqual(crop_count, manifest["recognition_crops"]) for split in ("train", "val"): for entry in (output / f"rec_{split}.txt").read_text().splitlines(): path, label = entry.split("\t", 1) self.assertTrue((output / path).is_file()) self.assertEqual("synthetic-00005" in path, split == "val") self.assertTrue(label) def test_seed_repeats_pixels_and_values(self): first, lines_a, values_a, hardness_a = _render(12345, None) second, lines_b, values_b, hardness_b = _render(12345, None) self.assertTrue(np.array_equal(first, second)) self.assertEqual(lines_a, lines_b) self.assertEqual(values_a, values_b) self.assertEqual(hardness_a, hardness_b) def test_requires_a_validation_document_and_never_overwrites(self): with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "dataset" with self.assertRaises(ValueError): generate(output, count=1, seed=1) self.assertFalse(output.exists()) generate(output, count=2, seed=1) with self.assertRaises(FileExistsError): generate(output, count=2, seed=2) def test_benchmark_separates_reading_from_assignment(self): with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "dataset" generate(output, count=2, seed=31415) row = json.loads((output / "ground_truth.jsonl").read_text().splitlines()[-1]) readings = [] for item in row["lines"]: if item["part"] != "value" or item["code"] not in TARGET_CODES: continue xs, ys = zip(*item["polygon"]) readings.append({"text": item["text"], "box": [min(xs), min(ys), max(xs), max(ys)]}) fields = [{"code": code, "value": value} for code, value in row["fields"].items() if code not in ("B", "2.2")] fields.append({"code": "B", "value": "wrong"}) result = {"status": "ok", "lines": readings, "fields": fields} with patch("ocr_compare.benchmark_synthetic.server._classic", return_value=result): counts = evaluate(output, limit=1) self.assertEqual(counts["value_lines"], 7) self.assertEqual(counts["value_text_read"], 7) self.assertEqual(counts["fields_exact"], 4) self.assertEqual(counts["fields_wrong"], 1) self.assertEqual(counts["fields_missing"], 1) def test_benchmark_rejects_non_synthetic_manifest(self): with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "dataset" generate(output, count=2, seed=31415) manifest_path = output / "manifest.json" manifest = json.loads(manifest_path.read_text()) manifest["training_on_real_documents"] = True manifest_path.write_text(json.dumps(manifest)) with self.assertRaisesRegex(ValueError, "generated datasets only"): evaluate(output, limit=1) def test_benchmark_counts_failed_documents_in_denominator(self): with tempfile.TemporaryDirectory() as directory: output = Path(directory) / "dataset" generate(output, count=2, seed=31415) with patch("ocr_compare.benchmark_synthetic.server._classic", return_value={"status": "error"}): counts = evaluate(output, limit=1) self.assertEqual(counts["documents"], 1) self.assertEqual(counts["failed_documents"], 1) self.assertEqual(counts["value_lines"], 7) self.assertEqual(counts["value_text_read"], 0) self.assertEqual(counts["fields"], 6) self.assertEqual(counts["fields_missing"], 6) if __name__ == "__main__": unittest.main()