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."""
from dataclasses import dataclass
from typing import Any
from typing import Protocol
@@ -28,6 +29,8 @@ class TranscriptionResult:
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
class TranscriptionProvider(Protocol):
+70
View File
@@ -66,6 +66,13 @@ class OpenRouterTranscriptionProvider:
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)
logger.info("OpenRouter transcription completed using model=%s", model)
return TranscriptionResult(
text=text,
@@ -76,8 +83,71 @@ class OpenRouterTranscriptionProvider:
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,
)
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:
image_b64 = base64.b64encode(image_bytes).decode("ascii")
data_url = f"data:{mime_type};base64,{image_b64}"
@@ -415,6 +415,8 @@ class TranscriptionService(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,
provider: str | None = None,
model: str | None = None,
prompt_name: str = DEFAULT_PROMPT_FILE,
@@ -462,11 +464,15 @@ class TranscriptionService(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,
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.error_detail = error_detail
job_source.status = JobSourceStatus.TRANSCRIBED if text is not None else JobSourceStatus.FAILED
job_source.executed_at = datetime.now(UTC)
+25 -2
View File
@@ -29,9 +29,13 @@ async def advance_job(
) -> Job | None:
"""Advance a single job by lifecycle status."""
settings = settings or get_settings()
match job.status:
current_status = _coerce_job_status(job.status)
match current_status:
case JobStatus.QUEUED:
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:
if job.retry_count < settings.worker_max_retries:
return await services.jobs.update_job_state(
@@ -56,11 +60,13 @@ async def process_queued_job(
) -> Job | None:
"""Process one complete transcription attempt for a queued job."""
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}")
return
# Transaction A: claim job for processing.
if current_status == JobStatus.QUEUED:
if session is None:
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING)
else:
@@ -376,6 +382,8 @@ async def _finalize_batch_outcome(
source_id=source.id,
text=result.text,
error_detail=None,
ai_metadata=result.ai_metadata,
raw_api_response=result.raw_api_response,
provider=result.provider,
model=result.model,
prompt_name=result.prompt_name,
@@ -402,6 +410,8 @@ async def _finalize_batch_outcome(
source_id=source.id,
text=result.text,
error_detail=None,
ai_metadata=result.ai_metadata,
raw_api_response=result.raw_api_response,
provider=result.provider,
model=result.model,
prompt_name=result.prompt_name,
@@ -470,3 +480,16 @@ def _line_count(text: str) -> int:
if not stripped:
return 0
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",
model="test-model",
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(
@@ -94,6 +96,11 @@ class TestPipelineSuccessFlow:
assert job is not None
assert job.status == JobStatus.TRANSCRIBED
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)
@pytest.mark.asyncio
+12 -1
View File
@@ -78,7 +78,13 @@ class TestOpenRouterProviderTranscribe:
"""Transcribe returns normalized text from a valid response payload."""
response = {
"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(
settings=Settings(openrouter_api_key="test-key"),
@@ -94,6 +100,11 @@ class TestOpenRouterProviderTranscribe:
assert result.text == "Line 1\nLine 2"
assert result.provider == "openrouter"
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
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):
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