Complete local hybrid search and improve agent reliability
This commit is contained in:
@@ -0,0 +1,266 @@
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
import base64
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from PIL import Image, ImageDraw
|
||||
from pypdf import PdfReader, PdfWriter
|
||||
from fastapi.testclient import TestClient
|
||||
from app.local_ocr import LocalOCRError
|
||||
from app.pdf_documents import extract_pdf_pages_text
|
||||
|
||||
_TEST_DATA_DIR = None
|
||||
if "SOVEREIGNAI_DATA_DIR" not in os.environ:
|
||||
_TEST_DATA_DIR = tempfile.TemporaryDirectory(prefix="sovereignai-pdf-tests-")
|
||||
os.environ["SOVEREIGNAI_DATA_DIR"] = _TEST_DATA_DIR.name
|
||||
|
||||
from app.main import app
|
||||
|
||||
|
||||
def make_pdf(text: str | None) -> bytes:
|
||||
stream = (
|
||||
b"BT /F1 12 Tf 72 720 Td (" + text.encode("ascii") + b") Tj ET"
|
||||
if text is not None
|
||||
else b""
|
||||
)
|
||||
objects = [
|
||||
b"<< /Type /Catalog /Pages 2 0 R >>",
|
||||
b"<< /Type /Pages /Kids [3 0 R] /Count 1 >>",
|
||||
b"<< /Type /Page /Parent 2 0 R /MediaBox [0 0 612 792] /Resources << /Font << /F1 4 0 R >> >> /Contents 5 0 R >>",
|
||||
b"<< /Type /Font /Subtype /Type1 /BaseFont /Helvetica >>",
|
||||
b"<< /Length " + str(len(stream)).encode("ascii") + b" >>\nstream\n" + stream + b"\nendstream",
|
||||
]
|
||||
data = bytearray(b"%PDF-1.4\n")
|
||||
offsets = [0]
|
||||
for number, body in enumerate(objects, 1):
|
||||
offsets.append(len(data))
|
||||
data.extend(f"{number} 0 obj\n".encode("ascii") + body + b"\nendobj\n")
|
||||
xref_offset = len(data)
|
||||
data.extend(f"xref\n0 {len(objects) + 1}\n".encode("ascii"))
|
||||
data.extend(b"0000000000 65535 f \n")
|
||||
for offset in offsets[1:]:
|
||||
data.extend(f"{offset:010d} 00000 n \n".encode("ascii"))
|
||||
data.extend(
|
||||
f"trailer\n<< /Size {len(objects) + 1} /Root 1 0 R >>\nstartxref\n{xref_offset}\n%%EOF\n".encode("ascii")
|
||||
)
|
||||
return bytes(data)
|
||||
|
||||
|
||||
def make_scanned_pdf(page_count: int = 1) -> bytes:
|
||||
pages = []
|
||||
for index in range(page_count):
|
||||
image = Image.new("RGB", (900, 1165), "white")
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.text((60, 80), f"Community Reading Meetup\nPlace: Amman\nPage: {index + 1}", fill="black")
|
||||
pages.append(image)
|
||||
output = BytesIO()
|
||||
pages[0].save(
|
||||
output,
|
||||
format="PDF",
|
||||
resolution=144,
|
||||
save_all=True,
|
||||
append_images=pages[1:],
|
||||
)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
def make_mixed_pdf() -> bytes:
|
||||
writer = PdfWriter()
|
||||
for page_data in (make_pdf("Digital page: IndexMarker"), make_scanned_pdf()):
|
||||
writer.append_pages_from_reader(PdfReader(BytesIO(page_data)))
|
||||
output = BytesIO()
|
||||
writer.write(output)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
class PdfAnalysisTests(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls) -> None:
|
||||
cls.client = TestClient(app)
|
||||
|
||||
def test_extracts_pdf_text_and_sends_page_number_to_model(self) -> None:
|
||||
completion = {"choices": [{"message": {"content": "يتحدث الملف عن لقاء مجتمعي في عمّان."}}]}
|
||||
with patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model:
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "لخص محتوى الملف"},
|
||||
files={"files": ("meetup.pdf", make_pdf("Community Meetup Amman"), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
self.assertEqual(response.json()["files"], ["meetup.pdf"])
|
||||
sent_prompt = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertIn("--- صفحة 1 ---", sent_prompt)
|
||||
self.assertIn("Community Meetup Amman", sent_prompt)
|
||||
self.assertIn("meetup.pdf", sent_prompt)
|
||||
|
||||
def test_extracts_text_per_page_and_marks_scanned_pages(self) -> None:
|
||||
parsed = extract_pdf_pages_text(make_mixed_pdf(), "mixed.pdf")
|
||||
|
||||
self.assertEqual(parsed["total_pages"], 2)
|
||||
self.assertTrue(parsed["pages"][0]["has_text"])
|
||||
self.assertIn("IndexMarker", parsed["pages"][0]["text"])
|
||||
self.assertFalse(parsed["pages"][1]["has_text"])
|
||||
|
||||
def test_mixed_pdf_analysis_combines_digital_text_ocr_and_page_image(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "يلخص المستند الرقمي والممسوح."}}]}
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "ScannedPageMarker Community Reading Meetup",
|
||||
"lines": [],
|
||||
"average_confidence": 0.8,
|
||||
"elapsed_seconds": 1.0,
|
||||
}
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", return_value=ocr),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "ما محتوى الملف؟"},
|
||||
files={"files": ("mixed.pdf", make_mixed_pdf(), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["pages_rendered"], 1)
|
||||
self.assertIn("ScannedPageMarker", body["result"])
|
||||
content = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertTrue(any("IndexMarker" in part.get("text", "") for part in content))
|
||||
self.assertTrue(any("ScannedPageMarker" in part.get("text", "") for part in content))
|
||||
self.assertEqual(len([part for part in content if part["type"] == "image_url"]), 1)
|
||||
|
||||
def test_scanned_pdf_is_rendered_and_routed_to_local_vision_model(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["gemma4:e2b", "ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "Community Reading Meetup — Amman."}}]}
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", side_effect=LocalOCRError("OCR weights absent")),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "اقرأ اسم اللقاء والمكان."},
|
||||
files={"files": ("meetup-scan.pdf", make_scanned_pdf(4), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["tool"], "local-pdf-vision-analysis")
|
||||
self.assertEqual(body["model"], "ministral-3:3b")
|
||||
self.assertTrue(body["auto_routed"])
|
||||
self.assertEqual(body["pages_rendered"], 3)
|
||||
self.assertIn("أول 3 صفحات فقط", body["result"])
|
||||
content = model.await_args.args[0]["messages"][1]["content"]
|
||||
system = model.await_args.args[0]["messages"][0]["content"]
|
||||
images = [part for part in content if part["type"] == "image_url"]
|
||||
self.assertEqual(len(images), 3)
|
||||
self.assertIn("دون ترجمتها", system)
|
||||
image_data = images[0]["image_url"].split(",", 1)[1]
|
||||
self.assertTrue(base64.b64decode(image_data).startswith(b"\xff\xd8\xff"))
|
||||
|
||||
def test_scanned_pdf_ocr_text_is_given_to_vision_model_and_returned(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "الإعلان عن لقاء قراءة."}}]}
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "لقاء القراءة المجتمعي\nCommunity Reading Meetup\nPlace: Amman",
|
||||
"lines": [],
|
||||
"average_confidence": 0.75,
|
||||
"elapsed_seconds": 1.2,
|
||||
}
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", return_value=ocr),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
data={"question": "انسخ عنوان الإعلان ومكانه."},
|
||||
files={"files": ("meetup-scan.pdf", make_scanned_pdf(), "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
body = response.json()
|
||||
self.assertEqual(body["ocr_engine"], "easyocr-local-ar-en")
|
||||
self.assertEqual(body["ocr_pages"], 1)
|
||||
self.assertIn("Community Reading Meetup", body["result"])
|
||||
content = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertTrue(any("Amman" in part.get("text", "") for part in content))
|
||||
|
||||
def test_image_ocr_text_is_fused_with_the_visual_model_request(self) -> None:
|
||||
provider = MagicMock()
|
||||
provider.default_model = "gemma4:e2b"
|
||||
provider.list_models = AsyncMock(return_value=["ministral-3:3b"])
|
||||
completion = {"choices": [{"message": {"content": "المكان عمّان."}}]}
|
||||
ocr = {
|
||||
"engine": "easyocr-local-ar-en",
|
||||
"text": "Place: Amman\nالمكان: عمان",
|
||||
"lines": [],
|
||||
"average_confidence": 0.82,
|
||||
"elapsed_seconds": 2.1,
|
||||
}
|
||||
image_output = BytesIO()
|
||||
Image.new("RGB", (100, 100), "white").save(image_output, format="PNG")
|
||||
with (
|
||||
patch("app.main.get_model_provider", return_value=provider),
|
||||
patch("app.main.recognize_image_text", return_value=ocr),
|
||||
patch("app.main.get_completion", new=AsyncMock(return_value=completion)) as model,
|
||||
):
|
||||
response = self.client.post(
|
||||
"/v1/agent/images/analyze",
|
||||
data={"question": "ما المكان المكتوب؟"},
|
||||
files={"files": ("notice.png", image_output.getvalue(), "image/png")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
self.assertEqual(response.json()["ocr_engine"], "easyocr-local-ar-en")
|
||||
self.assertIn("Place: Amman", response.json()["result"])
|
||||
user_parts = model.await_args.args[0]["messages"][1]["content"]
|
||||
self.assertIn("Place: Amman", user_parts[0]["text"])
|
||||
|
||||
def test_rejects_invalid_pdf_signature(self) -> None:
|
||||
response = self.client.post(
|
||||
"/v1/agent/files/analyze",
|
||||
files={"files": ("broken.pdf", b"not a pdf", "application/pdf")},
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 415)
|
||||
self.assertIn("ترويسة ملف PDF", response.json()["detail"])
|
||||
|
||||
def test_workspace_pdf_is_listed_indexed_and_searchable(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "meetup.pdf").write_bytes(make_pdf("Community Meetup Amman"))
|
||||
listed = self.client.post(
|
||||
"/v1/agent/workspace/files", json={"workspace_path": str(root)}
|
||||
)
|
||||
indexed = self.client.post(
|
||||
"/v1/agent/knowledge/index",
|
||||
json={"workspace_path": str(root), "files": ["meetup.pdf"]},
|
||||
)
|
||||
searched = self.client.post(
|
||||
"/v1/agent/knowledge/search",
|
||||
json={"workspace_path": str(root), "task": "Community Meetup"},
|
||||
)
|
||||
|
||||
self.assertEqual(listed.status_code, 200, listed.text)
|
||||
self.assertIn("meetup.pdf", listed.json()["files"])
|
||||
self.assertEqual(indexed.status_code, 200, indexed.text)
|
||||
self.assertEqual(searched.status_code, 200, searched.text)
|
||||
self.assertIn("Community Meetup Amman", searched.json()["results"][0]["text"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user