generated from john/python-template
Continue GC code review: Pydantic
This commit is contained in:
@@ -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`
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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"})
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"}}]},
|
||||
)
|
||||
|
||||
|
||||
@@ -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 +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(
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user