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
+76 -37
View File
@@ -2,8 +2,7 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime
from typing import Annotated
from uuid import UUID
from fastapi import APIRouter
@@ -11,7 +10,9 @@ from fastapi import Depends
from fastapi import Request
from fastapi import Response
from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field
from pydantic import model_validator
from transcription.db.models import Document
from transcription.db.models import DocumentPerson
@@ -23,28 +24,32 @@ from transcription.services import PeopleService
router = APIRouter(prefix="/api/v4", tags=["v4-documents"])
@dataclass(frozen=True)
class DocumentTypePayload:
id: UUID
code: str
label: str
is_active: bool
sort_order: int
created_at: datetime
updated_at: datetime
class ApiModel(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True)
@dataclass(frozen=True)
class PersonRolePayload:
id: UUID
code: str
label: str
is_active: bool
created_at: datetime
updated_at: datetime
class SelectorRequest(ApiModel):
@model_validator(mode="after")
def require_exactly_one_selector(self):
values = (self.selector_id, self.selector_code)
if sum(value is not None for value in values) != 1:
raise ValueError(f"Provide exactly one of {self.selector_names[0]} or {self.selector_names[1]}")
return self
@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
code: str
label: str
@@ -52,36 +57,66 @@ class DocumentTypeRead(BaseModel):
sort_order: int
class PersonRoleRead(BaseModel):
class PersonRoleRead(ApiModel):
id: UUID
code: str
label: str
is_active: bool
class DocumentTypeWriteRequest(BaseModel):
class DocumentTypeWriteRequest(SelectorRequest):
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_type_id: UUID | None
document_type_code: str | None
class DocumentPersonWriteRequest(BaseModel):
class DocumentPersonWriteRequest(ApiModel):
person_id: UUID
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_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
document_id: UUID
person_id: UUID
@@ -90,7 +125,7 @@ class DocumentPersonRead(BaseModel):
person_name: str | None = None
class DocumentPeopleResponse(BaseModel):
class DocumentPeopleResponse(ApiModel):
document_id: UUID
links: list[DocumentPersonRead] = Field(default_factory=list)
@@ -152,10 +187,14 @@ def get_people_service(request: Request) -> 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])
async def list_document_types(
service: DocumentServiceDependency,
active_only: bool = True,
service: DocumentService = Depends(get_document_service),
) -> list[DocumentTypeRead]:
items = await service.list_document_types(active_only=active_only)
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])
async def list_person_roles(
service: PeopleServiceDependency,
active_only: bool = True,
service: PeopleService = Depends(get_people_service),
) -> list[PersonRoleRead]:
items = await service.list_person_roles(active_only=active_only)
return [_person_role_to_read(item) for item in items]
@@ -174,7 +213,7 @@ async def list_person_roles(
async def set_document_type(
document_id: UUID,
payload: DocumentTypeWriteRequest,
service: DocumentService = Depends(get_document_service),
service: DocumentServiceDependency,
) -> DocumentTypeWriteResponse:
document = await service.set_document_type(
document_id=document_id,
@@ -187,7 +226,7 @@ async def set_document_type(
@router.get("/documents/{document_id}/people", response_model=DocumentPeopleResponse)
async def list_document_people(
document_id: UUID,
service: PeopleService = Depends(get_people_service),
service: PeopleServiceDependency,
) -> DocumentPeopleResponse:
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])
@@ -197,7 +236,7 @@ async def list_document_people(
async def add_document_person_link(
document_id: UUID,
payload: DocumentPersonWriteRequest,
service: PeopleService = Depends(get_people_service),
service: PeopleServiceDependency,
) -> DocumentPersonRead:
link = await service.add_document_person_link(
document_id=document_id,
@@ -212,7 +251,7 @@ async def add_document_person_link(
async def set_document_person_role(
document_person_id: UUID,
payload: DocumentPersonRoleUpdateRequest,
service: PeopleService = Depends(get_people_service),
service: PeopleServiceDependency,
) -> DocumentPersonRead:
link = await service.set_document_person_role(
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)
async def delete_document_person_link(
document_person_id: UUID,
service: PeopleService = Depends(get_people_service),
service: PeopleServiceDependency,
) -> Response:
await service.remove_document_person_link(document_person_id=document_person_id)
return Response(status_code=204)
+27 -14
View File
@@ -15,8 +15,10 @@ from typing import Any
from typing import Literal
from pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field
from pydantic import SecretStr
from pydantic import StringConstraints
from pydantic_settings import BaseSettings
from pydantic_settings import SettingsConfigDict
@@ -27,17 +29,27 @@ class Provider(StrEnum):
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):
model_config = ConfigDict(extra="forbid", frozen=True)
driver: Literal["sqlite"] = "sqlite"
path: str = "app.db"
path: NonEmptyStr = "app.db"
class PostgresSettings(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)
driver: Literal["postgres"] = "postgres"
host: str
port: int = 5432
database: str
user: str
host: NonEmptyStr
port: int = Field(default=5432, ge=1, le=65535)
database: NonEmptyStr
user: NonEmptyStr
password: SecretStr
@@ -55,6 +67,7 @@ class Settings(BaseSettings):
env_nested_delimiter="__",
cli_implicit_flags=True,
cli_kebab_case=True,
frozen=True,
)
# --- NiceGUI Server ---
@@ -65,13 +78,13 @@ class Settings(BaseSettings):
# --- AI provider ---
provider: Provider = Provider.OPENROUTER
openrouter_api_key: str
provider_model: str | None = None
openrouter_http_referer: str | None = None
openrouter_app_title: str | None = None
default_prompt_name: str = "transcribe_document.md"
transcription_temperature: float | None = None
transcription_top_p: float | None = None
openrouter_api_key: SecretStr
provider_model: NonEmptyStr | None = None
openrouter_http_referer: NonEmptyStr | None = None
openrouter_app_title: NonEmptyStr | None = None
default_prompt_name: PromptFilename = "transcribe_document.md"
transcription_temperature: Temperature | None = None
transcription_top_p: Probability | None = None
# --- runtime environment ---
environment: Literal["development", "test", "production"] = "development"
@@ -86,8 +99,8 @@ class Settings(BaseSettings):
prompt_dir: Path = Path("./prompts")
# --- worker reliability ---
worker_max_retries: int = 0
worker_retry_backoff_seconds: float = 0.0
worker_max_retries: int = Field(default=0, ge=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_min_transcription_chars: 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 datetime
from enum import StrEnum
from typing import Any
from typing import Optional
from uuid import UUID
from uuid import uuid4
from sqlalchemy import Column
from sqlalchemy import BigInteger
from sqlalchemy import Enum as SAEnum
from pydantic import JsonValue
from sqlalchemy import JSON
from sqlalchemy import BigInteger
from sqlalchemy import Column
from sqlalchemy import Enum as SAEnum
from sqlalchemy import UniqueConstraint
from sqlalchemy.dialects.postgresql import JSONB
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))
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):
@@ -82,7 +84,9 @@ class PersonRole(SQLModel, table=True):
created_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):
@@ -102,8 +106,12 @@ class Document(SQLModel, table=True):
jobs: list["Job"] = 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_type_ref: Optional["DocumentType"] = Relationship(back_populates="documents", sa_relationship_kwargs={"lazy": "selectin"})
document_people: list["DocumentPerson"] = Relationship(
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):
@@ -121,14 +129,16 @@ class Person(SQLModel, table=True):
death_place: str | None = None
biography: str | None = None
portrait_path: str | None = None
metadata_: dict[str, Any] | None = Field(
metadata_: dict[str, JsonValue] | None = Field(
default=None,
sa_column=Column("metadata", JSONBCompat(), nullable=True),
)
created_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):
@@ -159,9 +169,15 @@ class DocumentPerson(SQLModel, table=True):
UniqueConstraint("document_id", "person_id", "role", name="uq_document_person_role"),
)
document: Optional["Document"] = 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"})
document: Optional["Document"] = 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):
@@ -278,6 +294,7 @@ class Source(SQLModel, table=True):
"""Return the parent document name if loaded."""
return self.document.name if self.document else None
class JobSource(SQLModel, table=True):
"""A single AI execution record for one source page."""
@@ -298,12 +315,10 @@ class JobSource(SQLModel, table=True):
),
)
raw_transcription: str | None = None
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))
ai_metadata: dict[str, JsonValue] | 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
executed_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
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"})
+2
View File
@@ -6,6 +6,7 @@ from transcription.config import get_settings
from transcription.providers.base import ProviderAuthError
from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError
from transcription.providers.base import TranscriptionMetadata
from transcription.providers.base import TranscriptionProvider
from transcription.providers.base import TranscriptionResult
from transcription.providers.openrouter import OpenRouterTranscriptionProvider
@@ -25,6 +26,7 @@ __all__ = [
"ProviderAuthError",
"ProviderError",
"ProviderResponseError",
"TranscriptionMetadata",
"TranscriptionProvider",
"TranscriptionResult",
"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 pydantic import BaseModel
from pydantic import ConfigDict
from pydantic import Field
from pydantic import JsonValue
class ProviderError(RuntimeError):
"""Base error for provider failures."""
@@ -17,25 +20,64 @@ class ProviderResponseError(ProviderError):
"""Raised when provider responses are malformed or unusable."""
@dataclass(frozen=True)
class TranscriptionResult:
class ProviderUsage(BaseModel):
"""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."""
text: str
provider: str
model: str
model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True)
text: str = Field(min_length=1)
provider: str = Field(min_length=1)
model: str = Field(min_length=1)
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
user_prompt: str | None = None
temperature: float | None = None
top_p: float | None = None
finish_reason: str | None = None
usage_input_tokens: int | None = None
usage_output_tokens: int | None = None
usage_total_tokens: int | None = None
ai_metadata: dict[str, Any] | None = None
raw_api_response: dict[str, Any] | None = None
temperature: float | None = Field(default=None, ge=0.0, le=2.0)
top_p: float | None = Field(default=None, ge=0.0, le=1.0)
metadata: TranscriptionMetadata = Field(default_factory=TranscriptionMetadata)
raw_api_response: dict[str, JsonValue] | None = None
@property
def finish_reason(self) -> str | 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):
+142 -152
View File
@@ -4,18 +4,25 @@ from __future__ import annotations
import base64
import logging
from dataclasses import dataclass
from typing import Annotated
from typing import Any
from typing import cast
from typing import Literal
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 get_settings
from transcription.providers.base import ProviderAuthError
from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError
from transcription.providers.base import ProviderUsage
from transcription.providers.base import TranscriptionMetadata
from transcription.providers.base import TranscriptionResult
logger = logging.getLogger(__name__)
@@ -23,16 +30,92 @@ logger = logging.getLogger(__name__)
DEFAULT_OPENROUTER_MODEL = "google/gemini-2.5-flash"
@dataclass(frozen=True)
class OpenRouterRequest:
class _ProviderModel(BaseModel):
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."""
model: str
messages: list[dict[str, Any]]
model: str = Field(min_length=1)
messages: tuple[UserMessage, ...] = Field(min_length=1)
http_referer: str | None
x_open_router_title: str | None
temperature: float | None
top_p: float | None
temperature: float | None = Field(ge=0.0, le=2.0)
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:
@@ -41,7 +124,7 @@ class OpenRouterTranscriptionProvider:
def __init__(self, *, settings: Settings | None = None, client: OpenRouter | None = None):
self._settings = settings or get_settings()
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
def model(self) -> str:
@@ -66,31 +149,22 @@ class OpenRouterTranscriptionProvider:
top_p=top_p,
)
try:
response = await self._client.chat.send_async(
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,
)
response = await self._client.chat.send_async(**request.model_dump(mode="json", exclude_none=True))
except Exception as exc:
message = str(exc).lower()
if "401" in message or "auth" in message or "api key" in message:
raise ProviderAuthError("OpenRouter authentication 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)
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)
return TranscriptionResult(
text=text,
@@ -102,46 +176,37 @@ class OpenRouterTranscriptionProvider:
temperature=temperature,
top_p=top_p,
model=model,
finish_reason=finish_reason,
usage_input_tokens=usage_input_tokens,
usage_output_tokens=usage_output_tokens,
usage_total_tokens=usage_total_tokens,
ai_metadata=ai_metadata,
metadata=metadata,
raw_api_response=raw_api_response,
)
def _build_ai_metadata(
self,
*,
finish_reason: str | None,
usage_input_tokens: int | None,
usage_output_tokens: int | None,
usage_total_tokens: int | None,
) -> dict[str, Any] | None:
metadata: dict[str, Any] = {}
if finish_reason is not None:
metadata["finish_reason"] = finish_reason
def _build_metadata(self, response: OpenRouterResponse) -> TranscriptionMetadata:
choice = response.choices[0]
finish_reason = choice.finish_reason.strip() if choice.finish_reason and choice.finish_reason.strip() else None
normalized_usage = None
if response.usage is not None:
try:
usage = ResponseUsage.model_validate(response.usage)
except ValidationError as exc:
logger.warning("Ignoring invalid OpenRouter usage metadata: %s", exc)
else:
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] = {}
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:
def _coerce_raw_response(self, response: Any) -> dict[str, JsonValue]:
payload = self._to_json_compatible(response)
if payload is None:
return None
if isinstance(payload, dict):
return payload
return {"response": payload}
try:
return JSON_OBJECT_ADAPTER.validate_python(payload)
except ValidationError as exc:
raise ProviderResponseError("OpenRouter response is not a JSON object") from exc
def _to_json_compatible(self, value: Any) -> Any:
if value is None or isinstance(value, str | int | float | bool):
@@ -153,12 +218,13 @@ class OpenRouterTranscriptionProvider:
if isinstance(value, list | tuple | set):
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)
if callable(serializer):
try:
return self._to_json_compatible(serializer())
except Exception as exc: # noqa: BLE001
serialized = serializer(mode="json") if method_name == "model_dump" else serializer()
return self._to_json_compatible(serialized)
except (TypeError, ValueError) as exc:
logger.debug("OpenRouter response serializer %s failed: %s", method_name, exc)
continue
@@ -170,7 +236,7 @@ class OpenRouterTranscriptionProvider:
if not str(key).startswith("_")
}
return repr(value)
raise ProviderResponseError(f"OpenRouter response contains unsupported value type: {type(value).__name__}")
def _build_request(
self,
@@ -183,110 +249,34 @@ class OpenRouterTranscriptionProvider:
) -> OpenRouterRequest:
image_b64 = base64.b64encode(image_bytes).decode("ascii")
data_url = f"data:{mime_type};base64,{image_b64}"
media_content: dict[str, Any]
media_content: ImageContent | FileContent
if mime_type == "application/pdf":
media_content = {
"type": "file",
"file": {
"filename": "source.pdf",
"file_data": data_url,
},
}
media_content = FileContent(file=FileData(filename="source.pdf", file_data=data_url))
else:
media_content = {
"type": "image_url",
"image_url": {"url": data_url},
}
messages: list[dict[str, Any]] = [
{
"role": "user",
"content": [
{"type": "text", "text": prompt_text},
media_content,
],
}
]
media_content = ImageContent(image_url=ImageUrl(url=data_url))
return OpenRouterRequest(
model=self.model,
messages=messages,
messages=(UserMessage(content=(TextContent(text=prompt_text), media_content)),),
http_referer=self._settings.openrouter_http_referer,
x_open_router_title=self._settings.openrouter_app_title,
temperature=temperature,
top_p=top_p,
)
def _extract_text(self, response: Any) -> str:
choices = self._get_optional_attr(response, "choices")
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")
def _extract_text(self, response: OpenRouterResponse) -> str:
content = response.choices[0].message.content
text = self._normalize_content(content)
if not text:
raise ProviderResponseError("OpenRouter response contained no transcription text")
return text
def _extract_finish_reason(self, response: Any) -> str | None:
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:
def _normalize_content(self, content: str | tuple[ResponseContentPart, ...] | None) -> str:
if isinstance(content, str):
return content.strip()
if isinstance(content, list):
parts: list[str] = []
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())
if isinstance(content, tuple):
parts = [item.text.strip() for item in content if item.text and item.text.strip()]
return "\n".join(parts).strip()
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
from collections.abc import Sequence
from contextlib import contextmanager
from dataclasses import dataclass
from datetime import UTC
from datetime import datetime
from pathlib import Path
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.orm import selectinload
from sqlmodel import select
@@ -28,6 +33,7 @@ from transcription.errors import ErrorCategory
from transcription.providers import ProviderAuthError
from transcription.providers import ProviderError
from transcription.providers import ProviderResponseError
from transcription.providers import TranscriptionMetadata
from transcription.providers import TranscriptionProvider
from transcription.providers import TranscriptionResult
from transcription.providers import get_transcription_provider
@@ -46,18 +52,20 @@ SOURCE_MIME_TYPES = {
".pdf": "application/pdf",
}
SOURCE_EXTENSIONS = frozenset(SOURCE_MIME_TYPES)
JSON_OBJECT_ADAPTER = TypeAdapter(dict[str, JsonValue])
@dataclass(frozen=True)
class PromptExecution:
class PromptExecution(BaseModel):
"""Resolved prompt inputs captured for one page execution."""
prompt_name: str
prompt_hash: str
model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True)
prompt_name: str = Field(min_length=1, pattern=r"^[^/\\]+$")
prompt_hash: str = Field(pattern=r"^[0-9a-f]{64}$")
system_prompt: str | None
user_prompt: str
temperature: float | None
top_p: float | None
user_prompt: str = Field(min_length=1)
temperature: float | None = Field(ge=0.0, le=2.0)
top_p: float | None = Field(ge=0.0, le=1.0)
class PromptLoadError(AppError):
@@ -366,8 +374,8 @@ class SourceService(ServiceBase):
source_id: UUID,
text: str | None,
error_detail: str | None = None,
ai_metadata: dict[str, object] | None = None,
raw_api_response: dict[str, object] | None = None,
ai_metadata: TranscriptionMetadata | dict[str, JsonValue] | None = None,
raw_api_response: dict[str, JsonValue] | None = None,
provider: str | None = None,
model: str | 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.model = model or job.model or _resolve_transcript_model(provider=self.provider, settings=self.settings)
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:
source.raw_transcription = text
@@ -414,15 +424,15 @@ class SourceService(ServiceBase):
source_id=source_id,
status=JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED,
raw_transcription=text,
ai_metadata=ai_metadata,
raw_api_response=raw_api_response,
ai_metadata=metadata_payload,
raw_api_response=raw_response_payload,
error_detail=error_detail,
)
_session.add(job_source)
else:
job_source.raw_transcription = text
job_source.ai_metadata = ai_metadata
job_source.raw_api_response = raw_api_response
job_source.ai_metadata = metadata_payload
job_source.raw_api_response = raw_response_payload
job_source.error_detail = error_detail
job_source.status = JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED
job_source.executed_at = datetime.now(UTC)
@@ -492,6 +502,46 @@ def _resolve_transcript_model(*, provider: TranscriptionProvider, settings: Sett
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(
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()
prompt_execution = PromptExecution(
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,
user_prompt=prompt_text,
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,
top_p=result.top_p if result.top_p is not None else prompt_execution.top_p,
model=result.model,
finish_reason=result.finish_reason,
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,
metadata=result.metadata,
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)
return PromptExecution(
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,
user_prompt=user_prompt,
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:
"""Load and validate prompt text from PROMPT_DIR."""
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():
raise PromptLoadError(
+8 -5
View File
@@ -17,6 +17,7 @@ from ..providers import TranscriptionResult
from . import ServiceBundle
from .sources import PromptExecution
from .sources import build_prompt_execution
from .sources import hash_prompt_text
from .sources import transcribe_document_image
logger = logging.getLogger(__name__)
@@ -141,7 +142,8 @@ async def process_queued_job(
)
failed_pages.append((source, 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.document_id,
source.id,
@@ -157,7 +159,8 @@ async def process_queued_job(
failed_pages.append((source, 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.document_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:
return PromptExecution(
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,
user_prompt=source_job.user_prompt,
temperature=source_job.temperature,
@@ -269,7 +272,7 @@ async def _finalize_batch_outcome(
source_id=source.id,
text=result.text,
error_detail=None,
ai_metadata=result.ai_metadata,
ai_metadata=result.metadata_payload(),
raw_api_response=result.raw_api_response,
provider=result.provider,
model=result.model,
@@ -295,7 +298,7 @@ async def _finalize_batch_outcome(
source_id=source.id,
text=result.text,
error_detail=None,
ai_metadata=result.ai_metadata,
ai_metadata=result.metadata_payload(),
raw_api_response=result.raw_api_response,
provider=result.provider,
model=result.model,