AI metadata and api prompt results data capture now fixed

This commit is contained in:
Jim Lancaster
2026-08-08 14:59:45 -05:00
parent 89cac3c378
commit 58faa00d7b
7 changed files with 172 additions and 10 deletions
+3
View File
@@ -1,6 +1,7 @@
"""Provider interfaces and shared types for transcription adapters.""" """Provider interfaces and shared types for transcription adapters."""
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any
from typing import Protocol from typing import Protocol
@@ -28,6 +29,8 @@ class TranscriptionResult:
usage_input_tokens: int | None = None usage_input_tokens: int | None = None
usage_output_tokens: int | None = None usage_output_tokens: int | None = None
usage_total_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
class TranscriptionProvider(Protocol): class TranscriptionProvider(Protocol):
+70
View File
@@ -66,6 +66,13 @@ class OpenRouterTranscriptionProvider:
model = self._get_optional_attr(response, "model") or self.model model = self._get_optional_attr(response, "model") or self.model
finish_reason = self._extract_finish_reason(response) finish_reason = self._extract_finish_reason(response)
usage_input_tokens, usage_output_tokens, usage_total_tokens = self._extract_usage(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)
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,
@@ -76,8 +83,71 @@ class OpenRouterTranscriptionProvider:
usage_input_tokens=usage_input_tokens, usage_input_tokens=usage_input_tokens,
usage_output_tokens=usage_output_tokens, usage_output_tokens=usage_output_tokens,
usage_total_tokens=usage_total_tokens, usage_total_tokens=usage_total_tokens,
ai_metadata=ai_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
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:
payload = self._to_json_compatible(response)
if payload is None:
return None
if isinstance(payload, dict):
return payload
return {"response": payload}
def _to_json_compatible(self, value: Any) -> Any:
if value is None or isinstance(value, str | int | float | bool):
return value
if isinstance(value, dict):
return {str(key): self._to_json_compatible(item) for key, item in value.items()}
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"):
serializer = getattr(value, method_name, None)
if callable(serializer):
try:
return self._to_json_compatible(serializer())
except Exception: # noqa: BLE001
continue
object_dict = getattr(value, "__dict__", None)
if isinstance(object_dict, dict):
return {
str(key): self._to_json_compatible(item)
for key, item in object_dict.items()
if not str(key).startswith("_")
}
return repr(value)
def _build_request(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> OpenRouterRequest: def _build_request(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> 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}"
@@ -415,6 +415,8 @@ class TranscriptionService(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,
raw_api_response: dict[str, object] | None = None,
provider: str | None = None, provider: str | None = None,
model: str | None = None, model: str | None = None,
prompt_name: str = DEFAULT_PROMPT_FILE, prompt_name: str = DEFAULT_PROMPT_FILE,
@@ -462,11 +464,15 @@ class TranscriptionService(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,
raw_api_response=raw_api_response,
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.raw_api_response = raw_api_response
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)
+32 -9
View File
@@ -29,9 +29,13 @@ async def advance_job(
) -> Job | None: ) -> Job | None:
"""Advance a single job by lifecycle status.""" """Advance a single job by lifecycle status."""
settings = settings or get_settings() settings = settings or get_settings()
match job.status: current_status = _coerce_job_status(job.status)
match current_status:
case JobStatus.QUEUED: case JobStatus.QUEUED:
return await process_queued_job(job=job, services=services, settings=settings, session=session) return await process_queued_job(job=job, services=services, settings=settings, session=session)
case JobStatus.PROCESSING:
# Recover mid-flight jobs by continuing the queued processing path.
return await process_queued_job(job=job, services=services, settings=settings, session=session)
case JobStatus.FAILED: case JobStatus.FAILED:
if job.retry_count < settings.worker_max_retries: if job.retry_count < settings.worker_max_retries:
return await services.jobs.update_job_state( return await services.jobs.update_job_state(
@@ -56,18 +60,20 @@ async def process_queued_job(
) -> Job | None: ) -> Job | None:
"""Process one complete transcription attempt for a queued job.""" """Process one complete transcription attempt for a queued job."""
runtime_settings = settings or get_settings() runtime_settings = settings or get_settings()
if job.status != JobStatus.QUEUED: current_status = _coerce_job_status(job.status)
if current_status not in {JobStatus.QUEUED, JobStatus.PROCESSING}:
logger.warning(f"Job {job.id} is not queued. Current status: {job.status}") logger.warning(f"Job {job.id} is not queued. Current status: {job.status}")
return return
# Transaction A: claim job for processing. # Transaction A: claim job for processing.
if session is None: if current_status == JobStatus.QUEUED:
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING) if session is None:
else: job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING)
# If we're sharing the session, need to make sure setting the Job to PROCESSING is committed before we start else:
# the transcription, otherwise other workers may see the job as still QUEUED and try to process it. # If we're sharing the session, need to make sure setting the Job to PROCESSING is committed before we start
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING, session=session) # the transcription, otherwise other workers may see the job as still QUEUED and try to process it.
await session.commit() job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING, session=session)
await session.commit()
source_job = await services.jobs.read_job(job_id=job.id, session=session) source_job = await services.jobs.read_job(job_id=job.id, session=session)
sources = _resolve_job_sources(source_job) sources = _resolve_job_sources(source_job)
@@ -376,6 +382,8 @@ 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,
raw_api_response=result.raw_api_response,
provider=result.provider, provider=result.provider,
model=result.model, model=result.model,
prompt_name=result.prompt_name, prompt_name=result.prompt_name,
@@ -402,6 +410,8 @@ 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,
raw_api_response=result.raw_api_response,
provider=result.provider, provider=result.provider,
model=result.model, model=result.model,
prompt_name=result.prompt_name, prompt_name=result.prompt_name,
@@ -470,3 +480,16 @@ def _line_count(text: str) -> int:
if not stripped: if not stripped:
return 0 return 0
return sum(1 for line in stripped.splitlines() if line.strip()) return sum(1 for line in stripped.splitlines() if line.strip())
def _coerce_job_status(value: object) -> JobStatus | None:
if isinstance(value, JobStatus):
return value
if isinstance(value, str):
lowered = value.strip().lower()
for member in JobStatus:
if lowered in {member.value.lower(), member.name.lower()}:
return member
return None
+7
View File
@@ -76,6 +76,8 @@ 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}},
raw_api_response={"id": "resp_123", "choices": [{"message": {"content": "Pipeline transcript"}}]},
) )
monkeypatch.setattr( monkeypatch.setattr(
@@ -94,6 +96,11 @@ class TestPipelineSuccessFlow:
assert job is not None assert job is not None
assert job.status == JobStatus.TRANSCRIBED assert job.status == JobStatus.TRANSCRIBED
assert any(job_source.raw_transcription == "Pipeline transcript" for job_source in job.job_sources) assert any(job_source.raw_transcription == "Pipeline transcript" for job_source in job.job_sources)
assert any(job_source.ai_metadata == {"finish_reason": "stop", "usage": {"total_tokens": 42}} for job_source in job.job_sources)
assert any(
job_source.raw_api_response == {"id": "resp_123", "choices": [{"message": {"content": "Pipeline transcript"}}]}
for job_source in job.job_sources
)
assert all(job_source.error_detail is None for job_source in job.job_sources) assert all(job_source.error_detail is None for job_source in job.job_sources)
@pytest.mark.asyncio @pytest.mark.asyncio
+12 -1
View File
@@ -78,7 +78,13 @@ class TestOpenRouterProviderTranscribe:
"""Transcribe returns normalized text from a valid response payload.""" """Transcribe returns normalized text from a valid response payload."""
response = { response = {
"model": "vendor/model-b", "model": "vendor/model-b",
"choices": [{"message": {"content": [{"text": "Line 1"}, {"text": "Line 2"}]}}], "choices": [
{
"message": {"content": [{"text": "Line 1"}, {"text": "Line 2"}]},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 25, "total_tokens": 35},
} }
provider = OpenRouterTranscriptionProvider( provider = OpenRouterTranscriptionProvider(
settings=Settings(openrouter_api_key="test-key"), settings=Settings(openrouter_api_key="test-key"),
@@ -94,6 +100,11 @@ class TestOpenRouterProviderTranscribe:
assert result.text == "Line 1\nLine 2" assert result.text == "Line 1\nLine 2"
assert result.provider == "openrouter" assert result.provider == "openrouter"
assert result.model == "vendor/model-b" assert result.model == "vendor/model-b"
assert result.ai_metadata == {
"finish_reason": "stop",
"usage": {"input_tokens": 10, "output_tokens": 25, "total_tokens": 35},
}
assert result.raw_api_response == response
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_maps_sdk_exception_to_provider_error(self): async def test_maps_sdk_exception_to_provider_error(self):
+42
View File
@@ -211,3 +211,45 @@ async def test_source_delete_blocks_when_linked_to_multiple_jobs(default_session
with pytest.raises(SourceDeleteBlockedError): with pytest.raises(SourceDeleteBlockedError):
await transcriptions.delete_source_from_job_context(job_id=job_one.id, source_id=source.id) await transcriptions.delete_source_from_job_context(job_id=job_one.id, source_id=source.id)
@pytest.mark.asyncio
async def test_update_job_source_transcription_persists_provider_json_payloads(default_session_factory):
documents = DocumentService(session_factory=default_session_factory)
jobs = JobService(session_factory=default_session_factory)
transcriptions = TranscriptionService(session_factory=default_session_factory)
document = await documents.create_document(Document(id=uuid4(), name="provider-payloads-doc"))
job = await jobs.create_job(Job(document_id=document.id))
source = await transcriptions.create_source(
Source(
document_id=document.id,
page_number=1,
upload_name="provider.jpg",
filename="provider.jpg",
file_path="uploads/provider.jpg",
)
)
await transcriptions.create_job_source(
JobSource(job_id=job.id, source_id=source.id, status=JobSourceStatus.PENDING)
)
metadata = {"finish_reason": "stop", "usage": {"input_tokens": 11, "output_tokens": 22, "total_tokens": 33}}
raw_payload = {"id": "resp_xyz", "choices": [{"message": {"content": "provider transcript"}}]}
await transcriptions.update_job_source_transcription(
job_id=job.id,
source_id=source.id,
text="provider transcript",
ai_metadata=metadata,
raw_api_response=raw_payload,
provider="openrouter",
model="test-model",
prompt_name="transcribe_document.md",
)
stored_rows = await transcriptions.list_job_sources(job_id=job.id)
assert len(stored_rows) == 1
assert stored_rows[0].raw_transcription == "provider transcript"
assert stored_rows[0].ai_metadata == metadata
assert stored_rows[0].raw_api_response == raw_payload