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)