diff --git a/apps/worker/app/services/document_ingestion/processing_billing.py b/apps/worker/app/services/document_ingestion/processing_billing.py index 6fd6cc321..51c658ea6 100644 --- a/apps/worker/app/services/document_ingestion/processing_billing.py +++ b/apps/worker/app/services/document_ingestion/processing_billing.py @@ -4,7 +4,10 @@ from datetime import datetime from app.services.document_ingestion.page_estimator import WorkloadEstimate -from app.services.document_ingestion.processing_context import ParseJobContext +from app.services.document_ingestion.processing_context import ( + ParseJobContext, + persist_job_metadata_updates, +) from loguru import logger from sqlalchemy import select @@ -145,13 +148,8 @@ def record_processing_start( ) if extra_metadata is not None: metadata_updates.update(extra_metadata) - with get_sync_db_context() as db: - job_result = db.execute(select(Job).where(Job.job_id == job_id).with_for_update()) - job = job_result.scalar_one_or_none() - if job: - job.job_metadata = { - **dict(job.job_metadata or {}), - **metadata_updates, - } - job_context.metadata_service.update_metadata(job_id, metadata_updates) - job_context.job_metadata.update(metadata_updates) + persist_job_metadata_updates( + job_id=job_id, + job_context=job_context, + metadata_updates=metadata_updates, + ) diff --git a/apps/worker/app/services/document_ingestion/processing_context.py b/apps/worker/app/services/document_ingestion/processing_context.py index 9dfe87945..0d8e51728 100644 --- a/apps/worker/app/services/document_ingestion/processing_context.py +++ b/apps/worker/app/services/document_ingestion/processing_context.py @@ -27,6 +27,31 @@ class ParseJobContext: s3_key: str +def persist_job_metadata_updates( + *, + job_id: str, + job_context: ParseJobContext, + metadata_updates: dict[str, object], +) -> None: + """Merge metadata into Redis and the durable Job row. + + Redis remains the live working copy during processing. The Job row is + the durable record after Redis TTL expires, so token usage and other + stage fields must be written here as well. + """ + if not metadata_updates: + return + with get_sync_db_context() as db: + job = _select_job_row_for_update(db, job_id) + if job is not None: + job.job_metadata = { + **dict(job.job_metadata or {}), + **metadata_updates, + } + job_context.metadata_service.update_metadata(job_id, metadata_updates) + job_context.job_metadata.update(metadata_updates) + + def load_parse_job_context( job_id: str, requested_user_id: str | None, @@ -111,3 +136,9 @@ def _load_job_row(job_id: str) -> Job | None: def _select_job_row(db: Session, job_id: str) -> Job | None: return db.execute(select(Job).where(Job.job_id == job_id)).scalar_one_or_none() + + +def _select_job_row_for_update(db: Session, job_id: str) -> Job | None: + return db.execute( + select(Job).where(Job.job_id == job_id).with_for_update() + ).scalar_one_or_none() diff --git a/apps/worker/app/services/document_ingestion/processing_run.py b/apps/worker/app/services/document_ingestion/processing_run.py index cd5bcfce0..4a35f645b 100644 --- a/apps/worker/app/services/document_ingestion/processing_run.py +++ b/apps/worker/app/services/document_ingestion/processing_run.py @@ -17,6 +17,7 @@ from app.services.document_ingestion.processing_context import ( ParseJobContext, load_parse_job_context, + persist_job_metadata_updates, ) from app.services.document_ingestion.source_preparation import prepare_source_file from app.services.document_ingestion.success_finalization import finalize_parse_success @@ -199,10 +200,21 @@ def _run_parse_job( result_storage_factory=get_result_storage, ) finally: - job_context.job_metadata["stages"] = { + stages = { "timing_ms": dict(stage_timing_dict), "token_usage": dict(token_usage_dict), } + job_context.job_metadata["stages"] = stages + try: + persist_job_metadata_updates( + job_id=job_id, + job_context=job_context, + metadata_updates={"stages": stages}, + ) + except Exception as exc: + logger.warning( + f"Failed to persist processing stages: job_id={job_id}, error={exc}" + ) cleanup_llm_overrides() cleanup_token_tracker() cleanup_stage_tracker() diff --git a/apps/worker/app/services/document_ingestion/success_finalization.py b/apps/worker/app/services/document_ingestion/success_finalization.py index 9acc0e3ab..42330dbcc 100644 --- a/apps/worker/app/services/document_ingestion/success_finalization.py +++ b/apps/worker/app/services/document_ingestion/success_finalization.py @@ -18,7 +18,10 @@ build_generated_result_package, ) from app.services.document_ingestion.artifact_refs import collect_referenced_artifact_refs -from app.services.document_ingestion.processing_context import ParseJobContext +from app.services.document_ingestion.processing_context import ( + ParseJobContext, + persist_job_metadata_updates, +) from loguru import logger from shared.models.schemas.job_metadata import JobMetadataHelper @@ -210,8 +213,11 @@ def _record_processing_completion( _refresh_processing_stages(job_context) if "stages" in job_context.job_metadata: processing_timing_updates["stages"] = job_context.job_metadata["stages"] - job_context.metadata_service.update_metadata(job_id, processing_timing_updates) - job_context.job_metadata.update(processing_timing_updates) + persist_job_metadata_updates( + job_id=job_id, + job_context=job_context, + metadata_updates=processing_timing_updates, + ) def _generate_result_package( diff --git a/apps/worker/tests/unit/test_processing_metadata_persist.py b/apps/worker/tests/unit/test_processing_metadata_persist.py new file mode 100644 index 000000000..4cd3d8200 --- /dev/null +++ b/apps/worker/tests/unit/test_processing_metadata_persist.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +import os +from types import SimpleNamespace +from unittest.mock import Mock + +os.environ.setdefault("DATABASE_URL", "postgresql+asyncpg://test:test@localhost/test") +os.environ.setdefault("TMP_PATH", "/tmp/knowhere-test") +os.environ.setdefault("S3_BUCKET_NAME", "test-uploads") +os.environ.setdefault("S3_ACCESS_KEY_ID", "test") +os.environ.setdefault("S3_SECRET_ACCESS_KEY", "test") +os.environ.setdefault("S3_TEMP_PATH", "/tmp") + +import app.services.document_ingestion.processing_context as processing_context # noqa: E402 +import app.services.document_ingestion.success_finalization as success_finalization # noqa: E402 + + +class _FakeDbContext: + def __init__(self, session: object) -> None: + self.session = session + + def __enter__(self) -> object: + return self.session + + def __exit__(self, *_args: object) -> bool: + return False + + +def _job_context(*, metadata: dict[str, object] | None = None) -> processing_context.ParseJobContext: + return processing_context.ParseJobContext( + job_metadata=metadata or {"namespace": "default"}, + job_user_id="user-1", + metadata_service=Mock(), + redis_service=object(), + s3_key="uploads/job.pdf", + ) + + +def test_persist_job_metadata_updates_merges_stages_into_job_row( + monkeypatch: object, +) -> None: + job = SimpleNamespace(job_metadata={"namespace": "default", "page_count": 12}) + session = Mock() + monkeypatch.setattr( + processing_context, + "get_sync_db_context", + lambda: _FakeDbContext(session), + ) + monkeypatch.setattr( + processing_context, + "_select_job_row_for_update", + lambda _db, _job_id: job, + ) + job_context = _job_context() + stages = { + "token_usage": {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14, "calls": 1}, + "timing_ms": {"worker.parse.document": 1200}, + } + + processing_context.persist_job_metadata_updates( + job_id="job_abc", + job_context=job_context, + metadata_updates={"stages": stages}, + ) + + assert job.job_metadata["namespace"] == "default" + assert job.job_metadata["page_count"] == 12 + assert job.job_metadata["stages"]["token_usage"]["total_tokens"] == 14 + job_context.metadata_service.update_metadata.assert_called_once_with( + "job_abc", + {"stages": stages}, + ) + assert job_context.job_metadata["stages"] == stages + + +def test_record_processing_completion_persists_token_usage_to_job_row( + monkeypatch: object, +) -> None: + persist = Mock() + monkeypatch.setattr(success_finalization, "persist_job_metadata_updates", persist) + monkeypatch.setattr( + success_finalization, + "get_current_token_tracker", + lambda: {"prompt_tokens": 8, "completion_tokens": 2, "total_tokens": 10, "calls": 1}, + ) + monkeypatch.setattr( + success_finalization, + "get_current_stage_tracker", + lambda: {"worker.parse.document": 900}, + ) + from datetime import datetime, timezone + + job_context = _job_context() + started = datetime(2026, 8, 24, 10, 0, tzinfo=timezone.utc) + + success_finalization._record_processing_completion( + job_id="job_abc", + job_context=job_context, + processing_started_at=started, + ) + + persist.assert_called_once() + updates = persist.call_args.kwargs["metadata_updates"] + assert updates["stages"]["token_usage"]["total_tokens"] == 10 + assert "processing_completed_at" in updates + assert "processing_duration_ms" in updates diff --git a/packages/shared-python/shared/models/schemas/job_metadata.py b/packages/shared-python/shared/models/schemas/job_metadata.py index 4efb808b3..0162a257a 100644 --- a/packages/shared-python/shared/models/schemas/job_metadata.py +++ b/packages/shared-python/shared/models/schemas/job_metadata.py @@ -36,6 +36,10 @@ class JobMetadataBase(BaseModel): page_memory_config: Optional[Dict[str, Any]] = Field( None, description="Resolved page-memory worker configuration" ) + stages: Optional[Dict[str, Any]] = Field( + None, + description="Worker processing stages, including token_usage and timing_ms", + ) # result_mode was removed and is no longer supported. # Source-file fields.