Continue GC code review: Pydantic

This commit is contained in:
Jim Lancaster
2026-08-12 01:35:50 -05:00
parent 888a8c380a
commit 1e8d8572d4
18 changed files with 578 additions and 289 deletions
+5 -1
View File
@@ -76,6 +76,9 @@ DATABASE__PATH=app.db
SQLITE_CHECK_SAME_THREAD=false SQLITE_CHECK_SAME_THREAD=false
UPLOAD_DIR=./uploads UPLOAD_DIR=./uploads
PROMPT_DIR=./prompts PROMPT_DIR=./prompts
DEFAULT_PROMPT_NAME=transcribe_document.md
# TRANSCRIPTION_TEMPERATURE=0.2 # range: 0.0-2.0
# TRANSCRIPTION_TOP_P=0.9 # range: 0.0-1.0
``` ```
For PostgreSQL: For PostgreSQL:
@@ -138,7 +141,8 @@ Replace `localhost` with the server's hostname or IP address when connecting fro
## Prompt artifacts ## Prompt artifacts
Prompt files are stored in `prompts/` and loaded from `PROMPT_DIR` (default: `./prompts`). Prompt files are stored directly in `PROMPT_DIR` (default: `./prompts`). `DEFAULT_PROMPT_NAME` must be a filename,
not a path. Each job snapshots the validated prompt text, SHA-256 hash, and sampling values for reproducibility.
The canonical MVP prompt is: The canonical MVP prompt is:
- `prompts/transcribe_document.md` - `prompts/transcribe_document.md`
+2
View File
@@ -5,9 +5,11 @@ This directory stores transcription prompts as individual Markdown artifacts.
## Conventions ## Conventions
- Keep one prompt per file. - Keep one prompt per file.
- Use stable, descriptive snake_case file names. - Use stable, descriptive snake_case file names.
- Store prompt files directly in this directory; nested paths are rejected.
- Prefer incremental edits to a single prompt per change for clean history. - Prefer incremental edits to a single prompt per change for clean history.
- Keep prompts human-readable and policy-focused. - Keep prompts human-readable and policy-focused.
- Do not store secrets in prompt files. - Do not store secrets in prompt files.
- Runtime jobs snapshot prompt text, SHA-256 provenance, and sampling configuration.
## Current Prompt ## Current Prompt
- `transcribe_document.md`: baseline verbatim transcription policy for historical documents. - `transcribe_document.md`: baseline verbatim transcription policy for historical documents.
-10
View File
@@ -1,10 +0,0 @@
You are an assistant that may call tools.
Tool safety rules:
1) Tool arguments MUST be strict JSON matching the schema exactly.
2) Never place disallowed, sensitive, explicit, or policy-violating text directly into tool arguments.
3) If user content may be unsafe, first produce a brief neutral summary and pass only that summary.
4) Prefer IDs, enums, booleans, and short fields over raw free-form text.
5) Keep all string arguments <= 300 chars unless schema says otherwise.
6) If you cannot safely provide valid tool args, do not call the tool; respond with "NO_TOOL_CALL" and explain briefly.
7) Never include markdown/code fences in tool arguments.
+76 -37
View File
@@ -2,8 +2,7 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from typing import Annotated
from datetime import datetime
from uuid import UUID from uuid import UUID
from fastapi import APIRouter from fastapi import APIRouter
@@ -11,7 +10,9 @@ from fastapi import Depends
from fastapi import Request from fastapi import Request
from fastapi import Response from fastapi import Response
from pydantic import BaseModel from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field from pydantic import Field
from pydantic import model_validator
from transcription.db.models import Document from transcription.db.models import Document
from transcription.db.models import DocumentPerson from transcription.db.models import DocumentPerson
@@ -23,28 +24,32 @@ from transcription.services import PeopleService
router = APIRouter(prefix="/api/v4", tags=["v4-documents"]) router = APIRouter(prefix="/api/v4", tags=["v4-documents"])
@dataclass(frozen=True) class ApiModel(BaseModel):
class DocumentTypePayload: model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True)
id: UUID
code: str
label: str
is_active: bool
sort_order: int
created_at: datetime
updated_at: datetime
@dataclass(frozen=True) class SelectorRequest(ApiModel):
class PersonRolePayload: @model_validator(mode="after")
id: UUID def require_exactly_one_selector(self):
code: str values = (self.selector_id, self.selector_code)
label: str if sum(value is not None for value in values) != 1:
is_active: bool raise ValueError(f"Provide exactly one of {self.selector_names[0]} or {self.selector_names[1]}")
created_at: datetime return self
updated_at: datetime
@property
def selector_id(self) -> UUID | None:
raise NotImplementedError
@property
def selector_code(self) -> str | None:
raise NotImplementedError
@property
def selector_names(self) -> tuple[str, str]:
raise NotImplementedError
class DocumentTypeRead(BaseModel): class DocumentTypeRead(ApiModel):
id: UUID id: UUID
code: str code: str
label: str label: str
@@ -52,36 +57,66 @@ class DocumentTypeRead(BaseModel):
sort_order: int sort_order: int
class PersonRoleRead(BaseModel): class PersonRoleRead(ApiModel):
id: UUID id: UUID
code: str code: str
label: str label: str
is_active: bool is_active: bool
class DocumentTypeWriteRequest(BaseModel): class DocumentTypeWriteRequest(SelectorRequest):
document_type_id: UUID | None = None document_type_id: UUID | None = None
document_type_code: str | None = None document_type_code: str | None = Field(default=None, min_length=1, pattern=r"^[a-z0-9_]+$")
@property
def selector_id(self) -> UUID | None:
return self.document_type_id
@property
def selector_code(self) -> str | None:
return self.document_type_code
@property
def selector_names(self) -> tuple[str, str]:
return "document_type_id", "document_type_code"
class DocumentTypeWriteResponse(BaseModel): class DocumentTypeWriteResponse(ApiModel):
document_id: UUID document_id: UUID
document_type_id: UUID | None document_type_id: UUID | None
document_type_code: str | None document_type_code: str | None
class DocumentPersonWriteRequest(BaseModel): class DocumentPersonWriteRequest(ApiModel):
person_id: UUID person_id: UUID
role_id: UUID | None = None role_id: UUID | None = None
role_code: str | None = None role_code: str | None = Field(default=None, min_length=1, pattern=r"^[a-z0-9_]+$")
@model_validator(mode="after")
def reject_conflicting_role_selectors(self):
if self.role_id is not None and self.role_code is not None:
raise ValueError("Provide role_id or role_code, not both")
return self
class DocumentPersonRoleUpdateRequest(BaseModel): class DocumentPersonRoleUpdateRequest(SelectorRequest):
role_id: UUID | None = None role_id: UUID | None = None
role_code: str | None = None role_code: str | None = Field(default=None, min_length=1, pattern=r"^[a-z0-9_]+$")
@property
def selector_id(self) -> UUID | None:
return self.role_id
@property
def selector_code(self) -> str | None:
return self.role_code
@property
def selector_names(self) -> tuple[str, str]:
return "role_id", "role_code"
class DocumentPersonRead(BaseModel): class DocumentPersonRead(ApiModel):
id: UUID id: UUID
document_id: UUID document_id: UUID
person_id: UUID person_id: UUID
@@ -90,7 +125,7 @@ class DocumentPersonRead(BaseModel):
person_name: str | None = None person_name: str | None = None
class DocumentPeopleResponse(BaseModel): class DocumentPeopleResponse(ApiModel):
document_id: UUID document_id: UUID
links: list[DocumentPersonRead] = Field(default_factory=list) links: list[DocumentPersonRead] = Field(default_factory=list)
@@ -152,10 +187,14 @@ def get_people_service(request: Request) -> PeopleService:
return PeopleService() return PeopleService()
DocumentServiceDependency = Annotated[DocumentService, Depends(get_document_service)]
PeopleServiceDependency = Annotated[PeopleService, Depends(get_people_service)]
@router.get("/document-types", response_model=list[DocumentTypeRead]) @router.get("/document-types", response_model=list[DocumentTypeRead])
async def list_document_types( async def list_document_types(
service: DocumentServiceDependency,
active_only: bool = True, active_only: bool = True,
service: DocumentService = Depends(get_document_service),
) -> list[DocumentTypeRead]: ) -> list[DocumentTypeRead]:
items = await service.list_document_types(active_only=active_only) items = await service.list_document_types(active_only=active_only)
return [_document_type_to_read(item) for item in items] return [_document_type_to_read(item) for item in items]
@@ -163,8 +202,8 @@ async def list_document_types(
@router.get("/person-roles", response_model=list[PersonRoleRead]) @router.get("/person-roles", response_model=list[PersonRoleRead])
async def list_person_roles( async def list_person_roles(
service: PeopleServiceDependency,
active_only: bool = True, active_only: bool = True,
service: PeopleService = Depends(get_people_service),
) -> list[PersonRoleRead]: ) -> list[PersonRoleRead]:
items = await service.list_person_roles(active_only=active_only) items = await service.list_person_roles(active_only=active_only)
return [_person_role_to_read(item) for item in items] return [_person_role_to_read(item) for item in items]
@@ -174,7 +213,7 @@ async def list_person_roles(
async def set_document_type( async def set_document_type(
document_id: UUID, document_id: UUID,
payload: DocumentTypeWriteRequest, payload: DocumentTypeWriteRequest,
service: DocumentService = Depends(get_document_service), service: DocumentServiceDependency,
) -> DocumentTypeWriteResponse: ) -> DocumentTypeWriteResponse:
document = await service.set_document_type( document = await service.set_document_type(
document_id=document_id, document_id=document_id,
@@ -187,7 +226,7 @@ async def set_document_type(
@router.get("/documents/{document_id}/people", response_model=DocumentPeopleResponse) @router.get("/documents/{document_id}/people", response_model=DocumentPeopleResponse)
async def list_document_people( async def list_document_people(
document_id: UUID, document_id: UUID,
service: PeopleService = Depends(get_people_service), service: PeopleServiceDependency,
) -> DocumentPeopleResponse: ) -> DocumentPeopleResponse:
links = await service.list_document_people(document_id=document_id) links = await service.list_document_people(document_id=document_id)
return DocumentPeopleResponse(document_id=document_id, links=[_document_person_to_read(item) for item in links]) return DocumentPeopleResponse(document_id=document_id, links=[_document_person_to_read(item) for item in links])
@@ -197,7 +236,7 @@ async def list_document_people(
async def add_document_person_link( async def add_document_person_link(
document_id: UUID, document_id: UUID,
payload: DocumentPersonWriteRequest, payload: DocumentPersonWriteRequest,
service: PeopleService = Depends(get_people_service), service: PeopleServiceDependency,
) -> DocumentPersonRead: ) -> DocumentPersonRead:
link = await service.add_document_person_link( link = await service.add_document_person_link(
document_id=document_id, document_id=document_id,
@@ -212,7 +251,7 @@ async def add_document_person_link(
async def set_document_person_role( async def set_document_person_role(
document_person_id: UUID, document_person_id: UUID,
payload: DocumentPersonRoleUpdateRequest, payload: DocumentPersonRoleUpdateRequest,
service: PeopleService = Depends(get_people_service), service: PeopleServiceDependency,
) -> DocumentPersonRead: ) -> DocumentPersonRead:
link = await service.set_document_person_role( link = await service.set_document_person_role(
document_person_id=document_person_id, document_person_id=document_person_id,
@@ -225,7 +264,7 @@ async def set_document_person_role(
@router.delete("/document-people/{document_person_id}", status_code=204) @router.delete("/document-people/{document_person_id}", status_code=204)
async def delete_document_person_link( async def delete_document_person_link(
document_person_id: UUID, document_person_id: UUID,
service: PeopleService = Depends(get_people_service), service: PeopleServiceDependency,
) -> Response: ) -> Response:
await service.remove_document_person_link(document_person_id=document_person_id) await service.remove_document_person_link(document_person_id=document_person_id)
return Response(status_code=204) return Response(status_code=204)
+27 -14
View File
@@ -15,8 +15,10 @@ from typing import Any
from typing import Literal from typing import Literal
from pydantic import BaseModel from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field from pydantic import Field
from pydantic import SecretStr from pydantic import SecretStr
from pydantic import StringConstraints
from pydantic_settings import BaseSettings from pydantic_settings import BaseSettings
from pydantic_settings import SettingsConfigDict from pydantic_settings import SettingsConfigDict
@@ -27,17 +29,27 @@ class Provider(StrEnum):
OPENROUTER = "openrouter" OPENROUTER = "openrouter"
NonEmptyStr = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)]
PromptFilename = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1, pattern=r"^[^/\\]+$")]
Probability = Annotated[float, Field(ge=0.0, le=1.0)]
Temperature = Annotated[float, Field(ge=0.0, le=2.0)]
class SqliteSettings(BaseModel): class SqliteSettings(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
driver: Literal["sqlite"] = "sqlite" driver: Literal["sqlite"] = "sqlite"
path: str = "app.db" path: NonEmptyStr = "app.db"
class PostgresSettings(BaseModel): class PostgresSettings(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
driver: Literal["postgres"] = "postgres" driver: Literal["postgres"] = "postgres"
host: str host: NonEmptyStr
port: int = 5432 port: int = Field(default=5432, ge=1, le=65535)
database: str database: NonEmptyStr
user: str user: NonEmptyStr
password: SecretStr password: SecretStr
@@ -55,6 +67,7 @@ class Settings(BaseSettings):
env_nested_delimiter="__", env_nested_delimiter="__",
cli_implicit_flags=True, cli_implicit_flags=True,
cli_kebab_case=True, cli_kebab_case=True,
frozen=True,
) )
# --- NiceGUI Server --- # --- NiceGUI Server ---
@@ -65,13 +78,13 @@ class Settings(BaseSettings):
# --- AI provider --- # --- AI provider ---
provider: Provider = Provider.OPENROUTER provider: Provider = Provider.OPENROUTER
openrouter_api_key: str openrouter_api_key: SecretStr
provider_model: str | None = None provider_model: NonEmptyStr | None = None
openrouter_http_referer: str | None = None openrouter_http_referer: NonEmptyStr | None = None
openrouter_app_title: str | None = None openrouter_app_title: NonEmptyStr | None = None
default_prompt_name: str = "transcribe_document.md" default_prompt_name: PromptFilename = "transcribe_document.md"
transcription_temperature: float | None = None transcription_temperature: Temperature | None = None
transcription_top_p: float | None = None transcription_top_p: Probability | None = None
# --- runtime environment --- # --- runtime environment ---
environment: Literal["development", "test", "production"] = "development" environment: Literal["development", "test", "production"] = "development"
@@ -86,8 +99,8 @@ class Settings(BaseSettings):
prompt_dir: Path = Path("./prompts") prompt_dir: Path = Path("./prompts")
# --- worker reliability --- # --- worker reliability ---
worker_max_retries: int = 0 worker_max_retries: int = Field(default=0, ge=0)
worker_retry_backoff_seconds: float = 0.0 worker_retry_backoff_seconds: float = Field(default=0.0, ge=0.0)
worker_provider_timeout_seconds: float = Field(default=20.0, gt=0.0, le=20.0) worker_provider_timeout_seconds: float = Field(default=20.0, gt=0.0, le=20.0)
worker_min_transcription_chars: int = Field(default=0, ge=0) worker_min_transcription_chars: int = Field(default=0, ge=0)
worker_min_transcription_lines: int = Field(default=0, ge=0) worker_min_transcription_lines: int = Field(default=0, ge=0)
+32 -17
View File
@@ -4,15 +4,15 @@ from datetime import UTC
from datetime import date from datetime import date
from datetime import datetime from datetime import datetime
from enum import StrEnum from enum import StrEnum
from typing import Any
from typing import Optional from typing import Optional
from uuid import UUID from uuid import UUID
from uuid import uuid4 from uuid import uuid4
from sqlalchemy import Column from pydantic import JsonValue
from sqlalchemy import BigInteger
from sqlalchemy import Enum as SAEnum
from sqlalchemy import JSON from sqlalchemy import JSON
from sqlalchemy import BigInteger
from sqlalchemy import Column
from sqlalchemy import Enum as SAEnum
from sqlalchemy import UniqueConstraint from sqlalchemy import UniqueConstraint
from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm.exc import DetachedInstanceError from sqlalchemy.orm.exc import DetachedInstanceError
@@ -67,7 +67,9 @@ class DocumentType(SQLModel, table=True):
created_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) created_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
updated_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) updated_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
documents: list["Document"] = Relationship(back_populates="document_type_ref", sa_relationship_kwargs={"lazy": "selectin"}) documents: list["Document"] = Relationship(
back_populates="document_type_ref", sa_relationship_kwargs={"lazy": "selectin"}
)
class PersonRole(SQLModel, table=True): class PersonRole(SQLModel, table=True):
@@ -82,7 +84,9 @@ class PersonRole(SQLModel, table=True):
created_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) created_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
updated_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) updated_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
document_people: list["DocumentPerson"] = Relationship(back_populates="role_ref", sa_relationship_kwargs={"lazy": "selectin"}) document_people: list["DocumentPerson"] = Relationship(
back_populates="role_ref", sa_relationship_kwargs={"lazy": "selectin"}
)
class Document(SQLModel, table=True): class Document(SQLModel, table=True):
@@ -102,8 +106,12 @@ class Document(SQLModel, table=True):
jobs: list["Job"] = Relationship(back_populates="document", sa_relationship_kwargs={"lazy": "selectin"}) jobs: list["Job"] = Relationship(back_populates="document", sa_relationship_kwargs={"lazy": "selectin"})
sources: list["Source"] = Relationship(back_populates="document", sa_relationship_kwargs={"lazy": "selectin"}) sources: list["Source"] = Relationship(back_populates="document", sa_relationship_kwargs={"lazy": "selectin"})
document_people: list["DocumentPerson"] = Relationship(back_populates="document", sa_relationship_kwargs={"lazy": "selectin"}) document_people: list["DocumentPerson"] = Relationship(
document_type_ref: Optional["DocumentType"] = Relationship(back_populates="documents", sa_relationship_kwargs={"lazy": "selectin"}) back_populates="document", sa_relationship_kwargs={"lazy": "selectin"}
)
document_type_ref: Optional["DocumentType"] = Relationship(
back_populates="documents", sa_relationship_kwargs={"lazy": "selectin"}
)
class Person(SQLModel, table=True): class Person(SQLModel, table=True):
@@ -121,14 +129,16 @@ class Person(SQLModel, table=True):
death_place: str | None = None death_place: str | None = None
biography: str | None = None biography: str | None = None
portrait_path: str | None = None portrait_path: str | None = None
metadata_: dict[str, Any] | None = Field( metadata_: dict[str, JsonValue] | None = Field(
default=None, default=None,
sa_column=Column("metadata", JSONBCompat(), nullable=True), sa_column=Column("metadata", JSONBCompat(), nullable=True),
) )
created_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) created_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
updated_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) updated_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
document_people: list["DocumentPerson"] = Relationship(back_populates="person", sa_relationship_kwargs={"lazy": "selectin"}) document_people: list["DocumentPerson"] = Relationship(
back_populates="person", sa_relationship_kwargs={"lazy": "selectin"}
)
class DocumentPerson(SQLModel, table=True): class DocumentPerson(SQLModel, table=True):
@@ -159,9 +169,15 @@ class DocumentPerson(SQLModel, table=True):
UniqueConstraint("document_id", "person_id", "role", name="uq_document_person_role"), UniqueConstraint("document_id", "person_id", "role", name="uq_document_person_role"),
) )
document: Optional["Document"] = Relationship(back_populates="document_people", sa_relationship_kwargs={"lazy": "selectin"}) document: Optional["Document"] = Relationship(
person: Optional["Person"] = Relationship(back_populates="document_people", sa_relationship_kwargs={"lazy": "selectin"}) back_populates="document_people", sa_relationship_kwargs={"lazy": "selectin"}
role_ref: Optional["PersonRole"] = Relationship(back_populates="document_people", sa_relationship_kwargs={"lazy": "selectin"}) )
person: Optional["Person"] = Relationship(
back_populates="document_people", sa_relationship_kwargs={"lazy": "selectin"}
)
role_ref: Optional["PersonRole"] = Relationship(
back_populates="document_people", sa_relationship_kwargs={"lazy": "selectin"}
)
class Job(SQLModel, table=True): class Job(SQLModel, table=True):
@@ -278,6 +294,7 @@ class Source(SQLModel, table=True):
"""Return the parent document name if loaded.""" """Return the parent document name if loaded."""
return self.document.name if self.document else None return self.document.name if self.document else None
class JobSource(SQLModel, table=True): class JobSource(SQLModel, table=True):
"""A single AI execution record for one source page.""" """A single AI execution record for one source page."""
@@ -298,12 +315,10 @@ class JobSource(SQLModel, table=True):
), ),
) )
raw_transcription: str | None = None raw_transcription: str | None = None
ai_metadata: dict[str, Any] | None = Field(default=None, sa_column=Column(JSONBCompat(), nullable=True)) ai_metadata: dict[str, JsonValue] | 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, JsonValue] | None = Field(default=None, sa_column=Column(JSONBCompat(), nullable=True))
error_detail: str | None = None error_detail: str | None = None
executed_at: datetime = Field(default_factory=lambda: datetime.now(UTC)) executed_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
job: Optional["Job"] = Relationship(back_populates="job_sources", sa_relationship_kwargs={"lazy": "selectin"}) job: Optional["Job"] = Relationship(back_populates="job_sources", sa_relationship_kwargs={"lazy": "selectin"})
source: Optional["Source"] = Relationship(back_populates="job_sources", sa_relationship_kwargs={"lazy": "selectin"}) source: Optional["Source"] = Relationship(back_populates="job_sources", sa_relationship_kwargs={"lazy": "selectin"})
+2
View File
@@ -6,6 +6,7 @@ from transcription.config import get_settings
from transcription.providers.base import ProviderAuthError from transcription.providers.base import ProviderAuthError
from transcription.providers.base import ProviderError from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError from transcription.providers.base import ProviderResponseError
from transcription.providers.base import TranscriptionMetadata
from transcription.providers.base import TranscriptionProvider from transcription.providers.base import TranscriptionProvider
from transcription.providers.base import TranscriptionResult from transcription.providers.base import TranscriptionResult
from transcription.providers.openrouter import OpenRouterTranscriptionProvider from transcription.providers.openrouter import OpenRouterTranscriptionProvider
@@ -25,6 +26,7 @@ __all__ = [
"ProviderAuthError", "ProviderAuthError",
"ProviderError", "ProviderError",
"ProviderResponseError", "ProviderResponseError",
"TranscriptionMetadata",
"TranscriptionProvider", "TranscriptionProvider",
"TranscriptionResult", "TranscriptionResult",
"get_transcription_provider", "get_transcription_provider",
+59 -17
View File
@@ -1,9 +1,12 @@
"""Provider interfaces and shared types for transcription adapters.""" """Provider interfaces and validated shared contracts for transcription adapters."""
from dataclasses import dataclass
from typing import Any
from typing import Protocol from typing import Protocol
from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field
from pydantic import JsonValue
class ProviderError(RuntimeError): class ProviderError(RuntimeError):
"""Base error for provider failures.""" """Base error for provider failures."""
@@ -17,25 +20,64 @@ class ProviderResponseError(ProviderError):
"""Raised when provider responses are malformed or unusable.""" """Raised when provider responses are malformed or unusable."""
@dataclass(frozen=True) class ProviderUsage(BaseModel):
class TranscriptionResult: """Normalized provider token accounting."""
model_config = ConfigDict(extra="forbid", frozen=True)
input_tokens: int | None = Field(default=None, ge=0)
output_tokens: int | None = Field(default=None, ge=0)
total_tokens: int | None = Field(default=None, ge=0)
class TranscriptionMetadata(BaseModel):
"""Stable structured metadata persisted for one provider execution."""
model_config = ConfigDict(extra="forbid", frozen=True)
finish_reason: str | None = Field(default=None, min_length=1)
usage: ProviderUsage | None = None
def as_json_object(self) -> dict[str, JsonValue] | None:
payload = self.model_dump(mode="json", exclude_none=True)
return payload or None
class TranscriptionResult(BaseModel):
"""Normalized output returned by any transcription provider.""" """Normalized output returned by any transcription provider."""
text: str model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True)
provider: str
model: str text: str = Field(min_length=1)
provider: str = Field(min_length=1)
model: str = Field(min_length=1)
prompt_name: str | None = None prompt_name: str | None = None
prompt_hash: str | None = None prompt_hash: str | None = Field(default=None, pattern=r"^[0-9a-f]{64}$")
system_prompt: str | None = None system_prompt: str | None = None
user_prompt: str | None = None user_prompt: str | None = None
temperature: float | None = None temperature: float | None = Field(default=None, ge=0.0, le=2.0)
top_p: float | None = None top_p: float | None = Field(default=None, ge=0.0, le=1.0)
finish_reason: str | None = None metadata: TranscriptionMetadata = Field(default_factory=TranscriptionMetadata)
usage_input_tokens: int | None = None raw_api_response: dict[str, JsonValue] | None = None
usage_output_tokens: int | None = None
usage_total_tokens: int | None = None @property
ai_metadata: dict[str, Any] | None = None def finish_reason(self) -> str | None:
raw_api_response: dict[str, Any] | None = None return self.metadata.finish_reason
@property
def usage_input_tokens(self) -> int | None:
return self.metadata.usage.input_tokens if self.metadata.usage else None
@property
def usage_output_tokens(self) -> int | None:
return self.metadata.usage.output_tokens if self.metadata.usage else None
@property
def usage_total_tokens(self) -> int | None:
return self.metadata.usage.total_tokens if self.metadata.usage else None
def metadata_payload(self) -> dict[str, JsonValue] | None:
return self.metadata.as_json_object()
class TranscriptionProvider(Protocol): class TranscriptionProvider(Protocol):
+142 -152
View File
@@ -4,18 +4,25 @@ from __future__ import annotations
import base64 import base64
import logging import logging
from dataclasses import dataclass from typing import Annotated
from typing import Any from typing import Any
from typing import cast from typing import Literal
from openrouter import OpenRouter from openrouter import OpenRouter
from openrouter.components.chatmessages import ChatMessagesTypedDict from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field
from pydantic import JsonValue
from pydantic import TypeAdapter
from pydantic import ValidationError
from transcription.config import Settings from transcription.config import Settings
from transcription.config import get_settings from transcription.config import get_settings
from transcription.providers.base import ProviderAuthError from transcription.providers.base import ProviderAuthError
from transcription.providers.base import ProviderError from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError from transcription.providers.base import ProviderResponseError
from transcription.providers.base import ProviderUsage
from transcription.providers.base import TranscriptionMetadata
from transcription.providers.base import TranscriptionResult from transcription.providers.base import TranscriptionResult
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -23,16 +30,92 @@ logger = logging.getLogger(__name__)
DEFAULT_OPENROUTER_MODEL = "google/gemini-2.5-flash" DEFAULT_OPENROUTER_MODEL = "google/gemini-2.5-flash"
@dataclass(frozen=True) class _ProviderModel(BaseModel):
class OpenRouterRequest: model_config = ConfigDict(extra="forbid", frozen=True)
class TextContent(_ProviderModel):
type: Literal["text"] = "text"
text: str = Field(min_length=1)
class ImageUrl(_ProviderModel):
url: str = Field(min_length=1)
class ImageContent(_ProviderModel):
type: Literal["image_url"] = "image_url"
image_url: ImageUrl
class FileData(_ProviderModel):
filename: str = Field(min_length=1)
file_data: str = Field(min_length=1)
class FileContent(_ProviderModel):
type: Literal["file"] = "file"
file: FileData
MessageContent = Annotated[TextContent | ImageContent | FileContent, Field(discriminator="type")]
class UserMessage(_ProviderModel):
role: Literal["user"] = "user"
content: tuple[MessageContent, ...] = Field(min_length=2)
class OpenRouterRequest(_ProviderModel):
"""Normalized request payload fields for OpenRouter calls.""" """Normalized request payload fields for OpenRouter calls."""
model: str model: str = Field(min_length=1)
messages: list[dict[str, Any]] messages: tuple[UserMessage, ...] = Field(min_length=1)
http_referer: str | None http_referer: str | None
x_open_router_title: str | None x_open_router_title: str | None
temperature: float | None temperature: float | None = Field(ge=0.0, le=2.0)
top_p: float | None top_p: float | None = Field(ge=0.0, le=1.0)
class ResponseContentPart(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
text: str | None = None
class ResponseMessage(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
content: str | tuple[ResponseContentPart, ...] | None = None
class ResponseChoice(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
message: ResponseMessage
finish_reason: str | None = None
class ResponseUsage(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
prompt_tokens: int | None = Field(default=None, ge=0)
completion_tokens: int | None = Field(default=None, ge=0)
total_tokens: int | None = Field(default=None, ge=0)
input_tokens: int | None = Field(default=None, ge=0)
output_tokens: int | None = Field(default=None, ge=0)
total: int | None = Field(default=None, ge=0)
class OpenRouterResponse(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
model: str | None = None
choices: tuple[ResponseChoice, ...] = Field(min_length=1)
usage: dict[str, JsonValue] | None = None
JSON_OBJECT_ADAPTER = TypeAdapter(dict[str, JsonValue])
class OpenRouterTranscriptionProvider: class OpenRouterTranscriptionProvider:
@@ -41,7 +124,7 @@ class OpenRouterTranscriptionProvider:
def __init__(self, *, settings: Settings | None = None, client: OpenRouter | None = None): def __init__(self, *, settings: Settings | None = None, client: OpenRouter | None = None):
self._settings = settings or get_settings() self._settings = settings or get_settings()
self._model = self._settings.provider_model or DEFAULT_OPENROUTER_MODEL self._model = self._settings.provider_model or DEFAULT_OPENROUTER_MODEL
self._client = client or OpenRouter(api_key=self._settings.openrouter_api_key) self._client = client or OpenRouter(api_key=self._settings.openrouter_api_key.get_secret_value())
@property @property
def model(self) -> str: def model(self) -> str:
@@ -66,31 +149,22 @@ class OpenRouterTranscriptionProvider:
top_p=top_p, top_p=top_p,
) )
try: try:
response = await self._client.chat.send_async( response = await self._client.chat.send_async(**request.model_dump(mode="json", exclude_none=True))
messages=cast(list[ChatMessagesTypedDict], request.messages),
model=request.model,
http_referer=request.http_referer,
x_open_router_title=request.x_open_router_title,
temperature=request.temperature,
top_p=request.top_p,
)
except Exception as exc: except Exception as exc:
message = str(exc).lower() message = str(exc).lower()
if "401" in message or "auth" in message or "api key" in message: if "401" in message or "auth" in message or "api key" in message:
raise ProviderAuthError("OpenRouter authentication failed") from exc raise ProviderAuthError("OpenRouter authentication failed") from exc
raise ProviderError("OpenRouter request failed") from exc raise ProviderError("OpenRouter request failed") from exc
text = self._extract_text(response)
model = self._get_optional_attr(response, "model") or self.model
finish_reason = self._extract_finish_reason(response)
usage_input_tokens, usage_output_tokens, usage_total_tokens = self._extract_usage(response)
ai_metadata = self._build_ai_metadata(
finish_reason=finish_reason,
usage_input_tokens=usage_input_tokens,
usage_output_tokens=usage_output_tokens,
usage_total_tokens=usage_total_tokens,
)
raw_api_response = self._coerce_raw_response(response) raw_api_response = self._coerce_raw_response(response)
try:
validated_response = OpenRouterResponse.model_validate(raw_api_response)
except ValidationError as exc:
raise ProviderResponseError("OpenRouter response failed schema validation") from exc
text = self._extract_text(validated_response)
model = validated_response.model or self.model
metadata = self._build_metadata(validated_response)
logger.info("OpenRouter transcription completed using model=%s", model) logger.info("OpenRouter transcription completed using model=%s", model)
return TranscriptionResult( return TranscriptionResult(
text=text, text=text,
@@ -102,46 +176,37 @@ class OpenRouterTranscriptionProvider:
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
model=model, model=model,
finish_reason=finish_reason, metadata=metadata,
usage_input_tokens=usage_input_tokens,
usage_output_tokens=usage_output_tokens,
usage_total_tokens=usage_total_tokens,
ai_metadata=ai_metadata,
raw_api_response=raw_api_response, raw_api_response=raw_api_response,
) )
def _build_ai_metadata( def _build_metadata(self, response: OpenRouterResponse) -> TranscriptionMetadata:
self, choice = response.choices[0]
*, finish_reason = choice.finish_reason.strip() if choice.finish_reason and choice.finish_reason.strip() else None
finish_reason: str | None, normalized_usage = None
usage_input_tokens: int | None, if response.usage is not None:
usage_output_tokens: int | None, try:
usage_total_tokens: int | None, usage = ResponseUsage.model_validate(response.usage)
) -> dict[str, Any] | None: except ValidationError as exc:
metadata: dict[str, Any] = {} logger.warning("Ignoring invalid OpenRouter usage metadata: %s", exc)
if finish_reason is not None: else:
metadata["finish_reason"] = finish_reason normalized_usage = ProviderUsage(
input_tokens=usage.prompt_tokens if usage.prompt_tokens is not None else usage.input_tokens,
output_tokens=usage.completion_tokens
if usage.completion_tokens is not None
else usage.output_tokens,
total_tokens=usage.total_tokens if usage.total_tokens is not None else usage.total,
)
if normalized_usage.model_dump(exclude_none=True) == {}:
normalized_usage = None
return TranscriptionMetadata(finish_reason=finish_reason, usage=normalized_usage)
usage: dict[str, int] = {} def _coerce_raw_response(self, response: Any) -> dict[str, JsonValue]:
if usage_input_tokens is not None:
usage["input_tokens"] = usage_input_tokens
if usage_output_tokens is not None:
usage["output_tokens"] = usage_output_tokens
if usage_total_tokens is not None:
usage["total_tokens"] = usage_total_tokens
if usage:
metadata["usage"] = usage
return metadata or None
def _coerce_raw_response(self, response: Any) -> dict[str, Any] | None:
payload = self._to_json_compatible(response) payload = self._to_json_compatible(response)
if payload is None: try:
return None return JSON_OBJECT_ADAPTER.validate_python(payload)
if isinstance(payload, dict): except ValidationError as exc:
return payload raise ProviderResponseError("OpenRouter response is not a JSON object") from exc
return {"response": payload}
def _to_json_compatible(self, value: Any) -> Any: def _to_json_compatible(self, value: Any) -> Any:
if value is None or isinstance(value, str | int | float | bool): if value is None or isinstance(value, str | int | float | bool):
@@ -153,12 +218,13 @@ class OpenRouterTranscriptionProvider:
if isinstance(value, list | tuple | set): if isinstance(value, list | tuple | set):
return [self._to_json_compatible(item) for item in value] return [self._to_json_compatible(item) for item in value]
for method_name in ("model_dump", "dict", "to_dict"): for method_name in ("model_dump", "to_dict"):
serializer = getattr(value, method_name, None) serializer = getattr(value, method_name, None)
if callable(serializer): if callable(serializer):
try: try:
return self._to_json_compatible(serializer()) serialized = serializer(mode="json") if method_name == "model_dump" else serializer()
except Exception as exc: # noqa: BLE001 return self._to_json_compatible(serialized)
except (TypeError, ValueError) as exc:
logger.debug("OpenRouter response serializer %s failed: %s", method_name, exc) logger.debug("OpenRouter response serializer %s failed: %s", method_name, exc)
continue continue
@@ -170,7 +236,7 @@ class OpenRouterTranscriptionProvider:
if not str(key).startswith("_") if not str(key).startswith("_")
} }
return repr(value) raise ProviderResponseError(f"OpenRouter response contains unsupported value type: {type(value).__name__}")
def _build_request( def _build_request(
self, self,
@@ -183,110 +249,34 @@ class OpenRouterTranscriptionProvider:
) -> OpenRouterRequest: ) -> OpenRouterRequest:
image_b64 = base64.b64encode(image_bytes).decode("ascii") image_b64 = base64.b64encode(image_bytes).decode("ascii")
data_url = f"data:{mime_type};base64,{image_b64}" data_url = f"data:{mime_type};base64,{image_b64}"
media_content: dict[str, Any] media_content: ImageContent | FileContent
if mime_type == "application/pdf": if mime_type == "application/pdf":
media_content = { media_content = FileContent(file=FileData(filename="source.pdf", file_data=data_url))
"type": "file",
"file": {
"filename": "source.pdf",
"file_data": data_url,
},
}
else: else:
media_content = { media_content = ImageContent(image_url=ImageUrl(url=data_url))
"type": "image_url",
"image_url": {"url": data_url},
}
messages: list[dict[str, Any]] = [
{
"role": "user",
"content": [
{"type": "text", "text": prompt_text},
media_content,
],
}
]
return OpenRouterRequest( return OpenRouterRequest(
model=self.model, model=self.model,
messages=messages, messages=(UserMessage(content=(TextContent(text=prompt_text), media_content)),),
http_referer=self._settings.openrouter_http_referer, http_referer=self._settings.openrouter_http_referer,
x_open_router_title=self._settings.openrouter_app_title, x_open_router_title=self._settings.openrouter_app_title,
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
) )
def _extract_text(self, response: Any) -> str: def _extract_text(self, response: OpenRouterResponse) -> str:
choices = self._get_optional_attr(response, "choices") content = response.choices[0].message.content
if not choices:
raise ProviderResponseError("OpenRouter response missing choices")
first_choice = choices[0]
message = self._get_optional_attr(first_choice, "message")
if message is None:
raise ProviderResponseError("OpenRouter response missing assistant message")
content = self._get_optional_attr(message, "content")
text = self._normalize_content(content) text = self._normalize_content(content)
if not text: if not text:
raise ProviderResponseError("OpenRouter response contained no transcription text") raise ProviderResponseError("OpenRouter response contained no transcription text")
return text return text
def _extract_finish_reason(self, response: Any) -> str | None: def _normalize_content(self, content: str | tuple[ResponseContentPart, ...] | None) -> str:
choices = self._get_optional_attr(response, "choices")
if not choices:
return None
first_choice = choices[0]
finish_reason = self._get_optional_attr(first_choice, "finish_reason")
if isinstance(finish_reason, str) and finish_reason.strip():
return finish_reason.strip()
return None
def _extract_usage(self, response: Any) -> tuple[int | None, int | None, int | None]:
usage = self._get_optional_attr(response, "usage")
if usage is None:
return None, None, None
input_tokens = self._as_int(self._get_optional_attr(usage, "prompt_tokens"))
output_tokens = self._as_int(self._get_optional_attr(usage, "completion_tokens"))
total_tokens = self._as_int(self._get_optional_attr(usage, "total_tokens"))
if input_tokens is None:
input_tokens = self._as_int(self._get_optional_attr(usage, "input_tokens"))
if output_tokens is None:
output_tokens = self._as_int(self._get_optional_attr(usage, "output_tokens"))
if total_tokens is None:
total_tokens = self._as_int(self._get_optional_attr(usage, "total"))
return input_tokens, output_tokens, total_tokens
def _normalize_content(self, content: Any) -> str:
if isinstance(content, str): if isinstance(content, str):
return content.strip() return content.strip()
if isinstance(content, list): if isinstance(content, tuple):
parts: list[str] = [] parts = [item.text.strip() for item in content if item.text and item.text.strip()]
for item in content:
text_part = None
text_part = item.get("text") if isinstance(item, dict) else self._get_optional_attr(item, "text")
if isinstance(text_part, str) and text_part.strip():
parts.append(text_part.strip())
return "\n".join(parts).strip() return "\n".join(parts).strip()
return "" return ""
@staticmethod
def _get_optional_attr(obj: Any, key: str) -> Any:
if obj is None:
return None
if isinstance(obj, dict):
return obj.get(key)
return getattr(obj, key, None)
@staticmethod
def _as_int(value: Any) -> int | None:
if isinstance(value, int):
return value
return None
+75 -22
View File
@@ -6,12 +6,17 @@ import hashlib
import logging import logging
from collections.abc import Sequence from collections.abc import Sequence
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass
from datetime import UTC from datetime import UTC
from datetime import datetime from datetime import datetime
from pathlib import Path from pathlib import Path
from uuid import UUID from uuid import UUID
from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field
from pydantic import JsonValue
from pydantic import TypeAdapter
from pydantic import ValidationError
from sqlalchemy.ext.asyncio import async_sessionmaker from sqlalchemy.ext.asyncio import async_sessionmaker
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
from sqlmodel import select from sqlmodel import select
@@ -28,6 +33,7 @@ from transcription.errors import ErrorCategory
from transcription.providers import ProviderAuthError from transcription.providers import ProviderAuthError
from transcription.providers import ProviderError from transcription.providers import ProviderError
from transcription.providers import ProviderResponseError from transcription.providers import ProviderResponseError
from transcription.providers import TranscriptionMetadata
from transcription.providers import TranscriptionProvider from transcription.providers import TranscriptionProvider
from transcription.providers import TranscriptionResult from transcription.providers import TranscriptionResult
from transcription.providers import get_transcription_provider from transcription.providers import get_transcription_provider
@@ -46,18 +52,20 @@ SOURCE_MIME_TYPES = {
".pdf": "application/pdf", ".pdf": "application/pdf",
} }
SOURCE_EXTENSIONS = frozenset(SOURCE_MIME_TYPES) SOURCE_EXTENSIONS = frozenset(SOURCE_MIME_TYPES)
JSON_OBJECT_ADAPTER = TypeAdapter(dict[str, JsonValue])
@dataclass(frozen=True) class PromptExecution(BaseModel):
class PromptExecution:
"""Resolved prompt inputs captured for one page execution.""" """Resolved prompt inputs captured for one page execution."""
prompt_name: str model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True)
prompt_hash: str
prompt_name: str = Field(min_length=1, pattern=r"^[^/\\]+$")
prompt_hash: str = Field(pattern=r"^[0-9a-f]{64}$")
system_prompt: str | None system_prompt: str | None
user_prompt: str user_prompt: str = Field(min_length=1)
temperature: float | None temperature: float | None = Field(ge=0.0, le=2.0)
top_p: float | None top_p: float | None = Field(ge=0.0, le=1.0)
class PromptLoadError(AppError): class PromptLoadError(AppError):
@@ -366,8 +374,8 @@ class SourceService(ServiceBase):
source_id: UUID, source_id: UUID,
text: str | None, text: str | None,
error_detail: str | None = None, error_detail: str | None = None,
ai_metadata: dict[str, object] | None = None, ai_metadata: TranscriptionMetadata | dict[str, JsonValue] | None = None,
raw_api_response: dict[str, object] | None = None, raw_api_response: dict[str, JsonValue] | None = None,
provider: str | None = None, provider: str | None = None,
model: str | None = None, model: str | None = None,
session: AsyncSession | None = None, session: AsyncSession | None = None,
@@ -400,6 +408,8 @@ class SourceService(ServiceBase):
job.provider = provider or job.provider or self.settings.provider.value job.provider = provider or job.provider or self.settings.provider.value
job.model = model or job.model or _resolve_transcript_model(provider=self.provider, settings=self.settings) job.model = model or job.model or _resolve_transcript_model(provider=self.provider, settings=self.settings)
job.date_updated = datetime.now(UTC) job.date_updated = datetime.now(UTC)
metadata_payload = _validate_transcription_metadata(ai_metadata)
raw_response_payload = _validate_json_object(raw_api_response, field_name="raw_api_response")
if text is not None: if text is not None:
source.raw_transcription = text source.raw_transcription = text
@@ -414,15 +424,15 @@ class SourceService(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,
ai_metadata=ai_metadata, ai_metadata=metadata_payload,
raw_api_response=raw_api_response, raw_api_response=raw_response_payload,
error_detail=error_detail, error_detail=error_detail,
) )
_session.add(job_source) _session.add(job_source)
else: else:
job_source.raw_transcription = text job_source.raw_transcription = text
job_source.ai_metadata = ai_metadata job_source.ai_metadata = metadata_payload
job_source.raw_api_response = raw_api_response job_source.raw_api_response = raw_response_payload
job_source.error_detail = error_detail job_source.error_detail = error_detail
job_source.status = JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED job_source.status = JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED
job_source.executed_at = datetime.now(UTC) job_source.executed_at = datetime.now(UTC)
@@ -492,6 +502,46 @@ def _resolve_transcript_model(*, provider: TranscriptionProvider, settings: Sett
return "unknown" return "unknown"
def _validate_transcription_metadata(
metadata: TranscriptionMetadata | dict[str, JsonValue] | None,
) -> dict[str, JsonValue] | None:
if metadata is None:
return None
try:
validated = (
metadata if isinstance(metadata, TranscriptionMetadata) else TranscriptionMetadata.model_validate(metadata)
)
except ValidationError as exc:
raise TranscriptionError(
"Transcription metadata failed validation",
category=ErrorCategory.VALIDATION,
suggestion="Persist only normalized provider execution metadata.",
) from exc
return validated.as_json_object()
def _validate_json_object(
payload: dict[str, JsonValue] | None,
*,
field_name: str,
) -> dict[str, JsonValue] | None:
if payload is None:
return None
try:
return JSON_OBJECT_ADAPTER.validate_python(payload)
except ValidationError as exc:
raise TranscriptionError(
f"{field_name} must be a JSON-compatible object",
category=ErrorCategory.VALIDATION,
suggestion="Remove non-JSON values before persisting provider diagnostics.",
) from exc
def hash_prompt_text(prompt_text: str) -> str:
"""Return the canonical SHA-256 provenance hash for prompt text."""
return hashlib.sha256(prompt_text.encode("utf-8")).hexdigest()
async def transcribe_document_image( async def transcribe_document_image(
image_path: str | Path, image_path: str | Path,
*, *,
@@ -510,7 +560,7 @@ async def transcribe_document_image(
effective_prompt_name = (prompt_name or runtime_settings.default_prompt_name or DEFAULT_PROMPT_FILE).strip() effective_prompt_name = (prompt_name or runtime_settings.default_prompt_name or DEFAULT_PROMPT_FILE).strip()
prompt_execution = PromptExecution( prompt_execution = PromptExecution(
prompt_name=effective_prompt_name, prompt_name=effective_prompt_name,
prompt_hash=hashlib.sha256(prompt_text.encode("utf-8")).hexdigest(), prompt_hash=hash_prompt_text(prompt_text),
system_prompt=None, system_prompt=None,
user_prompt=prompt_text, user_prompt=prompt_text,
temperature=temperature if temperature is not None else runtime_settings.transcription_temperature, temperature=temperature if temperature is not None else runtime_settings.transcription_temperature,
@@ -540,11 +590,7 @@ async def transcribe_document_image(
temperature=result.temperature if result.temperature is not None else prompt_execution.temperature, temperature=result.temperature if result.temperature is not None else prompt_execution.temperature,
top_p=result.top_p if result.top_p is not None else prompt_execution.top_p, top_p=result.top_p if result.top_p is not None else prompt_execution.top_p,
model=result.model, model=result.model,
finish_reason=result.finish_reason, metadata=result.metadata,
usage_input_tokens=result.usage_input_tokens,
usage_output_tokens=result.usage_output_tokens,
usage_total_tokens=result.usage_total_tokens,
ai_metadata=result.ai_metadata,
raw_api_response=result.raw_api_response, raw_api_response=result.raw_api_response,
) )
@@ -556,7 +602,7 @@ def build_prompt_execution(*, prompt_name: str | None = None, settings: Settings
user_prompt = load_prompt_text(prompt_name=effective_prompt_name, settings=runtime_settings) user_prompt = load_prompt_text(prompt_name=effective_prompt_name, settings=runtime_settings)
return PromptExecution( return PromptExecution(
prompt_name=effective_prompt_name, prompt_name=effective_prompt_name,
prompt_hash=hashlib.sha256(user_prompt.encode("utf-8")).hexdigest(), prompt_hash=hash_prompt_text(user_prompt),
system_prompt=None, system_prompt=None,
user_prompt=user_prompt, user_prompt=user_prompt,
temperature=runtime_settings.transcription_temperature, temperature=runtime_settings.transcription_temperature,
@@ -567,7 +613,14 @@ def build_prompt_execution(*, prompt_name: str | None = None, settings: Settings
def load_prompt_text(*, prompt_name: str = DEFAULT_PROMPT_FILE, settings: Settings | None = None) -> str: def load_prompt_text(*, prompt_name: str = DEFAULT_PROMPT_FILE, settings: Settings | None = None) -> str:
"""Load and validate prompt text from PROMPT_DIR.""" """Load and validate prompt text from PROMPT_DIR."""
runtime_settings = settings or get_settings() runtime_settings = settings or get_settings()
prompt_path = runtime_settings.prompt_dir / prompt_name prompt_root = runtime_settings.prompt_dir.resolve()
prompt_path = (prompt_root / prompt_name).resolve()
if prompt_path.parent != prompt_root:
raise PromptLoadError(
f"Prompt file must be directly inside PROMPT_DIR: {prompt_name}",
category=ErrorCategory.VALIDATION,
suggestion="Configure a prompt filename without directory components.",
)
if not prompt_path.exists() or not prompt_path.is_file(): if not prompt_path.exists() or not prompt_path.is_file():
raise PromptLoadError( raise PromptLoadError(
+8 -5
View File
@@ -17,6 +17,7 @@ from ..providers import TranscriptionResult
from . import ServiceBundle from . import ServiceBundle
from .sources import PromptExecution from .sources import PromptExecution
from .sources import build_prompt_execution from .sources import build_prompt_execution
from .sources import hash_prompt_text
from .sources import transcribe_document_image from .sources import transcribe_document_image
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -141,7 +142,8 @@ async def process_queued_job(
) )
failed_pages.append((source, error)) failed_pages.append((source, error))
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,
job.document_id, job.document_id,
source.id, source.id,
@@ -157,7 +159,8 @@ async def process_queued_job(
failed_pages.append((source, error)) failed_pages.append((source, error))
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,
job.document_id, job.document_id,
source.id, source.id,
@@ -230,7 +233,7 @@ def _resolve_job_prompt_execution(*, source_job: Job, settings: Settings) -> Pro
if source_job.user_prompt and source_job.prompt_name: if source_job.user_prompt and source_job.prompt_name:
return PromptExecution( return PromptExecution(
prompt_name=source_job.prompt_name, prompt_name=source_job.prompt_name,
prompt_hash=source_job.prompt_hash or "", prompt_hash=source_job.prompt_hash or hash_prompt_text(source_job.user_prompt),
system_prompt=source_job.system_prompt, system_prompt=source_job.system_prompt,
user_prompt=source_job.user_prompt, user_prompt=source_job.user_prompt,
temperature=source_job.temperature, temperature=source_job.temperature,
@@ -269,7 +272,7 @@ async def _finalize_batch_outcome(
source_id=source.id, source_id=source.id,
text=result.text, text=result.text,
error_detail=None, error_detail=None,
ai_metadata=result.ai_metadata, ai_metadata=result.metadata_payload(),
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,
@@ -295,7 +298,7 @@ async def _finalize_batch_outcome(
source_id=source.id, source_id=source.id,
text=result.text, text=result.text,
error_detail=None, error_detail=None,
ai_metadata=result.ai_metadata, ai_metadata=result.metadata_payload(),
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,
+32
View File
@@ -117,6 +117,25 @@ def test_set_document_type_by_code_updates_canonical_fields(tmp_path):
assert payload["document_type_code"] == "record" assert payload["document_type_code"] == "record"
def test_document_type_payload_requires_exactly_one_selector(tmp_path):
with _v4_api_client(tmp_path, db_filename="api-doc-type-validation.db") as (client, db_url):
document_id, _ = _seed_document_and_person(db_url=db_url)
missing = client.put(f"/api/v4/documents/{document_id}/type", json={})
conflicting = client.put(
f"/api/v4/documents/{document_id}/type",
json={"document_type_id": str(UUID(int=1)), "document_type_code": "record"},
)
unexpected = client.put(
f"/api/v4/documents/{document_id}/type",
json={"document_type_code": "record", "ignored": True},
)
assert missing.status_code == 422
assert conflicting.status_code == 422
assert unexpected.status_code == 422
def test_document_people_role_aware_write_read_and_delete(tmp_path): def test_document_people_role_aware_write_read_and_delete(tmp_path):
with _v4_api_client(tmp_path, db_filename="api-links.db") as (client, db_url): with _v4_api_client(tmp_path, db_filename="api-links.db") as (client, db_url):
document_id, person_id = _seed_document_and_person(db_url=db_url) document_id, person_id = _seed_document_and_person(db_url=db_url)
@@ -155,6 +174,19 @@ def test_document_people_role_aware_write_read_and_delete(tmp_path):
assert list_after_delete.json()["links"] == [] assert list_after_delete.json()["links"] == []
def test_document_person_link_defaults_to_author_when_role_is_omitted(tmp_path):
with _v4_api_client(tmp_path, db_filename="api-default-role.db") as (client, db_url):
document_id, person_id = _seed_document_and_person(db_url=db_url)
response = client.post(
f"/api/v4/documents/{document_id}/people",
json={"person_id": str(person_id)},
)
assert response.status_code == 200
assert response.json()["role_code"] == "author"
def test_duplicate_document_person_link_returns_conflict_envelope(tmp_path): def test_duplicate_document_person_link_returns_conflict_envelope(tmp_path):
with _v4_api_client(tmp_path, db_filename="api-dup.db") as (client, db_url): with _v4_api_client(tmp_path, db_filename="api-dup.db") as (client, db_url):
document_id, person_id = _seed_document_and_person(db_url=db_url) document_id, person_id = _seed_document_and_person(db_url=db_url)
+6 -1
View File
@@ -9,6 +9,8 @@ from transcription.config import Settings
from transcription.db.models import Document from transcription.db.models import Document
from transcription.db.models import JobSourceStatus from transcription.db.models import JobSourceStatus
from transcription.db.models import JobStatus from transcription.db.models import JobStatus
from transcription.providers.base import ProviderUsage
from transcription.providers.base import TranscriptionMetadata
from transcription.providers.base import TranscriptionResult from transcription.providers.base import TranscriptionResult
from transcription.services import ServiceBundle from transcription.services import ServiceBundle
from transcription.services.store import create_document_job from transcription.services.store import create_document_job
@@ -84,7 +86,10 @@ class TestPipelineSuccessFlow:
provider="openrouter", provider="openrouter",
model="test-model", model="test-model",
prompt_name="transcribe_document.md", prompt_name="transcribe_document.md",
ai_metadata={"finish_reason": "stop", "usage": {"total_tokens": 42}}, metadata=TranscriptionMetadata(
finish_reason="stop",
usage=ProviderUsage(total_tokens=42),
),
raw_api_response={"id": "resp_123", "choices": [{"message": {"content": "Pipeline transcript"}}]}, raw_api_response={"id": "resp_123", "choices": [{"message": {"content": "Pipeline transcript"}}]},
) )
+23 -1
View File
@@ -122,7 +122,7 @@ class TestOpenRouterProviderTranscribe:
assert result.provider == "openrouter" assert result.provider == "openrouter"
assert result.prompt_name is None assert result.prompt_name is None
assert result.model == "vendor/model-b" assert result.model == "vendor/model-b"
assert result.ai_metadata == { assert result.metadata_payload() == {
"finish_reason": "stop", "finish_reason": "stop",
"usage": {"input_tokens": 10, "output_tokens": 25, "total_tokens": 35}, "usage": {"input_tokens": 10, "output_tokens": 25, "total_tokens": 35},
} }
@@ -178,3 +178,25 @@ class TestOpenRouterProviderTranscribe:
image_bytes=b"img-bytes", image_bytes=b"img-bytes",
mime_type="image/png", mime_type="image/png",
) )
@pytest.mark.asyncio
async def test_ignores_invalid_token_metadata_without_discarding_transcript(self):
response = {
"model": "vendor/model-a",
"choices": [{"message": {"content": "Transcript text"}}],
"usage": {"prompt_tokens": -1},
}
provider = OpenRouterTranscriptionProvider(
settings=Settings(openrouter_api_key="test-key"),
client=_FakeClient(response=response),
)
result = await provider.transcribe(
prompt_text="Prompt body",
image_bytes=b"img-bytes",
mime_type="image/png",
)
assert result.text == "Transcript text"
assert result.metadata_payload() is None
assert result.raw_api_response == response
+7 -6
View File
@@ -7,6 +7,7 @@ from uuid import uuid4
import pytest import pytest
from sqlmodel import select from sqlmodel import select
from transcription.config import Settings
from transcription.db.models import Document from transcription.db.models import Document
from transcription.db.models import DocumentPerson from transcription.db.models import DocumentPerson
from transcription.db.models import DocumentPersonRole from transcription.db.models import DocumentPersonRole
@@ -98,8 +99,8 @@ async def test_delete_document_blocks_when_dependencies_exist(default_session_fa
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_document_succeeds_when_unlinked(default_session_factory, tmp_path): async def test_delete_document_succeeds_when_unlinked(default_session_factory, tmp_path):
service = DocumentService(session_factory=default_session_factory) settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
service.settings.upload_dir = tmp_path service = DocumentService(session_factory=default_session_factory, settings=settings)
document = await service.create_document( document = await service.create_document(
Document( Document(
@@ -123,9 +124,9 @@ async def test_delete_document_succeeds_when_unlinked(default_session_factory, t
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_document_removes_person_links(default_session_factory, tmp_path): async def test_delete_document_removes_person_links(default_session_factory, tmp_path):
service = DocumentService(session_factory=default_session_factory) settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
service = DocumentService(session_factory=default_session_factory, settings=settings)
people_service = PeopleService(session_factory=default_session_factory) people_service = PeopleService(session_factory=default_session_factory)
service.settings.upload_dir = tmp_path
document = await service.create_document( document = await service.create_document(
Document( Document(
@@ -163,8 +164,8 @@ async def test_delete_document_removes_person_links(default_session_factory, tmp
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_document_removes_populated_storage_tree(default_session_factory, tmp_path): async def test_delete_document_removes_populated_storage_tree(default_session_factory, tmp_path):
service = DocumentService(session_factory=default_session_factory) settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
service.settings.upload_dir = tmp_path service = DocumentService(session_factory=default_session_factory, settings=settings)
document = await service.create_document( document = await service.create_document(
Document( Document(
+5 -4
View File
@@ -4,6 +4,7 @@ from uuid import uuid4
import pytest import pytest
from transcription.config import Settings
from transcription.db.models import Document from transcription.db.models import Document
from transcription.db.models import Job from transcription.db.models import Job
from transcription.db.models import JobSource from transcription.db.models import JobSource
@@ -102,8 +103,8 @@ class TestSourceServiceRevisionUpsert:
): ):
documents = DocumentService(session_factory=default_session_factory) documents = DocumentService(session_factory=default_session_factory)
jobs = JobService(session_factory=default_session_factory) jobs = JobService(session_factory=default_session_factory)
transcriptions = SourceService(session_factory=default_session_factory) settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
transcriptions.settings.upload_dir = tmp_path transcriptions = SourceService(session_factory=default_session_factory, settings=settings)
document = Document(id=uuid4(), name="delete-source-success") document = Document(id=uuid4(), name="delete-source-success")
await documents.create_document(document=document) await documents.create_document(document=document)
@@ -174,8 +175,8 @@ class TestSourceServiceRevisionUpsert:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_delete_unlinked_source_succeeds(self, default_session_factory, tmp_path): async def test_delete_unlinked_source_succeeds(self, default_session_factory, tmp_path):
documents = DocumentService(session_factory=default_session_factory) documents = DocumentService(session_factory=default_session_factory)
transcriptions = SourceService(session_factory=default_session_factory) settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
transcriptions.settings.upload_dir = tmp_path transcriptions = SourceService(session_factory=default_session_factory, settings=settings)
document = Document(id=uuid4(), name="delete-unlinked-source") document = Document(id=uuid4(), name="delete-unlinked-source")
await documents.create_document(document=document) await documents.create_document(document=document)
+25 -2
View File
@@ -24,7 +24,7 @@ class TestSettingsLoading:
"""Settings constructs when OPENROUTER_API_KEY is provided.""" """Settings constructs when OPENROUTER_API_KEY is provided."""
monkeypatch.setenv("OPENROUTER_API_KEY", "test-key-xyz") monkeypatch.setenv("OPENROUTER_API_KEY", "test-key-xyz")
settings = Settings() settings = Settings()
assert settings.openrouter_api_key == "test-key-xyz" assert settings.openrouter_api_key.get_secret_value() == "test-key-xyz"
def test_requires_api_key(self, monkeypatch): def test_requires_api_key(self, monkeypatch):
"""Settings raises ValidationError when OPENROUTER_API_KEY is missing.""" """Settings raises ValidationError when OPENROUTER_API_KEY is missing."""
@@ -52,7 +52,7 @@ class TestSettingsLoading:
] ]
) )
assert settings.openrouter_api_key == "test-key" assert settings.openrouter_api_key.get_secret_value() == "test-key"
assert settings.port == 8123 assert settings.port == 8123
assert settings.reload is True assert settings.reload is True
@@ -77,6 +77,29 @@ class TestProviderSettings:
assert settings.openrouter_http_referer is None assert settings.openrouter_http_referer is None
assert settings.openrouter_app_title is None assert settings.openrouter_app_title is None
@pytest.mark.parametrize(
("field", "value"),
[
("transcription_temperature", -0.1),
("transcription_temperature", 2.1),
("transcription_top_p", -0.1),
("transcription_top_p", 1.1),
],
)
def test_rejects_sampling_values_outside_provider_ranges(self, field, value):
with pytest.raises(ValidationError):
_make_settings(**{field: value})
def test_rejects_prompt_paths_outside_prompt_directory(self):
with pytest.raises(ValidationError):
_make_settings(default_prompt_name="../secret.md")
def test_settings_are_immutable_runtime_snapshots(self):
settings = _make_settings()
with pytest.raises(ValidationError):
settings.port = 9000
def test_provider_model_accepts_env_default(self, monkeypatch): def test_provider_model_accepts_env_default(self, monkeypatch):
"""provider_model is sourced when provided through environment configuration.""" """provider_model is sourced when provided through environment configuration."""
monkeypatch.setenv("PROVIDER_MODEL", "google/gemini-2.5-flash") monkeypatch.setenv("PROVIDER_MODEL", "google/gemini-2.5-flash")
+52
View File
@@ -1,7 +1,17 @@
"""Tests for prompt artifacts in prompts/.""" """Tests for prompt artifacts in prompts/."""
import hashlib
from pathlib import Path from pathlib import Path
import pytest
from pydantic import ValidationError
from transcription.config import Settings
from transcription.services.sources import PromptExecution
from transcription.services.sources import PromptLoadError
from transcription.services.sources import build_prompt_execution
from transcription.services.sources import load_prompt_text
PROMPT_PATH = Path("prompts/transcribe_document.md") PROMPT_PATH = Path("prompts/transcribe_document.md")
@@ -39,3 +49,45 @@ class TestPromptArtifact:
text = _prompt_text().lower() text = _prompt_text().lower()
assert "[deleted:" in text assert "[deleted:" in text
assert "[inserted:" in text assert "[inserted:" in text
class TestPromptConfiguration:
def test_builds_validated_immutable_prompt_provenance(self, tmp_path):
prompt_text = "Transcribe this document verbatim."
(tmp_path / "custom.md").write_text(prompt_text, encoding="utf-8")
settings = Settings(
openrouter_api_key="test-key",
prompt_dir=tmp_path,
default_prompt_name="custom.md",
transcription_temperature=0.2,
transcription_top_p=0.9,
)
execution = build_prompt_execution(settings=settings)
assert execution.prompt_hash == hashlib.sha256(prompt_text.encode()).hexdigest()
assert execution.temperature == 0.2
assert execution.top_p == 0.9
with pytest.raises(ValidationError):
execution.prompt_name = "changed.md"
def test_rejects_prompt_path_traversal_even_with_direct_loader_call(self, tmp_path):
outside_prompt = tmp_path / "outside.md"
prompt_dir = tmp_path / "prompts"
prompt_dir.mkdir()
outside_prompt.write_text("secret", encoding="utf-8")
settings = Settings(openrouter_api_key="test-key", prompt_dir=prompt_dir)
with pytest.raises(PromptLoadError):
load_prompt_text(prompt_name="../outside.md", settings=settings)
def test_prompt_execution_rejects_invalid_provenance_hash(self):
with pytest.raises(ValidationError):
PromptExecution(
prompt_name="prompt.md",
prompt_hash="not-a-sha256",
system_prompt=None,
user_prompt="text",
temperature=None,
top_p=None,
)