test_ocr_service.py 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100
  1. from __future__ import annotations
  2. import pytest
  3. class FakeProvider:
  4. def __init__(self, responses):
  5. self.responses = list(responses)
  6. self.calls = []
  7. def extract(self, image_bytes):
  8. self.calls.append(image_bytes)
  9. return self.responses.pop(0)
  10. def block(text, bbox, confidence=0.9):
  11. from app.core.data_research.ocr.base import OcrBlock
  12. return OcrBlock(text=text, bbox=bbox, page=99, confidence=confidence)
  13. def test_ocr_service_preserves_scanned_pdf_page_order_and_review_flags():
  14. from app.core.data_research.ocr.service import OcrService
  15. provider = FakeProvider(
  16. [
  17. [block("第一页", (0.1, 0.2, 0.8, 0.3), 0.95)],
  18. [block("第二页低置信", (0.0, 0.0, 1.0, 1.0), 0.4)],
  19. ]
  20. )
  21. results = OcrService(provider, review_threshold=0.75).extract_pages(
  22. [b"page-one", b"page-two"]
  23. )
  24. assert [result.page for result in results] == [1, 2]
  25. assert results[0].bbox == (0.1, 0.2, 0.8, 0.3)
  26. assert results[0].review_required is False
  27. assert results[1].review_required is True
  28. def test_ocr_service_rejects_malformed_provider_response():
  29. from app.core.data_research.ocr.service import OcrResponseInvalid, OcrService
  30. provider = FakeProvider([[block("bad", (-1.0, 0.0, 2.0, 1.0))]])
  31. with pytest.raises(OcrResponseInvalid, match="bounding box"):
  32. OcrService(provider).extract_pages([b"image"])
  33. def test_ocr_is_fail_closed_when_provider_is_disabled():
  34. from app.core.data_research.ocr.service import OcrDisabled, OcrService
  35. with pytest.raises(OcrDisabled, match="not configured"):
  36. OcrService(None).extract_pages([b"image"])
  37. class FakeResponse:
  38. def __init__(self, payload, *, content=b"{}", headers=None):
  39. self._payload = payload
  40. self.content = content
  41. self.headers = headers or {}
  42. def raise_for_status(self):
  43. return None
  44. def json(self):
  45. return self._payload
  46. def test_http_provider_enforces_timeout_tls_and_response_limit():
  47. from app.core.data_research.ocr.http_provider import HttpOcrProvider
  48. from app.core.data_research.ocr.service import OcrResponseInvalid
  49. calls = []
  50. def post(*args, **kwargs):
  51. calls.append((args, kwargs))
  52. return FakeResponse(
  53. {"blocks": [{"text": "编号", "bbox": [0, 0, 1, 1], "confidence": 0.9}]},
  54. content=b"x" * 20,
  55. )
  56. provider = HttpOcrProvider(
  57. "https://ocr.internal/v1/extract",
  58. token="secret-token",
  59. timeout_seconds=3,
  60. verify_tls=True,
  61. max_response_bytes=100,
  62. post=post,
  63. )
  64. blocks = provider.extract(b"png")
  65. assert blocks[0].text == "编号"
  66. assert calls[0][1]["timeout"] == 3
  67. assert calls[0][1]["verify"] is True
  68. provider.max_response_bytes = 10
  69. with pytest.raises(OcrResponseInvalid, match="response size") as error:
  70. provider.extract(b"png-secret-source")
  71. assert "secret-token" not in str(error.value)
  72. assert "png-secret-source" not in str(error.value)