generated from john/python-template
AI metadata and api prompt results data capture now fixed
This commit is contained in:
@@ -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):
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,11 +60,13 @@ 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 current_status == JobStatus.QUEUED:
|
||||||
if session is None:
|
if session is None:
|
||||||
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING)
|
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING)
|
||||||
else:
|
else:
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user