Merge pull request #131 from scub-france/fix/clean-architecture-audit

fix(arch): inject repositories and extract domain logic
This commit is contained in:
Pier-Jean Malandrino 2026-04-10 13:41:09 +02:00 committed by GitHub
commit 34d906d4b9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 836 additions and 716 deletions

View file

@ -3,13 +3,13 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
from typing import Annotated
from fastapi import APIRouter, HTTPException, Query, UploadFile from fastapi import APIRouter, Depends, HTTPException, Query, Request, UploadFile
from fastapi.responses import Response from fastapi.responses import Response
from api.schemas import DocumentResponse from api.schemas import DocumentResponse
from infra.settings import settings from services.document_service import DocumentService
from services import document_service
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/documents", tags=["documents"]) router = APIRouter(prefix="/api/documents", tags=["documents"])
@ -17,6 +17,13 @@ router = APIRouter(prefix="/api/documents", tags=["documents"])
_READ_CHUNK_SIZE = 64 * 1024 # 64 KB _READ_CHUNK_SIZE = 64 * 1024 # 64 KB
def _get_service(request: Request) -> DocumentService:
return request.app.state.document_service
ServiceDep = Annotated[DocumentService, Depends(_get_service)]
def _to_response(doc) -> DocumentResponse: def _to_response(doc) -> DocumentResponse:
return DocumentResponse( return DocumentResponse(
id=doc.id, id=doc.id,
@ -29,14 +36,14 @@ def _to_response(doc) -> DocumentResponse:
@router.post("/upload", response_model=DocumentResponse, status_code=200) @router.post("/upload", response_model=DocumentResponse, status_code=200)
async def upload(file: UploadFile) -> DocumentResponse: async def upload(file: UploadFile, service: ServiceDep) -> DocumentResponse:
"""Upload a PDF document.""" """Upload a PDF document."""
if not file.filename: if not file.filename:
raise HTTPException(status_code=400, detail="No filename provided") raise HTTPException(status_code=400, detail="No filename provided")
# Reject early if Content-Length exceeds limit (before reading body) # Reject early if Content-Length exceeds limit (before reading body)
_max = document_service.MAX_FILE_SIZE _max = service.max_file_size
_detail = f"File too large (max {settings.max_file_size_mb} MB)" _detail = f"File too large (max {service.max_file_size_mb} MB)"
if _max > 0 and file.size and file.size > _max: if _max > 0 and file.size and file.size > _max:
raise HTTPException(status_code=413, detail=_detail) raise HTTPException(status_code=413, detail=_detail)
@ -51,7 +58,7 @@ async def upload(file: UploadFile) -> DocumentResponse:
content = b"".join(chunks) content = b"".join(chunks)
try: try:
doc = await document_service.upload( doc = await service.upload(
filename=file.filename, filename=file.filename,
content_type=file.content_type or "application/pdf", content_type=file.content_type or "application/pdf",
file_content=content, file_content=content,
@ -63,25 +70,25 @@ async def upload(file: UploadFile) -> DocumentResponse:
@router.get("", response_model=list[DocumentResponse]) @router.get("", response_model=list[DocumentResponse])
async def list_documents() -> list[DocumentResponse]: async def list_documents(service: ServiceDep) -> list[DocumentResponse]:
"""List all documents.""" """List all documents."""
docs = await document_service.find_all() docs = await service.find_all()
return [_to_response(d) for d in docs] return [_to_response(d) for d in docs]
@router.get("/{doc_id}", response_model=DocumentResponse) @router.get("/{doc_id}", response_model=DocumentResponse)
async def get_document(doc_id: str) -> DocumentResponse: async def get_document(doc_id: str, service: ServiceDep) -> DocumentResponse:
"""Get a single document.""" """Get a single document."""
doc = await document_service.find_by_id(doc_id) doc = await service.find_by_id(doc_id)
if not doc: if not doc:
raise HTTPException(status_code=404, detail="Document not found") raise HTTPException(status_code=404, detail="Document not found")
return _to_response(doc) return _to_response(doc)
@router.delete("/{doc_id}", status_code=204) @router.delete("/{doc_id}", status_code=204)
async def delete_document(doc_id: str) -> None: async def delete_document(doc_id: str, service: ServiceDep) -> None:
"""Delete a document and its file.""" """Delete a document and its file."""
deleted = await document_service.delete(doc_id) deleted = await service.delete(doc_id)
if not deleted: if not deleted:
raise HTTPException(status_code=404, detail="Document not found") raise HTTPException(status_code=404, detail="Document not found")
@ -89,11 +96,12 @@ async def delete_document(doc_id: str) -> None:
@router.get("/{doc_id}/preview") @router.get("/{doc_id}/preview")
async def preview( async def preview(
doc_id: str, doc_id: str,
service: ServiceDep,
page: int = Query(1, ge=1), page: int = Query(1, ge=1),
dpi: int = Query(150, ge=72, le=300), dpi: int = Query(150, ge=72, le=300),
) -> Response: ) -> Response:
"""Generate a PNG preview of a specific PDF page.""" """Generate a PNG preview of a specific PDF page."""
doc = await document_service.find_by_id(doc_id) doc = await service.find_by_id(doc_id)
if not doc: if not doc:
raise HTTPException(status_code=404, detail="Document not found") raise HTTPException(status_code=404, detail="Document not found")
@ -106,7 +114,7 @@ async def preview(
try: try:
with open(doc.storage_path, "rb") as f: with open(doc.storage_path, "rb") as f:
file_content = f.read() file_content = f.read()
png_bytes = document_service.generate_preview(file_content, page=page, dpi=dpi) png_bytes = DocumentService.generate_preview(file_content, page=page, dpi=dpi)
return Response(content=png_bytes, media_type="image/png") return Response(content=png_bytes, media_type="image/png")
except ValueError as e: except ValueError as e:
raise HTTPException(status_code=404, detail=str(e)) from e raise HTTPException(status_code=404, detail=str(e)) from e

View file

@ -9,6 +9,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Protocol from typing import TYPE_CHECKING, Protocol
if TYPE_CHECKING: if TYPE_CHECKING:
from domain.models import AnalysisJob, Document
from domain.value_objects import ( from domain.value_objects import (
ChunkingOptions, ChunkingOptions,
ChunkResult, ChunkResult,
@ -44,3 +45,37 @@ class DocumentChunker(Protocol):
document_json: str, document_json: str,
options: ChunkingOptions, options: ChunkingOptions,
) -> list[ChunkResult]: ... ) -> list[ChunkResult]: ...
class DocumentRepository(Protocol):
"""Port for document persistence."""
async def insert(self, doc: Document) -> None: ...
async def find_all(self, *, limit: int = 200, offset: int = 0) -> list[Document]: ...
async def find_by_id(self, doc_id: str) -> Document | None: ...
async def update_page_count(self, doc_id: str, page_count: int) -> None: ...
async def delete(self, doc_id: str) -> bool: ...
class AnalysisRepository(Protocol):
"""Port for analysis job persistence."""
async def insert(self, job: AnalysisJob) -> None: ...
async def find_all(self, *, limit: int = 200, offset: int = 0) -> list[AnalysisJob]: ...
async def find_by_id(self, job_id: str) -> AnalysisJob | None: ...
async def update_status(self, job: AnalysisJob) -> None: ...
async def update_progress(self, job_id: str, current: int, total: int) -> None: ...
async def update_chunks(self, job_id: str, chunks_json: str) -> bool: ...
async def delete(self, job_id: str) -> bool: ...
async def delete_by_document(self, document_id: str) -> int: ...

View file

@ -0,0 +1,82 @@
"""Domain services — pure business logic with no infrastructure dependencies."""
from __future__ import annotations
import re
from domain.value_objects import ConversionResult, PageDetail
# Regex to extract <body> content from Docling's well-formed HTML output.
_BODY_RE = re.compile(r"<body[^>]*>(.*)</body>", re.DOTALL | re.IGNORECASE)
def extract_html_body(html: str) -> str:
"""Extract content between <body> tags.
Docling produces well-formed HTML regex is safe for this controlled output.
Returns raw html as fallback if no <body> tag is found.
"""
match = _BODY_RE.search(html)
return match.group(1).strip() if match else html
def merge_results(results: list[ConversionResult]) -> ConversionResult:
"""Merge multiple batch ConversionResults into a single consolidated result.
document_json is intentionally set to None: merging DoclingDocument's internal
tree structure across batches is error-prone. Re-chunking is disabled for
batched conversions (robustness decision for 0.3.1).
"""
if not results:
return ConversionResult(page_count=0, content_markdown="", content_html="", pages=[])
all_pages: list[PageDetail] = []
all_md: list[str] = []
html_bodies: list[str] = []
total_skipped = 0
for r in results:
all_pages.extend(r.pages)
all_md.append(r.content_markdown)
html_bodies.append(extract_html_body(r.content_html))
total_skipped += r.skipped_items
merged_body = "\n".join(html_bodies)
merged_html = (
f'<!DOCTYPE html><html><head><meta charset="utf-8"></head><body>{merged_body}</body></html>'
)
return ConversionResult(
page_count=sum(r.page_count for r in results),
content_markdown="\n\n".join(all_md),
content_html=merged_html,
pages=all_pages,
skipped_items=total_skipped,
document_json=None,
)
def classify_error(exc: Exception) -> str:
"""Return a user-friendly error message based on the exception type/content."""
msg = str(exc).lower()
if "invalidcxxcompiler" in msg or "no working c++ compiler" in msg:
return "Missing C++ compiler — set TORCHDYNAMO_DISABLE=1 to work around this"
if "out of memory" in msg or "oom" in msg:
return "Out of memory — try a smaller document or disable table structure analysis"
if "could not acquire converter lock" in msg:
return "Server busy — a previous conversion is still running. Please retry later"
if "pipeline" in msg and "failed" in msg:
return "Document processing failed — the document may be corrupted or unsupported"
if "timeout" in msg:
return "Processing took too long — try with fewer pages or simpler options"
# Fallback: truncate raw error to something reasonable
raw = str(exc)
if len(raw) > 200:
raw = raw[:200] + ""
return raw

View file

@ -22,8 +22,11 @@ from api.analyses import router as analyses_router
from api.documents import router as documents_router from api.documents import router as documents_router
from infra.rate_limiter import RateLimiterMiddleware from infra.rate_limiter import RateLimiterMiddleware
from infra.settings import settings from infra.settings import settings
from persistence.analysis_repo import SqliteAnalysisRepository
from persistence.database import get_connection, init_db from persistence.database import get_connection, init_db
from services.analysis_service import AnalysisService from persistence.document_repo import SqliteDocumentRepository
from services.analysis_service import AnalysisConfig, AnalysisService
from services.document_service import DocumentConfig, DocumentService
logging.basicConfig( logging.basicConfig(
level=logging.INFO, level=logging.INFO,
@ -58,14 +61,44 @@ def _build_chunker():
return None return None
def _build_analysis_service() -> AnalysisService: def _build_repos() -> tuple[SqliteDocumentRepository, SqliteAnalysisRepository]:
return SqliteDocumentRepository(), SqliteAnalysisRepository()
def _build_analysis_service(
document_repo: SqliteDocumentRepository,
analysis_repo: SqliteAnalysisRepository,
) -> AnalysisService:
converter = _build_converter() converter = _build_converter()
chunker = _build_chunker() chunker = _build_chunker()
config = AnalysisConfig(
default_table_mode=settings.default_table_mode,
batch_page_size=settings.batch_page_size,
)
return AnalysisService( return AnalysisService(
converter=converter, converter=converter,
analysis_repo=analysis_repo,
document_repo=document_repo,
chunker=chunker, chunker=chunker,
conversion_timeout=settings.conversion_timeout, conversion_timeout=settings.conversion_timeout,
max_concurrent=settings.max_concurrent_analyses, max_concurrent=settings.max_concurrent_analyses,
config=config,
)
def _build_document_service(
document_repo: SqliteDocumentRepository,
analysis_repo: SqliteAnalysisRepository,
) -> DocumentService:
config = DocumentConfig(
upload_dir=settings.upload_dir,
max_file_size_mb=settings.max_file_size_mb,
max_page_count=settings.max_page_count,
)
return DocumentService(
document_repo=document_repo,
analysis_repo=analysis_repo,
config=config,
) )
@ -77,7 +110,9 @@ def _build_analysis_service() -> AnalysisService:
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]: async def lifespan(app: FastAPI) -> AsyncIterator[None]:
await init_db() await init_db()
app.state.analysis_service = _build_analysis_service() document_repo, analysis_repo = _build_repos()
app.state.analysis_service = _build_analysis_service(document_repo, analysis_repo)
app.state.document_service = _build_document_service(document_repo, analysis_repo)
logger.info("Docling Studio backend ready (engine=%s)", settings.conversion_engine) logger.info("Docling Studio backend ready (engine=%s)", settings.conversion_engine)
yield yield

View file

@ -43,96 +43,94 @@ _SELECT_WITH_DOC = """
""" """
async def insert(job: AnalysisJob) -> None: class SqliteAnalysisRepository:
"""Persist a new analysis job record.""" """SQLite implementation of the AnalysisRepository port."""
async with get_connection() as db:
await db.execute(
"""INSERT INTO analysis_jobs (id, document_id, status, created_at)
VALUES (?, ?, ?, ?)""",
(job.id, job.document_id, job.status.value, str(job.created_at)),
)
await db.commit()
async def insert(self, job: AnalysisJob) -> None:
"""Persist a new analysis job record."""
async with get_connection() as db:
await db.execute(
"""INSERT INTO analysis_jobs (id, document_id, status, created_at)
VALUES (?, ?, ?, ?)""",
(job.id, job.document_id, job.status.value, str(job.created_at)),
)
await db.commit()
async def find_all(*, limit: int = 200, offset: int = 0) -> list[AnalysisJob]: async def find_all(self, *, limit: int = 200, offset: int = 0) -> list[AnalysisJob]:
"""Return analysis jobs with document info, newest first.""" """Return analysis jobs with document info, newest first."""
async with get_connection() as db: async with get_connection() as db:
cursor = await db.execute( cursor = await db.execute(
f"{_SELECT_WITH_DOC} ORDER BY aj.created_at DESC LIMIT ? OFFSET ?", f"{_SELECT_WITH_DOC} ORDER BY aj.created_at DESC LIMIT ? OFFSET ?",
(limit, offset), (limit, offset),
) )
rows = await cursor.fetchall() rows = await cursor.fetchall()
return [_row_to_job(r) for r in rows] return [_row_to_job(r) for r in rows]
async def find_by_id(self, job_id: str) -> AnalysisJob | None:
"""Find an analysis job by ID (with document filename), or return None."""
async with get_connection() as db:
cursor = await db.execute(f"{_SELECT_WITH_DOC} WHERE aj.id = ?", (job_id,))
row = await cursor.fetchone()
return _row_to_job(row) if row else None
async def find_by_id(job_id: str) -> AnalysisJob | None: async def update_status(self, job: AnalysisJob) -> None:
"""Find an analysis job by ID (with document filename), or return None.""" """Persist all mutable fields of an analysis job (status, results, timestamps)."""
async with get_connection() as db: async with get_connection() as db:
cursor = await db.execute(f"{_SELECT_WITH_DOC} WHERE aj.id = ?", (job_id,)) await db.execute(
row = await cursor.fetchone() """UPDATE analysis_jobs
return _row_to_job(row) if row else None SET status = ?, content_markdown = ?, content_html = ?,
pages_json = ?, document_json = ?, chunks_json = ?,
error_message = ?, progress_current = ?, progress_total = ?,
started_at = ?, completed_at = ?
WHERE id = ?""",
(
job.status.value,
job.content_markdown,
job.content_html,
job.pages_json,
job.document_json,
job.chunks_json,
job.error_message,
job.progress_current,
job.progress_total,
str(job.started_at) if job.started_at else None,
str(job.completed_at) if job.completed_at else None,
job.id,
),
)
await db.commit()
async def update_progress(self, job_id: str, current: int, total: int) -> None:
"""Update only the progress columns for a running analysis."""
async with get_connection() as db:
await db.execute(
"UPDATE analysis_jobs SET progress_current = ?, progress_total = ? WHERE id = ?",
(current, total, job_id),
)
await db.commit()
async def update_status(job: AnalysisJob) -> None: async def update_chunks(self, job_id: str, chunks_json: str) -> bool:
"""Persist all mutable fields of an analysis job (status, results, timestamps).""" """Update only the chunks_json column for a completed analysis."""
async with get_connection() as db: async with get_connection() as db:
await db.execute( cursor = await db.execute(
"""UPDATE analysis_jobs "UPDATE analysis_jobs SET chunks_json = ? WHERE id = ?",
SET status = ?, content_markdown = ?, content_html = ?, (chunks_json, job_id),
pages_json = ?, document_json = ?, chunks_json = ?, )
error_message = ?, progress_current = ?, progress_total = ?, await db.commit()
started_at = ?, completed_at = ? return cursor.rowcount > 0
WHERE id = ?""",
(
job.status.value,
job.content_markdown,
job.content_html,
job.pages_json,
job.document_json,
job.chunks_json,
job.error_message,
job.progress_current,
job.progress_total,
str(job.started_at) if job.started_at else None,
str(job.completed_at) if job.completed_at else None,
job.id,
),
)
await db.commit()
async def delete(self, job_id: str) -> bool:
"""Delete an analysis job by ID. Returns True if a row was removed."""
async with get_connection() as db:
cursor = await db.execute("DELETE FROM analysis_jobs WHERE id = ?", (job_id,))
await db.commit()
return cursor.rowcount > 0
async def update_progress(job_id: str, current: int, total: int) -> None: async def delete_by_document(self, document_id: str) -> int:
"""Update only the progress columns for a running analysis.""" """Delete all analysis jobs for a given document. Returns count deleted."""
async with get_connection() as db: async with get_connection() as db:
await db.execute( cursor = await db.execute(
"UPDATE analysis_jobs SET progress_current = ?, progress_total = ? WHERE id = ?", "DELETE FROM analysis_jobs WHERE document_id = ?", (document_id,)
(current, total, job_id), )
) await db.commit()
await db.commit() return cursor.rowcount
async def update_chunks(job_id: str, chunks_json: str) -> bool:
"""Update only the chunks_json column for a completed analysis."""
async with get_connection() as db:
cursor = await db.execute(
"UPDATE analysis_jobs SET chunks_json = ? WHERE id = ?",
(chunks_json, job_id),
)
await db.commit()
return cursor.rowcount > 0
async def delete(job_id: str) -> bool:
"""Delete an analysis job by ID. Returns True if a row was removed."""
async with get_connection() as db:
cursor = await db.execute("DELETE FROM analysis_jobs WHERE id = ?", (job_id,))
await db.commit()
return cursor.rowcount > 0
async def delete_by_document(document_id: str) -> int:
"""Delete all analysis jobs for a given document. Returns count deleted."""
async with get_connection() as db:
cursor = await db.execute("DELETE FROM analysis_jobs WHERE document_id = ?", (document_id,))
await db.commit()
return cursor.rowcount

View file

@ -25,57 +25,56 @@ def _row_to_document(row) -> Document:
) )
async def insert(doc: Document) -> None: class SqliteDocumentRepository:
"""Persist a new document record.""" """SQLite implementation of the DocumentRepository port."""
async with get_connection() as db:
await db.execute(
"""INSERT INTO documents (id, filename, content_type, file_size, page_count, storage_path, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?)""",
(
doc.id,
doc.filename,
doc.content_type,
doc.file_size,
doc.page_count,
doc.storage_path,
str(doc.created_at),
),
)
await db.commit()
async def insert(self, doc: Document) -> None:
"""Persist a new document record."""
async with get_connection() as db:
await db.execute(
"""INSERT INTO documents (id, filename, content_type, file_size, page_count, storage_path, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?)""",
(
doc.id,
doc.filename,
doc.content_type,
doc.file_size,
doc.page_count,
doc.storage_path,
str(doc.created_at),
),
)
await db.commit()
async def find_all(*, limit: int = 200, offset: int = 0) -> list[Document]: async def find_all(self, *, limit: int = 200, offset: int = 0) -> list[Document]:
"""Return documents ordered by creation date (newest first).""" """Return documents ordered by creation date (newest first)."""
async with get_connection() as db: async with get_connection() as db:
cursor = await db.execute( cursor = await db.execute(
"SELECT * FROM documents ORDER BY created_at DESC LIMIT ? OFFSET ?", "SELECT * FROM documents ORDER BY created_at DESC LIMIT ? OFFSET ?",
(limit, offset), (limit, offset),
) )
rows = await cursor.fetchall() rows = await cursor.fetchall()
return [_row_to_document(r) for r in rows] return [_row_to_document(r) for r in rows]
async def find_by_id(self, doc_id: str) -> Document | None:
"""Find a document by its ID, or return None."""
async with get_connection() as db:
cursor = await db.execute("SELECT * FROM documents WHERE id = ?", (doc_id,))
row = await cursor.fetchone()
return _row_to_document(row) if row else None
async def find_by_id(doc_id: str) -> Document | None: async def update_page_count(self, doc_id: str, page_count: int) -> None:
"""Find a document by its ID, or return None.""" """Update the page count after conversion has determined it."""
async with get_connection() as db: async with get_connection() as db:
cursor = await db.execute("SELECT * FROM documents WHERE id = ?", (doc_id,)) await db.execute(
row = await cursor.fetchone() "UPDATE documents SET page_count = ? WHERE id = ?",
return _row_to_document(row) if row else None (page_count, doc_id),
)
await db.commit()
async def delete(self, doc_id: str) -> bool:
async def update_page_count(doc_id: str, page_count: int) -> None: """Delete a document by ID. Returns True if a row was removed."""
"""Update the page count after conversion has determined it.""" async with get_connection() as db:
async with get_connection() as db: cursor = await db.execute("DELETE FROM documents WHERE id = ?", (doc_id,))
await db.execute( await db.commit()
"UPDATE documents SET page_count = ? WHERE id = ?", return cursor.rowcount > 0
(page_count, doc_id),
)
await db.commit()
async def delete(doc_id: str) -> bool:
"""Delete a document by ID. Returns True if a row was removed."""
async with get_connection() as db:
cursor = await db.execute("DELETE FROM documents WHERE id = ?", (doc_id,))
await db.commit()
return cursor.rowcount > 0

View file

@ -1,7 +1,7 @@
"""Analysis service — async document parsing orchestration. """Analysis service — async document parsing orchestration.
Uses an injected DocumentConverter (port) so the service is decoupled Uses injected ports (converter, chunker, repositories) so the service is
from the conversion implementation (local Docling lib vs remote Docling Serve). decoupled from infrastructure implementations.
""" """
from __future__ import annotations from __future__ import annotations
@ -11,25 +11,27 @@ import functools
import json import json
import logging import logging
import math import math
import re from dataclasses import asdict, dataclass
from dataclasses import asdict
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import pypdfium2 as pdfium import pypdfium2 as pdfium
from domain.models import AnalysisJob, AnalysisStatus from domain.models import AnalysisJob, AnalysisStatus
from domain.services import classify_error, merge_results
from domain.value_objects import ( from domain.value_objects import (
ChunkingOptions, ChunkingOptions,
ChunkResult, ChunkResult,
ConversionOptions, ConversionOptions,
ConversionResult, ConversionResult,
PageDetail,
) )
from infra.settings import settings
if TYPE_CHECKING: if TYPE_CHECKING:
from domain.ports import DocumentChunker, DocumentConverter from domain.ports import (
from persistence import analysis_repo, document_repo AnalysisRepository,
DocumentChunker,
DocumentConverter,
DocumentRepository,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -48,9 +50,6 @@ def _chunk_to_dict(c: ChunkResult) -> dict:
# Maximum number of concurrent analysis jobs to prevent resource exhaustion. # Maximum number of concurrent analysis jobs to prevent resource exhaustion.
_DEFAULT_MAX_CONCURRENT = 3 _DEFAULT_MAX_CONCURRENT = 3
# Regex to extract <body> content from Docling's well-formed HTML output.
_BODY_RE = re.compile(r"<body[^>]*>(.*)</body>", re.DOTALL | re.IGNORECASE)
def _count_pdf_pages(file_path: str) -> int: def _count_pdf_pages(file_path: str) -> int:
"""Count pages in a PDF. Returns 0 if the file is not a valid PDF.""" """Count pages in a PDF. Returns 0 if the file is not a valid PDF."""
@ -64,67 +63,36 @@ def _count_pdf_pages(file_path: str) -> int:
return 0 return 0
def _extract_html_body(html: str) -> str: @dataclass
"""Extract content between <body> tags. class AnalysisConfig:
"""Configuration values needed by AnalysisService, extracted from settings."""
Docling produces well-formed HTML regex is safe for this controlled output. default_table_mode: str = "accurate"
Returns raw html as fallback if no <body> tag is found. batch_page_size: int = 0
"""
match = _BODY_RE.search(html)
return match.group(1).strip() if match else html
def _merge_results(results: list[ConversionResult]) -> ConversionResult:
"""Merge multiple batch ConversionResults into a single consolidated result.
document_json is intentionally set to None: merging DoclingDocument's internal
tree structure across batches is error-prone. Re-chunking is disabled for
batched conversions (robustness decision for 0.3.1).
"""
if not results:
return ConversionResult(page_count=0, content_markdown="", content_html="", pages=[])
all_pages: list[PageDetail] = []
all_md: list[str] = []
html_bodies: list[str] = []
total_skipped = 0
for r in results:
all_pages.extend(r.pages)
all_md.append(r.content_markdown)
html_bodies.append(_extract_html_body(r.content_html))
total_skipped += r.skipped_items
merged_body = "\n".join(html_bodies)
merged_html = (
f'<!DOCTYPE html><html><head><meta charset="utf-8"></head><body>{merged_body}</body></html>'
)
return ConversionResult(
page_count=sum(r.page_count for r in results),
content_markdown="\n\n".join(all_md),
content_html=merged_html,
pages=all_pages,
skipped_items=total_skipped,
document_json=None,
)
class AnalysisService: class AnalysisService:
"""Orchestrates document analysis using an injected converter.""" """Orchestrates document analysis using injected ports."""
def __init__( def __init__(
self, self,
converter: DocumentConverter, converter: DocumentConverter,
analysis_repo: AnalysisRepository,
document_repo: DocumentRepository,
chunker: DocumentChunker | None = None, chunker: DocumentChunker | None = None,
conversion_timeout: int = 600, conversion_timeout: int = 600,
max_concurrent: int = _DEFAULT_MAX_CONCURRENT, max_concurrent: int = _DEFAULT_MAX_CONCURRENT,
config: AnalysisConfig | None = None,
): ):
self._converter = converter self._converter = converter
self._chunker = chunker self._chunker = chunker
self._analysis_repo = analysis_repo
self._document_repo = document_repo
self._conversion_timeout = conversion_timeout self._conversion_timeout = conversion_timeout
self._semaphore = asyncio.Semaphore(max_concurrent) self._semaphore = asyncio.Semaphore(max_concurrent)
self._running_tasks: dict[str, asyncio.Task] = {} self._running_tasks: dict[str, asyncio.Task] = {}
self._background_tasks: set[asyncio.Task] = set()
self._config = config or AnalysisConfig()
async def create( async def create(
self, self,
@ -134,13 +102,13 @@ class AnalysisService:
chunking_options: dict | None = None, chunking_options: dict | None = None,
) -> AnalysisJob: ) -> AnalysisJob:
"""Create a new analysis job and launch background processing.""" """Create a new analysis job and launch background processing."""
doc = await document_repo.find_by_id(document_id) doc = await self._document_repo.find_by_id(document_id)
if not doc: if not doc:
raise ValueError(f"Document not found: {document_id}") raise ValueError(f"Document not found: {document_id}")
job = AnalysisJob(document_id=document_id) job = AnalysisJob(document_id=document_id)
job.document_filename = doc.filename job.document_filename = doc.filename
await analysis_repo.insert(job) await self._analysis_repo.insert(job)
task = asyncio.create_task( task = asyncio.create_task(
self._run_analysis( self._run_analysis(
@ -158,11 +126,11 @@ class AnalysisService:
async def find_all(self) -> list[AnalysisJob]: async def find_all(self) -> list[AnalysisJob]:
"""Return all analysis jobs, newest first.""" """Return all analysis jobs, newest first."""
return await analysis_repo.find_all() return await self._analysis_repo.find_all()
async def find_by_id(self, job_id: str) -> AnalysisJob | None: async def find_by_id(self, job_id: str) -> AnalysisJob | None:
"""Find an analysis job by ID, or return None.""" """Find an analysis job by ID, or return None."""
return await analysis_repo.find_by_id(job_id) return await self._analysis_repo.find_by_id(job_id)
async def delete(self, job_id: str) -> bool: async def delete(self, job_id: str) -> bool:
"""Delete an analysis job, cancelling any running task. Returns True if it existed.""" """Delete an analysis job, cancelling any running task. Returns True if it existed."""
@ -170,11 +138,11 @@ class AnalysisService:
if task and not task.done(): if task and not task.done():
task.cancel() task.cancel()
logger.info("Cancelled running task for job %s", job_id) logger.info("Cancelled running task for job %s", job_id)
return await analysis_repo.delete(job_id) return await self._analysis_repo.delete(job_id)
async def rechunk(self, job_id: str, chunking_options: dict) -> list[ChunkResult]: async def rechunk(self, job_id: str, chunking_options: dict) -> list[ChunkResult]:
"""Re-chunk an existing completed analysis with new options.""" """Re-chunk an existing completed analysis with new options."""
job = await analysis_repo.find_by_id(job_id) job = await self._analysis_repo.find_by_id(job_id)
if not job: if not job:
raise ValueError(f"Analysis not found: {job_id}") raise ValueError(f"Analysis not found: {job_id}")
if job.status != AnalysisStatus.COMPLETED: if job.status != AnalysisStatus.COMPLETED:
@ -188,7 +156,7 @@ class AnalysisService:
chunks = await self._chunker.chunk(job.document_json, options) chunks = await self._chunker.chunk(job.document_json, options)
chunks_json = json.dumps([_chunk_to_dict(c) for c in chunks]) chunks_json = json.dumps([_chunk_to_dict(c) for c in chunks])
await analysis_repo.update_chunks(job_id, chunks_json) await self._analysis_repo.update_chunks(job_id, chunks_json)
return chunks return chunks
@ -206,7 +174,7 @@ class AnalysisService:
Raises on batch failure (fail-fast: entire job fails). Raises on batch failure (fail-fast: entire job fails).
""" """
num_batches = math.ceil(total_pages / batch_size) num_batches = math.ceil(total_pages / batch_size)
await analysis_repo.update_progress(job_id, 0, total_pages) await self._analysis_repo.update_progress(job_id, 0, total_pages)
logger.info( logger.info(
"Batched conversion: %d pages in %d batches of %d for job %s", "Batched conversion: %d pages in %d batches of %d for job %s",
total_pages, total_pages,
@ -220,7 +188,7 @@ class AnalysisService:
start = batch_idx * batch_size + 1 start = batch_idx * batch_size + 1
end = min(start + batch_size - 1, total_pages) end = min(start + batch_size - 1, total_pages)
if not await analysis_repo.find_by_id(job_id): if not await self._analysis_repo.find_by_id(job_id):
logger.info( logger.info(
"Job %s deleted during batch %d/%d, aborting", "Job %s deleted during batch %d/%d, aborting",
job_id, job_id,
@ -240,7 +208,7 @@ class AnalysisService:
) from exc ) from exc
results.append(batch_result) results.append(batch_result)
await analysis_repo.update_progress(job_id, end, total_pages) await self._analysis_repo.update_progress(job_id, end, total_pages)
logger.info( logger.info(
"Batch %d/%d done (pages %d-%d) for job %s", "Batch %d/%d done (pages %d-%d) for job %s",
batch_idx + 1, batch_idx + 1,
@ -250,12 +218,37 @@ class AnalysisService:
job_id, job_id,
) )
return _merge_results(results) return merge_results(results)
def _on_task_done(self, task: asyncio.Task, *, job_id: str) -> None: def _on_task_done(self, task: asyncio.Task, *, job_id: str) -> None:
"""Cleanup running tasks and delegate to module-level handler.""" """Cleanup running tasks and handle failures."""
self._running_tasks.pop(job_id, None) self._running_tasks.pop(job_id, None)
_on_task_done(task, job_id=job_id) if task.cancelled():
logger.warning("Analysis task was cancelled: %s", job_id)
self._schedule_mark_failed(job_id, "Task was cancelled")
return
exc = task.exception()
if exc:
logger.error("Unhandled exception in analysis task %s: %s", job_id, exc, exc_info=True)
self._schedule_mark_failed(job_id, classify_error(exc))
def _schedule_mark_failed(self, job_id: str, error: str) -> None:
"""Schedule _mark_failed as a tracked background task."""
t = asyncio.ensure_future(self._mark_failed(job_id, error))
self._background_tasks.add(t)
t.add_done_callback(self._background_tasks.discard)
async def _mark_failed(self, job_id: str, error: str) -> None:
"""Safely mark a job as failed, handling DB errors gracefully."""
try:
job = await self._analysis_repo.find_by_id(job_id)
if job:
job.mark_failed(error)
await self._analysis_repo.update_status(job)
except OSError:
logger.exception("Database I/O error marking job %s as failed", job_id)
except Exception:
logger.exception("Unexpected error marking job %s as failed", job_id)
async def _run_analysis( async def _run_analysis(
self, self,
@ -285,22 +278,22 @@ class AnalysisService:
) -> None: ) -> None:
"""Inner analysis logic — called under the concurrency semaphore.""" """Inner analysis logic — called under the concurrency semaphore."""
try: try:
job = await analysis_repo.find_by_id(job_id) job = await self._analysis_repo.find_by_id(job_id)
if not job: if not job:
logger.error("Analysis job %s not found", job_id) logger.error("Analysis job %s not found", job_id)
return return
job.mark_running() job.mark_running()
await analysis_repo.update_status(job) await self._analysis_repo.update_status(job)
logger.info("Analysis started: %s (file: %s)", job_id, filename) logger.info("Analysis started: %s (file: %s)", job_id, filename)
opts_dict = pipeline_options or {} opts_dict = pipeline_options or {}
if "table_mode" not in opts_dict: if "table_mode" not in opts_dict:
opts_dict = {**opts_dict, "table_mode": settings.default_table_mode} opts_dict = {**opts_dict, "table_mode": self._config.default_table_mode}
options = ConversionOptions(**opts_dict) options = ConversionOptions(**opts_dict)
total_pages = _count_pdf_pages(file_path) total_pages = _count_pdf_pages(file_path)
batch_size = settings.batch_page_size batch_size = self._config.batch_page_size
if batch_size > 0 and total_pages > batch_size: if batch_size > 0 and total_pages > batch_size:
result = await self._run_batched_conversion( result = await self._run_batched_conversion(
@ -329,7 +322,7 @@ class AnalysisService:
# Re-read the job so we don't lose progress_current/progress_total # Re-read the job so we don't lose progress_current/progress_total
# written to the DB during batched conversion. # written to the DB during batched conversion.
job = await analysis_repo.find_by_id(job_id) or job job = await self._analysis_repo.find_by_id(job_id) or job
job.mark_completed( job.mark_completed(
markdown=result.content_markdown, markdown=result.content_markdown,
html=result.content_html, html=result.content_html,
@ -337,81 +330,19 @@ class AnalysisService:
document_json=result.document_json, document_json=result.document_json,
chunks_json=chunks_json, chunks_json=chunks_json,
) )
await analysis_repo.update_status(job) await self._analysis_repo.update_status(job)
if result.page_count: if result.page_count:
await document_repo.update_page_count(job.document_id, result.page_count) await self._document_repo.update_page_count(job.document_id, result.page_count)
logger.info("Analysis completed: %s (%d pages)", job_id, result.page_count) logger.info("Analysis completed: %s (%d pages)", job_id, result.page_count)
except TimeoutError: except TimeoutError:
logger.error("Analysis timed out after %ds: %s", self._conversion_timeout, job_id) logger.error("Analysis timed out after %ds: %s", self._conversion_timeout, job_id)
await _mark_failed(job_id, f"Conversion timed out after {self._conversion_timeout}s") await self._mark_failed(
job_id, f"Conversion timed out after {self._conversion_timeout}s"
)
except Exception as e: except Exception as e:
logger.exception("Analysis failed: %s", job_id) logger.exception("Analysis failed: %s", job_id)
await _mark_failed(job_id, _classify_error(e)) await self._mark_failed(job_id, classify_error(e))
def _classify_error(exc: Exception) -> str:
"""Return a user-friendly error message based on the exception type/content."""
msg = str(exc).lower()
if "invalidcxxcompiler" in msg or "no working c++ compiler" in msg:
return "Missing C++ compiler — set TORCHDYNAMO_DISABLE=1 to work around this"
if "out of memory" in msg or "oom" in msg:
return "Out of memory — try a smaller document or disable table structure analysis"
if "could not acquire converter lock" in msg:
return "Server busy — a previous conversion is still running. Please retry later"
if "pipeline" in msg and "failed" in msg:
return "Document processing failed — the document may be corrupted or unsupported"
if "timeout" in msg:
return "Processing took too long — try with fewer pages or simpler options"
# Fallback: truncate raw error to something reasonable
raw = str(exc)
if len(raw) > 200:
raw = raw[:200] + ""
return raw
_background_tasks: set[asyncio.Task] = set()
def _on_task_done(task: asyncio.Task, *, job_id: str) -> None:
"""Log unhandled exceptions from background analysis tasks and mark job as FAILED."""
if task.cancelled():
logger.warning("Analysis task was cancelled: %s", job_id)
_schedule_mark_failed(job_id, "Task was cancelled")
return
exc = task.exception()
if exc:
logger.error("Unhandled exception in analysis task %s: %s", job_id, exc, exc_info=True)
_schedule_mark_failed(job_id, _classify_error(exc))
# Keep the module-level function as the default, but AnalysisService uses its own method.
def _schedule_mark_failed(job_id: str, error: str) -> None:
"""Schedule _mark_failed as a tracked background task."""
t = asyncio.ensure_future(_mark_failed(job_id, error))
_background_tasks.add(t)
t.add_done_callback(_background_tasks.discard)
async def _mark_failed(job_id: str, error: str) -> None:
"""Safely mark a job as failed, handling DB errors gracefully."""
try:
job = await analysis_repo.find_by_id(job_id)
if job:
job.mark_failed(error)
await analysis_repo.update_status(job)
except OSError:
logger.exception("Database I/O error marking job %s as failed", job_id)
except Exception:
logger.exception("Unexpected error marking job %s as failed", job_id)

View file

@ -6,113 +6,149 @@ import io
import logging import logging
import os import os
import uuid import uuid
from dataclasses import dataclass
from typing import TYPE_CHECKING
from pdf2image import convert_from_bytes, pdfinfo_from_bytes from pdf2image import convert_from_bytes, pdfinfo_from_bytes
from domain.models import Document from domain.models import Document
from infra.settings import settings
from persistence import analysis_repo, document_repo if TYPE_CHECKING:
from domain.ports import AnalysisRepository, DocumentRepository
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
UPLOAD_DIR = settings.upload_dir
MAX_FILE_SIZE = settings.max_file_size_mb * 1024 * 1024 if settings.max_file_size_mb > 0 else 0
MAX_PAGE_COUNT = settings.max_page_count # 0 = unlimited
# PDF magic bytes: %PDF # PDF magic bytes: %PDF
_PDF_MAGIC = b"%PDF" _PDF_MAGIC = b"%PDF"
_UPLOAD_CHUNK_SIZE = 64 * 1024 # 64 KB chunks for streaming writes _UPLOAD_CHUNK_SIZE = 64 * 1024 # 64 KB chunks for streaming writes
async def upload(filename: str, content_type: str, file_content: bytes) -> Document: @dataclass
"""Save uploaded file to disk and persist metadata. class DocumentConfig:
"""Configuration values needed by DocumentService, extracted from settings."""
Writes the file in fixed-size chunks to keep peak memory usage low. upload_dir: str = "uploads"
""" max_file_size_mb: int = 0
if MAX_FILE_SIZE > 0 and len(file_content) > MAX_FILE_SIZE: max_page_count: int = 0
raise ValueError(f"File too large (max {settings.max_file_size_mb} MB)")
if not file_content[:4].startswith(_PDF_MAGIC):
raise ValueError("Invalid file: not a PDF document")
os.makedirs(UPLOAD_DIR, exist_ok=True)
ext = ".pdf" # Content already validated as PDF
safe_name = f"{uuid.uuid4()}{ext}"
file_path = os.path.join(UPLOAD_DIR, safe_name)
# Write in chunks to avoid doubling memory usage for large files
with open(file_path, "wb") as f:
for offset in range(0, len(file_content), _UPLOAD_CHUNK_SIZE):
f.write(file_content[offset : offset + _UPLOAD_CHUNK_SIZE])
# Count PDF pages
page_count = _count_pages(file_content)
if MAX_PAGE_COUNT > 0 and page_count is not None and page_count > MAX_PAGE_COUNT:
os.unlink(file_path)
raise ValueError(f"Too many pages ({page_count}). Maximum allowed: {MAX_PAGE_COUNT}")
doc = Document(
filename=filename,
content_type=content_type,
file_size=len(file_content),
page_count=page_count,
storage_path=os.path.abspath(file_path),
)
await document_repo.insert(doc)
return doc
async def find_all() -> list[Document]: class DocumentService:
"""Return all documents, newest first.""" """Orchestrates document upload, storage, and preview."""
return await document_repo.find_all()
def __init__(
self,
document_repo: DocumentRepository,
analysis_repo: AnalysisRepository,
config: DocumentConfig,
):
self._document_repo = document_repo
self._analysis_repo = analysis_repo
self._config = config
self._upload_dir = config.upload_dir
self._max_file_size = (
config.max_file_size_mb * 1024 * 1024 if config.max_file_size_mb > 0 else 0
)
self._max_page_count = config.max_page_count
async def find_by_id(doc_id: str) -> Document | None: @property
"""Find a document by its ID, or return None.""" def max_file_size(self) -> int:
return await document_repo.find_by_id(doc_id) return self._max_file_size
@property
def max_file_size_mb(self) -> int:
return self._config.max_file_size_mb
async def delete(doc_id: str) -> bool: async def upload(self, filename: str, content_type: str, file_content: bytes) -> Document:
"""Delete document file, associated analyses, and database record.""" """Save uploaded file to disk and persist metadata.
doc = await document_repo.find_by_id(doc_id)
if not doc:
return False
# Delete associated analyses first (cascade) Writes the file in fixed-size chunks to keep peak memory usage low.
await analysis_repo.delete_by_document(doc_id) """
if self._max_file_size > 0 and len(file_content) > self._max_file_size:
raise ValueError(f"File too large (max {self._config.max_file_size_mb} MB)")
# Delete file from disk (only if inside UPLOAD_DIR) if not file_content[:4].startswith(_PDF_MAGIC):
try: raise ValueError("Invalid file: not a PDF document")
real_path = os.path.realpath(doc.storage_path)
real_upload_dir = os.path.realpath(UPLOAD_DIR)
if real_path.startswith(real_upload_dir + os.sep) and os.path.exists(real_path):
os.unlink(real_path)
elif os.path.exists(doc.storage_path):
logger.warning("Refused to delete file outside upload dir: %s", doc.storage_path)
except FileNotFoundError:
logger.info("File already removed: %s", doc.storage_path)
except PermissionError:
logger.error("Permission denied deleting file: %s", doc.storage_path)
except OSError:
logger.warning("Could not delete file: %s", doc.storage_path, exc_info=True)
return await document_repo.delete(doc_id) os.makedirs(self._upload_dir, exist_ok=True)
ext = ".pdf" # Content already validated as PDF
safe_name = f"{uuid.uuid4()}{ext}"
file_path = os.path.join(self._upload_dir, safe_name)
def generate_preview(file_content: bytes, page: int = 1, dpi: int = 150) -> bytes: # Write in chunks to avoid doubling memory usage for large files
"""Generate a PNG preview of a specific PDF page.""" with open(file_path, "wb") as f:
images = convert_from_bytes(file_content, first_page=page, last_page=page, dpi=dpi) for offset in range(0, len(file_content), _UPLOAD_CHUNK_SIZE):
if not images: f.write(file_content[offset : offset + _UPLOAD_CHUNK_SIZE])
raise ValueError(f"Page {page} not found")
buf = io.BytesIO() # Count PDF pages
images[0].save(buf, format="PNG") page_count = _count_pages(file_content)
return buf.getvalue()
if (
self._max_page_count > 0
and page_count is not None
and page_count > self._max_page_count
):
os.unlink(file_path)
raise ValueError(
f"Too many pages ({page_count}). Maximum allowed: {self._max_page_count}"
)
doc = Document(
filename=filename,
content_type=content_type,
file_size=len(file_content),
page_count=page_count,
storage_path=os.path.abspath(file_path),
)
await self._document_repo.insert(doc)
return doc
async def find_all(self) -> list[Document]:
"""Return all documents, newest first."""
return await self._document_repo.find_all()
async def find_by_id(self, doc_id: str) -> Document | None:
"""Find a document by its ID, or return None."""
return await self._document_repo.find_by_id(doc_id)
async def delete(self, doc_id: str) -> bool:
"""Delete document file, associated analyses, and database record."""
doc = await self._document_repo.find_by_id(doc_id)
if not doc:
return False
# Delete associated analyses first (cascade)
await self._analysis_repo.delete_by_document(doc_id)
# Delete file from disk (only if inside upload dir)
try:
real_path = os.path.realpath(doc.storage_path)
real_upload_dir = os.path.realpath(self._upload_dir)
if real_path.startswith(real_upload_dir + os.sep) and os.path.exists(real_path):
os.unlink(real_path)
elif os.path.exists(doc.storage_path):
logger.warning("Refused to delete file outside upload dir: %s", doc.storage_path)
except FileNotFoundError:
logger.info("File already removed: %s", doc.storage_path)
except PermissionError:
logger.error("Permission denied deleting file: %s", doc.storage_path)
except OSError:
logger.warning("Could not delete file: %s", doc.storage_path, exc_info=True)
return await self._document_repo.delete(doc_id)
@staticmethod
def generate_preview(file_content: bytes, page: int = 1, dpi: int = 150) -> bytes:
"""Generate a PNG preview of a specific PDF page."""
images = convert_from_bytes(file_content, first_page=page, last_page=page, dpi=dpi)
if not images:
raise ValueError(f"Page {page} not found")
buf = io.BytesIO()
images[0].save(buf, format="PNG")
return buf.getvalue()
def _count_pages(file_content: bytes) -> int | None: def _count_pages(file_content: bytes) -> int | None:

View file

@ -8,14 +8,20 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from domain.services import extract_html_body, merge_results
from domain.value_objects import ConversionResult, PageDetail from domain.value_objects import ConversionResult, PageDetail
from services.analysis_service import ( from services.analysis_service import AnalysisConfig, AnalysisService, _count_pdf_pages
AnalysisService,
_count_pdf_pages,
_extract_html_body, def _make_service(**kwargs) -> AnalysisService:
_merge_results, """Create an AnalysisService with mock repos for testing."""
_on_task_done, defaults = {
) "converter": MagicMock(),
"analysis_repo": MagicMock(),
"document_repo": MagicMock(),
}
defaults.update(kwargs)
return AnalysisService(**defaults)
class TestOnTaskDone: class TestOnTaskDone:
@ -25,25 +31,25 @@ class TestOnTaskDone:
async def test_exception_marks_job_failed(self): async def test_exception_marks_job_failed(self):
"""When a background task raises, the job should be marked FAILED.""" """When a background task raises, the job should be marked FAILED."""
job_id = "job-123" job_id = "job-123"
service = _make_service()
# Create a task that raises
async def failing_task(): async def failing_task():
raise RuntimeError("unexpected failure") raise RuntimeError("unexpected failure")
task = asyncio.create_task(failing_task()) task = asyncio.create_task(failing_task())
await asyncio.sleep(0) # let the task fail await asyncio.sleep(0)
with patch("services.analysis_service._mark_failed", new_callable=AsyncMock) as mock_mark: with patch.object(service, "_mark_failed", new_callable=AsyncMock) as mock_mark:
_on_task_done(task, job_id=job_id) service._on_task_done(task, job_id=job_id)
# ensure_future schedules it; give the event loop a tick
await asyncio.sleep(0) await asyncio.sleep(0)
mock_mark.assert_called_once_with(job_id, "unexpected failure") mock_mark.assert_called_once_with(job_id, "unexpected failure")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_exception_uses_classify_error(self): async def test_exception_uses_classify_error(self):
"""_on_task_done should route exceptions through _classify_error.""" """_on_task_done should route exceptions through classify_error."""
job_id = "job-classify" job_id = "job-classify"
service = _make_service()
async def timeout_task(): async def timeout_task():
raise TimeoutError("timeout exceeded while processing") raise TimeoutError("timeout exceeded while processing")
@ -51,8 +57,8 @@ class TestOnTaskDone:
task = asyncio.create_task(timeout_task()) task = asyncio.create_task(timeout_task())
await asyncio.sleep(0) await asyncio.sleep(0)
with patch("services.analysis_service._mark_failed", new_callable=AsyncMock) as mock_mark: with patch.object(service, "_mark_failed", new_callable=AsyncMock) as mock_mark:
_on_task_done(task, job_id=job_id) service._on_task_done(task, job_id=job_id)
await asyncio.sleep(0) await asyncio.sleep(0)
mock_mark.assert_called_once_with( mock_mark.assert_called_once_with(
@ -63,6 +69,7 @@ class TestOnTaskDone:
async def test_cancelled_task_marks_job_failed(self): async def test_cancelled_task_marks_job_failed(self):
"""When a background task is cancelled, the job should be marked FAILED.""" """When a background task is cancelled, the job should be marked FAILED."""
job_id = "job-456" job_id = "job-456"
service = _make_service()
async def slow_task(): async def slow_task():
await asyncio.sleep(999) await asyncio.sleep(999)
@ -74,8 +81,8 @@ class TestOnTaskDone:
with contextlib.suppress(asyncio.CancelledError): with contextlib.suppress(asyncio.CancelledError):
await task await task
with patch("services.analysis_service._mark_failed", new_callable=AsyncMock) as mock_mark: with patch.object(service, "_mark_failed", new_callable=AsyncMock) as mock_mark:
_on_task_done(task, job_id=job_id) service._on_task_done(task, job_id=job_id)
await asyncio.sleep(0) await asyncio.sleep(0)
mock_mark.assert_called_once_with(job_id, "Task was cancelled") mock_mark.assert_called_once_with(job_id, "Task was cancelled")
@ -84,6 +91,7 @@ class TestOnTaskDone:
async def test_successful_task_does_not_mark_failed(self): async def test_successful_task_does_not_mark_failed(self):
"""When a background task succeeds, _mark_failed should not be called.""" """When a background task succeeds, _mark_failed should not be called."""
job_id = "job-789" job_id = "job-789"
service = _make_service()
async def ok_task(): async def ok_task():
return "done" return "done"
@ -91,8 +99,8 @@ class TestOnTaskDone:
task = asyncio.create_task(ok_task()) task = asyncio.create_task(ok_task())
await task await task
with patch("services.analysis_service._mark_failed", new_callable=AsyncMock) as mock_mark: with patch.object(service, "_mark_failed", new_callable=AsyncMock) as mock_mark:
_on_task_done(task, job_id=job_id) service._on_task_done(task, job_id=job_id)
await asyncio.sleep(0) await asyncio.sleep(0)
mock_mark.assert_not_called() mock_mark.assert_not_called()
@ -104,8 +112,9 @@ class TestAnalysisServiceCancellation:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_cancels_running_task(self): async def test_delete_cancels_running_task(self):
"""Deleting a job while running should cancel its task.""" """Deleting a job while running should cancel its task."""
converter = MagicMock() mock_analysis_repo = MagicMock()
service = AnalysisService(converter=converter) mock_analysis_repo.delete = AsyncMock(return_value=True)
service = _make_service(analysis_repo=mock_analysis_repo)
blocker = asyncio.Event() blocker = asyncio.Event()
@ -115,9 +124,7 @@ class TestAnalysisServiceCancellation:
task = asyncio.create_task(slow_analysis()) task = asyncio.create_task(slow_analysis())
service._running_tasks["j1"] = task service._running_tasks["j1"] = task
with patch("services.analysis_service.analysis_repo") as mock_repo: result = await service.delete("j1")
mock_repo.delete = AsyncMock(return_value=True)
result = await service.delete("j1")
assert result is True assert result is True
assert task.cancelling() or task.cancelled() assert task.cancelling() or task.cancelled()
@ -126,20 +133,18 @@ class TestAnalysisServiceCancellation:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_completed_job_no_error(self): async def test_delete_completed_job_no_error(self):
"""Deleting a completed job should not raise even if no task tracked.""" """Deleting a completed job should not raise even if no task tracked."""
converter = MagicMock() mock_analysis_repo = MagicMock()
service = AnalysisService(converter=converter) mock_analysis_repo.delete = AsyncMock(return_value=True)
service = _make_service(analysis_repo=mock_analysis_repo)
with patch("services.analysis_service.analysis_repo") as mock_repo: result = await service.delete("j-gone")
mock_repo.delete = AsyncMock(return_value=True)
result = await service.delete("j-gone")
assert result is True assert result is True
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_task_cleaned_from_running_on_completion(self): async def test_task_cleaned_from_running_on_completion(self):
"""After a task completes, it should be removed from _running_tasks.""" """After a task completes, it should be removed from _running_tasks."""
converter = MagicMock() service = _make_service()
service = AnalysisService(converter=converter)
async def instant(): async def instant():
pass pass
@ -156,13 +161,11 @@ class TestAnalysisServiceConcurrency:
"""Verify that the semaphore limits concurrent analysis jobs.""" """Verify that the semaphore limits concurrent analysis jobs."""
def test_semaphore_initialized_with_max_concurrent(self): def test_semaphore_initialized_with_max_concurrent(self):
converter = MagicMock() service = _make_service(max_concurrent=5)
service = AnalysisService(converter=converter, max_concurrent=5)
assert service._semaphore._value == 5 assert service._semaphore._value == 5
def test_default_max_concurrent(self): def test_default_max_concurrent(self):
converter = MagicMock() service = _make_service()
service = AnalysisService(converter=converter)
assert service._semaphore._value == 3 assert service._semaphore._value == 3
@pytest.mark.asyncio @pytest.mark.asyncio
@ -171,8 +174,7 @@ class TestAnalysisServiceConcurrency:
call_order: list[str] = [] call_order: list[str] = []
blocker = asyncio.Event() blocker = asyncio.Event()
converter = MagicMock() service = _make_service(max_concurrent=1)
service = AnalysisService(converter=converter, max_concurrent=1)
async def fake_inner(self, *args, **kwargs): async def fake_inner(self, *args, **kwargs):
call_order.append("start") call_order.append("start")
@ -203,7 +205,6 @@ class TestAnalysisServiceConcurrency:
class TestCountPdfPages: class TestCountPdfPages:
def test_valid_pdf(self, tmp_path): def test_valid_pdf(self, tmp_path):
"""A real (minimal) PDF should return its page count.""" """A real (minimal) PDF should return its page count."""
# Create a minimal valid 1-page PDF using pypdfium2
import pypdfium2 as pdfium import pypdfium2 as pdfium
pdf = pdfium.PdfDocument.new() pdf = pdfium.PdfDocument.new()
@ -234,20 +235,20 @@ class TestCountPdfPages:
class TestExtractHtmlBody: class TestExtractHtmlBody:
def test_extracts_body(self): def test_extracts_body(self):
html = '<html><head></head><body class="x"><p>Hello</p></body></html>' html = '<html><head></head><body class="x"><p>Hello</p></body></html>'
assert _extract_html_body(html) == "<p>Hello</p>" assert extract_html_body(html) == "<p>Hello</p>"
def test_no_body_tag_returns_raw(self): def test_no_body_tag_returns_raw(self):
html = "<p>No body tag</p>" html = "<p>No body tag</p>"
assert _extract_html_body(html) == html assert extract_html_body(html) == html
def test_empty_body(self): def test_empty_body(self):
html = "<html><body></body></html>" html = "<html><body></body></html>"
assert _extract_html_body(html) == "" assert extract_html_body(html) == ""
class TestMergeResults: class TestMergeResults:
def test_empty_list(self): def test_empty_list(self):
result = _merge_results([]) result = merge_results([])
assert result.page_count == 0 assert result.page_count == 0
assert result.content_markdown == "" assert result.content_markdown == ""
assert result.pages == [] assert result.pages == []
@ -261,7 +262,7 @@ class TestMergeResults:
pages=[PageDetail(page_number=1, width=612, height=792)], pages=[PageDetail(page_number=1, width=612, height=792)],
document_json='{"pages": {}}', document_json='{"pages": {}}',
) )
merged = _merge_results([r]) merged = merge_results([r])
assert merged.page_count == 3 assert merged.page_count == 3
assert merged.content_markdown == "# Page 1" assert merged.content_markdown == "# Page 1"
assert merged.pages == [PageDetail(page_number=1, width=612, height=792)] assert merged.pages == [PageDetail(page_number=1, width=612, height=792)]
@ -288,7 +289,7 @@ class TestMergeResults:
], ],
skipped_items=2, skipped_items=2,
) )
merged = _merge_results([r1, r2]) merged = merge_results([r1, r2])
assert merged.page_count == 4 assert merged.page_count == 4
assert merged.content_markdown == "# Batch 1\n\n# Batch 2" assert merged.content_markdown == "# Batch 1\n\n# Batch 2"
assert len(merged.pages) == 4 assert len(merged.pages) == 4
@ -330,19 +331,23 @@ class TestBatchedConversion:
), ),
] ]
service = AnalysisService(converter=converter, conversion_timeout=60) mock_analysis_repo = MagicMock()
mock_analysis_repo.find_by_id = AsyncMock(return_value=MagicMock())
mock_analysis_repo.update_progress = AsyncMock()
with patch("services.analysis_service.analysis_repo") as mock_repo: service = _make_service(
mock_repo.find_by_id = AsyncMock(return_value=MagicMock()) converter=converter,
mock_repo.update_progress = AsyncMock() analysis_repo=mock_analysis_repo,
conversion_timeout=60,
)
result = await service._run_batched_conversion( result = await service._run_batched_conversion(
"job-1", "job-1",
"/fake.pdf", "/fake.pdf",
MagicMock(), MagicMock(),
total_pages=8, total_pages=8,
batch_size=5, batch_size=5,
) )
assert result is not None assert result is not None
assert result.page_count == 8 assert result.page_count == 8
@ -370,20 +375,24 @@ class TestBatchedConversion:
RuntimeError("OOM"), RuntimeError("OOM"),
] ]
service = AnalysisService(converter=converter, conversion_timeout=60) mock_analysis_repo = MagicMock()
mock_analysis_repo.find_by_id = AsyncMock(return_value=MagicMock())
mock_analysis_repo.update_progress = AsyncMock()
with patch("services.analysis_service.analysis_repo") as mock_repo: service = _make_service(
mock_repo.find_by_id = AsyncMock(return_value=MagicMock()) converter=converter,
mock_repo.update_progress = AsyncMock() analysis_repo=mock_analysis_repo,
conversion_timeout=60,
)
with pytest.raises(RuntimeError, match=r"Batch 2/2 \(pages 6-8\) failed: OOM"): with pytest.raises(RuntimeError, match=r"Batch 2/2 \(pages 6-8\) failed: OOM"):
await service._run_batched_conversion( await service._run_batched_conversion(
"job-fail", "job-fail",
"/fake.pdf", "/fake.pdf",
MagicMock(), MagicMock(),
total_pages=8, total_pages=8,
batch_size=5, batch_size=5,
) )
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_progress_preserved_through_full_analysis_flow(self): async def test_progress_preserved_through_full_analysis_flow(self):
@ -412,8 +421,6 @@ class TestBatchedConversion:
), ),
] ]
service = AnalysisService(converter=converter, conversion_timeout=60)
# Simulated DB state: find_by_id is called 4 times: # Simulated DB state: find_by_id is called 4 times:
# 1) start of _run_analysis_inner (initial load) # 1) start of _run_analysis_inner (initial load)
# 2) batch 1 mid-flight deletion check # 2) batch 1 mid-flight deletion check
@ -442,21 +449,27 @@ class TestBatchedConversion:
async def capture_update_status(job): async def capture_update_status(job):
saved_jobs.append(job) saved_jobs.append(job)
with ( mock_analysis_repo = MagicMock()
patch("services.analysis_service.analysis_repo") as mock_repo, mock_analysis_repo.find_by_id = AsyncMock(
patch("services.analysis_service.document_repo") as mock_doc_repo, side_effect=[initial_job, batch_check_job, batch_check_job, refreshed_job]
patch("services.analysis_service._count_pdf_pages", return_value=8), )
patch("services.analysis_service.settings") as mock_settings, mock_analysis_repo.update_status = AsyncMock(side_effect=capture_update_status)
): mock_analysis_repo.update_progress = AsyncMock()
mock_settings.batch_page_size = 5
mock_settings.default_table_mode = "accurate"
mock_repo.find_by_id = AsyncMock(
side_effect=[initial_job, batch_check_job, batch_check_job, refreshed_job]
)
mock_repo.update_status = AsyncMock(side_effect=capture_update_status)
mock_repo.update_progress = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock()
mock_document_repo = MagicMock()
mock_document_repo.update_page_count = AsyncMock()
config = AnalysisConfig(default_table_mode="accurate", batch_page_size=5)
service = AnalysisService(
converter=converter,
analysis_repo=mock_analysis_repo,
document_repo=mock_document_repo,
conversion_timeout=60,
config=config,
)
with patch("services.analysis_service._count_pdf_pages", return_value=8):
await service._run_analysis_inner("job-progress", "/fake.pdf", "fake.pdf") await service._run_analysis_inner("job-progress", "/fake.pdf", "fake.pdf")
# The last update_status call is the COMPLETED one # The last update_status call is the COMPLETED one
@ -480,20 +493,24 @@ class TestBatchedConversion:
pages=[PageDetail(page_number=i, width=612, height=792) for i in range(1, 6)], pages=[PageDetail(page_number=i, width=612, height=792) for i in range(1, 6)],
) )
service = AnalysisService(converter=converter, conversion_timeout=60) mock_analysis_repo = MagicMock()
# First check returns job, second returns None (deleted) # First check returns job, second returns None (deleted)
with patch("services.analysis_service.analysis_repo") as mock_repo: mock_analysis_repo.find_by_id = AsyncMock(side_effect=[MagicMock(), None])
mock_repo.find_by_id = AsyncMock(side_effect=[MagicMock(), None]) mock_analysis_repo.update_progress = AsyncMock()
mock_repo.update_progress = AsyncMock()
result = await service._run_batched_conversion( service = _make_service(
"job-del", converter=converter,
"/fake.pdf", analysis_repo=mock_analysis_repo,
MagicMock(), conversion_timeout=60,
total_pages=10, )
batch_size=5,
) result = await service._run_batched_conversion(
"job-del",
"/fake.pdf",
MagicMock(),
total_pages=10,
batch_size=5,
)
assert result is None assert result is None
# Only first batch should have been converted # Only first batch should have been converted

View file

@ -1,6 +1,6 @@
"""Tests for FastAPI API endpoints using TestClient.""" """Tests for FastAPI API endpoints using TestClient."""
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
@ -24,6 +24,18 @@ def mock_analysis_service(client):
app.state.analysis_service = original app.state.analysis_service = original
@pytest.fixture
def mock_document_service(client):
"""Inject a mock DocumentService into app.state for the duration of the test."""
mock_svc = MagicMock()
mock_svc.max_file_size = 50 * 1024 * 1024
mock_svc.max_file_size_mb = 50
original = getattr(app.state, "document_service", None)
app.state.document_service = mock_svc
yield mock_svc
app.state.document_service = original
class TestHealthEndpoint: class TestHealthEndpoint:
def test_health(self, client): def test_health(self, client):
resp = client.get("/api/health") resp = client.get("/api/health")
@ -41,12 +53,13 @@ class TestHealthEndpoint:
class TestDocumentEndpoints: class TestDocumentEndpoints:
@patch("services.document_service.find_all", new_callable=AsyncMock) def test_list_documents(self, client, mock_document_service):
def test_list_documents(self, mock_find_all, client): mock_document_service.find_all = AsyncMock(
mock_find_all.return_value = [ return_value=[
Document(id="d1", filename="a.pdf", storage_path="/tmp/a"), Document(id="d1", filename="a.pdf", storage_path="/tmp/a"),
Document(id="d2", filename="b.pdf", storage_path="/tmp/b"), Document(id="d2", filename="b.pdf", storage_path="/tmp/b"),
] ]
)
resp = client.get("/api/documents") resp = client.get("/api/documents")
assert resp.status_code == 200 assert resp.status_code == 200
@ -57,15 +70,16 @@ class TestDocumentEndpoints:
# Verify camelCase # Verify camelCase
assert "createdAt" in data[0] assert "createdAt" in data[0]
@patch("services.document_service.find_by_id", new_callable=AsyncMock) def test_get_document(self, client, mock_document_service):
def test_get_document(self, mock_find, client): mock_document_service.find_by_id = AsyncMock(
mock_find.return_value = Document( return_value=Document(
id="d1", id="d1",
filename="test.pdf", filename="test.pdf",
content_type="application/pdf", content_type="application/pdf",
file_size=2048, file_size=2048,
page_count=3, page_count=3,
storage_path="/tmp/test", storage_path="/tmp/test",
)
) )
resp = client.get("/api/documents/d1") resp = client.get("/api/documents/d1")
@ -75,21 +89,21 @@ class TestDocumentEndpoints:
assert data["fileSize"] == 2048 assert data["fileSize"] == 2048
assert data["pageCount"] == 3 assert data["pageCount"] == 3
@patch("services.document_service.find_by_id", new_callable=AsyncMock) def test_get_document_not_found(self, client, mock_document_service):
def test_get_document_not_found(self, mock_find, client): mock_document_service.find_by_id = AsyncMock(return_value=None)
mock_find.return_value = None
resp = client.get("/api/documents/missing") resp = client.get("/api/documents/missing")
assert resp.status_code == 404 assert resp.status_code == 404
@patch("services.document_service.upload", new_callable=AsyncMock) def test_upload_document(self, client, mock_document_service):
def test_upload_document(self, mock_upload, client): mock_document_service.upload = AsyncMock(
mock_upload.return_value = Document( return_value=Document(
id="new-1", id="new-1",
filename="uploaded.pdf", filename="uploaded.pdf",
content_type="application/pdf", content_type="application/pdf",
file_size=512, file_size=512,
storage_path="/tmp/uploaded", storage_path="/tmp/uploaded",
)
) )
resp = client.post( resp = client.post(
@ -101,9 +115,10 @@ class TestDocumentEndpoints:
assert data["id"] == "new-1" assert data["id"] == "new-1"
assert data["filename"] == "uploaded.pdf" assert data["filename"] == "uploaded.pdf"
@patch("services.document_service.upload", new_callable=AsyncMock) def test_upload_too_large(self, client, mock_document_service):
def test_upload_too_large(self, mock_upload, client): mock_document_service.upload = AsyncMock(
mock_upload.side_effect = ValueError("File too large (max 5 MB)") side_effect=ValueError("File too large (max 5 MB)")
)
resp = client.post( resp = client.post(
"/api/documents/upload", "/api/documents/upload",
@ -111,29 +126,28 @@ class TestDocumentEndpoints:
) )
assert resp.status_code == 400 assert resp.status_code == 400
@patch("services.document_service.find_by_id", new_callable=AsyncMock) def test_preview_page_out_of_range(self, client, mock_document_service):
def test_preview_page_out_of_range(self, mock_find, client): mock_document_service.find_by_id = AsyncMock(
mock_find.return_value = Document( return_value=Document(
id="d1", id="d1",
filename="test.pdf", filename="test.pdf",
page_count=3, page_count=3,
storage_path="/tmp/test.pdf", storage_path="/tmp/test.pdf",
)
) )
resp = client.get("/api/documents/d1/preview?page=10") resp = client.get("/api/documents/d1/preview?page=10")
assert resp.status_code == 400 assert resp.status_code == 400
assert "out of range" in resp.json()["detail"] assert "out of range" in resp.json()["detail"]
@patch("services.document_service.delete", new_callable=AsyncMock) def test_delete_document(self, client, mock_document_service):
def test_delete_document(self, mock_delete, client): mock_document_service.delete = AsyncMock(return_value=True)
mock_delete.return_value = True
resp = client.delete("/api/documents/d1") resp = client.delete("/api/documents/d1")
assert resp.status_code == 204 assert resp.status_code == 204
@patch("services.document_service.delete", new_callable=AsyncMock) def test_delete_document_not_found(self, client, mock_document_service):
def test_delete_document_not_found(self, mock_delete, client): mock_document_service.delete = AsyncMock(return_value=False)
mock_delete.return_value = False
resp = client.delete("/api/documents/missing") resp = client.delete("/api/documents/missing")
assert resp.status_code == 404 assert resp.status_code == 404

View file

@ -1,4 +1,4 @@
"""Tests for document_service — upload, preview, page counting, and deletion.""" """Tests for DocumentService — upload, preview, page counting, and deletion."""
from __future__ import annotations from __future__ import annotations
@ -8,82 +8,87 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from domain.models import Document from domain.models import Document
from services import document_service from services.document_service import DocumentConfig, DocumentService, _count_pages
def _make_service(
upload_dir: str = "/tmp/uploads",
max_file_size_mb: int = 50,
max_page_count: int = 0,
) -> DocumentService:
"""Create a DocumentService with mock repos for testing."""
config = DocumentConfig(
upload_dir=upload_dir,
max_file_size_mb=max_file_size_mb,
max_page_count=max_page_count,
)
return DocumentService(
document_repo=AsyncMock(),
analysis_repo=AsyncMock(),
config=config,
)
class TestUploadValidation: class TestUploadValidation:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_rejects_oversized_file(self): async def test_rejects_oversized_file(self):
content = b"x" * (document_service.MAX_FILE_SIZE + 1) service = _make_service(max_file_size_mb=1)
content = b"x" * (1 * 1024 * 1024 + 1)
with pytest.raises(ValueError, match="File too large"): with pytest.raises(ValueError, match="File too large"):
await document_service.upload("big.pdf", "application/pdf", content) await service.upload("big.pdf", "application/pdf", content)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_rejects_non_pdf(self): async def test_rejects_non_pdf(self):
service = _make_service()
content = b"NOT-A-PDF-FILE" content = b"NOT-A-PDF-FILE"
with pytest.raises(ValueError, match="not a PDF"): with pytest.raises(ValueError, match="not a PDF"):
await document_service.upload("fake.pdf", "application/pdf", content) await service.upload("fake.pdf", "application/pdf", content)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_rejects_too_many_pages(self, tmp_path, monkeypatch): async def test_rejects_too_many_pages(self, tmp_path):
monkeypatch.setattr(document_service, "UPLOAD_DIR", str(tmp_path)) service = _make_service(upload_dir=str(tmp_path), max_page_count=20)
monkeypatch.setattr(document_service, "MAX_PAGE_COUNT", 20)
with patch.object(document_service, "_count_pages", return_value=40): with patch("services.document_service._count_pages", return_value=40):
content = b"%PDF-1.4 fake pdf content" content = b"%PDF-1.4 fake pdf content"
with pytest.raises(ValueError, match="Too many pages"): with pytest.raises(ValueError, match="Too many pages"):
await document_service.upload("big.pdf", "application/pdf", content) await service.upload("big.pdf", "application/pdf", content)
# Verify temp file was cleaned up # Verify temp file was cleaned up
assert len(os.listdir(tmp_path)) == 0 assert len(os.listdir(tmp_path)) == 0
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_allows_pdf_under_page_limit(self, tmp_path, monkeypatch): async def test_allows_pdf_under_page_limit(self, tmp_path):
monkeypatch.setattr(document_service, "UPLOAD_DIR", str(tmp_path)) service = _make_service(upload_dir=str(tmp_path), max_page_count=20)
monkeypatch.setattr(document_service, "MAX_PAGE_COUNT", 20)
mock_insert = AsyncMock() with patch("services.document_service._count_pages", return_value=15):
with (
patch("persistence.document_repo.insert", mock_insert),
patch.object(document_service, "_count_pages", return_value=15),
):
content = b"%PDF-1.4 fake pdf content" content = b"%PDF-1.4 fake pdf content"
doc = await document_service.upload("ok.pdf", "application/pdf", content) doc = await service.upload("ok.pdf", "application/pdf", content)
assert doc.page_count == 15 assert doc.page_count == 15
mock_insert.assert_called_once() service._document_repo.insert.assert_called_once()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_unlimited_pages_when_zero(self, tmp_path, monkeypatch): async def test_unlimited_pages_when_zero(self, tmp_path):
monkeypatch.setattr(document_service, "UPLOAD_DIR", str(tmp_path)) service = _make_service(upload_dir=str(tmp_path), max_page_count=0)
monkeypatch.setattr(document_service, "MAX_PAGE_COUNT", 0)
mock_insert = AsyncMock() with patch("services.document_service._count_pages", return_value=100):
with (
patch("persistence.document_repo.insert", mock_insert),
patch.object(document_service, "_count_pages", return_value=100),
):
content = b"%PDF-1.4 fake pdf content" content = b"%PDF-1.4 fake pdf content"
doc = await document_service.upload("big.pdf", "application/pdf", content) doc = await service.upload("big.pdf", "application/pdf", content)
assert doc.page_count == 100 assert doc.page_count == 100
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_accepts_valid_pdf(self, tmp_path, monkeypatch): async def test_accepts_valid_pdf(self, tmp_path):
monkeypatch.setattr(document_service, "UPLOAD_DIR", str(tmp_path)) service = _make_service(upload_dir=str(tmp_path))
mock_insert = AsyncMock() with patch("services.document_service._count_pages", return_value=5):
with (
patch("persistence.document_repo.insert", mock_insert),
patch.object(document_service, "_count_pages", return_value=5),
):
content = b"%PDF-1.4 fake pdf content" content = b"%PDF-1.4 fake pdf content"
doc = await document_service.upload("test.pdf", "application/pdf", content) doc = await service.upload("test.pdf", "application/pdf", content)
assert doc.filename == "test.pdf" assert doc.filename == "test.pdf"
assert doc.file_size == len(content) assert doc.file_size == len(content)
assert doc.page_count == 5 assert doc.page_count == 5
mock_insert.assert_called_once() service._document_repo.insert.assert_called_once()
# Verify file was actually written to disk # Verify file was actually written to disk
assert os.path.exists(doc.storage_path) assert os.path.exists(doc.storage_path)
@ -98,7 +103,7 @@ class TestGeneratePreview:
patch("services.document_service.convert_from_bytes", return_value=[]), patch("services.document_service.convert_from_bytes", return_value=[]),
pytest.raises(ValueError, match="Page 1 not found"), pytest.raises(ValueError, match="Page 1 not found"),
): ):
document_service.generate_preview(b"%PDF-fake", page=1) DocumentService.generate_preview(b"%PDF-fake", page=1)
def test_returns_png_bytes(self): def test_returns_png_bytes(self):
"""generate_preview should return PNG bytes from pdf2image.""" """generate_preview should return PNG bytes from pdf2image."""
@ -106,7 +111,7 @@ class TestGeneratePreview:
mock_image.save = MagicMock(side_effect=lambda buf, format: buf.write(b"PNG-DATA")) mock_image.save = MagicMock(side_effect=lambda buf, format: buf.write(b"PNG-DATA"))
with patch("services.document_service.convert_from_bytes", return_value=[mock_image]): with patch("services.document_service.convert_from_bytes", return_value=[mock_image]):
result = document_service.generate_preview(b"%PDF-fake", page=1, dpi=72) result = DocumentService.generate_preview(b"%PDF-fake", page=1, dpi=72)
assert result == b"PNG-DATA" assert result == b"PNG-DATA"
@ -117,21 +122,19 @@ class TestCountPages:
"services.document_service.pdfinfo_from_bytes", "services.document_service.pdfinfo_from_bytes",
return_value={"Pages": 42}, return_value={"Pages": 42},
): ):
assert document_service._count_pages(b"pdf") == 42 assert _count_pages(b"pdf") == 42
def test_returns_none_on_error(self): def test_returns_none_on_error(self):
with patch( with patch(
"services.document_service.pdfinfo_from_bytes", "services.document_service.pdfinfo_from_bytes",
side_effect=FileNotFoundError("poppler not found"), side_effect=FileNotFoundError("poppler not found"),
): ):
assert document_service._count_pages(b"pdf") is None assert _count_pages(b"pdf") is None
class TestDelete: class TestDelete:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_removes_file_and_records(self, tmp_path, monkeypatch): async def test_delete_removes_file_and_records(self, tmp_path):
monkeypatch.setattr(document_service, "UPLOAD_DIR", str(tmp_path))
# Create a fake file # Create a fake file
fake_file = tmp_path / "test.pdf" fake_file = tmp_path / "test.pdf"
fake_file.write_bytes(b"content") fake_file.write_bytes(b"content")
@ -142,21 +145,21 @@ class TestDelete:
storage_path=str(fake_file), storage_path=str(fake_file),
) )
with ( service = _make_service(upload_dir=str(tmp_path))
patch("persistence.document_repo.find_by_id", AsyncMock(return_value=doc)), service._document_repo.find_by_id = AsyncMock(return_value=doc)
patch("persistence.analysis_repo.delete_by_document", AsyncMock(return_value=2)), service._document_repo.delete = AsyncMock(return_value=True)
patch("persistence.document_repo.delete", AsyncMock(return_value=True)), service._analysis_repo.delete_by_document = AsyncMock(return_value=2)
):
result = await document_service.delete("doc-1") result = await service.delete("doc-1")
assert result is True assert result is True
assert not fake_file.exists() assert not fake_file.exists()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_refuses_file_outside_upload_dir(self, tmp_path, monkeypatch): async def test_delete_refuses_file_outside_upload_dir(self, tmp_path):
"""Files outside UPLOAD_DIR should not be deleted (path traversal protection).""" """Files outside upload dir should not be deleted (path traversal protection)."""
monkeypatch.setattr(document_service, "UPLOAD_DIR", str(tmp_path / "uploads")) uploads = tmp_path / "uploads"
os.makedirs(tmp_path / "uploads", exist_ok=True) os.makedirs(uploads, exist_ok=True)
# File is outside the upload dir # File is outside the upload dir
outside_file = tmp_path / "secret.txt" outside_file = tmp_path / "secret.txt"
@ -164,18 +167,20 @@ class TestDelete:
doc = Document(id="doc-1", filename="x.pdf", storage_path=str(outside_file)) doc = Document(id="doc-1", filename="x.pdf", storage_path=str(outside_file))
with ( service = _make_service(upload_dir=str(uploads))
patch("persistence.document_repo.find_by_id", AsyncMock(return_value=doc)), service._document_repo.find_by_id = AsyncMock(return_value=doc)
patch("persistence.analysis_repo.delete_by_document", AsyncMock(return_value=0)), service._document_repo.delete = AsyncMock(return_value=True)
patch("persistence.document_repo.delete", AsyncMock(return_value=True)), service._analysis_repo.delete_by_document = AsyncMock(return_value=0)
):
await document_service.delete("doc-1") await service.delete("doc-1")
# File should NOT have been deleted # File should NOT have been deleted
assert outside_file.exists() assert outside_file.exists()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_not_found_returns_false(self): async def test_delete_not_found_returns_false(self):
with patch("persistence.document_repo.find_by_id", AsyncMock(return_value=None)): service = _make_service()
result = await document_service.delete("missing") service._document_repo.find_by_id = AsyncMock(return_value=None)
result = await service.delete("missing")
assert result is False assert result is False

View file

@ -325,6 +325,22 @@ class TestConvertDocumentRouting:
class TestServiceForwardsPipelineOptions: class TestServiceForwardsPipelineOptions:
"""Verify analysis_service.create and _run_analysis forward pipeline options.""" """Verify analysis_service.create and _run_analysis forward pipeline options."""
def _make_service(self, converter):
from services.analysis_service import AnalysisService
mock_analysis_repo = MagicMock()
mock_analysis_repo.find_by_id = AsyncMock()
mock_analysis_repo.insert = AsyncMock()
mock_analysis_repo.update_status = AsyncMock()
mock_document_repo = MagicMock()
mock_document_repo.find_by_id = AsyncMock()
mock_document_repo.update_page_count = AsyncMock()
return AnalysisService(
converter=converter,
analysis_repo=mock_analysis_repo,
document_repo=mock_document_repo,
)
@pytest.fixture @pytest.fixture
def mock_doc(self): def mock_doc(self):
from domain.models import Document from domain.models import Document
@ -337,22 +353,11 @@ class TestServiceForwardsPipelineOptions:
return AnalysisJob(id="j1", document_id="d1", document_filename="test.pdf") return AnalysisJob(id="j1", document_id="d1", document_filename="test.pdf")
@patch("services.analysis_service.document_repo")
@patch("services.analysis_service.analysis_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_create_passes_pipeline_options_to_run( async def test_create_passes_pipeline_options_to_run(self, mock_doc):
self,
mock_analysis_repo,
mock_doc_repo,
mock_doc,
):
mock_doc_repo.find_by_id = AsyncMock(return_value=mock_doc)
mock_analysis_repo.insert = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
from services.analysis_service import AnalysisService svc = self._make_service(mock_converter)
svc._document_repo.find_by_id = AsyncMock(return_value=mock_doc)
svc = AnalysisService(converter=mock_converter)
opts = {"do_ocr": False, "table_mode": "fast"} opts = {"do_ocr": False, "table_mode": "fast"}
@ -360,42 +365,20 @@ class TestServiceForwardsPipelineOptions:
await svc.create("d1", pipeline_options=opts) await svc.create("d1", pipeline_options=opts)
mock_task.assert_called_once() mock_task.assert_called_once()
@patch("services.analysis_service.document_repo")
@patch("services.analysis_service.analysis_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_create_passes_none_when_no_options( async def test_create_passes_none_when_no_options(self, mock_doc):
self,
mock_analysis_repo,
mock_doc_repo,
mock_doc,
):
mock_doc_repo.find_by_id = AsyncMock(return_value=mock_doc)
mock_analysis_repo.insert = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
from services.analysis_service import AnalysisService svc = self._make_service(mock_converter)
svc._document_repo.find_by_id = AsyncMock(return_value=mock_doc)
svc = AnalysisService(converter=mock_converter)
with patch("services.analysis_service.asyncio.create_task") as mock_task: with patch("services.analysis_service.asyncio.create_task") as mock_task:
await svc.create("d1") await svc.create("d1")
mock_task.assert_called_once() mock_task.assert_called_once()
@patch("services.analysis_service.analysis_repo")
@patch("services.analysis_service.document_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_analysis_forwards_options_to_convert( async def test_run_analysis_forwards_options_to_convert(self, mock_job):
self,
mock_doc_repo,
mock_analysis_repo,
mock_job,
):
from domain.value_objects import ConversionResult, PageDetail from domain.value_objects import ConversionResult, PageDetail
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.return_value = ConversionResult( mock_converter.convert.return_value = ConversionResult(
page_count=1, page_count=1,
@ -404,9 +387,8 @@ class TestServiceForwardsPipelineOptions:
pages=[PageDetail(page_number=1, width=612.0, height=792.0)], pages=[PageDetail(page_number=1, width=612.0, height=792.0)],
) )
from services.analysis_service import AnalysisService svc = self._make_service(mock_converter)
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
opts = { opts = {
"do_ocr": False, "do_ocr": False,
@ -432,21 +414,10 @@ class TestServiceForwardsPipelineOptions:
assert conv_opts.generate_picture_images is True assert conv_opts.generate_picture_images is True
assert conv_opts.images_scale == 2.0 assert conv_opts.images_scale == 2.0
@patch("services.analysis_service.analysis_repo")
@patch("services.analysis_service.document_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_analysis_uses_defaults_when_no_options( async def test_run_analysis_uses_defaults_when_no_options(self, mock_job):
self,
mock_doc_repo,
mock_analysis_repo,
mock_job,
):
from domain.value_objects import ConversionResult, PageDetail from domain.value_objects import ConversionResult, PageDetail
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.return_value = ConversionResult( mock_converter.convert.return_value = ConversionResult(
page_count=1, page_count=1,
@ -455,9 +426,8 @@ class TestServiceForwardsPipelineOptions:
pages=[PageDetail(page_number=1, width=612.0, height=792.0)], pages=[PageDetail(page_number=1, width=612.0, height=792.0)],
) )
from services.analysis_service import AnalysisService svc = self._make_service(mock_converter)
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", None) await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", None)
@ -466,30 +436,19 @@ class TestServiceForwardsPipelineOptions:
assert call_args[0][0] == "/tmp/test.pdf" assert call_args[0][0] == "/tmp/test.pdf"
assert call_args[0][1] == ConversionOptions() assert call_args[0][1] == ConversionOptions()
@patch("services.analysis_service.analysis_repo")
@patch("services.analysis_service.document_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_analysis_marks_failed_on_error( async def test_run_analysis_marks_failed_on_error(self, mock_job):
self,
mock_doc_repo,
mock_analysis_repo,
mock_job,
):
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.side_effect = RuntimeError("Docling crashed") mock_converter.convert.side_effect = RuntimeError("Docling crashed")
from services.analysis_service import AnalysisService svc = self._make_service(mock_converter)
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {"do_ocr": False}) await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {"do_ocr": False})
# Should have called update_status twice: RUNNING then FAILED # Should have called update_status twice: RUNNING then FAILED
assert mock_analysis_repo.update_status.call_count == 2 assert svc._analysis_repo.update_status.call_count == 2
last_job = mock_analysis_repo.update_status.call_args_list[-1][0][0] last_job = svc._analysis_repo.update_status.call_args_list[-1][0][0]
assert last_job.status.value == "FAILED" assert last_job.status.value == "FAILED"
assert "Docling crashed" in last_job.error_message assert "Docling crashed" in last_job.error_message

View file

@ -3,8 +3,9 @@
import pytest import pytest
from domain.models import AnalysisJob, AnalysisStatus, Document from domain.models import AnalysisJob, AnalysisStatus, Document
from persistence import analysis_repo, document_repo from persistence.analysis_repo import SqliteAnalysisRepository
from persistence.database import init_db from persistence.database import init_db
from persistence.document_repo import SqliteDocumentRepository
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
@ -16,8 +17,18 @@ async def setup_db(monkeypatch, tmp_path):
yield yield
@pytest.fixture
def document_repo():
return SqliteDocumentRepository()
@pytest.fixture
def analysis_repo():
return SqliteAnalysisRepository()
class TestDocumentRepo: class TestDocumentRepo:
async def test_insert_and_find_by_id(self): async def test_insert_and_find_by_id(self, document_repo):
doc = Document( doc = Document(
id="doc-1", id="doc-1",
filename="test.pdf", filename="test.pdf",
@ -33,11 +44,11 @@ class TestDocumentRepo:
assert found.filename == "test.pdf" assert found.filename == "test.pdf"
assert found.file_size == 1024 assert found.file_size == 1024
async def test_find_by_id_not_found(self): async def test_find_by_id_not_found(self, document_repo):
found = await document_repo.find_by_id("nonexistent") found = await document_repo.find_by_id("nonexistent")
assert found is None assert found is None
async def test_find_all(self): async def test_find_all(self, document_repo):
for i in range(3): for i in range(3):
doc = Document(id=f"doc-{i}", filename=f"file{i}.pdf", storage_path=f"/tmp/{i}") doc = Document(id=f"doc-{i}", filename=f"file{i}.pdf", storage_path=f"/tmp/{i}")
await document_repo.insert(doc) await document_repo.insert(doc)
@ -45,7 +56,7 @@ class TestDocumentRepo:
all_docs = await document_repo.find_all() all_docs = await document_repo.find_all()
assert len(all_docs) == 3 assert len(all_docs) == 3
async def test_update_page_count(self): async def test_update_page_count(self, document_repo):
doc = Document(id="doc-1", filename="test.pdf", storage_path="/tmp/test.pdf") doc = Document(id="doc-1", filename="test.pdf", storage_path="/tmp/test.pdf")
await document_repo.insert(doc) await document_repo.insert(doc)
@ -54,7 +65,7 @@ class TestDocumentRepo:
updated = await document_repo.find_by_id("doc-1") updated = await document_repo.find_by_id("doc-1")
assert updated.page_count == 10 assert updated.page_count == 10
async def test_delete(self): async def test_delete(self, document_repo):
doc = Document(id="doc-1", filename="test.pdf", storage_path="/tmp/test.pdf") doc = Document(id="doc-1", filename="test.pdf", storage_path="/tmp/test.pdf")
await document_repo.insert(doc) await document_repo.insert(doc)
@ -64,19 +75,19 @@ class TestDocumentRepo:
found = await document_repo.find_by_id("doc-1") found = await document_repo.find_by_id("doc-1")
assert found is None assert found is None
async def test_delete_nonexistent(self): async def test_delete_nonexistent(self, document_repo):
deleted = await document_repo.delete("nonexistent") deleted = await document_repo.delete("nonexistent")
assert deleted is False assert deleted is False
class TestAnalysisRepo: class TestAnalysisRepo:
async def _insert_doc(self): async def _insert_doc(self, document_repo):
doc = Document(id="doc-1", filename="test.pdf", storage_path="/tmp/test.pdf") doc = Document(id="doc-1", filename="test.pdf", storage_path="/tmp/test.pdf")
await document_repo.insert(doc) await document_repo.insert(doc)
return doc return doc
async def test_insert_and_find_by_id(self): async def test_insert_and_find_by_id(self, document_repo, analysis_repo):
await self._insert_doc() await self._insert_doc(document_repo)
job = AnalysisJob(id="job-1", document_id="doc-1") job = AnalysisJob(id="job-1", document_id="doc-1")
await analysis_repo.insert(job) await analysis_repo.insert(job)
@ -87,12 +98,12 @@ class TestAnalysisRepo:
assert found.status == AnalysisStatus.PENDING assert found.status == AnalysisStatus.PENDING
assert found.document_filename == "test.pdf" assert found.document_filename == "test.pdf"
async def test_find_by_id_not_found(self): async def test_find_by_id_not_found(self, analysis_repo):
found = await analysis_repo.find_by_id("nonexistent") found = await analysis_repo.find_by_id("nonexistent")
assert found is None assert found is None
async def test_find_all(self): async def test_find_all(self, document_repo, analysis_repo):
await self._insert_doc() await self._insert_doc(document_repo)
for i in range(3): for i in range(3):
job = AnalysisJob(id=f"job-{i}", document_id="doc-1") job = AnalysisJob(id=f"job-{i}", document_id="doc-1")
await analysis_repo.insert(job) await analysis_repo.insert(job)
@ -100,8 +111,8 @@ class TestAnalysisRepo:
all_jobs = await analysis_repo.find_all() all_jobs = await analysis_repo.find_all()
assert len(all_jobs) == 3 assert len(all_jobs) == 3
async def test_update_status(self): async def test_update_status(self, document_repo, analysis_repo):
await self._insert_doc() await self._insert_doc(document_repo)
job = AnalysisJob(id="job-1", document_id="doc-1") job = AnalysisJob(id="job-1", document_id="doc-1")
await analysis_repo.insert(job) await analysis_repo.insert(job)
@ -112,8 +123,8 @@ class TestAnalysisRepo:
assert found.status == AnalysisStatus.RUNNING assert found.status == AnalysisStatus.RUNNING
assert found.started_at is not None assert found.started_at is not None
async def test_update_status_completed(self): async def test_update_status_completed(self, document_repo, analysis_repo):
await self._insert_doc() await self._insert_doc(document_repo)
job = AnalysisJob(id="job-1", document_id="doc-1") job = AnalysisJob(id="job-1", document_id="doc-1")
await analysis_repo.insert(job) await analysis_repo.insert(job)
@ -127,8 +138,8 @@ class TestAnalysisRepo:
assert found.content_html == "<h1>Test</h1>" assert found.content_html == "<h1>Test</h1>"
assert found.pages_json == "[]" assert found.pages_json == "[]"
async def test_delete(self): async def test_delete(self, document_repo, analysis_repo):
await self._insert_doc() await self._insert_doc(document_repo)
job = AnalysisJob(id="job-1", document_id="doc-1") job = AnalysisJob(id="job-1", document_id="doc-1")
await analysis_repo.insert(job) await analysis_repo.insert(job)
@ -138,12 +149,12 @@ class TestAnalysisRepo:
found = await analysis_repo.find_by_id("job-1") found = await analysis_repo.find_by_id("job-1")
assert found is None assert found is None
async def test_delete_nonexistent(self): async def test_delete_nonexistent(self, analysis_repo):
deleted = await analysis_repo.delete("nonexistent") deleted = await analysis_repo.delete("nonexistent")
assert deleted is False assert deleted is False
async def test_delete_by_document(self): async def test_delete_by_document(self, document_repo, analysis_repo):
await self._insert_doc() await self._insert_doc(document_repo)
for i in range(3): for i in range(3):
job = AnalysisJob(id=f"job-{i}", document_id="doc-1") job = AnalysisJob(id=f"job-{i}", document_id="doc-1")
await analysis_repo.insert(job) await analysis_repo.insert(job)

View file

@ -11,9 +11,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
from domain.models import AnalysisJob from domain.models import AnalysisJob
from domain.services import classify_error
from domain.value_objects import ConversionOptions, ConversionResult, PageDetail from domain.value_objects import ConversionOptions, ConversionResult, PageDetail
from infra.settings import Settings from infra.settings import Settings
from services.analysis_service import AnalysisService, _classify_error from services.analysis_service import AnalysisConfig, AnalysisService
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# M1 — _classify_error: user-friendly error messages # M1 — _classify_error: user-friendly error messages
@ -25,40 +26,40 @@ class TestClassifyError:
def test_cxx_compiler_error(self): def test_cxx_compiler_error(self):
exc = RuntimeError("InvalidCxxCompiler: No working C++ compiler found") exc = RuntimeError("InvalidCxxCompiler: No working C++ compiler found")
assert "Missing C++ compiler" in _classify_error(exc) assert "Missing C++ compiler" in classify_error(exc)
def test_no_working_compiler(self): def test_no_working_compiler(self):
exc = RuntimeError("no working c++ compiler found in torch") exc = RuntimeError("no working c++ compiler found in torch")
assert "Missing C++ compiler" in _classify_error(exc) assert "Missing C++ compiler" in classify_error(exc)
def test_out_of_memory(self): def test_out_of_memory(self):
exc = MemoryError("Out of memory allocating tensor") exc = MemoryError("Out of memory allocating tensor")
assert "Out of memory" in _classify_error(exc) assert "Out of memory" in classify_error(exc)
def test_oom_shorthand(self): def test_oom_shorthand(self):
exc = RuntimeError("OOM during inference on page 5") exc = RuntimeError("OOM during inference on page 5")
assert "Out of memory" in _classify_error(exc) assert "Out of memory" in classify_error(exc)
def test_lock_timeout(self): def test_lock_timeout(self):
exc = TimeoutError("Could not acquire converter lock within 300s") exc = TimeoutError("Could not acquire converter lock within 300s")
assert "Server busy" in _classify_error(exc) assert "Server busy" in classify_error(exc)
def test_pipeline_failed(self): def test_pipeline_failed(self):
exc = RuntimeError("Pipeline StandardPdfPipeline failed on page 3") exc = RuntimeError("Pipeline StandardPdfPipeline failed on page 3")
assert "Document processing failed" in _classify_error(exc) assert "Document processing failed" in classify_error(exc)
def test_timeout_generic(self): def test_timeout_generic(self):
exc = TimeoutError("timeout exceeded while processing") exc = TimeoutError("timeout exceeded while processing")
assert "Processing took too long" in _classify_error(exc) assert "Processing took too long" in classify_error(exc)
def test_unknown_short_error(self): def test_unknown_short_error(self):
exc = ValueError("something weird happened") exc = ValueError("something weird happened")
assert _classify_error(exc) == "something weird happened" assert classify_error(exc) == "something weird happened"
def test_unknown_long_error_truncated(self): def test_unknown_long_error_truncated(self):
long_msg = "x" * 300 long_msg = "x" * 300
exc = ValueError(long_msg) exc = ValueError(long_msg)
result = _classify_error(exc) result = classify_error(exc)
assert len(result) <= 201 assert len(result) <= 201
assert result.endswith("") assert result.endswith("")
@ -83,74 +84,63 @@ class TestDefaultTableMode:
pages=[PageDetail(page_number=1, width=612.0, height=792.0)], pages=[PageDetail(page_number=1, width=612.0, height=792.0)],
) )
@patch("services.analysis_service.analysis_repo") def _make_service(self, converter, *, default_table_mode="accurate"):
@patch("services.analysis_service.document_repo") mock_analysis_repo = MagicMock()
@pytest.mark.asyncio mock_analysis_repo.find_by_id = AsyncMock()
async def test_default_table_mode_injected_when_missing(
self, mock_doc_repo, mock_analysis_repo, mock_job
):
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock() mock_analysis_repo.update_status = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock() mock_document_repo = MagicMock()
mock_document_repo.update_page_count = AsyncMock()
config = AnalysisConfig(default_table_mode=default_table_mode)
return AnalysisService(
converter=converter,
analysis_repo=mock_analysis_repo,
document_repo=mock_document_repo,
config=config,
)
@pytest.mark.asyncio
async def test_default_table_mode_injected_when_missing(self, mock_job):
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.return_value = self._make_result() mock_converter.convert.return_value = self._make_result()
svc = self._make_service(mock_converter)
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {}) await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {})
opts = mock_converter.convert.call_args[0][1] opts = mock_converter.convert.call_args[0][1]
assert opts.table_mode == "accurate" assert opts.table_mode == "accurate"
@patch("services.analysis_service.analysis_repo")
@patch("services.analysis_service.document_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_default_table_mode_injected_when_none( async def test_default_table_mode_injected_when_none(self, mock_job):
self, mock_doc_repo, mock_analysis_repo, mock_job
):
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.return_value = self._make_result() mock_converter.convert.return_value = self._make_result()
svc = self._make_service(mock_converter)
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", None) await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", None)
opts = mock_converter.convert.call_args[0][1] opts = mock_converter.convert.call_args[0][1]
assert opts.table_mode == "accurate" assert opts.table_mode == "accurate"
@patch("services.analysis_service.analysis_repo")
@patch("services.analysis_service.document_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_user_table_mode_preserved(self, mock_doc_repo, mock_analysis_repo, mock_job): async def test_user_table_mode_preserved(self, mock_job):
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.return_value = self._make_result() mock_converter.convert.return_value = self._make_result()
svc = self._make_service(mock_converter)
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {"table_mode": "fast"}) await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {"table_mode": "fast"})
opts = mock_converter.convert.call_args[0][1] opts = mock_converter.convert.call_args[0][1]
assert opts.table_mode == "fast" assert opts.table_mode == "fast"
@patch("services.analysis_service.settings", Settings(default_table_mode="fast"))
@patch("services.analysis_service.analysis_repo")
@patch("services.analysis_service.document_repo")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_custom_default_from_settings(self, mock_doc_repo, mock_analysis_repo, mock_job): async def test_custom_default_from_settings(self, mock_job):
mock_analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
mock_analysis_repo.update_status = AsyncMock()
mock_doc_repo.update_page_count = AsyncMock()
mock_converter = AsyncMock() mock_converter = AsyncMock()
mock_converter.convert.return_value = self._make_result() mock_converter.convert.return_value = self._make_result()
svc = self._make_service(mock_converter, default_table_mode="fast")
svc._analysis_repo.find_by_id = AsyncMock(return_value=mock_job)
svc = AnalysisService(converter=mock_converter)
await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {}) await svc._run_analysis("j1", "/tmp/test.pdf", "test.pdf", {})
opts = mock_converter.convert.call_args[0][1] opts = mock_converter.convert.call_args[0][1]