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