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
UPLOAD_DIR=./uploads
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:
@@ -138,7 +141,8 @@ Replace `localhost` with the server's hostname or IP address when connecting fro
## 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:
- `prompts/transcribe_document.md`
+2
View File
@@ -5,9 +5,11 @@ This directory stores transcription prompts as individual Markdown artifacts.
## Conventions
- Keep one prompt per file.
- 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.
- Keep prompts human-readable and policy-focused.
- Do not store secrets in prompt files.
- Runtime jobs snapshot prompt text, SHA-256 provenance, and sampling configuration.
## Current Prompt
- `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 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,
+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"
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):
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)
@@ -155,6 +174,19 @@ def test_document_people_role_aware_write_read_and_delete(tmp_path):
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):
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)
+6 -1
View File
@@ -9,6 +9,8 @@ from transcription.config import Settings
from transcription.db.models import Document
from transcription.db.models import JobSourceStatus
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.services import ServiceBundle
from transcription.services.store import create_document_job
@@ -84,7 +86,10 @@ class TestPipelineSuccessFlow:
provider="openrouter",
model="test-model",
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"}}]},
)
+23 -1
View File
@@ -122,7 +122,7 @@ class TestOpenRouterProviderTranscribe:
assert result.provider == "openrouter"
assert result.prompt_name is None
assert result.model == "vendor/model-b"
assert result.ai_metadata == {
assert result.metadata_payload() == {
"finish_reason": "stop",
"usage": {"input_tokens": 10, "output_tokens": 25, "total_tokens": 35},
}
@@ -178,3 +178,25 @@ class TestOpenRouterProviderTranscribe:
image_bytes=b"img-bytes",
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
from sqlmodel import select
from transcription.config import Settings
from transcription.db.models import Document
from transcription.db.models import DocumentPerson
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
async def test_delete_document_succeeds_when_unlinked(default_session_factory, tmp_path):
service = DocumentService(session_factory=default_session_factory)
service.settings.upload_dir = tmp_path
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
service = DocumentService(session_factory=default_session_factory, settings=settings)
document = await service.create_document(
Document(
@@ -123,9 +124,9 @@ async def test_delete_document_succeeds_when_unlinked(default_session_factory, t
@pytest.mark.asyncio
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)
service.settings.upload_dir = tmp_path
document = await service.create_document(
Document(
@@ -163,8 +164,8 @@ async def test_delete_document_removes_person_links(default_session_factory, tmp
@pytest.mark.asyncio
async def test_delete_document_removes_populated_storage_tree(default_session_factory, tmp_path):
service = DocumentService(session_factory=default_session_factory)
service.settings.upload_dir = tmp_path
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
service = DocumentService(session_factory=default_session_factory, settings=settings)
document = await service.create_document(
Document(
+5 -4
View File
@@ -4,6 +4,7 @@ from uuid import uuid4
import pytest
from transcription.config import Settings
from transcription.db.models import Document
from transcription.db.models import Job
from transcription.db.models import JobSource
@@ -102,8 +103,8 @@ class TestSourceServiceRevisionUpsert:
):
documents = DocumentService(session_factory=default_session_factory)
jobs = JobService(session_factory=default_session_factory)
transcriptions = SourceService(session_factory=default_session_factory)
transcriptions.settings.upload_dir = tmp_path
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
transcriptions = SourceService(session_factory=default_session_factory, settings=settings)
document = Document(id=uuid4(), name="delete-source-success")
await documents.create_document(document=document)
@@ -174,8 +175,8 @@ class TestSourceServiceRevisionUpsert:
@pytest.mark.asyncio
async def test_delete_unlinked_source_succeeds(self, default_session_factory, tmp_path):
documents = DocumentService(session_factory=default_session_factory)
transcriptions = SourceService(session_factory=default_session_factory)
transcriptions.settings.upload_dir = tmp_path
settings = Settings(openrouter_api_key="test-key", upload_dir=tmp_path)
transcriptions = SourceService(session_factory=default_session_factory, settings=settings)
document = Document(id=uuid4(), name="delete-unlinked-source")
await documents.create_document(document=document)
+25 -2
View File
@@ -24,7 +24,7 @@ class TestSettingsLoading:
"""Settings constructs when OPENROUTER_API_KEY is provided."""
monkeypatch.setenv("OPENROUTER_API_KEY", "test-key-xyz")
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):
"""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.reload is True
@@ -77,6 +77,29 @@ class TestProviderSettings:
assert settings.openrouter_http_referer 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):
"""provider_model is sourced when provided through environment configuration."""
monkeypatch.setenv("PROVIDER_MODEL", "google/gemini-2.5-flash")
+52
View File
@@ -1,7 +1,17 @@
"""Tests for prompt artifacts in prompts/."""
import hashlib
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")
@@ -39,3 +49,45 @@ class TestPromptArtifact:
text = _prompt_text().lower()
assert "[deleted:" 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,
)