diff --git a/README.md b/README.md index 0b1e826..800cbe8 100644 --- a/README.md +++ b/README.md @@ -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` diff --git a/prompts/README.md b/prompts/README.md index 0b3e454..e098784 100644 --- a/prompts/README.md +++ b/prompts/README.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. diff --git a/prompts/system.txt b/prompts/system.txt deleted file mode 100644 index 5e777ef..0000000 --- a/prompts/system.txt +++ /dev/null @@ -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. \ No newline at end of file diff --git a/src/transcription/api/v4_documents.py b/src/transcription/api/v4_documents.py index 778b1e5..020df16 100644 --- a/src/transcription/api/v4_documents.py +++ b/src/transcription/api/v4_documents.py @@ -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) diff --git a/src/transcription/config.py b/src/transcription/config.py index e9f7b1a..e546563 100644 --- a/src/transcription/config.py +++ b/src/transcription/config.py @@ -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) diff --git a/src/transcription/db/models.py b/src/transcription/db/models.py index 3fe49ba..b05519c 100644 --- a/src/transcription/db/models.py +++ b/src/transcription/db/models.py @@ -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"}) - - diff --git a/src/transcription/providers/__init__.py b/src/transcription/providers/__init__.py index 318e40a..2463e48 100644 --- a/src/transcription/providers/__init__.py +++ b/src/transcription/providers/__init__.py @@ -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", diff --git a/src/transcription/providers/base.py b/src/transcription/providers/base.py index d86095f..2488738 100644 --- a/src/transcription/providers/base.py +++ b/src/transcription/providers/base.py @@ -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): diff --git a/src/transcription/providers/openrouter.py b/src/transcription/providers/openrouter.py index fedc292..2d6644c 100644 --- a/src/transcription/providers/openrouter.py +++ b/src/transcription/providers/openrouter.py @@ -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 diff --git a/src/transcription/services/sources.py b/src/transcription/services/sources.py index 2a87c5e..a88375d 100644 --- a/src/transcription/services/sources.py +++ b/src/transcription/services/sources.py @@ -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( diff --git a/src/transcription/services/workflows.py b/src/transcription/services/workflows.py index ecd2de8..8abfa14 100644 --- a/src/transcription/services/workflows.py +++ b/src/transcription/services/workflows.py @@ -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, diff --git a/tests/api/test_v4_documents.py b/tests/api/test_v4_documents.py index e9c68e4..f7f336c 100644 --- a/tests/api/test_v4_documents.py +++ b/tests/api/test_v4_documents.py @@ -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) diff --git a/tests/integration/test_pipeline_flow.py b/tests/integration/test_pipeline_flow.py index 1e611b2..9b7d4eb 100644 --- a/tests/integration/test_pipeline_flow.py +++ b/tests/integration/test_pipeline_flow.py @@ -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"}}]}, ) diff --git a/tests/providers/test_openrouter.py b/tests/providers/test_openrouter.py index 882d382..587a219 100644 --- a/tests/providers/test_openrouter.py +++ b/tests/providers/test_openrouter.py @@ -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 diff --git a/tests/services/test_document_service.py b/tests/services/test_document_service.py index c70391f..049f2ca 100644 --- a/tests/services/test_document_service.py +++ b/tests/services/test_document_service.py @@ -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( diff --git a/tests/services/test_transcription_service.py b/tests/services/test_transcription_service.py index a10c5a8..1ca898f 100644 --- a/tests/services/test_transcription_service.py +++ b/tests/services/test_transcription_service.py @@ -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) diff --git a/tests/test_config.py b/tests/test_config.py index fdc1e6b..da6d565 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -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") diff --git a/tests/test_prompts.py b/tests/test_prompts.py index 1b27ed3..7a41184 100644 --- a/tests/test_prompts.py +++ b/tests/test_prompts.py @@ -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, + )