generated from john/python-template
73 lines
2.5 KiB
Python
73 lines
2.5 KiB
Python
from uuid import uuid4
|
|
|
|
import pytest
|
|
from sqlmodel import select
|
|
|
|
from transcription.config import Settings
|
|
from transcription.db.models import Document
|
|
from transcription.db.models import Job
|
|
from transcription.db.models import JobSource
|
|
from transcription.db.models import Source
|
|
from transcription.services.store import UploadError
|
|
from transcription.services.store import create_job_for_document
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_create_job_for_document_requires_at_least_one_upload(async_session, tmp_path):
|
|
document = Document(id=uuid4(), name="needs-upload")
|
|
async_session.add(document)
|
|
await async_session.commit()
|
|
|
|
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
|
|
|
|
with pytest.raises(UploadError):
|
|
await create_job_for_document(
|
|
document_id=document.id,
|
|
uploads=[],
|
|
session=async_session,
|
|
settings=settings,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_create_job_for_document_sorts_uploads_and_creates_links(async_session, tmp_path):
|
|
document = Document(id=uuid4(), name="ordered-upload-doc")
|
|
async_session.add(document)
|
|
await async_session.commit()
|
|
|
|
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
|
|
|
|
result = await create_job_for_document(
|
|
document_id=document.id,
|
|
uploads=[
|
|
("folder/b_page.pdf", b"b"),
|
|
("folder/A_page.pdf", b"a"),
|
|
],
|
|
provider="openrouter",
|
|
model="test-model",
|
|
prompt_name="transcribe_document.md",
|
|
session=async_session,
|
|
settings=settings,
|
|
)
|
|
|
|
created_job = await async_session.get(Job, result.job_id)
|
|
assert created_job is not None
|
|
assert created_job.provider == "openrouter"
|
|
assert created_job.model == "test-model"
|
|
assert created_job.prompt_name == "transcribe_document.md"
|
|
|
|
sources = (
|
|
await async_session.exec(
|
|
select(Source)
|
|
.where(Source.document_id == document.id)
|
|
.order_by(Source.page_number) # pyright: ignore[reportArgumentType]
|
|
)
|
|
).all()
|
|
assert [source.upload_name for source in sources] == ["A_page.pdf", "b_page.pdf"]
|
|
assert all(source.filename.endswith(".pdf") for source in sources)
|
|
assert all("A_page" not in source.filename and "b_page" not in source.filename for source in sources)
|
|
|
|
job_sources = (await async_session.exec(select(JobSource).where(JobSource.job_id == result.job_id))).all()
|
|
assert len(job_sources) == 2
|
|
assert set(result.source_ids) == {job_source.source_id for job_source in job_sources}
|