279 lines
9.9 KiB
Python
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)
|