generated from john/python-template
V3 post step 2 refinement: add temperature & top-p settings to config (and .env), add prompt fields back to job table so that the prompt settings get frozen at runtime for all sources being processed.
This commit is contained in:
@@ -126,6 +126,12 @@ class Job(SQLModel, table=True):
|
|||||||
date_updated: datetime = Field(default_factory=lambda: datetime.now(UTC))
|
date_updated: datetime = Field(default_factory=lambda: datetime.now(UTC))
|
||||||
provider: str | None = None
|
provider: str | None = None
|
||||||
model: str | None = None
|
model: str | None = None
|
||||||
|
prompt_name: str | None = None
|
||||||
|
prompt_hash: str | None = None
|
||||||
|
system_prompt: str | None = None
|
||||||
|
user_prompt: str | None = None
|
||||||
|
temperature: float | None = None
|
||||||
|
top_p: float | None = None
|
||||||
|
|
||||||
document: Optional["Document"] = Relationship(back_populates="jobs", sa_relationship_kwargs={"lazy": "selectin"})
|
document: Optional["Document"] = Relationship(back_populates="jobs", sa_relationship_kwargs={"lazy": "selectin"})
|
||||||
job_sources: list["JobSource"] = Relationship(back_populates="job", sa_relationship_kwargs={"lazy": "selectin"})
|
job_sources: list["JobSource"] = Relationship(back_populates="job", sa_relationship_kwargs={"lazy": "selectin"})
|
||||||
@@ -223,12 +229,6 @@ class JobSource(SQLModel, table=True):
|
|||||||
source_id: UUID = Field(foreign_key="source.id")
|
source_id: UUID = Field(foreign_key="source.id")
|
||||||
status: JobSourceStatus = Field(default=JobSourceStatus.PENDING)
|
status: JobSourceStatus = Field(default=JobSourceStatus.PENDING)
|
||||||
raw_transcription: str | None = None
|
raw_transcription: str | None = None
|
||||||
prompt_name: str | None = None
|
|
||||||
prompt_hash: str | None = None
|
|
||||||
system_prompt: str | None = None
|
|
||||||
user_prompt: str | None = None
|
|
||||||
temperature: float | None = None
|
|
||||||
top_p: float | None = None
|
|
||||||
ai_metadata: dict[str, Any] | None = Field(default=None, sa_column=Column(JSONBCompat(), nullable=True))
|
ai_metadata: dict[str, Any] | None = Field(default=None, sa_column=Column(JSONBCompat(), nullable=True))
|
||||||
raw_api_response: dict[str, Any] | None = Field(default=None, sa_column=Column(JSONBCompat(), nullable=True))
|
raw_api_response: dict[str, Any] | None = Field(default=None, sa_column=Column(JSONBCompat(), nullable=True))
|
||||||
error_detail: str | None = None
|
error_detail: str | None = None
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from ..db.models import JobSource
|
|||||||
from ..db.models import JobSourceStatus
|
from ..db.models import JobSourceStatus
|
||||||
from ..db.models import Source
|
from ..db.models import Source
|
||||||
from .documents import UploadJobResult
|
from .documents import UploadJobResult
|
||||||
|
from .transcription import build_prompt_execution
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -62,6 +63,7 @@ async def create_upload_job(
|
|||||||
) -> UploadJobResult:
|
) -> UploadJobResult:
|
||||||
"""Create upload-backed document and queued job records."""
|
"""Create upload-backed document and queued job records."""
|
||||||
runtime_settings = settings or get_settings()
|
runtime_settings = settings or get_settings()
|
||||||
|
prompt_execution = build_prompt_execution(settings=runtime_settings)
|
||||||
document_id = uuid4()
|
document_id = uuid4()
|
||||||
source_id = uuid4()
|
source_id = uuid4()
|
||||||
stored_path = store_file(
|
stored_path = store_file(
|
||||||
@@ -81,6 +83,7 @@ async def create_upload_job(
|
|||||||
stored_path=stored_path,
|
stored_path=stored_path,
|
||||||
file_hash=file_hash,
|
file_hash=file_hash,
|
||||||
file_size_bytes=file_size_bytes,
|
file_size_bytes=file_size_bytes,
|
||||||
|
prompt_execution=prompt_execution,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
_best_effort_delete(stored_path)
|
_best_effort_delete(stored_path)
|
||||||
@@ -118,6 +121,7 @@ async def create_job_for_document(
|
|||||||
)
|
)
|
||||||
|
|
||||||
runtime_settings = settings or get_settings()
|
runtime_settings = settings or get_settings()
|
||||||
|
prompt_execution = build_prompt_execution(settings=runtime_settings)
|
||||||
sorted_uploads = sorted(uploads, key=lambda item: Path(item[0]).name.casefold())
|
sorted_uploads = sorted(uploads, key=lambda item: Path(item[0]).name.casefold())
|
||||||
stored_uploads: list[PendingStoredUpload] = []
|
stored_uploads: list[PendingStoredUpload] = []
|
||||||
for filename, file_bytes in sorted_uploads:
|
for filename, file_bytes in sorted_uploads:
|
||||||
@@ -145,6 +149,7 @@ async def create_job_for_document(
|
|||||||
stored_uploads=stored_uploads,
|
stored_uploads=stored_uploads,
|
||||||
provider=provider,
|
provider=provider,
|
||||||
model=model,
|
model=model,
|
||||||
|
prompt_execution=prompt_execution,
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
for upload in stored_uploads:
|
for upload in stored_uploads:
|
||||||
@@ -173,6 +178,7 @@ async def _create_upload_records(
|
|||||||
stored_path: Path,
|
stored_path: Path,
|
||||||
file_hash: str,
|
file_hash: str,
|
||||||
file_size_bytes: int,
|
file_size_bytes: int,
|
||||||
|
prompt_execution,
|
||||||
) -> tuple[Document, Job]:
|
) -> tuple[Document, Job]:
|
||||||
document = Document(
|
document = Document(
|
||||||
id=document_id,
|
id=document_id,
|
||||||
@@ -181,7 +187,15 @@ async def _create_upload_records(
|
|||||||
session.add(document)
|
session.add(document)
|
||||||
await session.flush()
|
await session.flush()
|
||||||
|
|
||||||
job = Job(document_id=document.id)
|
job = Job(
|
||||||
|
document_id=document.id,
|
||||||
|
prompt_name=prompt_execution.prompt_name,
|
||||||
|
prompt_hash=prompt_execution.prompt_hash,
|
||||||
|
system_prompt=prompt_execution.system_prompt,
|
||||||
|
user_prompt=prompt_execution.user_prompt,
|
||||||
|
temperature=prompt_execution.temperature,
|
||||||
|
top_p=prompt_execution.top_p,
|
||||||
|
)
|
||||||
session.add(job)
|
session.add(job)
|
||||||
await session.flush()
|
await session.flush()
|
||||||
|
|
||||||
@@ -219,6 +233,7 @@ async def _create_job_for_document_records(
|
|||||||
stored_uploads: Sequence[PendingStoredUpload],
|
stored_uploads: Sequence[PendingStoredUpload],
|
||||||
provider: str | None,
|
provider: str | None,
|
||||||
model: str | None,
|
model: str | None,
|
||||||
|
prompt_execution,
|
||||||
) -> tuple[Job, list[UUID]]:
|
) -> tuple[Job, list[UUID]]:
|
||||||
document = await session.get(Document, document_id)
|
document = await session.get(Document, document_id)
|
||||||
if document is None:
|
if document is None:
|
||||||
@@ -237,6 +252,12 @@ async def _create_job_for_document_records(
|
|||||||
document_id=document_id,
|
document_id=document_id,
|
||||||
provider=(provider or None),
|
provider=(provider or None),
|
||||||
model=(model or None),
|
model=(model or None),
|
||||||
|
prompt_name=prompt_execution.prompt_name,
|
||||||
|
prompt_hash=prompt_execution.prompt_hash,
|
||||||
|
system_prompt=prompt_execution.system_prompt,
|
||||||
|
user_prompt=prompt_execution.user_prompt,
|
||||||
|
temperature=prompt_execution.temperature,
|
||||||
|
top_p=prompt_execution.top_p,
|
||||||
)
|
)
|
||||||
session.add(job)
|
session.add(job)
|
||||||
await session.flush()
|
await session.flush()
|
||||||
|
|||||||
@@ -424,12 +424,6 @@ class TranscriptionService(ServiceBase):
|
|||||||
error_detail=error_detail,
|
error_detail=error_detail,
|
||||||
provider=provider,
|
provider=provider,
|
||||||
model=model,
|
model=model,
|
||||||
prompt_name=prompt_name,
|
|
||||||
prompt_hash=prompt_hash,
|
|
||||||
system_prompt=system_prompt,
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
session=_session,
|
session=_session,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -447,12 +441,6 @@ class TranscriptionService(ServiceBase):
|
|||||||
raw_api_response: dict[str, object] | None = None,
|
raw_api_response: dict[str, object] | None = None,
|
||||||
provider: str | None = None,
|
provider: str | None = None,
|
||||||
model: str | None = None,
|
model: str | None = None,
|
||||||
prompt_name: str = DEFAULT_PROMPT_FILE,
|
|
||||||
prompt_hash: str | None = None,
|
|
||||||
system_prompt: str | None = None,
|
|
||||||
user_prompt: str | None = None,
|
|
||||||
temperature: float | None = None,
|
|
||||||
top_p: float | None = None,
|
|
||||||
session: AsyncSession | None = None,
|
session: AsyncSession | None = None,
|
||||||
) -> JobSource:
|
) -> JobSource:
|
||||||
"""Persist transcription fields for one source within a specific job."""
|
"""Persist transcription fields for one source within a specific job."""
|
||||||
@@ -497,12 +485,6 @@ class TranscriptionService(ServiceBase):
|
|||||||
source_id=source_id,
|
source_id=source_id,
|
||||||
status=JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED,
|
status=JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED,
|
||||||
raw_transcription=text,
|
raw_transcription=text,
|
||||||
prompt_name=prompt_name,
|
|
||||||
prompt_hash=prompt_hash,
|
|
||||||
system_prompt=system_prompt,
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
ai_metadata=ai_metadata,
|
ai_metadata=ai_metadata,
|
||||||
raw_api_response=raw_api_response,
|
raw_api_response=raw_api_response,
|
||||||
error_detail=error_detail,
|
error_detail=error_detail,
|
||||||
@@ -510,12 +492,6 @@ class TranscriptionService(ServiceBase):
|
|||||||
_session.add(job_source)
|
_session.add(job_source)
|
||||||
else:
|
else:
|
||||||
job_source.raw_transcription = text
|
job_source.raw_transcription = text
|
||||||
job_source.prompt_name = prompt_name
|
|
||||||
job_source.prompt_hash = prompt_hash
|
|
||||||
job_source.system_prompt = system_prompt
|
|
||||||
job_source.user_prompt = user_prompt
|
|
||||||
job_source.temperature = temperature
|
|
||||||
job_source.top_p = top_p
|
|
||||||
job_source.ai_metadata = ai_metadata
|
job_source.ai_metadata = ai_metadata
|
||||||
job_source.raw_api_response = raw_api_response
|
job_source.raw_api_response = raw_api_response
|
||||||
job_source.error_detail = error_detail
|
job_source.error_detail = error_detail
|
||||||
@@ -591,12 +567,26 @@ async def transcribe_document_image(
|
|||||||
image_path: str | Path,
|
image_path: str | Path,
|
||||||
*,
|
*,
|
||||||
prompt_name: str | None = None,
|
prompt_name: str | None = None,
|
||||||
|
prompt_text: str | None = None,
|
||||||
|
temperature: float | None = None,
|
||||||
|
top_p: float | None = None,
|
||||||
settings: Settings | None = None,
|
settings: Settings | None = None,
|
||||||
provider: TranscriptionProvider | None = None,
|
provider: TranscriptionProvider | None = None,
|
||||||
) -> TranscriptionResult:
|
) -> TranscriptionResult:
|
||||||
"""Transcribe a local image using the configured prompt and provider."""
|
"""Transcribe a local image using the configured prompt and provider."""
|
||||||
runtime_settings = settings or get_settings()
|
runtime_settings = settings or get_settings()
|
||||||
prompt_execution = build_prompt_execution(prompt_name=prompt_name, settings=runtime_settings)
|
if prompt_text is None:
|
||||||
|
prompt_execution = build_prompt_execution(prompt_name=prompt_name, settings=runtime_settings)
|
||||||
|
else:
|
||||||
|
effective_prompt_name = (prompt_name or runtime_settings.default_prompt_name or DEFAULT_PROMPT_FILE).strip()
|
||||||
|
prompt_execution = PromptExecution(
|
||||||
|
prompt_name=effective_prompt_name,
|
||||||
|
prompt_hash=hashlib.sha256(prompt_text.encode("utf-8")).hexdigest(),
|
||||||
|
system_prompt=None,
|
||||||
|
user_prompt=prompt_text,
|
||||||
|
temperature=temperature if temperature is not None else runtime_settings.transcription_temperature,
|
||||||
|
top_p=top_p if top_p is not None else runtime_settings.transcription_top_p,
|
||||||
|
)
|
||||||
image_bytes, mime_type = load_image_payload(image_path)
|
image_bytes, mime_type = load_image_payload(image_path)
|
||||||
|
|
||||||
adapter = provider or get_transcription_provider(settings=runtime_settings)
|
adapter = provider or get_transcription_provider(settings=runtime_settings)
|
||||||
|
|||||||
@@ -15,8 +15,8 @@ from ..errors import classify_unexpected_error
|
|||||||
from ..errors import format_error_detail
|
from ..errors import format_error_detail
|
||||||
from ..providers import TranscriptionResult
|
from ..providers import TranscriptionResult
|
||||||
from . import ServiceBundle
|
from . import ServiceBundle
|
||||||
from .transcription import DEFAULT_PROMPT_FILE
|
|
||||||
from .transcription import build_prompt_execution
|
from .transcription import build_prompt_execution
|
||||||
|
from .transcription import PromptExecution
|
||||||
from .transcription import transcribe_document_image
|
from .transcription import transcribe_document_image
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -85,13 +85,13 @@ async def process_queued_job(
|
|||||||
if not sources:
|
if not sources:
|
||||||
return await services.jobs.mark_job_status(job.id, JobStatus.TRANSCRIBED, session=session)
|
return await services.jobs.mark_job_status(job.id, JobStatus.TRANSCRIBED, session=session)
|
||||||
|
|
||||||
successful_pages: list[tuple[Source, TranscriptionResult, str, str | None, str, float | None, float | None]] = []
|
successful_pages: list[tuple[Source, TranscriptionResult]] = []
|
||||||
failed_pages: list[tuple[Source, AppError, str, str, str | None, str, float | None, float | None]] = []
|
failed_pages: list[tuple[Source, AppError]] = []
|
||||||
externally_stopped = False
|
externally_stopped = False
|
||||||
|
|
||||||
for source in sources:
|
prompt_execution = _resolve_job_prompt_execution(source_job=source_job, settings=runtime_settings)
|
||||||
prompt_execution = build_prompt_execution(settings=runtime_settings)
|
|
||||||
|
|
||||||
|
for source in sources:
|
||||||
if await _job_no_longer_processing(job_id=job.id, services=services, session=session):
|
if await _job_no_longer_processing(job_id=job.id, services=services, session=session):
|
||||||
externally_stopped = True
|
externally_stopped = True
|
||||||
break
|
break
|
||||||
@@ -101,6 +101,10 @@ async def process_queued_job(
|
|||||||
result = await asyncio.wait_for(
|
result = await asyncio.wait_for(
|
||||||
transcribe_document_image(
|
transcribe_document_image(
|
||||||
source.file_path,
|
source.file_path,
|
||||||
|
prompt_name=prompt_execution.prompt_name,
|
||||||
|
prompt_text=prompt_execution.user_prompt,
|
||||||
|
temperature=prompt_execution.temperature,
|
||||||
|
top_p=prompt_execution.top_p,
|
||||||
settings=runtime_settings,
|
settings=runtime_settings,
|
||||||
provider=services.transcriptions.provider,
|
provider=services.transcriptions.provider,
|
||||||
),
|
),
|
||||||
@@ -127,17 +131,7 @@ async def process_queued_job(
|
|||||||
)
|
)
|
||||||
|
|
||||||
_validate_transcription_quality(result=result, settings=runtime_settings)
|
_validate_transcription_quality(result=result, settings=runtime_settings)
|
||||||
successful_pages.append(
|
successful_pages.append((source, result))
|
||||||
(
|
|
||||||
source,
|
|
||||||
result,
|
|
||||||
prompt_execution.prompt_hash,
|
|
||||||
prompt_execution.system_prompt,
|
|
||||||
prompt_execution.user_prompt,
|
|
||||||
prompt_execution.temperature,
|
|
||||||
prompt_execution.top_p,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
except TimeoutError:
|
except TimeoutError:
|
||||||
error = AppError(
|
error = AppError(
|
||||||
f"Provider call timed out after {runtime_settings.worker_provider_timeout_seconds:.1f}s",
|
f"Provider call timed out after {runtime_settings.worker_provider_timeout_seconds:.1f}s",
|
||||||
@@ -145,18 +139,7 @@ async def process_queued_job(
|
|||||||
suggestion="Retry the job. If this repeats, verify provider latency and request payload size.",
|
suggestion="Retry the job. If this repeats, verify provider latency and request payload size.",
|
||||||
retriable=True,
|
retriable=True,
|
||||||
)
|
)
|
||||||
failed_pages.append(
|
failed_pages.append((source, error))
|
||||||
(
|
|
||||||
source,
|
|
||||||
error,
|
|
||||||
prompt_execution.prompt_name,
|
|
||||||
prompt_execution.prompt_hash,
|
|
||||||
prompt_execution.system_prompt,
|
|
||||||
prompt_execution.user_prompt,
|
|
||||||
prompt_execution.temperature,
|
|
||||||
prompt_execution.top_p,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Source failed operation=worker.process_job job_id=%s document_id=%s source_id=%s error_id=%s category=%s",
|
"Source failed operation=worker.process_job job_id=%s document_id=%s source_id=%s error_id=%s category=%s",
|
||||||
job.id,
|
job.id,
|
||||||
@@ -172,18 +155,7 @@ async def process_queued_job(
|
|||||||
case _:
|
case _:
|
||||||
error = classify_unexpected_error(exc, operation="worker.process_job")
|
error = classify_unexpected_error(exc, operation="worker.process_job")
|
||||||
|
|
||||||
failed_pages.append(
|
failed_pages.append((source, error))
|
||||||
(
|
|
||||||
source,
|
|
||||||
error,
|
|
||||||
prompt_execution.prompt_name,
|
|
||||||
prompt_execution.prompt_hash,
|
|
||||||
prompt_execution.system_prompt,
|
|
||||||
prompt_execution.user_prompt,
|
|
||||||
prompt_execution.temperature,
|
|
||||||
prompt_execution.top_p,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Source failed operation=worker.process_job job_id=%s document_id=%s source_id=%s error_id=%s category=%s",
|
"Source failed operation=worker.process_job job_id=%s document_id=%s source_id=%s error_id=%s category=%s",
|
||||||
job.id,
|
job.id,
|
||||||
@@ -392,6 +364,20 @@ def _resolve_job_sources(job: Job) -> list[Source]:
|
|||||||
return list(sorted(sources, key=lambda item: (item.page_number, item.upload_name.casefold())))
|
return list(sorted(sources, key=lambda item: (item.page_number, item.upload_name.casefold())))
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_job_prompt_execution(*, source_job: Job, settings: Settings) -> PromptExecution:
|
||||||
|
if source_job.user_prompt and source_job.prompt_name:
|
||||||
|
return PromptExecution(
|
||||||
|
prompt_name=source_job.prompt_name,
|
||||||
|
prompt_hash=source_job.prompt_hash or "",
|
||||||
|
system_prompt=source_job.system_prompt,
|
||||||
|
user_prompt=source_job.user_prompt,
|
||||||
|
temperature=source_job.temperature,
|
||||||
|
top_p=source_job.top_p,
|
||||||
|
)
|
||||||
|
|
||||||
|
return build_prompt_execution(settings=settings)
|
||||||
|
|
||||||
|
|
||||||
async def _job_no_longer_processing(
|
async def _job_no_longer_processing(
|
||||||
*,
|
*,
|
||||||
job_id,
|
job_id,
|
||||||
@@ -407,15 +393,15 @@ async def _finalize_batch_outcome(
|
|||||||
*,
|
*,
|
||||||
job: Job,
|
job: Job,
|
||||||
services: ServiceBundle,
|
services: ServiceBundle,
|
||||||
successful_pages: list[tuple[Source, TranscriptionResult, str, str | None, str, float | None, float | None]],
|
successful_pages: list[tuple[Source, TranscriptionResult]],
|
||||||
failed_pages: list[tuple[Source, AppError, str, str, str | None, str, float | None, float | None]],
|
failed_pages: list[tuple[Source, AppError]],
|
||||||
status: JobStatus,
|
status: JobStatus,
|
||||||
session: AsyncSession | None = None,
|
session: AsyncSession | None = None,
|
||||||
) -> Job:
|
) -> Job:
|
||||||
"""Transaction B: write per-source outcomes and terminal job status atomically."""
|
"""Transaction B: write per-source outcomes and terminal job status atomically."""
|
||||||
if session is None:
|
if session is None:
|
||||||
async with services.jobs._session_scope() as local_session:
|
async with services.jobs._session_scope() as local_session:
|
||||||
for source, result, prompt_hash, system_prompt, user_prompt, temperature, top_p in successful_pages:
|
for source, result in successful_pages:
|
||||||
await services.transcriptions.update_job_source_transcription(
|
await services.transcriptions.update_job_source_transcription(
|
||||||
job_id=job.id,
|
job_id=job.id,
|
||||||
source_id=source.id,
|
source_id=source.id,
|
||||||
@@ -425,27 +411,15 @@ async def _finalize_batch_outcome(
|
|||||||
raw_api_response=result.raw_api_response,
|
raw_api_response=result.raw_api_response,
|
||||||
provider=result.provider,
|
provider=result.provider,
|
||||||
model=result.model,
|
model=result.model,
|
||||||
prompt_name=result.prompt_name,
|
|
||||||
prompt_hash=result.prompt_hash or prompt_hash,
|
|
||||||
system_prompt=result.system_prompt if result.system_prompt is not None else system_prompt,
|
|
||||||
user_prompt=result.user_prompt if result.user_prompt is not None else user_prompt,
|
|
||||||
temperature=result.temperature if result.temperature is not None else temperature,
|
|
||||||
top_p=result.top_p if result.top_p is not None else top_p,
|
|
||||||
session=local_session,
|
session=local_session,
|
||||||
)
|
)
|
||||||
|
|
||||||
for source, error, prompt_name, prompt_hash, system_prompt, user_prompt, temperature, top_p in failed_pages:
|
for source, error in failed_pages:
|
||||||
await services.transcriptions.update_job_source_transcription(
|
await services.transcriptions.update_job_source_transcription(
|
||||||
job_id=job.id,
|
job_id=job.id,
|
||||||
source_id=source.id,
|
source_id=source.id,
|
||||||
text=None,
|
text=None,
|
||||||
error_detail=format_error_detail(error),
|
error_detail=format_error_detail(error),
|
||||||
prompt_name=prompt_name,
|
|
||||||
prompt_hash=prompt_hash,
|
|
||||||
system_prompt=system_prompt,
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
session=local_session,
|
session=local_session,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -453,7 +427,7 @@ async def _finalize_batch_outcome(
|
|||||||
await local_session.commit()
|
await local_session.commit()
|
||||||
return updated_job
|
return updated_job
|
||||||
|
|
||||||
for source, result, prompt_hash, system_prompt, user_prompt, temperature, top_p in successful_pages:
|
for source, result in successful_pages:
|
||||||
await services.transcriptions.update_job_source_transcription(
|
await services.transcriptions.update_job_source_transcription(
|
||||||
job_id=job.id,
|
job_id=job.id,
|
||||||
source_id=source.id,
|
source_id=source.id,
|
||||||
@@ -463,27 +437,15 @@ async def _finalize_batch_outcome(
|
|||||||
raw_api_response=result.raw_api_response,
|
raw_api_response=result.raw_api_response,
|
||||||
provider=result.provider,
|
provider=result.provider,
|
||||||
model=result.model,
|
model=result.model,
|
||||||
prompt_name=result.prompt_name,
|
|
||||||
prompt_hash=result.prompt_hash or prompt_hash,
|
|
||||||
system_prompt=result.system_prompt if result.system_prompt is not None else system_prompt,
|
|
||||||
user_prompt=result.user_prompt if result.user_prompt is not None else user_prompt,
|
|
||||||
temperature=result.temperature if result.temperature is not None else temperature,
|
|
||||||
top_p=result.top_p if result.top_p is not None else top_p,
|
|
||||||
session=session,
|
session=session,
|
||||||
)
|
)
|
||||||
|
|
||||||
for source, error, prompt_name, prompt_hash, system_prompt, user_prompt, temperature, top_p in failed_pages:
|
for source, error in failed_pages:
|
||||||
await services.transcriptions.update_job_source_transcription(
|
await services.transcriptions.update_job_source_transcription(
|
||||||
job_id=job.id,
|
job_id=job.id,
|
||||||
source_id=source.id,
|
source_id=source.id,
|
||||||
text=None,
|
text=None,
|
||||||
error_detail=format_error_detail(error),
|
error_detail=format_error_detail(error),
|
||||||
prompt_name=prompt_name,
|
|
||||||
prompt_hash=prompt_hash,
|
|
||||||
system_prompt=system_prompt,
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
session=session,
|
session=session,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -116,8 +116,8 @@ def _latest_job_error_detail(job: Job) -> str | None:
|
|||||||
|
|
||||||
def _latest_job_prompt(job: Job) -> str | None:
|
def _latest_job_prompt(job: Job) -> str | None:
|
||||||
for job_source in sorted(job.job_sources, key=lambda item: item.executed_at, reverse=True):
|
for job_source in sorted(job.job_sources, key=lambda item: item.executed_at, reverse=True):
|
||||||
if job_source.prompt_name:
|
if job_source.job and job_source.job.prompt_name:
|
||||||
return job_source.prompt_name
|
return job_source.job.prompt_name
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -452,6 +452,6 @@ def _parse_uuid(value: str | None) -> UUID | None:
|
|||||||
def _latest_prompt_name(job: Job) -> str | None:
|
def _latest_prompt_name(job: Job) -> str | None:
|
||||||
ordered = sorted(job.job_sources, key=lambda item: item.executed_at, reverse=True)
|
ordered = sorted(job.job_sources, key=lambda item: item.executed_at, reverse=True)
|
||||||
for job_source in ordered:
|
for job_source in ordered:
|
||||||
if job_source.prompt_name:
|
if job_source.job and job_source.job.prompt_name:
|
||||||
return job_source.prompt_name
|
return job_source.job.prompt_name
|
||||||
return None
|
return None
|
||||||
@@ -281,7 +281,7 @@ def _render_source_job_metadata_zone(latest_job_source: JobSource | None) -> Non
|
|||||||
)
|
)
|
||||||
metadata_row(
|
metadata_row(
|
||||||
"Prompt:",
|
"Prompt:",
|
||||||
latest_job_source.prompt_name or "unknown",
|
latest_job_source.job.prompt_name if latest_job_source.job and latest_job_source.job.prompt_name else "unknown",
|
||||||
)
|
)
|
||||||
|
|
||||||
if latest_job_source.error_detail:
|
if latest_job_source.error_detail:
|
||||||
|
|||||||
@@ -73,10 +73,13 @@ class TestPipelineSuccessFlow:
|
|||||||
image_path,
|
image_path,
|
||||||
*,
|
*,
|
||||||
prompt_name="transcribe_document.md",
|
prompt_name="transcribe_document.md",
|
||||||
|
prompt_text=None,
|
||||||
|
temperature=None,
|
||||||
|
top_p=None,
|
||||||
settings=None,
|
settings=None,
|
||||||
provider=None,
|
provider=None,
|
||||||
) -> TranscriptionResult:
|
) -> TranscriptionResult:
|
||||||
_ = (image_path, prompt_name, settings, provider)
|
_ = (image_path, prompt_name, prompt_text, temperature, top_p, settings, provider)
|
||||||
return TranscriptionResult(
|
return TranscriptionResult(
|
||||||
text="Pipeline transcript",
|
text="Pipeline transcript",
|
||||||
provider="openrouter",
|
provider="openrouter",
|
||||||
@@ -102,10 +105,10 @@ class TestPipelineSuccessFlow:
|
|||||||
assert job is not None
|
assert job is not None
|
||||||
assert job.status == JobStatus.TRANSCRIBED
|
assert job.status == JobStatus.TRANSCRIBED
|
||||||
assert any(job_source.raw_transcription == "Pipeline transcript" for job_source in job.job_sources)
|
assert any(job_source.raw_transcription == "Pipeline transcript" for job_source in job.job_sources)
|
||||||
assert any(job_source.prompt_name == "transcribe_document.md" for job_source in job.job_sources)
|
assert job.prompt_name == "transcribe_document.md"
|
||||||
assert any(job_source.user_prompt is not None for job_source in job.job_sources)
|
assert job.user_prompt is not None
|
||||||
assert any(job_source.temperature == 0.2 for job_source in job.job_sources)
|
assert job.temperature == 0.2
|
||||||
assert any(job_source.top_p == 0.85 for job_source in job.job_sources)
|
assert job.top_p == 0.85
|
||||||
assert any(job_source.ai_metadata == {"finish_reason": "stop", "usage": {"total_tokens": 42}} for job_source in job.job_sources)
|
assert any(job_source.ai_metadata == {"finish_reason": "stop", "usage": {"total_tokens": 42}} for job_source in job.job_sources)
|
||||||
assert any(
|
assert any(
|
||||||
job_source.raw_api_response == {"id": "resp_123", "choices": [{"message": {"content": "Pipeline transcript"}}]}
|
job_source.raw_api_response == {"id": "resp_123", "choices": [{"message": {"content": "Pipeline transcript"}}]}
|
||||||
@@ -142,11 +145,14 @@ class TestPipelineSuccessFlow:
|
|||||||
image_path,
|
image_path,
|
||||||
*,
|
*,
|
||||||
prompt_name="transcribe_document.md",
|
prompt_name="transcribe_document.md",
|
||||||
|
prompt_text=None,
|
||||||
|
temperature=None,
|
||||||
|
top_p=None,
|
||||||
settings=None,
|
settings=None,
|
||||||
provider=None,
|
provider=None,
|
||||||
) -> TranscriptionResult:
|
) -> TranscriptionResult:
|
||||||
page_name = Path(image_path).name
|
page_name = Path(image_path).name
|
||||||
_ = (prompt_name, settings, provider)
|
_ = (prompt_name, prompt_text, temperature, top_p, settings, provider)
|
||||||
return TranscriptionResult(
|
return TranscriptionResult(
|
||||||
text=f"Transcript for {page_name}",
|
text=f"Transcript for {page_name}",
|
||||||
provider="openrouter",
|
provider="openrouter",
|
||||||
@@ -171,6 +177,7 @@ class TestPipelineSuccessFlow:
|
|||||||
assert all(job_source.status == JobSourceStatus.TRANSCRIBED for job_source in job.job_sources)
|
assert all(job_source.status == JobSourceStatus.TRANSCRIBED for job_source in job.job_sources)
|
||||||
assert all(job_source.raw_transcription for job_source in job.job_sources)
|
assert all(job_source.raw_transcription for job_source in job.job_sources)
|
||||||
assert all(job_source.source is not None and job_source.source.raw_transcription for job_source in job.job_sources)
|
assert all(job_source.source is not None and job_source.source.raw_transcription for job_source in job.job_sources)
|
||||||
|
assert job.prompt_name == "transcribe_document.md"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_worker_marks_partial_success_when_some_sources_fail(
|
async def test_worker_marks_partial_success_when_some_sources_fail(
|
||||||
@@ -202,12 +209,15 @@ class TestPipelineSuccessFlow:
|
|||||||
image_path,
|
image_path,
|
||||||
*,
|
*,
|
||||||
prompt_name="transcribe_document.md",
|
prompt_name="transcribe_document.md",
|
||||||
|
prompt_text=None,
|
||||||
|
temperature=None,
|
||||||
|
top_p=None,
|
||||||
settings=None,
|
settings=None,
|
||||||
provider=None,
|
provider=None,
|
||||||
) -> TranscriptionResult:
|
) -> TranscriptionResult:
|
||||||
nonlocal call_count
|
nonlocal call_count
|
||||||
call_count += 1
|
call_count += 1
|
||||||
_ = (prompt_name, settings, provider)
|
_ = (image_path, prompt_name, prompt_text, temperature, top_p, settings, provider)
|
||||||
if call_count == 2:
|
if call_count == 2:
|
||||||
raise RuntimeError("simulated page failure")
|
raise RuntimeError("simulated page failure")
|
||||||
return TranscriptionResult(
|
return TranscriptionResult(
|
||||||
@@ -279,11 +289,14 @@ class TestPipelineSuccessFlow:
|
|||||||
image_path,
|
image_path,
|
||||||
*,
|
*,
|
||||||
prompt_name="transcribe_document.md",
|
prompt_name="transcribe_document.md",
|
||||||
|
prompt_text=None,
|
||||||
|
temperature=None,
|
||||||
|
top_p=None,
|
||||||
settings=None,
|
settings=None,
|
||||||
provider=None,
|
provider=None,
|
||||||
) -> TranscriptionResult:
|
) -> TranscriptionResult:
|
||||||
nonlocal call_count
|
nonlocal call_count
|
||||||
_ = (image_path, prompt_name, settings, provider)
|
_ = (image_path, prompt_name, prompt_text, temperature, top_p, settings, provider)
|
||||||
call_count += 1
|
call_count += 1
|
||||||
return TranscriptionResult(
|
return TranscriptionResult(
|
||||||
text="new transcript",
|
text="new transcript",
|
||||||
@@ -332,10 +345,13 @@ class TestPipelineFailureFlow:
|
|||||||
image_path,
|
image_path,
|
||||||
*,
|
*,
|
||||||
prompt_name="transcribe_document.md",
|
prompt_name="transcribe_document.md",
|
||||||
|
prompt_text=None,
|
||||||
|
temperature=None,
|
||||||
|
top_p=None,
|
||||||
settings=None,
|
settings=None,
|
||||||
provider=None,
|
provider=None,
|
||||||
) -> TranscriptionResult:
|
) -> TranscriptionResult:
|
||||||
_ = (image_path, prompt_name, settings, provider)
|
_ = (image_path, prompt_name, prompt_text, temperature, top_p, settings, provider)
|
||||||
raise RuntimeError("pipeline provider failure")
|
raise RuntimeError("pipeline provider failure")
|
||||||
|
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
|
|||||||
@@ -56,6 +56,8 @@ async def test_create_job_for_document_sorts_uploads_and_creates_links(async_ses
|
|||||||
assert created_job is not None
|
assert created_job is not None
|
||||||
assert created_job.provider == "openrouter"
|
assert created_job.provider == "openrouter"
|
||||||
assert created_job.model == "test-model"
|
assert created_job.model == "test-model"
|
||||||
|
assert created_job.prompt_name == "transcribe_document.md"
|
||||||
|
assert created_job.user_prompt is not None
|
||||||
|
|
||||||
sources = (
|
sources = (
|
||||||
await async_session.exec(
|
await async_session.exec(
|
||||||
@@ -78,7 +80,6 @@ async def test_create_job_for_document_sorts_uploads_and_creates_links(async_ses
|
|||||||
job_sources = (await async_session.exec(select(JobSource).where(JobSource.job_id == result.job_id))).all()
|
job_sources = (await async_session.exec(select(JobSource).where(JobSource.job_id == result.job_id))).all()
|
||||||
assert len(job_sources) == 2
|
assert len(job_sources) == 2
|
||||||
assert set(result.source_ids) == {job_source.source_id for job_source in job_sources}
|
assert set(result.source_ids) == {job_source.source_id for job_source in job_sources}
|
||||||
assert {job_source.prompt_name for job_source in job_sources} == {None}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -110,6 +111,10 @@ async def test_create_upload_job_stores_source_under_document_id_directory(async
|
|||||||
assert source.file_hash == "2c8648d103e3dd7ad87660da0f126a1443b6d21ac1bd3ec000c5e24e2373a90c"
|
assert source.file_hash == "2c8648d103e3dd7ad87660da0f126a1443b6d21ac1bd3ec000c5e24e2373a90c"
|
||||||
assert source.file_size_bytes == len(b"image-bytes")
|
assert source.file_size_bytes == len(b"image-bytes")
|
||||||
|
|
||||||
|
created_job = await async_session.get(Job, result.job_id)
|
||||||
|
assert created_job is not None
|
||||||
|
assert created_job.prompt_name == "transcribe_document.md"
|
||||||
|
|
||||||
|
|
||||||
def test_store_person_portrait_stores_file_under_person_id_directory(tmp_path):
|
def test_store_person_portrait_stores_file_under_person_id_directory(tmp_path):
|
||||||
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
|
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
|
||||||
|
|||||||
@@ -255,12 +255,10 @@ async def test_update_job_source_transcription_persists_provider_json_payloads(d
|
|||||||
raw_api_response=raw_payload,
|
raw_api_response=raw_payload,
|
||||||
provider="openrouter",
|
provider="openrouter",
|
||||||
model="test-model",
|
model="test-model",
|
||||||
prompt_name="transcribe_document.md",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
stored_rows = await transcriptions.list_job_sources(job_id=job.id)
|
stored_rows = await transcriptions.list_job_sources(job_id=job.id)
|
||||||
assert len(stored_rows) == 1
|
assert len(stored_rows) == 1
|
||||||
assert stored_rows[0].raw_transcription == "provider transcript"
|
assert stored_rows[0].raw_transcription == "provider transcript"
|
||||||
assert stored_rows[0].prompt_name == "transcribe_document.md"
|
|
||||||
assert stored_rows[0].ai_metadata == metadata
|
assert stored_rows[0].ai_metadata == metadata
|
||||||
assert stored_rows[0].raw_api_response == raw_payload
|
assert stored_rows[0].raw_api_response == raw_payload
|
||||||
|
|||||||
@@ -65,8 +65,17 @@ class TestWorkflowReliability:
|
|||||||
|
|
||||||
loaded = await services.jobs.read_job(job_id=job.id, session=session)
|
loaded = await services.jobs.read_job(job_id=job.id, session=session)
|
||||||
|
|
||||||
async def _never_returns(image_path, *, prompt_name="transcribe_document.md", settings=None, provider=None):
|
async def _never_returns(
|
||||||
_ = (image_path, prompt_name, settings, provider)
|
image_path,
|
||||||
|
*,
|
||||||
|
prompt_name="transcribe_document.md",
|
||||||
|
prompt_text=None,
|
||||||
|
temperature=None,
|
||||||
|
top_p=None,
|
||||||
|
settings=None,
|
||||||
|
provider=None,
|
||||||
|
):
|
||||||
|
_ = (image_path, prompt_name, prompt_text, temperature, top_p, settings, provider)
|
||||||
raise TimeoutError("simulated provider timeout")
|
raise TimeoutError("simulated provider timeout")
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.services.workflows.transcribe_document_image", _never_returns)
|
monkeypatch.setattr("transcription.services.workflows.transcribe_document_image", _never_returns)
|
||||||
|
|||||||
@@ -95,6 +95,7 @@ async def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., Awai
|
|||||||
retry_count=0,
|
retry_count=0,
|
||||||
provider="openrouter",
|
provider="openrouter",
|
||||||
model="google/gemini-2.5-flash",
|
model="google/gemini-2.5-flash",
|
||||||
|
prompt_name="transcribe_document.md",
|
||||||
)
|
)
|
||||||
session.add(job)
|
session.add(job)
|
||||||
await session.flush()
|
await session.flush()
|
||||||
@@ -121,7 +122,6 @@ async def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., Awai
|
|||||||
if transcription_text is not None
|
if transcription_text is not None
|
||||||
else JobSourceStatus.FAILED
|
else JobSourceStatus.FAILED
|
||||||
),
|
),
|
||||||
prompt_name="transcribe_document.md",
|
|
||||||
raw_transcription=transcription_text,
|
raw_transcription=transcription_text,
|
||||||
error_detail=error_detail,
|
error_detail=error_detail,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user