Files
OCR/docengine/tests/repositories/test_document_repository.py
2026-06-01 21:49:53 +05:30

279 lines
9.9 KiB
Python

from __future__ import annotations
import uuid
import pytest
from sqlalchemy.orm import Session
from app.models.document import (
Document,
DocumentImage,
DocumentPage,
DocumentTable,
DocumentTextBlock,
TemplateMatch,
)
from app.repositories.document_repository import (
DocumentPageRepository,
DocumentRepository,
DocumentTextBlockRepository,
TemplateMatchRepository,
)
from tests.conftest import (
create_test_document,
create_test_fingerprint,
create_test_page,
create_test_template,
create_test_text_block,
create_test_user,
)
class TestDocumentRepository:
"""Tests for DocumentRepository CRUD."""
def test_create_document(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user, filename="created.pdf")
assert doc.id is not None
assert doc.original_filename == "created.pdf"
assert doc.status == "pending"
def test_get_by_id(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
repo = DocumentRepository(db)
found = repo.get_by_id(doc.id)
assert found is not None
assert found.id == doc.id
def test_get_by_id_not_found(self, db: Session) -> None:
repo = DocumentRepository(db)
assert repo.get_by_id(uuid.uuid4()) is None
def test_get_with_pages(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
page1 = create_test_page(db, doc, page_number=1)
page2 = create_test_page(db, doc, page_number=2)
create_test_text_block(db, page1, text="Page 1 text")
repo = DocumentRepository(db)
result = repo.get_with_pages(doc.id)
assert result is not None
assert len(result.pages) == 2
def test_update_status(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user, status="pending")
repo = DocumentRepository(db)
repo.update_status(doc.id, "processing")
db.flush()
updated = repo.get_by_id(doc.id)
assert updated is not None
assert updated.status == "processing"
def test_update_status_with_error(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user, status="processing")
repo = DocumentRepository(db)
repo.update_status(doc.id, "failed", error_message="OCR engine crashed")
db.flush()
updated = repo.get_by_id(doc.id)
assert updated is not None
assert updated.status == "failed"
assert updated.error_message == "OCR engine crashed"
def test_get_all_with_pagination(self, db: Session) -> None:
user = create_test_user(db)
for i in range(5):
create_test_document(db, user=user, filename=f"pagdoc_{i}.pdf")
repo = DocumentRepository(db)
page1 = repo.get_all(offset=0, limit=3)
assert len(page1) == 3
page2 = repo.get_all(offset=3, limit=3)
assert len(page2) == 2
def test_get_all_with_status_filter(self, db: Session) -> None:
user = create_test_user(db)
create_test_document(db, user=user, filename="pend.pdf", status="pending")
create_test_document(db, user=user, filename="comp.pdf", status="completed")
create_test_document(db, user=user, filename="fail.pdf", status="failed")
repo = DocumentRepository(db)
pending = repo.get_all(offset=0, limit=100, filters={"status": "pending"})
assert all(d.status == "pending" for d in pending)
def test_count(self, db: Session) -> None:
user = create_test_user(db)
create_test_document(db, user=user, filename="cnt1.pdf")
create_test_document(db, user=user, filename="cnt2.pdf")
repo = DocumentRepository(db)
total = repo.count()
assert total >= 2
def test_count_with_filter(self, db: Session) -> None:
user = create_test_user(db)
create_test_document(db, user=user, filename="cnt_p.pdf", status="pending")
create_test_document(db, user=user, filename="cnt_c.pdf", status="completed")
repo = DocumentRepository(db)
pending_count = repo.count(filters={"status": "pending"})
assert pending_count >= 1
def test_delete(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
repo = DocumentRepository(db)
repo.delete(doc)
db.flush()
assert repo.get_by_id(doc.id) is None
def test_get_by_checksum(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
repo = DocumentRepository(db)
found = repo.get_by_checksum(doc.checksum)
assert found is not None
assert found.id == doc.id
class TestDocumentPageRepository:
"""Tests for DocumentPageRepository."""
def test_create_page(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
page = create_test_page(db, doc, page_number=1)
assert page.id is not None
assert page.document_id == doc.id
assert page.page_number == 1
def test_get_pages_by_document(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
create_test_page(db, doc, page_number=1)
create_test_page(db, doc, page_number=2)
create_test_page(db, doc, page_number=3)
repo = DocumentPageRepository(db)
pages = repo.get_document_pages(doc.id)
assert len(pages) == 3
assert [p.page_number for p in pages] == [1, 2, 3]
class TestDocumentTextBlockRepository:
"""Tests for DocumentTextBlockRepository."""
def test_create_text_block(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
page = create_test_page(db, doc)
block = create_test_text_block(db, page, text="Hello World")
assert block.id is not None
assert block.text == "Hello World"
assert block.page_id == page.id
def test_get_blocks_by_page(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
page = create_test_page(db, doc)
create_test_text_block(db, page, text="First", sequence=0)
create_test_text_block(db, page, text="Second", sequence=1, y=120.0)
create_test_text_block(db, page, text="Third", sequence=2, y=140.0)
repo = DocumentTextBlockRepository(db)
blocks = repo.get_page_text_blocks(page.id)
assert len(blocks) == 3
def test_get_blocks_by_type(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user)
page = create_test_page(db, doc)
create_test_text_block(db, page, text="Header", block_type="header")
create_test_text_block(db, page, text="Body", block_type="text", y=200.0)
create_test_text_block(db, page, text="Footer", block_type="footer", y=700.0)
repo = DocumentTextBlockRepository(db)
headers = repo.get_by_block_type(page.id, "header")
assert len(headers) == 1
assert headers[0].text == "Header"
class TestTemplateMatchRepository:
"""Tests for TemplateMatchRepository."""
def test_create_match(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user, status="completed")
template = create_test_template(db, created_by=user)
match = TemplateMatch(
id=uuid.uuid4(),
document_id=doc.id,
format_id=template.id,
confidence_score=0.85,
selected=True,
)
db.add(match)
db.flush()
repo = TemplateMatchRepository(db)
matches = repo.get_document_matches(doc.id)
assert len(matches) == 1
assert matches[0].confidence_score == 0.85
def test_get_selected_match(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user, status="completed")
t1 = create_test_template(db, name="Low", created_by=user)
t2 = create_test_template(db, name="High", created_by=user)
db.add(TemplateMatch(
id=uuid.uuid4(), document_id=doc.id, format_id=t1.id,
confidence_score=0.5, selected=False,
))
db.add(TemplateMatch(
id=uuid.uuid4(), document_id=doc.id, format_id=t2.id,
confidence_score=0.95, selected=True,
))
db.flush()
repo = TemplateMatchRepository(db)
best = repo.get_selected_match(doc.id)
assert best is not None
assert best.format_id == t2.id
assert best.confidence_score == 0.95
def test_get_document_matches_ordered(self, db: Session) -> None:
user = create_test_user(db)
doc = create_test_document(db, user=user, status="completed")
t1 = create_test_template(db, name="T1", created_by=user)
t2 = create_test_template(db, name="T2", created_by=user)
t3 = create_test_template(db, name="T3", created_by=user)
db.add(TemplateMatch(
id=uuid.uuid4(), document_id=doc.id, format_id=t1.id,
confidence_score=0.3, selected=False,
))
db.add(TemplateMatch(
id=uuid.uuid4(), document_id=doc.id, format_id=t2.id,
confidence_score=0.9, selected=True,
))
db.add(TemplateMatch(
id=uuid.uuid4(), document_id=doc.id, format_id=t3.id,
confidence_score=0.6, selected=False,
))
db.flush()
repo = TemplateMatchRepository(db)
matches = repo.get_document_matches(doc.id)
scores = [m.confidence_score for m in matches]
assert scores == sorted(scores, reverse=True)