diff --git a/apps/worker/app/services/document_ingestion/parse_execution.py b/apps/worker/app/services/document_ingestion/parse_execution.py index 594ff60fe..a65856d33 100644 --- a/apps/worker/app/services/document_ingestion/parse_execution.py +++ b/apps/worker/app/services/document_ingestion/parse_execution.py @@ -17,7 +17,7 @@ def execute_document_parse( prepared_source: PreparedSourceFile, output_dir: str, ) -> ParseOutput: - """Run the parser adapter for a prepared source file.""" + """Run worker document parsing for a prepared local source file.""" doc_type = JobMetadataHelper.get_parsing_param( job_context.job_metadata, "doc_type", diff --git a/apps/worker/app/services/document_ingestion/processing_billing.py b/apps/worker/app/services/document_ingestion/processing_billing.py index f0d313fa4..94a3d77ad 100644 --- a/apps/worker/app/services/document_ingestion/processing_billing.py +++ b/apps/worker/app/services/document_ingestion/processing_billing.py @@ -24,6 +24,27 @@ class ParseJobBillingSnapshot: billing_status: str +def record_skipped_parse_job_billing( + *, + job_id: str, + workload_estimate: WorkloadEstimate, +) -> ParseJobBillingSnapshot: + page_count = workload_estimate.page_count + 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.page_count = page_count + job.credits_charged = 0 + job.billing_status = "skipped" + + return ParseJobBillingSnapshot( + billing_amount_micro_dollars=0, + billing_credits=0.0, + billing_status="skipped", + ) + + def charge_parse_job_pages( *, job_id: str, @@ -121,5 +142,13 @@ def record_processing_start( metadata_updates["workload_estimate_fallback_reason"] = ( workload_estimate.fallback_reason ) + 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) diff --git a/apps/worker/app/services/document_ingestion/processing_context.py b/apps/worker/app/services/document_ingestion/processing_context.py index 51a40cd43..9dfe87945 100644 --- a/apps/worker/app/services/document_ingestion/processing_context.py +++ b/apps/worker/app/services/document_ingestion/processing_context.py @@ -1,24 +1,21 @@ from __future__ import annotations -import os from dataclasses import dataclass from typing import Any from loguru import logger from sqlalchemy import select +from sqlalchemy.orm import Session -from shared.core.config import settings from shared.core.database_sync import get_sync_db_context from shared.core.exceptions.domain_exceptions import ( NotFoundException, - ValidationException, ) from shared.models.database.job import Job from shared.services.redis.redis_sync_service import ( SyncJobInfoRedisService, SyncJobMetadataService, ) -from shared.services.storage.job_file_storage import JobFileStorage @dataclass(frozen=True) @@ -37,15 +34,13 @@ def load_parse_job_context( ) -> ParseJobContext: job_info_service = SyncJobInfoRedisService(redis_service) job_info = job_info_service.get_job_info(job_id) + job_row: Job | None = None if not job_info: logger.warning( f"JobInfo not found in Redis for job_id={job_id}; falling back to database" ) - with get_sync_db_context() as fallback_db: - job_row = fallback_db.execute( - select(Job).where(Job.job_id == job_id) - ).scalar_one_or_none() + job_row = _load_job_row(job_id) if not job_row or not job_row.s3_key: raise NotFoundException( @@ -58,6 +53,17 @@ def load_parse_job_context( job_user_id: str | None = ( str(job_row.user_id) if job_row.user_id else requested_user_id ) + job_info_service.save_job_info( + job_id, + { + "job_id": job_id, + "s3_key": s3_key, + "user_id": job_user_id, + "webhook_enabled": bool(job_row.webhook_enabled), + "job_type": "document_ingestion", + "source_type": job_row.source_type, + }, + ) logger.info(f"Recovered JobInfo from database: job_id={job_id}, s3_key={s3_key}") else: raw_s3_key = job_info.get("s3_key") @@ -77,11 +83,17 @@ def load_parse_job_context( metadata_service = SyncJobMetadataService(redis_service) raw_job_metadata = metadata_service.get_metadata(job_id) if not isinstance(raw_job_metadata, dict) or not raw_job_metadata: - raise NotFoundException( - resource="JobMetadata", - resource_id=job_id, - internal_message=f"Job metadata not found for job_id={job_id}", - ) + job_row = job_row or _load_job_row(job_id) + raw_job_metadata = job_row.job_metadata if job_row else None + if isinstance(raw_job_metadata, dict) and raw_job_metadata: + metadata_service.save_metadata(job_id, raw_job_metadata) + logger.info(f"Recovered JobMetadata from database: job_id={job_id}") + else: + raise NotFoundException( + resource="JobMetadata", + resource_id=job_id, + internal_message=f"Job metadata not found for job_id={job_id}", + ) return ParseJobContext( job_metadata=dict(raw_job_metadata), @@ -92,30 +104,10 @@ def load_parse_job_context( ) -def assert_source_file_within_size_limit(s3_key: str) -> None: - file_info = JobFileStorage().verify_upload_exists(s3_key) - if not file_info.get("exists"): - raise NotFoundException( - resource="S3File", - resource_id=s3_key, - internal_message=f"S3 file not found: {s3_key}", - ) +def _load_job_row(job_id: str) -> Job | None: + with get_sync_db_context() as fallback_db: + return _select_job_row(fallback_db, job_id) - logger.info(f"S3 file verified: {s3_key}") - - file_size = file_info.get("size", 0) - file_extension = os.path.splitext(s3_key)[1].lower() - if file_size > settings.MAX_FILE_SIZE: - limit_mb = settings.MAX_FILE_SIZE // (1024 * 1024) - raise ValidationException( - user_message=f"File size exceeds limit (max {limit_mb}MB for {file_extension})", - violations=[ - { - "field": "file_size", - "description": ( - f"Size {file_size} bytes exceeds limit of " - f"{settings.MAX_FILE_SIZE} bytes" - ), - } - ], - ) + +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() diff --git a/apps/worker/app/services/document_ingestion/processing_run.py b/apps/worker/app/services/document_ingestion/processing_run.py index 995990379..69072d51d 100644 --- a/apps/worker/app/services/document_ingestion/processing_run.py +++ b/apps/worker/app/services/document_ingestion/processing_run.py @@ -11,22 +11,21 @@ from app.services.document_ingestion.parse_execution import execute_document_parse from app.services.document_ingestion.processing_billing import ( charge_parse_job_pages, + record_skipped_parse_job_billing, record_processing_start, ) from app.services.document_ingestion.processing_context import ( ParseJobContext, - assert_source_file_within_size_limit, load_parse_job_context, ) from app.services.document_ingestion.source_preparation import prepare_source_file from app.services.document_ingestion.success_finalization import finalize_parse_success from app.services.document_ingestion.workspace import ( TemporaryParseWorkspace, - cleanup_task_workspace, - download_s3_file_to_temp, ) from loguru import logger +from shared.core.exceptions.domain_exceptions import ValidationException from shared.services.jobs.lifecycle.service import get_sync_job_lifecycle_service from shared.services.redis.distributed_lock import RedisJobLock from shared.services.redis.redis_sync_service import ( @@ -44,7 +43,6 @@ def execute(self, job_id: str, user_id: str | None) -> dict[str, object]: redis_service = SyncRedisServiceFactory.get_service() job_context = load_parse_job_context(job_id, user_id, redis_service) - assert_source_file_within_size_limit(job_context.s3_key) should_process = mark_job_running(job_id, job_context.redis_service) if not should_process: @@ -65,7 +63,7 @@ def execute(self, job_id: str, user_id: str | None) -> dict[str, object]: task_workspace=task_workspace, ) finally: - task_workspace.cleanup(cleanup_task_workspace) + task_workspace.cleanup() return result @@ -83,7 +81,6 @@ def _run_parse_job( job_id=job_id, job_context=job_context, input_dir=task_workspace.input_dir, - download_source_file=download_s3_file_to_temp, ) workload_estimate = PageEstimator.estimate_workload(prepared_source.local_file_path) @@ -96,6 +93,23 @@ def _run_parse_job( ) processing_started_at = datetime.now(timezone.utc) + if _is_pdf_page_limit_exceeded( + file_extension=prepared_source.file_extension, + page_count=page_count, + ): + billing_snapshot = record_skipped_parse_job_billing( + job_id=job_id, + workload_estimate=workload_estimate, + ) + record_processing_start( + job_id=job_id, + job_context=job_context, + billing_snapshot=billing_snapshot, + processing_started_at=processing_started_at, + workload_estimate=workload_estimate, + ) + _raise_pdf_page_limit_exceeded(page_count) + billing_snapshot = charge_parse_job_pages( job_id=job_id, filename=prepared_source.source_file_name, @@ -146,3 +160,27 @@ def _run_parse_job( task_workspace_dir=task_workspace.root_dir, result_storage_factory=get_result_storage, ) + + +def _is_pdf_page_limit_exceeded(*, file_extension: str, page_count: int) -> bool: + from shared.core.config import settings + + return file_extension == ".pdf" and page_count > settings.MAX_PDF_PAGE_LIMIT + + +def _raise_pdf_page_limit_exceeded(page_count: int) -> None: + from shared.core.config import settings + + pdf_page_limit = settings.MAX_PDF_PAGE_LIMIT + raise ValidationException( + user_message=( + f"Document too large: {page_count} pages exceeds the {pdf_page_limit}-page limit. " + "Please split the document and upload in smaller batches." + ), + violations=[ + { + "field": "page_count", + "description": f"PDF has {page_count} pages, limit is {pdf_page_limit}", + } + ], + ) diff --git a/apps/worker/app/services/document_ingestion/source_preparation.py b/apps/worker/app/services/document_ingestion/source_preparation.py index cc3c1bd79..0b542d022 100644 --- a/apps/worker/app/services/document_ingestion/source_preparation.py +++ b/apps/worker/app/services/document_ingestion/source_preparation.py @@ -1,19 +1,21 @@ from __future__ import annotations import os -from collections.abc import Callable from dataclasses import dataclass from app.services.document_ingestion.processing_context import ParseJobContext -from app.services.document_ingestion.workspace import download_s3_file_to_temp from app.services.document_parser.support.internal_parse_name import ( prepare_internal_parse_input, ) from loguru import logger +from shared.core.config import settings +from shared.core.exceptions.domain_exceptions import ( + NotFoundException, + ValidationException, +) from shared.models.schemas.job_metadata import JobMetadataHelper - -DownloadSourceFile = Callable[[str, str, str], str] +from shared.services.storage.job_file_storage import JobFileStorage @dataclass(frozen=True) @@ -29,9 +31,8 @@ def prepare_source_file( job_id: str, job_context: ParseJobContext, input_dir: str, - download_source_file: DownloadSourceFile = download_s3_file_to_temp, ) -> PreparedSourceFile: - """Download and normalize the source file before parser execution.""" + """Verify, download, and normalize the uploaded source file.""" source_file_name = JobMetadataHelper.get_source_file_name( job_context.job_metadata, ) or os.path.basename(job_context.s3_key) @@ -39,10 +40,12 @@ def prepare_source_file( os.path.splitext(job_context.s3_key)[1].lower() if job_context.s3_key else "" ) - local_file_path = download_source_file( + storage = JobFileStorage() + _assert_source_file_within_size_limit(storage, job_context.s3_key) + local_file_path = storage.download_upload_to_temp( job_context.s3_key, - file_extension, - input_dir, + suffix=file_extension, + temp_dir=input_dir, ) logger.info(f"File downloaded: job_id={job_id}, local_path={local_file_path}") @@ -64,3 +67,36 @@ def prepare_source_file( local_file_path=prepared_parse_input.file_path, file_extension=file_extension, ) + + +def _assert_source_file_within_size_limit( + storage: JobFileStorage, + s3_key: str, +) -> None: + file_info = storage.verify_upload_exists(s3_key) + if not file_info.get("exists"): + raise NotFoundException( + resource="S3File", + resource_id=s3_key, + internal_message=f"S3 file not found: {s3_key}", + ) + + logger.info(f"S3 file verified: {s3_key}") + + raw_file_size = file_info.get("size", 0) + file_size = raw_file_size if isinstance(raw_file_size, int) else 0 + file_extension = os.path.splitext(s3_key)[1].lower() + if file_size > settings.MAX_FILE_SIZE: + limit_mb = settings.MAX_FILE_SIZE // (1024 * 1024) + raise ValidationException( + user_message=f"File size exceeds limit (max {limit_mb}MB for {file_extension})", + violations=[ + { + "field": "file_size", + "description": ( + f"Size {file_size} bytes exceeds limit of " + f"{settings.MAX_FILE_SIZE} bytes" + ), + } + ], + ) diff --git a/apps/worker/app/services/document_ingestion/workspace.py b/apps/worker/app/services/document_ingestion/workspace.py index e8a0188dc..3dc65290b 100644 --- a/apps/worker/app/services/document_ingestion/workspace.py +++ b/apps/worker/app/services/document_ingestion/workspace.py @@ -3,7 +3,6 @@ import os import shutil import tempfile -from collections.abc import Callable from dataclasses import dataclass from loguru import logger @@ -14,9 +13,6 @@ SystemSettingInvalidException, SystemSettingMissingException, ) -from shared.services.storage.job_file_storage import JobFileStorage - -CleanupTaskWorkspace = Callable[[str | None], bool] @dataclass(frozen=True) @@ -37,12 +33,8 @@ def create(cls, job_id: str) -> "TemporaryParseWorkspace": logger.info(f"Task workspace ready: job_id={job_id}, workspace={root_dir}") return cls(root_dir=root_dir, input_dir=input_dir, output_dir=output_dir) - def cleanup( - self, - cleanup_workspace: CleanupTaskWorkspace | None = None, - ) -> bool: - resolved_cleanup = cleanup_workspace or cleanup_task_workspace - return resolved_cleanup(self.root_dir) + def cleanup(self) -> bool: + return cleanup_task_workspace(self.root_dir) def cleanup_temp_file(file_path: str | None) -> None: @@ -99,13 +91,3 @@ def create_task_workspace(job_id: str) -> str: internal_message=f"Failed to create task workspace in {temp_root}", original_exception=exc, ) from exc - - -def download_s3_file_to_temp(s3_key: str, file_ext: str, temp_dir: str) -> str: - """Download the source file from object storage into the task workspace.""" - storage = JobFileStorage() - return storage.download_upload_to_temp( - s3_key, - suffix=file_ext, - temp_dir=temp_dir, - ) diff --git a/apps/worker/app/services/workload/url_upload_context.py b/apps/worker/app/services/workload/url_upload_context.py index af7fb807b..bcb400662 100644 --- a/apps/worker/app/services/workload/url_upload_context.py +++ b/apps/worker/app/services/workload/url_upload_context.py @@ -3,7 +3,13 @@ from dataclasses import dataclass from typing import Any +from loguru import logger +from sqlalchemy import select +from sqlalchemy.orm import Session + +from shared.core.database_sync import get_sync_db_context from shared.core.exceptions.domain_exceptions import NotFoundException +from shared.models.database.job import Job from shared.services.redis.redis_sync_service import ( SyncJobInfoRedisService, SyncJobMetadataService, @@ -18,25 +24,49 @@ class UrlUploadContext: def load_url_upload_context(job_id: str, redis_service: Any) -> UrlUploadContext: job_info_service = SyncJobInfoRedisService(redis_service) job_info = job_info_service.get_job_info(job_id) + metadata_service = SyncJobMetadataService(redis_service) if job_info: raw_s3_key = job_info.get("s3_key") else: - metadata_service = SyncJobMetadataService(redis_service) job_metadata = metadata_service.get_metadata(job_id) - if not job_metadata: - raise NotFoundException( - resource="JobInfo", - resource_id=job_id, - internal_message="Job info not found in Redis or Metadata", + raw_s3_key = job_metadata.get("s3_key") if job_metadata else None + + if not raw_s3_key: + logger.warning( + f"URL upload JobInfo missing in Redis for job_id={job_id}; falling back to database" + ) + job_row = _load_job_row(job_id) + if job_row and job_row.s3_key: + raw_s3_key = str(job_row.s3_key) + job_info_service.save_job_info( + job_id, + { + "job_id": job_id, + "s3_key": raw_s3_key, + "user_id": str(job_row.user_id) if job_row.user_id else None, + "webhook_enabled": bool(job_row.webhook_enabled), + "job_type": "document_ingestion", + "source_type": job_row.source_type, + }, ) - raw_s3_key = job_metadata.get("s3_key") + if isinstance(job_row.job_metadata, dict) and job_row.job_metadata: + metadata_service.save_metadata(job_id, job_row.job_metadata) if not raw_s3_key: raise NotFoundException( resource="JobInfo", resource_id="s3_key", - internal_message=f"Missing s3_key in Redis job info for job_id={job_id}", + internal_message=f"Missing s3_key in Redis or database for job_id={job_id}", ) return UrlUploadContext(s3_key=str(raw_s3_key)) + + +def _load_job_row(job_id: str) -> Job | None: + with get_sync_db_context() as db: + return _select_job_row(db, job_id) + + +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() diff --git a/apps/worker/app/services/workload/url_upload_service.py b/apps/worker/app/services/workload/url_upload_service.py index 27cf81616..e5dec5979 100644 --- a/apps/worker/app/services/workload/url_upload_service.py +++ b/apps/worker/app/services/workload/url_upload_service.py @@ -30,18 +30,27 @@ def upload_url_file( upload_context = load_url_upload_context(job_id, redis_service) lifecycle_service.update_progress( - job_id, progress=3, message="Validating URL file type..." + job_id, + progress=3, + message="Validating URL file type...", + redis_service=redis_service, ) file_extension = resolve_supported_url_extension(source_url) lifecycle_service.update_progress( - job_id, progress=10, message="Downloading file from URL..." + job_id, + progress=10, + message="Downloading file from URL...", + redis_service=redis_service, ) temp_file_path = download_source_url_to_temp(source_url) try: lifecycle_service.update_progress( - job_id, progress=30, message="Validating file size..." + job_id, + progress=30, + message="Validating file size...", + redis_service=redis_service, ) assert_temp_file_within_size_limit( temp_file_path=temp_file_path, @@ -49,7 +58,10 @@ def upload_url_file( ) lifecycle_service.update_progress( - job_id, progress=50, message="Uploading file to S3..." + job_id, + progress=50, + message="Uploading file to S3...", + redis_service=redis_service, ) upload_temp_file_to_source_storage( temp_file_path=temp_file_path, @@ -60,7 +72,10 @@ def upload_url_file( cleanup_temp_file(temp_file_path) lifecycle_service.update_progress( - job_id, progress=80, message="Verifying upload result..." + job_id, + progress=80, + message="Verifying upload result...", + redis_service=redis_service, ) file_info = verify_source_upload(upload_context.s3_key) @@ -68,6 +83,7 @@ def upload_url_file( job_id, progress=100, message="URL file upload complete, waiting for processing...", + redis_service=redis_service, ) logger.info( "URL file upload complete, waiting for S3 webhook: " diff --git a/apps/worker/tests/contract/conftest.py b/apps/worker/tests/contract/conftest.py index 300d1c2ba..c1550b145 100644 --- a/apps/worker/tests/contract/conftest.py +++ b/apps/worker/tests/contract/conftest.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import importlib import os import sys from collections.abc import Generator @@ -58,6 +59,13 @@ def worker_contract_environment( sys.path.insert(0, worker_root_value) contract_runtime.clear_application_modules() + from shared.core.celery_app import get_celery_app + + celery_app = get_celery_app() + monkeypatch.setattr(celery_app.conf, "task_always_eager", True) + monkeypatch.setattr(celery_app.conf, "task_eager_propagates", False) + importlib.import_module("app.core.tasks.document_ingestion_tasks") + try: yield finally: diff --git a/apps/worker/tests/contract/test_parse_task_contract.py b/apps/worker/tests/contract/test_parse_task_contract.py index 7fd6f5684..494c6f8ef 100644 --- a/apps/worker/tests/contract/test_parse_task_contract.py +++ b/apps/worker/tests/contract/test_parse_task_contract.py @@ -1,26 +1,16 @@ from __future__ import annotations -import json -import shutil -import zipfile from concurrent.futures import ThreadPoolExecutor from pathlib import Path -from threading import Barrier -from types import SimpleNamespace -from typing import Any from uuid import uuid4 -import pandas as pd import pytest -from pytest import MonkeyPatch -from sqlalchemy import text -from sqlalchemy.engine import Engine -from support.contract_database import insert_contract_job, insert_contract_user +from support.worker_parse_contract import WorkerParseContract _REPO_ROOT: Path = Path(__file__).resolve().parents[4] _FIXTURES_ROOT: Path = _REPO_ROOT / "apps" / "worker" / "tests" / "fixtures" -_SAMPLE_PDF_PATH: Path = _FIXTURES_ROOT / "sample_3pages.pdf" +_SAMPLE_XLSX_PATH: Path = _FIXTURES_ROOT / "sample_100rows.xlsx" def _write_blank_pdf(file_path: Path, page_count: int) -> None: @@ -34,1460 +24,347 @@ def _write_blank_pdf(file_path: Path, page_count: int) -> None: writer.write(pdf_file) -def _build_pending_file_job_metadata(source_file_name: str) -> dict[str, Any]: - job_metadata: dict[str, Any] = { - "namespace": "worker-contract", - "source_type": "file", - "source_file_name": source_file_name, - "parsing_params": {"kb_dir": "legacy-ignored"}, - } - return job_metadata - - -def _load_parse_task_modules() -> tuple[Any, Any, Any, Engine, Any, Any, Any]: - import app.core.tasks.document_ingestion_tasks as document_ingestion_tasks - import app.services.document_ingestion.processing_run as parse_job_service - import app.services.document_parser.parse_service as parse_service - from shared.core.database_sync import get_sync_engine - from shared.services.redis.redis_sync_service import ( - SyncJobInfoRedisService, - SyncJobMetadataService, - SyncRedisServiceFactory, - ) +def test_parse_task_should_process_uploaded_file_through_real_contract_boundaries( + worker_contract_environment: None, + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=False) - return ( - document_ingestion_tasks, - parse_service, - parse_job_service, - get_sync_engine(), - SyncJobInfoRedisService, - SyncJobMetadataService, - SyncRedisServiceFactory, + job = contract.create_file_job( + source_file_name="contract-real.xlsx", + job_id_prefix="job_parse_real", + ) + contract.upload_source_file( + local_file_path=_SAMPLE_XLSX_PATH, + s3_key=job["s3_key"], ) + celery_result = contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=job["user_id"], + ) -def _load_worker_settings() -> Any: - from shared.core.config import settings - - return settings - - -def _save_worker_task_cache( - *, - job_id: str, - user_id: str, - s3_key: str, - metadata: dict[str, Any], - sync_job_info_service_cls: Any, - sync_job_metadata_service_cls: Any, - sync_redis_service_factory: Any, -) -> Any: - redis_service = sync_redis_service_factory.get_service() - sync_job_info_service = sync_job_info_service_cls(redis_service) - sync_job_metadata_service = sync_job_metadata_service_cls(redis_service) - - sync_job_info_service.save_job_info( - job_id, - { - "job_id": job_id, - "s3_key": s3_key, - "user_id": user_id, - "webhook_enabled": False, - "job_type": "document_ingestion", - "source_type": "file", - }, + assert celery_result.successful() + assert celery_result.result["status"] == "success" + assert celery_result.result["job_id"] == job["job_id"] + assert celery_result.result["delivery_mode"] == "url" + assert celery_result.result["result_s3_key"] == contract.storage.build_result_zip_key( + job_id=job["job_id"] ) - sync_job_metadata_service.save_metadata(job_id, metadata) - return redis_service + observed = contract.observe_successful_job(job["job_id"]) + job_row = observed["job"] + result_row = observed["result"] + job_chunks = observed["job_chunks"] + document_chunks = observed["document_chunks"] -def _patch_verify_upload_exists( - monkeypatch: MonkeyPatch, - file_info_for_storage_key: Any, -) -> None: - from shared.services.storage.job_file_storage import JobFileStorage + assert job_row["status"] == "done" + assert job_row["billing_status"] == "skipped" + assert job_row["page_count"] and job_row["page_count"] > 0 + assert job_row["error_message"] is None - monkeypatch.setattr( - JobFileStorage, - "verify_upload_exists", - lambda self, storage_key: file_info_for_storage_key(storage_key), + assert result_row["document_id"] + assert result_row["result_s3_key"] == contract.storage.build_result_zip_key( + job_id=job["job_id"] ) - - -def _find_task_workspaces(root: Path, job_id: str) -> list[Path]: - return sorted( - path - for path in root.iterdir() - if path.is_dir() and path.name.startswith(f"document_ingestion_task_{job_id}_") + assert result_row["result_size"] and result_row["result_size"] > 0 + + assert len(job_chunks) > 0 + assert len(document_chunks) == len(job_chunks) + assert observed["document_sections_count"] > 0 + assert all(row["chunk_type"] == "table" for row in job_chunks) + assert any("tables/" in str(row["path"]) for row in job_chunks) + + assert contract.get_task_status(job["job_id"]) == "done" + task_progress = contract.get_task_progress(job["job_id"]) + assert task_progress["progress"] == 100 + assert task_progress["message"] == "Task complete!" + + result_file_info = contract.verify_result_zip_object(result_row["result_s3_key"]) + assert result_file_info["exists"] is True + assert result_file_info["size"] == result_row["result_size"] + + result_zip = contract.read_result_zip( + result_s3_key=result_row["result_s3_key"], + tmp_path=tmp_path, ) + assert {"chunks.json", "doc_nav.json", "manifest.json"}.issubset( + result_zip["members"] + ) + assert any(member.startswith("tables/") for member in result_zip["members"]) + chunks_payload = result_zip["chunks"] + assert len(chunks_payload["chunks"]) == len(job_chunks) + assert all(chunk["type"] == "table" for chunk in chunks_payload["chunks"]) + assert any( + "Table summary:" in chunk["content"] for chunk in chunks_payload["chunks"] + ) -def _build_fake_parse_output(*, output_dir: Path, rows: list[dict[str, Any]]) -> Any: - from app.services.document_parser.orchestration.parse_output import ParseOutput - - return ParseOutput(output_dir=str(output_dir), parsed_df=pd.DataFrame(rows)) - + manifest_payload = result_zip["manifest"] + assert manifest_payload["source_file_name"] == job["source_file_name"] + assert manifest_payload["statistics"]["total_chunks"] == len(job_chunks) -def _bind_parse_task_to_current_module( - monkeypatch: MonkeyPatch, - *, - document_ingestion_tasks: Any, -) -> None: - monkeypatch.setitem( - document_ingestion_tasks.parse_task._orig_run.__globals__, - "_parse", - document_ingestion_tasks._parse, - ) - monkeypatch.setattr(document_ingestion_tasks.parse_task, "__trace__", None, raising=False) + assert contract.find_task_workspaces(tmp_path, job["job_id"]) == [] -@pytest.mark.parametrize( - ("billing_enabled", "expected_billing_status", "expected_transaction_types"), - [ - (True, "charged", ["initial_grant", "usage"]), - (False, "skipped", []), - ], -) -def test_should_parse_a_pending_file_job_and_persist_the_published_result_state( +def test_parse_task_should_charge_user_when_billing_is_enabled( worker_contract_environment: None, - monkeypatch: MonkeyPatch, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, - billing_enabled: bool, - expected_billing_status: str, - expected_transaction_types: list[str], ) -> None: - monkeypatch.setenv("BILLING_ENABLED", "true" if billing_enabled else "false") - ( - document_ingestion_tasks, - parse_service, - parse_job_service, - engine, - sync_job_info_service_cls, - sync_job_metadata_service_cls, - sync_redis_service_factory, - ) = _load_parse_task_modules() - settings = _load_worker_settings() - - user_id: str = f"worker-user-{uuid4().hex[:12]}" - job_id: str = f"job_parse_success_{uuid4().hex[:12]}" - source_file_name: str = "contract-parse.pdf" - s3_key: str = f"uploads/{job_id}.pdf" - text_content_with_refs: str = ( - "chunk-1 embeds [images/page-1.png] and [tables/table-1.html]" - ) - captured_artifacts: dict[str, Any] = {} - - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - job_metadata = _build_pending_file_job_metadata(source_file_name) - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - status="pending", - source_type="file", - s3_key=s3_key, - webhook_enabled=False, - job_metadata=job_metadata, - billing_status="pending", - ) + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=True) - redis_service = _save_worker_task_cache( - job_id=job_id, - user_id=user_id, - s3_key=s3_key, - metadata=job_metadata, - sync_job_info_service_cls=sync_job_info_service_cls, - sync_job_metadata_service_cls=sync_job_metadata_service_cls, - sync_redis_service_factory=sync_redis_service_factory, + job = contract.create_file_job( + source_file_name="contract-billing.xlsx", + job_id_prefix="job_parse_billing", + ) + contract.upload_source_file( + local_file_path=_SAMPLE_XLSX_PATH, + s3_key=job["s3_key"], ) - _bind_parse_task_to_current_module(monkeypatch, document_ingestion_tasks=document_ingestion_tasks) - monkeypatch.setattr(settings, "TMP_PATH", str(tmp_path)) - monkeypatch.setattr(settings, "BILLING_ENABLED", billing_enabled) - - def fake_verify_s3_file_exists(storage_key: str) -> dict[str, Any]: - return { - "exists": storage_key == s3_key, - "size": _SAMPLE_PDF_PATH.stat().st_size, - } - - _patch_verify_upload_exists(monkeypatch, fake_verify_s3_file_exists) - - def fake_download_s3_file_to_temp( - storage_key: str, file_ext: str, temp_dir: str - ) -> str: - assert storage_key == s3_key - assert file_ext == ".pdf" - downloaded_path = Path(temp_dir) / f"downloaded{file_ext}" - shutil.copy2(_SAMPLE_PDF_PATH, downloaded_path) - return str(downloaded_path) - - def fake_checkerboard_parse_output(**kwargs: Any) -> Any: - captured_artifacts["parse_kwargs"] = kwargs - output_dir = ( - Path(str(kwargs["output_dir"])) - / str(kwargs["internal_output_filename"]) - ) - images_dir = output_dir / "images" - tables_dir = output_dir / "tables" - images_dir.mkdir(parents=True, exist_ok=True) - tables_dir.mkdir(parents=True, exist_ok=True) - (output_dir / "full.md").write_text("body", encoding="utf-8") - (images_dir / "page-1.png").write_bytes(b"png") - (tables_dir / "table-1.html").write_text("
", encoding="utf-8") - - file_root = str(kwargs["internal_output_filename"]) - parsed_rows: list[dict[str, Any]] = [ - { - "content": text_content_with_refs, - "path": f"{file_root}/公司研究/自主可控加强,寒武纪或迎来营收快速放量周期", - "type": "text", - "length": len(text_content_with_refs), - "keywords": "", - "summary": "", - "know_id": "kid-1", - "tokens": "", - "connectto": json.dumps( - [ - { - "target": "table-1", - "relation": "embeds", - "ref": "[tables/table-1.html]", - } - ] - ), - "addtime": "now", - "page_nums": "1", - }, - { - "content": "chunk-2", - "path": f"{file_root}/相关研报/要点", - "type": "text", - "length": 7, - "keywords": "", - "summary": "", - "know_id": "kid-2", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "2", - }, - { - "content": "image caption", - "path": f"{file_root}/images/page-1.png", - "type": "image", - "length": 13, - "keywords": "", - "summary": "", - "know_id": "image-1", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "3", - }, - { - "content": "table content", - "path": f"{file_root}/tables/table-1.html", - "type": "table", - "length": 13, - "keywords": "", - "summary": "", - "know_id": "table-1", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "3", - }, - ] - return _build_fake_parse_output(output_dir=output_dir, rows=parsed_rows) - - class FakeResultStorage: - def upload(self, *, job_id: str, result_dir: str, zip_file_path: str) -> Any: - result_dir_path = Path(result_dir) - zip_path = Path(zip_file_path) - captured_artifacts["result_dir"] = result_dir - captured_artifacts["zip_file_path"] = zip_file_path - captured_artifacts["raw_entries"] = sorted( - path.relative_to(result_dir_path).as_posix() - for path in result_dir_path.rglob("*") - if path.is_file() - ) - captured_artifacts["doc_nav"] = json.loads( - (result_dir_path / "doc_nav.json").read_text(encoding="utf-8") - ) + celery_result = contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=job["user_id"], + ) - with zipfile.ZipFile(zip_path) as zip_file: - captured_artifacts["zip_entries"] = sorted(zip_file.namelist()) - captured_artifacts["zip_chunks"] = json.loads( - zip_file.read("chunks.json") - )["chunks"] - captured_artifacts["manifest"] = json.loads( - zip_file.read("manifest.json") - ) - - return SimpleNamespace( - zip_key=f"results/{job_id}.zip", - raw_prefix=f"results/{job_id}/", - raw_files={}, - ) + assert celery_result.successful() + observed = contract.observe_successful_job(job["job_id"]) + job_row = observed["job"] + expected_charge = job_row["page_count"] * int( + contract.settings.MICRO_DOLLARS_PER_PAGE + ) - monkeypatch.setattr(parse_job_service, "download_s3_file_to_temp", fake_download_s3_file_to_temp) - monkeypatch.setattr(parse_service, "checkerboard_parse_output", fake_checkerboard_parse_output) - monkeypatch.setattr(parse_job_service, "get_result_storage", lambda: FakeResultStorage()) - - result = document_ingestion_tasks.parse_task.run(job_id, user_id, "document_ingestion") - - expected_summary = "This document includes: 公司研究, 相关研报" - expected_connect_to = [ - { - "target": "image-1", - "relation": "embeds", - "ref": "[images/page-1.png]", - "position": { - "start": text_content_with_refs.index("[images/page-1.png]"), - "end": text_content_with_refs.index("[images/page-1.png]") - + len("[images/page-1.png]"), - }, - }, - { - "target": "table-1", - "relation": "embeds", - "ref": "[tables/table-1.html]", - "position": { - "start": text_content_with_refs.index("[tables/table-1.html]"), - "end": text_content_with_refs.index("[tables/table-1.html]") - + len("[tables/table-1.html]"), - }, - }, - ] - expected_credits_charged = 3 * int(settings.MICRO_DOLLARS_PER_PAGE) - expected_initial_balance = int(settings.FREE_PLAN_INITIAL_CREDITS) * 1_000_000 - - assert result == { - "status": "success", - "job_id": job_id, - "add_dir": None, - "vectors_count": 0, - "contents_count": 4, - "stored_count": 0, - "delivery_mode": "url", - "result_s3_key": f"results/{job_id}.zip", - } - assert captured_artifacts["parse_kwargs"]["filename"] == source_file_name - assert captured_artifacts["parse_kwargs"]["internal_output_filename"] == source_file_name - assert Path(str(captured_artifacts["parse_kwargs"]["file_full_path"])).name == source_file_name - assert "namespace" not in captured_artifacts["parse_kwargs"] - assert "kb_dir" not in captured_artifacts["parse_kwargs"] - assert captured_artifacts["result_dir"].endswith("contract-parse.pdf") - assert captured_artifacts["doc_nav"]["file_name"] == source_file_name - assert captured_artifacts["doc_nav"]["sections"][0]["title"] == "公司研究" - assert captured_artifacts["manifest"]["HIERARCHY"] == { - "公司研究": { - "自主可控加强,寒武纪或迎来营收快速放量周期": {}, - }, - "相关研报": { - "要点": {}, - }, - } - assert "doc_nav.json" in captured_artifacts["raw_entries"] - assert "hierarchy.json" not in captured_artifacts["raw_entries"] - assert "hierarchy_slim.json" not in captured_artifacts["raw_entries"] - assert "chunks.json" in captured_artifacts["zip_entries"] - assert "full.md" in captured_artifacts["zip_entries"] - assert "doc_nav.json" in captured_artifacts["zip_entries"] - assert "chunks_slim.json" not in captured_artifacts["zip_entries"] - assert "hierarchy.json" not in captured_artifacts["zip_entries"] - assert "hierarchy_slim.json" not in captured_artifacts["zip_entries"] - assert "images/page-1.png" in captured_artifacts["zip_entries"] - assert "tables/table-1.html" in captured_artifacts["zip_entries"] - assert captured_artifacts["zip_chunks"][0]["metadata"]["document_top_summary"] == expected_summary - assert captured_artifacts["zip_chunks"][0]["metadata"]["connect_to"] == expected_connect_to - assert captured_artifacts["zip_chunks"][2]["metadata"]["file_path"] == "images/page-1.png" - assert captured_artifacts["zip_chunks"][3]["metadata"]["file_path"] == "tables/table-1.html" - assert _find_task_workspaces(tmp_path, job_id) == [] - - progress = redis_service.hgetall(f"task:{job_id}:progress") - assert progress["progress"] == 100 - assert progress["message"] == "Task complete!" - assert progress["timestamp"] - - metadata = sync_job_metadata_service_cls(redis_service).get_metadata(job_id) - assert metadata is not None - assert metadata["page_count"] == 3 - assert metadata["workload_estimate_method"] == "pdf_metadata" - assert "workload_estimate_fallback_reason" not in metadata - assert metadata["billing_status"] == expected_billing_status - if billing_enabled: - assert metadata["billing_amount_micro_dollars"] == expected_credits_charged - assert metadata["billing_credits"] == expected_credits_charged / 1_000_000 - else: - assert metadata["billing_amount_micro_dollars"] == 0 - assert metadata["billing_credits"] == 0.0 - assert metadata["processing_started_at"] - assert metadata["processing_completed_at"] - assert metadata["processing_duration_ms"] >= 0 - - with engine.begin() as connection: - job_row = ( - connection.execute( - text( - """ - SELECT - status, - billing_status, - page_count, - credits_charged, - error_code, - error_message - FROM jobs - WHERE job_id = :job_id - """ - ), - {"job_id": job_id}, - ) - .mappings() - .one() - ) - job_result_row = ( - connection.execute( - text( - """ - SELECT id, delivery_mode, result_s3_key, result_size, inline_payload - FROM job_results - WHERE job_id = :job_id - """ - ), - {"job_id": job_id}, - ) - .mappings() - .one() - ) - job_chunks = list( - connection.execute( - text( - """ - SELECT chunk_type, text, path, sort_order - FROM job_chunks - WHERE job_result_id = :job_result_id - ORDER BY sort_order - """ - ), - {"job_result_id": job_result_row["id"]}, - ) - .mappings() - .all() - ) - document_row = ( - connection.execute( - text( - """ - SELECT document_id, namespace, status, current_job_result_id, source_file_name - FROM documents - WHERE user_id = :user_id - """ - ), - {"user_id": user_id}, - ) - .mappings() - .one() - ) - document_chunks = list( - connection.execute( - text( - """ - SELECT chunk_type, file_path, source_chunk_path, chunk_metadata - FROM document_chunks - WHERE document_id = :document_id - ORDER BY sort_order - """ - ), - {"document_id": document_row["document_id"]}, - ) - .mappings() - .all() - ) - graph_node_row = ( - connection.execute( - text( - """ - SELECT properties - FROM graph_nodes - WHERE owner_document_id = :document_id - """ - ), - {"document_id": document_row["document_id"]}, - ) - .mappings() - .one() - ) - balance_row = ( - connection.execute( - text( - """ - SELECT credits_balance - FROM user_balances - WHERE user_id = :user_id - """ - ), - {"user_id": user_id}, - ) - .mappings() - .one_or_none() - ) - transaction_types = list( - connection.execute( - text( - """ - SELECT transaction_type - FROM credits_transactions - WHERE user_id = :user_id - ORDER BY created_at ASC - """ - ), - {"user_id": user_id}, - ) - .scalars() - .all() - ) - audit_transitions = list( - connection.execute( - text( - """ - SELECT transition_reason, to_state - FROM job_state_audit_logs - WHERE job_id = :job_id - ORDER BY id ASC - """ - ), - {"job_id": job_id}, - ) - .mappings() - .all() - ) + assert job_row["status"] == "done" + assert job_row["billing_status"] == "charged" + assert job_row["credits_charged"] == expected_charge - graph_properties = dict(graph_node_row["properties"]) + billing = contract.observe_user_billing(job["user_id"]) + expected_initial_balance = int(contract.settings.FREE_PLAN_INITIAL_CREDITS) * 1_000_000 + assert billing["balance"] == expected_initial_balance - expected_charge + assert billing["transaction_types"] == ["initial_grant", "usage"] - assert job_row["status"] == "done" - assert job_row["billing_status"] == expected_billing_status - assert job_row["page_count"] == 3 - if billing_enabled: - assert job_row["credits_charged"] == expected_credits_charged - else: - assert job_row["credits_charged"] == 0 - assert job_row["error_code"] is None - assert job_row["error_message"] is None - assert job_result_row["delivery_mode"] == "url" - assert job_result_row["result_s3_key"] == f"results/{job_id}.zip" - assert job_result_row["result_size"] > 0 - assert dict(job_result_row["inline_payload"])["checksum"] - assert [chunk["chunk_type"] for chunk in job_chunks] == [ - "text", - "text", - "image", - "table", - ] - assert job_chunks[0]["text"] == text_content_with_refs - assert job_chunks[0]["path"].endswith( - "公司研究/自主可控加强,寒武纪或迎来营收快速放量周期" - ) - assert job_chunks[0]["sort_order"] == 0 - assert document_row["namespace"] == "worker-contract" - assert document_row["status"] == "active" - assert document_row["source_file_name"] == source_file_name - assert document_row["current_job_result_id"] - assert len(document_chunks) == 4 - assert document_chunks[0]["chunk_type"] == "text" - assert dict(document_chunks[0]["chunk_metadata"])["document_top_summary"] == expected_summary - assert dict(document_chunks[0]["chunk_metadata"])["connect_to"] == expected_connect_to - assert document_chunks[2]["chunk_type"] == "image" - assert document_chunks[2]["file_path"] == "images/page-1.png" - assert document_chunks[3]["chunk_type"] == "table" - assert document_chunks[3]["file_path"] == "tables/table-1.html" - assert graph_properties["chunks_count"] == 4 - assert graph_properties["top_summary"] == expected_summary - if billing_enabled: - assert balance_row is not None - assert ( - balance_row["credits_balance"] - == expected_initial_balance - expected_credits_charged - ) - else: - assert balance_row is None - assert transaction_types == expected_transaction_types - assert [(row["transition_reason"], row["to_state"]) for row in audit_transitions] == [ + assert contract.observe_job_state_transitions(job["job_id"]) == [ ("start_processing", "running"), ("mark_completed", "done"), ] -def test_should_export_full_result_when_publication_deduplicates_existing_chunks( +def test_parse_task_should_export_full_result_when_same_content_was_already_published( worker_contract_environment: None, - monkeypatch: MonkeyPatch, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: - monkeypatch.setenv("BILLING_ENABLED", "false") - ( - document_ingestion_tasks, - parse_service, - parse_job_service, - engine, - sync_job_info_service_cls, - sync_job_metadata_service_cls, - sync_redis_service_factory, - ) = _load_parse_task_modules() - settings = _load_worker_settings() - - user_id: str = f"worker-user-{uuid4().hex[:12]}" - existing_job_id: str = f"job_existing_{uuid4().hex[:12]}" - existing_result_id: str = str(uuid4()) - existing_document_id: str = f"doc_{uuid4().hex[:12]}" - job_id: str = f"job_parse_dedup_{uuid4().hex[:12]}" - source_file_name: str = "dedup-export.pdf" - s3_key: str = f"uploads/{job_id}.pdf" - job_metadata = _build_pending_file_job_metadata(source_file_name) - captured_artifacts: dict[str, Any] = {} - - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - insert_contract_job( - connection, - job_id=existing_job_id, - user_id=user_id, - status="done", - source_type="file", - s3_key=f"uploads/{existing_job_id}.pdf", - webhook_enabled=False, - job_metadata=_build_pending_file_job_metadata("existing.pdf"), - billing_status="skipped", - ) - connection.execute( - text( - """ - INSERT INTO job_results ( - id, - job_id, - delivery_mode, - inline_payload, - result_s3_key, - result_size, - created_at, - updated_at - ) VALUES ( - :result_id, - :job_id, - 'url', - CAST(:inline_payload AS JSON), - :result_s3_key, - :result_size, - NOW(), - NOW() - ) - """ - ), - { - "result_id": existing_result_id, - "job_id": existing_job_id, - "inline_payload": json.dumps({"checksum": "existing"}), - "result_s3_key": f"results/{existing_job_id}.zip", - "result_size": 123, - }, - ) - connection.execute( - text( - """ - INSERT INTO documents ( - document_id, - user_id, - namespace, - status, - current_job_result_id, - source_file_name, - created_at, - updated_at - ) VALUES ( - :document_id, - :user_id, - 'worker-contract', - 'active', - :result_id, - 'existing.pdf', - NOW(), - NOW() - ) - """ - ), - { - "document_id": existing_document_id, - "user_id": user_id, - "result_id": existing_result_id, - }, - ) - connection.execute( - text( - """ - UPDATE job_results - SET document_id = :document_id - WHERE id = :result_id - """ - ), - { - "document_id": existing_document_id, - "result_id": existing_result_id, - }, - ) - connection.execute( - text( - """ - INSERT INTO document_chunks ( - id, - chunk_id, - user_id, - namespace, - document_id, - job_result_id, - chunk_type, - content, - source_chunk_path, - file_path, - chunk_metadata, - sort_order, - created_at - ) VALUES - ( - :text_id, - 'duplicate-text', - :user_id, - 'worker-contract', - :document_id, - :result_id, - 'text', - 'already published text', - 'existing.pdf/Section/Duplicate text', - NULL, - CAST(:text_metadata AS JSON), - 0, - NOW() - ), - ( - :image_id, - 'duplicate-image', - :user_id, - 'worker-contract', - :document_id, - :result_id, - 'image', - 'already published image', - 'existing.pdf/images/duplicate.png', - 'images/duplicate.png', - CAST(:image_metadata AS JSON), - 1, - NOW() - ) - """ - ), - { - "text_id": f"dchk_{uuid4().hex[:12]}", - "image_id": f"dchk_{uuid4().hex[:12]}", - "user_id": user_id, - "document_id": existing_document_id, - "result_id": existing_result_id, - "text_metadata": json.dumps({}), - "image_metadata": json.dumps({"file_path": "images/duplicate.png"}), - }, - ) - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - status="pending", - source_type="file", - s3_key=s3_key, - webhook_enabled=False, - job_metadata=job_metadata, - billing_status="pending", - ) + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=False) - _save_worker_task_cache( - job_id=job_id, + user_id = f"worker-contract-user-{uuid4().hex[:12]}" + first_job = contract.create_file_job( user_id=user_id, - s3_key=s3_key, - metadata=job_metadata, - sync_job_info_service_cls=sync_job_info_service_cls, - sync_job_metadata_service_cls=sync_job_metadata_service_cls, - sync_redis_service_factory=sync_redis_service_factory, + source_file_name="contract-original.xlsx", + job_id_prefix="job_parse_original", ) - - _bind_parse_task_to_current_module(monkeypatch, document_ingestion_tasks=document_ingestion_tasks) - monkeypatch.setattr(settings, "TMP_PATH", str(tmp_path)) - monkeypatch.setattr(settings, "BILLING_ENABLED", False) - - def fake_cleanup_task_workspace(workspace_dir: str | None) -> bool: - captured_artifacts["workspace_dir"] = workspace_dir - return True - - def fake_verify_s3_file_exists(storage_key: str) -> dict[str, Any]: - return { - "exists": storage_key == s3_key, - "size": _SAMPLE_PDF_PATH.stat().st_size, - } - - def fake_download_s3_file_to_temp( - storage_key: str, file_ext: str, temp_dir: str - ) -> str: - assert storage_key == s3_key - downloaded_path = Path(temp_dir) / f"downloaded{file_ext}" - shutil.copy2(_SAMPLE_PDF_PATH, downloaded_path) - return str(downloaded_path) - - def fake_checkerboard_parse_output(**kwargs: Any) -> Any: - output_dir = ( - Path(str(kwargs["output_dir"])) - / str(kwargs["internal_output_filename"]) + second_job = contract.create_file_job( + user_id=user_id, + source_file_name="contract-duplicate.xlsx", + job_id_prefix="job_parse_duplicate", + ) + for job in [first_job, second_job]: + contract.upload_source_file( + local_file_path=_SAMPLE_XLSX_PATH, + s3_key=job["s3_key"], ) - images_dir = output_dir / "images" - images_dir.mkdir(parents=True, exist_ok=True) - (output_dir / "full.md").write_text("body", encoding="utf-8") - (images_dir / "duplicate.png").write_bytes(b"png") - file_root = str(kwargs["internal_output_filename"]) - parsed_rows: list[dict[str, Any]] = [ - { - "content": "duplicate text", - "path": f"{file_root}/Section/Duplicate text", - "type": "text", - "length": 14, - "keywords": "", - "summary": "", - "know_id": "duplicate-text", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "1", - }, - { - "content": "duplicate image", - "path": f"{file_root}/images/duplicate.png", - "type": "image", - "length": 15, - "keywords": "", - "summary": "", - "know_id": "duplicate-image", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "2", - }, - { - "content": "new text", - "path": f"{file_root}/Section/New text", - "type": "text", - "length": 8, - "keywords": "", - "summary": "", - "know_id": "new-text", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "3", - }, - ] - return _build_fake_parse_output(output_dir=output_dir, rows=parsed_rows) - - class FakeResultStorage: - def upload(self, *, job_id: str, result_dir: str, zip_file_path: str) -> Any: - result_dir_path = Path(result_dir) - zip_path = Path(zip_file_path) - captured_artifacts["raw_entries"] = sorted( - path.relative_to(result_dir_path).as_posix() - for path in result_dir_path.rglob("*") - if path.is_file() - ) - with zipfile.ZipFile(zip_path) as zip_file: - captured_artifacts["zip_entries"] = sorted(zip_file.namelist()) - captured_artifacts["zip_chunks"] = json.loads( - zip_file.read("chunks.json") - )["chunks"] - captured_artifacts["manifest"] = json.loads( - zip_file.read("manifest.json") - ) - - return SimpleNamespace( - zip_key=f"results/{job_id}.zip", - raw_prefix=f"results/{job_id}/", - raw_files={}, - ) + first_result = contract.enqueue_parse_task( + job_id=first_job["job_id"], + user_id=user_id, + ) + second_result = contract.enqueue_parse_task( + job_id=second_job["job_id"], + user_id=user_id, + ) - _patch_verify_upload_exists(monkeypatch, fake_verify_s3_file_exists) - monkeypatch.setattr(parse_job_service, "download_s3_file_to_temp", fake_download_s3_file_to_temp) - monkeypatch.setattr(parse_service, "checkerboard_parse_output", fake_checkerboard_parse_output) - monkeypatch.setattr(parse_job_service, "get_result_storage", lambda: FakeResultStorage()) - monkeypatch.setattr(parse_job_service, "cleanup_task_workspace", fake_cleanup_task_workspace) - - result = document_ingestion_tasks.parse_task.run(job_id, user_id, "document_ingestion") - - assert result["contents_count"] == 3 - assert "images/duplicate.png" in captured_artifacts["zip_entries"] - assert "images/duplicate.png" in captured_artifacts["raw_entries"] - assert [chunk["chunk_id"] for chunk in captured_artifacts["zip_chunks"]] == [ - "duplicate-text", - "duplicate-image", - "new-text", - ] - assert captured_artifacts["manifest"]["statistics"] == { - "total_chunks": 3, - "text_chunks": 2, - "image_chunks": 1, - "table_chunks": 0, - "total_pages": None, - } - workspace_dir = Path(str(captured_artifacts["workspace_dir"])) - assert list(workspace_dir.rglob("images/duplicate.png")) - - with engine.begin() as connection: - job_result_row = ( - connection.execute( - text( - """ - SELECT id, document_metadata - FROM job_results - WHERE job_id = :job_id - """ - ), - {"job_id": job_id}, - ) - .mappings() - .one() - ) - job_chunk_ids = list( - connection.execute( - text( - """ - SELECT chunk_id - FROM job_chunks - WHERE job_result_id = :job_result_id - ORDER BY sort_order - """ - ), - {"job_result_id": job_result_row["id"]}, - ) - .scalars() - .all() - ) - document_chunk_ids = list( - connection.execute( - text( - """ - SELECT document_chunks.chunk_id - FROM document_chunks - JOIN documents - ON documents.document_id = document_chunks.document_id - WHERE documents.current_job_result_id = :job_result_id - ORDER BY document_chunks.sort_order - """ - ), - {"job_result_id": job_result_row["id"]}, - ) - .scalars() - .all() - ) + assert first_result.successful() + assert second_result.successful() + + observed = contract.observe_successful_job(second_job["job_id"]) + result_row = observed["result"] + job_chunks = observed["job_chunks"] + document_chunks = observed["document_chunks"] + result_zip = contract.read_result_zip( + result_s3_key=result_row["result_s3_key"], + tmp_path=tmp_path, + ) - assert job_chunk_ids == ["duplicate-text", "duplicate-image", "new-text"] - assert document_chunk_ids == ["duplicate-text", "duplicate-image", "new-text"] - assert "chunk_overlap" not in dict(job_result_row["document_metadata"] or {}) + assert len(job_chunks) > 0 + assert len(document_chunks) == len(job_chunks) + assert len(result_zip["chunks"]["chunks"]) == len(job_chunks) + assert any(member.startswith("tables/") for member in result_zip["members"]) + assert "chunk_overlap" not in dict(result_row["document_metadata"] or {}) -def test_should_initialize_billing_once_for_concurrent_parse_tasks( +def test_parse_task_should_initialize_billing_once_for_concurrent_parse_tasks( worker_contract_environment: None, - monkeypatch: MonkeyPatch, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: - ( - document_ingestion_tasks, - parse_service, - parse_job_service, - engine, - sync_job_info_service_cls, - sync_job_metadata_service_cls, - sync_redis_service_factory, - ) = _load_parse_task_modules() - settings = _load_worker_settings() - - user_id: str = f"worker-concurrent-user-{uuid4().hex[:12]}" - job_ids: list[str] = [f"job_cb_{index}_{uuid4().hex[:12]}" for index in range(2)] - source_file_name: str = "contract-concurrent.pdf" - s3_keys: dict[str, str] = { - job_id: f"uploads/{job_id}.pdf" for job_id in job_ids - } - job_metadata_by_id: dict[str, dict[str, Any]] = { - job_id: _build_pending_file_job_metadata(source_file_name) - for job_id in job_ids - } + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=True) - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - for job_id in job_ids: - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - status="pending", - source_type="file", - s3_key=s3_keys[job_id], - webhook_enabled=False, - job_metadata=job_metadata_by_id[job_id], - billing_status="pending", - ) - - redis_service = sync_redis_service_factory.get_service() - for job_id in job_ids: - _save_worker_task_cache( - job_id=job_id, + user_id = f"worker-contract-user-{uuid4().hex[:12]}" + jobs = [ + contract.create_file_job( user_id=user_id, - s3_key=s3_keys[job_id], - metadata=job_metadata_by_id[job_id], - sync_job_info_service_cls=sync_job_info_service_cls, - sync_job_metadata_service_cls=sync_job_metadata_service_cls, - sync_redis_service_factory=sync_redis_service_factory, - ) - - _bind_parse_task_to_current_module(monkeypatch, document_ingestion_tasks=document_ingestion_tasks) - monkeypatch.setattr(settings, "TMP_PATH", str(tmp_path)) - monkeypatch.setattr(settings, "BILLING_ENABLED", True) - - def fake_verify_s3_file_exists(storage_key: str) -> dict[str, Any]: - return { - "exists": storage_key in s3_keys.values(), - "size": _SAMPLE_PDF_PATH.stat().st_size, - } - - def fake_download_s3_file_to_temp( - storage_key: str, file_ext: str, temp_dir: str - ) -> str: - assert storage_key in s3_keys.values() - assert file_ext == ".pdf" - downloaded_path = Path(temp_dir) / f"downloaded{file_ext}" - shutil.copy2(_SAMPLE_PDF_PATH, downloaded_path) - return str(downloaded_path) - - billing_start_barrier = Barrier(len(job_ids)) - - def fake_estimate_workload(file_path: str) -> Any: - from app.services.document_ingestion.page_estimator import WorkloadEstimate - - billing_start_barrier.wait(timeout=10) - return WorkloadEstimate(page_count=1, method="contract_fake") - - def fake_checkerboard_parse_output(**kwargs: Any) -> Any: - output_dir = ( - Path(str(kwargs["output_dir"])) - / str(kwargs["internal_output_filename"]) + source_file_name=f"contract-concurrent-{index}.xlsx", + job_id_prefix=f"job_parse_concurrent_{index}", ) - output_dir.mkdir(parents=True, exist_ok=True) - (output_dir / "full.md").write_text("body", encoding="utf-8") - - file_root = str(kwargs["internal_output_filename"]) - parsed_rows: list[dict[str, Any]] = [ - { - "content": "chunk body", - "path": f"{file_root}/Section/Point", - "type": "text", - "length": 10, - "keywords": "", - "summary": "", - "know_id": f"{kwargs['job_id']}-chunk-1", - "tokens": "", - "connectto": "", - "addtime": "now", - "page_nums": "1", - } - ] - return _build_fake_parse_output(output_dir=output_dir, rows=parsed_rows) - - class FakeResultStorage: - def upload(self, *, job_id: str, result_dir: str, zip_file_path: str) -> Any: - return SimpleNamespace( - zip_key=f"results/{job_id}.zip", - raw_prefix=f"results/{job_id}/", - raw_files={}, - ) - - _patch_verify_upload_exists(monkeypatch, fake_verify_s3_file_exists) - monkeypatch.setattr(parse_job_service, "download_s3_file_to_temp", fake_download_s3_file_to_temp) - monkeypatch.setattr( - parse_job_service.PageEstimator, - "estimate_workload", - fake_estimate_workload, - ) - monkeypatch.setattr(parse_service, "checkerboard_parse_output", fake_checkerboard_parse_output) - monkeypatch.setattr(parse_job_service, "get_result_storage", lambda: FakeResultStorage()) - - def run_parse_task(job_id: str) -> dict[str, Any]: - return dict(document_ingestion_tasks.parse_task.run(job_id, user_id, "document_ingestion")) - - with ThreadPoolExecutor(max_workers=len(job_ids)) as executor: - results = list(executor.map(run_parse_task, job_ids)) - - expected_credits_charged = int(settings.MICRO_DOLLARS_PER_PAGE) - expected_initial_balance = ( - int(settings.FREE_PLAN_INITIAL_CREDITS) * 1_000_000 - ) - - with engine.begin() as connection: - job_rows = list( - connection.execute( - text( - """ - SELECT job_id, status, billing_status, page_count, credits_charged - FROM jobs - WHERE job_id = ANY(:job_ids) - ORDER BY job_id - """ - ), - {"job_ids": job_ids}, - ) - .mappings() - .all() - ) - balance_row = ( - connection.execute( - text( - """ - SELECT credits_balance - FROM user_balances - WHERE user_id = :user_id - """ - ), - {"user_id": user_id}, - ) - .mappings() - .one() - ) - transaction_rows = list( - connection.execute( - text( - """ - SELECT transaction_type, COUNT(*) AS count - FROM credits_transactions - WHERE user_id = :user_id - GROUP BY transaction_type - ORDER BY transaction_type - """ - ), - {"user_id": user_id}, - ) - .mappings() - .all() - ) - payment_count_row = ( - connection.execute( - text( - """ - SELECT COUNT(*) AS count - FROM payment_records - WHERE user_id = :user_id - AND payment_type = 'system_grant' - """ + for index in range(2) + ] + for job in jobs: + contract.upload_source_file( + local_file_path=_SAMPLE_XLSX_PATH, + s3_key=job["s3_key"], + ) + + with ThreadPoolExecutor(max_workers=len(jobs)) as executor: + celery_results = list( + executor.map( + lambda job: contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=user_id, ), - {"user_id": user_id}, + jobs, ) - .mappings() - .one() ) - metadata_by_job_id = { - job_id: sync_job_metadata_service_cls(redis_service).get_metadata(job_id) - for job_id in job_ids - } - result_job_ids = {str(result["job_id"]) for result in results} - transaction_counts = { - str(row["transaction_type"]): int(row["count"]) for row in transaction_rows - } - - assert result_job_ids == set(job_ids) - assert all(result["status"] == "success" for result in results) - assert [ - { - "job_id": row["job_id"], - "status": row["status"], - "billing_status": row["billing_status"], - "page_count": row["page_count"], - "credits_charged": row["credits_charged"], - } - for row in job_rows - ] == [ - { - "job_id": job_id, - "status": "done", - "billing_status": "charged", - "page_count": 1, - "credits_charged": expected_credits_charged, - } - for job_id in sorted(job_ids) + assert all(result.successful() for result in celery_results) + observed_jobs = [ + contract.observe_successful_job(job["job_id"])["job"] for job in jobs ] - assert all( - metadata_by_job_id[job_id] is not None - and metadata_by_job_id[job_id]["billing_status"] == "charged" - and metadata_by_job_id[job_id]["billing_amount_micro_dollars"] - == expected_credits_charged - for job_id in job_ids - ) - assert balance_row == { - "credits_balance": expected_initial_balance - - (expected_credits_charged * len(job_ids)) + assert all(row["billing_status"] == "charged" for row in observed_jobs) + + billing = contract.observe_user_billing(user_id) + expected_total_charge = sum(row["credits_charged"] for row in observed_jobs) + expected_initial_balance = int(contract.settings.FREE_PLAN_INITIAL_CREDITS) * 1_000_000 + assert billing["balance"] == expected_initial_balance - expected_total_charge + assert billing["transaction_counts"] == { + "initial_grant": 1, + "usage": len(jobs), } - assert transaction_counts == {"initial_grant": 1, "usage": len(job_ids)} - assert payment_count_row == {"count": 1} + assert billing["system_grant_payment_count"] == 1 -def test_should_skip_parse_task_when_the_job_is_already_terminal( +def test_parse_task_should_skip_terminal_job_without_creating_outputs( worker_contract_environment: None, - monkeypatch: MonkeyPatch, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: - ( - document_ingestion_tasks, - parse_service, - parse_job_service, - engine, - sync_job_info_service_cls, - sync_job_metadata_service_cls, - sync_redis_service_factory, - ) = _load_parse_task_modules() - settings = _load_worker_settings() - - user_id: str = f"worker-user-{uuid4().hex[:12]}" - job_id: str = f"job_parse_skipped_{uuid4().hex[:12]}" - source_file_name: str = "contract-skip.pdf" - s3_key: str = f"uploads/{job_id}.pdf" - - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - job_metadata = _build_pending_file_job_metadata(source_file_name) - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - s3_key=s3_key, - status="done", - source_type="file", - webhook_enabled=False, - job_metadata=job_metadata, - billing_status="charged", - ) - - redis_service = _save_worker_task_cache( - job_id=job_id, - user_id=user_id, - s3_key=s3_key, - metadata=job_metadata, - sync_job_info_service_cls=sync_job_info_service_cls, - sync_job_metadata_service_cls=sync_job_metadata_service_cls, - sync_redis_service_factory=sync_redis_service_factory, + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + + job = contract.create_file_job( + source_file_name="contract-skip.xlsx", + status="done", + billing_status="charged", + job_id_prefix="job_skip", ) - _bind_parse_task_to_current_module(monkeypatch, document_ingestion_tasks=document_ingestion_tasks) - monkeypatch.setattr(settings, "TMP_PATH", str(tmp_path)) - def fake_verify_s3_file_exists(storage_key: str) -> dict[str, Any]: - return {"exists": storage_key == s3_key, "size": 1024} - - _patch_verify_upload_exists(monkeypatch, fake_verify_s3_file_exists) - monkeypatch.setattr( - parse_service, - "checkerboard_parse_output", - lambda **_kwargs: (_ for _ in ()).throw( - AssertionError("terminal parse task should not invoke the parser") - ), + celery_result = contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=job["user_id"], ) - result = document_ingestion_tasks.parse_task.run(job_id, user_id, "document_ingestion") - - assert result == { + assert celery_result.successful() + assert celery_result.result == { "status": "skipped", - "job_id": job_id, + "job_id": job["job_id"], "reason": "job_already_terminal", } - assert redis_service.hgetall(f"task:{job_id}:progress") == {} - assert _find_task_workspaces(tmp_path, job_id) == [] - - with engine.begin() as connection: - job_result_count = int( - connection.execute( - text("SELECT COUNT(*) FROM job_results WHERE job_id = :job_id"), - {"job_id": job_id}, - ).scalar_one() - ) - audit_transition_count = int( - connection.execute( - text( - "SELECT COUNT(*) FROM job_state_audit_logs WHERE job_id = :job_id" - ), - {"job_id": job_id}, - ).scalar_one() - ) + assert contract.get_task_progress(job["job_id"]) == {} + assert contract.find_task_workspaces(tmp_path, job["job_id"]) == [] - assert job_result_count == 0 - assert audit_transition_count == 0 + job_row = contract.observe_job_status(job["job_id"]) + assert job_row["status"] == "done" + assert job_row["billing_status"] == "charged" + assert contract.count_job_results(job["job_id"]) == 0 -def test_should_mark_the_job_failed_and_cleanup_the_workspace_when_parse_execution_raises( +def test_parse_task_should_mark_failed_and_cleanup_when_uploaded_source_is_missing( worker_contract_environment: None, - monkeypatch: MonkeyPatch, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: - ( - document_ingestion_tasks, - parse_service, - parse_job_service, - engine, - sync_job_info_service_cls, - sync_job_metadata_service_cls, - sync_redis_service_factory, - ) = _load_parse_task_modules() - settings = _load_worker_settings() - - user_id: str = f"worker-user-{uuid4().hex[:12]}" - job_id: str = f"job_parse_failure_{uuid4().hex[:12]}" - source_file_name: str = "contract-failure.pdf" - s3_key: str = f"uploads/{job_id}.pdf" - - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - job_metadata = _build_pending_file_job_metadata(source_file_name) - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - status="pending", - source_type="file", - s3_key=s3_key, - webhook_enabled=False, - job_metadata=job_metadata, - billing_status="pending", - ) + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=False) - _save_worker_task_cache( - job_id=job_id, - user_id=user_id, - s3_key=s3_key, - metadata=job_metadata, - sync_job_info_service_cls=sync_job_info_service_cls, - sync_job_metadata_service_cls=sync_job_metadata_service_cls, - sync_redis_service_factory=sync_redis_service_factory, + job = contract.create_file_job( + source_file_name="contract-missing-source.xlsx", + job_id_prefix="job_missing", ) - _bind_parse_task_to_current_module(monkeypatch, document_ingestion_tasks=document_ingestion_tasks) - monkeypatch.setattr(settings, "TMP_PATH", str(tmp_path)) - def fake_verify_s3_file_exists(storage_key: str) -> dict[str, Any]: - return { - "exists": storage_key == s3_key, - "size": _SAMPLE_PDF_PATH.stat().st_size, - } - - _patch_verify_upload_exists(monkeypatch, fake_verify_s3_file_exists) - - def fake_download_s3_file_to_temp( - storage_key: str, file_ext: str, temp_dir: str - ) -> str: - assert storage_key == s3_key - downloaded_path = Path(temp_dir) / f"downloaded{file_ext}" - shutil.copy2(_SAMPLE_PDF_PATH, downloaded_path) - return str(downloaded_path) - - monkeypatch.setattr(parse_job_service, "download_s3_file_to_temp", fake_download_s3_file_to_temp) - monkeypatch.setattr( - parse_service, - "checkerboard_parse_output", - lambda **_kwargs: (_ for _ in ()).throw(RuntimeError("parse failed")), + celery_result = contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=job["user_id"], ) - monkeypatch.setattr( - parse_job_service, - "get_result_storage", - lambda: (_ for _ in ()).throw( - AssertionError("result storage should not run after parser failure") - ), + + assert celery_result.failed() + assert contract.find_task_workspaces(tmp_path, job["job_id"]) == [] + + job_row = contract.observe_job_status(job["job_id"]) + assert job_row["status"] == "failed" + assert job_row["billing_status"] in {"pending", "skipped"} + assert job_row["error_code"] + assert job_row["error_message"] + assert contract.count_job_results(job["job_id"]) == 0 + + +def test_parse_task_should_refund_charged_job_when_uploaded_file_cannot_be_parsed( + worker_contract_environment: None, + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=True) + + invalid_xlsx_path = tmp_path / "invalid.xlsx" + invalid_xlsx_path.write_bytes(b"this is not an xlsx workbook") + job = contract.create_file_job( + source_file_name="contract-invalid.xlsx", + job_id_prefix="job_invalid_parse", + ) + contract.upload_source_file( + local_file_path=invalid_xlsx_path, + s3_key=job["s3_key"], ) - result = document_ingestion_tasks.parse_task.apply( - args=[job_id, user_id, "document_ingestion"], - throw=False, + celery_result = contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=job["user_id"], ) - assert result.status == "FAILURE" - assert _find_task_workspaces(tmp_path, job_id) == [] - - expected_credits_charged = 3 * int(settings.MICRO_DOLLARS_PER_PAGE) - expected_initial_balance = int(settings.FREE_PLAN_INITIAL_CREDITS) * 1_000_000 - - with engine.begin() as connection: - job_row = ( - connection.execute( - text( - """ - SELECT status, billing_status, page_count, credits_charged, error_code, error_message - FROM jobs - WHERE job_id = :job_id - """ - ), - {"job_id": job_id}, - ) - .mappings() - .one() - ) - balance_row = ( - connection.execute( - text( - """ - SELECT credits_balance - FROM user_balances - WHERE user_id = :user_id - """ - ), - {"user_id": user_id}, - ) - .mappings() - .one() - ) - transaction_types = list( - connection.execute( - text( - """ - SELECT transaction_type - FROM credits_transactions - WHERE user_id = :user_id - ORDER BY created_at ASC - """ - ), - {"user_id": user_id}, - ) - .scalars() - .all() - ) - audit_transitions = list( - connection.execute( - text( - """ - SELECT transition_reason, to_state - FROM job_state_audit_logs - WHERE job_id = :job_id - ORDER BY id ASC - """ - ), - {"job_id": job_id}, - ) - .mappings() - .all() - ) + assert celery_result.failed() + assert contract.find_task_workspaces(tmp_path, job["job_id"]) == [] + job_row = contract.observe_job_status(job["job_id"]) assert job_row["status"] == "failed" assert job_row["billing_status"] == "refunded" - assert job_row["page_count"] == 3 - assert job_row["credits_charged"] == expected_credits_charged - assert job_row["error_code"] == "UNKNOWN" - assert job_row["error_message"] == "An unexpected error occurred" - assert balance_row["credits_balance"] == expected_initial_balance - assert transaction_types == ["initial_grant", "usage", "refund"] - assert [(row["transition_reason"], row["to_state"]) for row in audit_transitions] == [ + assert job_row["credits_charged"] == int(contract.settings.MICRO_DOLLARS_PER_PAGE) + assert contract.count_job_results(job["job_id"]) == 0 + + billing = contract.observe_user_billing(job["user_id"]) + expected_initial_balance = int(contract.settings.FREE_PLAN_INITIAL_CREDITS) * 1_000_000 + assert billing["balance"] == expected_initial_balance + assert billing["transaction_types"] == ["initial_grant", "usage", "refund"] + assert contract.observe_job_state_transitions(job["job_id"]) == [ ("start_processing", "running"), ("mark_failed", "failed"), ] @@ -1495,120 +372,42 @@ def fake_download_s3_file_to_temp( def test_should_reject_pdf_when_page_count_exceeds_configured_limit( worker_contract_environment: None, - monkeypatch: MonkeyPatch, + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: - monkeypatch.setenv("BILLING_ENABLED", "false") - ( - document_ingestion_tasks, - _parse_service, - parse_job_service, - engine, - sync_job_info_service_cls, - sync_job_metadata_service_cls, - sync_redis_service_factory, - ) = _load_parse_task_modules() - settings = _load_worker_settings() - - user_id: str = f"worker-user-{uuid4().hex[:12]}" - job_id: str = f"job_pdf_page_limit_{uuid4().hex[:12]}" - source_file_name: str = "oversized-contract.pdf" - s3_key: str = f"uploads/{job_id}.pdf" + contract = WorkerParseContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.use_billing(monkeypatch, is_enabled=False) + max_pdf_page_limit: int = 1 actual_page_count: int = 2 + source_file_name: str = "oversized-contract.pdf" pdf_path = tmp_path / source_file_name _write_blank_pdf(pdf_path, actual_page_count) - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - job_metadata = _build_pending_file_job_metadata(source_file_name) - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - status="pending", - source_type="file", - s3_key=s3_key, - webhook_enabled=False, - job_metadata=job_metadata, - billing_status="pending", - ) - - redis_service = _save_worker_task_cache( - job_id=job_id, - user_id=user_id, - s3_key=s3_key, - metadata=job_metadata, - sync_job_info_service_cls=sync_job_info_service_cls, - sync_job_metadata_service_cls=sync_job_metadata_service_cls, - sync_redis_service_factory=sync_redis_service_factory, + contract.use_pdf_page_limit(monkeypatch, max_pdf_page_limit) + job = contract.create_file_job( + source_file_name=source_file_name, + job_id_prefix="job_pdf_page_limit", ) - - _bind_parse_task_to_current_module( - monkeypatch, - document_ingestion_tasks=document_ingestion_tasks, - ) - monkeypatch.setattr(settings, "TMP_PATH", str(tmp_path)) - monkeypatch.setattr(settings, "BILLING_ENABLED", False) - monkeypatch.setattr(settings, "MAX_PDF_PAGE_LIMIT", max_pdf_page_limit) - - def fake_verify_s3_file_exists(storage_key: str) -> dict[str, Any]: - return { - "exists": storage_key == s3_key, - "size": pdf_path.stat().st_size, - } - - _patch_verify_upload_exists(monkeypatch, fake_verify_s3_file_exists) - - def fake_download_s3_file_to_temp( - storage_key: str, - file_ext: str, - temp_dir: str, - ) -> str: - assert storage_key == s3_key - assert file_ext == ".pdf" - downloaded_path = Path(temp_dir) / f"downloaded{file_ext}" - shutil.copy2(pdf_path, downloaded_path) - return str(downloaded_path) - - monkeypatch.setattr(parse_job_service, "download_s3_file_to_temp", fake_download_s3_file_to_temp) - monkeypatch.setattr( - parse_job_service, - "get_result_storage", - lambda: (_ for _ in ()).throw( - AssertionError("result storage should not run after page-limit rejection") - ), + contract.upload_source_file( + local_file_path=pdf_path, + s3_key=job["s3_key"], ) - result = document_ingestion_tasks.parse_task.apply( - args=[job_id, user_id, "document_ingestion"], - throw=False, + celery_result = contract.enqueue_parse_task( + job_id=job["job_id"], + user_id=job["user_id"], ) - assert result.status == "FAILURE" - assert _find_task_workspaces(tmp_path, job_id) == [] + assert celery_result.failed() + assert contract.find_task_workspaces(tmp_path, job["job_id"]) == [] - metadata = sync_job_metadata_service_cls(redis_service).get_metadata(job_id) - assert metadata is not None + metadata = contract.get_job_metadata(job["job_id"]) assert metadata["page_count"] == actual_page_count assert metadata["billing_status"] == "skipped" - with engine.begin() as connection: - job_row = ( - connection.execute( - text( - """ - SELECT status, billing_status, page_count, credits_charged, - error_code, error_message, job_metadata - FROM jobs - WHERE job_id = :job_id - """ - ), - {"job_id": job_id}, - ) - .mappings() - .one() - ) + job_row = contract.observe_job_status(job["job_id"]) assert job_row["status"] == "failed" assert job_row["billing_status"] == "skipped" @@ -1619,7 +418,7 @@ def fake_download_s3_file_to_temp( job_row["error_message"] == "Document too large: 2 pages exceeds the 1-page limit. Please split the document and upload in smaller batches." ) - assert job_row["job_metadata"]["error_details"] == { + assert metadata["error_details"] == { "violations": [ { "field": "page_count", diff --git a/apps/worker/tests/contract/test_url_upload_contract.py b/apps/worker/tests/contract/test_url_upload_contract.py index 4765472cd..4a2bf32ed 100644 --- a/apps/worker/tests/contract/test_url_upload_contract.py +++ b/apps/worker/tests/contract/test_url_upload_contract.py @@ -1,52 +1,34 @@ from __future__ import annotations -import os -import socket +from collections.abc import Iterator +from contextlib import contextmanager +from functools import partial +from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path -from typing import Any -from uuid import uuid4 +from threading import Thread from pytest import MonkeyPatch -from sqlalchemy import text -from sqlalchemy.engine import Engine -from support.contract_database import insert_contract_job, insert_contract_user +from support.worker_url_upload_contract import WorkerUrlUploadContract -def _load_upload_task_modules() -> tuple[Any, Any, Engine, Any, Any]: - import app.core.tasks.document_ingestion_tasks as document_ingestion_tasks - import app.services.workload.url_upload_service as url_upload_service - from shared.core.database_sync import get_sync_engine - from shared.services.redis.redis_sync_service import ( - SyncJobInfoRedisService, - SyncRedisServiceFactory, - ) +class _SilentStaticFileHandler(SimpleHTTPRequestHandler): + def log_message(self, format: str, *args: object) -> None: + return - return ( - document_ingestion_tasks, - url_upload_service, - get_sync_engine(), - SyncJobInfoRedisService, - SyncRedisServiceFactory, - ) - -def _bind_upload_task_to_current_module( - monkeypatch: MonkeyPatch, - *, - document_ingestion_tasks: Any, -) -> None: - monkeypatch.setitem( - document_ingestion_tasks.upload_url_file_task._orig_run.__globals__, - "_upload_url_file", - document_ingestion_tasks._upload_url_file, - ) - monkeypatch.setattr( - document_ingestion_tasks.upload_url_file_task, - "__trace__", - None, - raising=False, - ) +@contextmanager +def _serve_directory(directory: Path) -> Iterator[str]: + handler = partial(_SilentStaticFileHandler, directory=str(directory)) + server = ThreadingHTTPServer(("127.0.0.1", 0), handler) + thread = Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_port}" + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) def test_should_upload_a_url_job_to_the_expected_storage_key_and_publish_progress( @@ -54,142 +36,44 @@ def test_should_upload_a_url_job_to_the_expected_storage_key_and_publish_progres monkeypatch: MonkeyPatch, tmp_path: Path, ) -> None: - ( - document_ingestion_tasks, - url_upload_service, - engine, - sync_job_info_service_cls, - sync_redis_service_factory, - ) = _load_upload_task_modules() - from shared.core.config import settings - - user_id = f"worker-user-{uuid4().hex[:12]}" - job_id = f"job_url_upload_{uuid4().hex[:12]}" - source_url = "https://example.test/files/contract-source.pdf" - s3_key = f"uploads/{job_id}.pdf" - downloaded_path = tmp_path / "downloaded-contract-source.pdf" - uploaded_calls: list[tuple[str, str, str]] = [] - - def resolve_public_address( - host: str, - port: int | None, - *args: object, - **kwargs: object, - ) -> list[tuple[socket.AddressFamily, socket.SocketKind, int, str, tuple[str, int]]]: - return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 0))] - - monkeypatch.setattr( - url_upload_service, - "download_source_url_to_temp", - lambda _source_url: str(downloaded_path), - ) - monkeypatch.setitem( - document_ingestion_tasks.upload_url_file.__globals__, - "download_source_url_to_temp", - lambda _source_url: str(downloaded_path), - ) - monkeypatch.setattr( - url_upload_service, - "upload_temp_file_to_source_storage", - lambda *, temp_file_path, s3_key: uploaded_calls.append( - (temp_file_path, s3_key, settings.S3_BUCKET_NAME) - ), - ) - monkeypatch.setitem( - document_ingestion_tasks.upload_url_file.__globals__, - "upload_temp_file_to_source_storage", - lambda *, temp_file_path, s3_key: uploaded_calls.append( - (temp_file_path, s3_key, settings.S3_BUCKET_NAME) - ), - ) - monkeypatch.setattr( - url_upload_service, - "verify_source_upload", - lambda storage_key: {"exists": storage_key == s3_key, "size": 3}, - ) - monkeypatch.setitem( - document_ingestion_tasks.upload_url_file.__globals__, - "verify_source_upload", - lambda storage_key: {"exists": storage_key == s3_key, "size": 3}, - ) - monkeypatch.setattr(socket, "getaddrinfo", resolve_public_address) - _bind_upload_task_to_current_module( - monkeypatch, - document_ingestion_tasks=document_ingestion_tasks, - ) - - downloaded_path.write_bytes(b"pdf") - - with engine.begin() as connection: - insert_contract_user(connection, user_id=user_id) - insert_contract_job( - connection, - job_id=job_id, - user_id=user_id, - status="waiting-file", - source_type="url", - s3_key=s3_key, - job_metadata={ - "namespace": "worker-contract", - "source_type": "url", - "source_url": source_url, - "source_file_name": "contract-source.pdf", - }, + source_dir = tmp_path / "source" + source_dir.mkdir() + source_file = source_dir / "contract-source.pdf" + source_file.write_bytes(b"pdf") + + contract = WorkerUrlUploadContract.create() + contract.use_workspace_root(monkeypatch, tmp_path) + contract.allow_private_url_sources(monkeypatch) + + with _serve_directory(source_dir) as server_url: + source_url = f"{server_url}/{source_file.name}" + job = contract.create_url_job(source_url=source_url) + + celery_result = contract.enqueue_upload_url_task( + job_id=job["job_id"], + source_url=source_url, + user_id=job["user_id"], ) - redis_service = sync_redis_service_factory.get_service() - sync_job_info_service = sync_job_info_service_cls(redis_service) - sync_job_info_service.save_job_info( - job_id, - { - "job_id": job_id, - "s3_key": s3_key, - "user_id": user_id, - "webhook_enabled": False, - "job_type": "document_ingestion", - "source_type": "url", - }, - ) - - result = document_ingestion_tasks.upload_url_file_task.run( - job_id, - source_url, - user_id, - "document_ingestion", - ) - - assert result == { + assert celery_result.successful() + assert celery_result.result == { "status": "success", - "job_id": job_id, - "s3_key": s3_key, - "file_size": 3, + "job_id": job["job_id"], + "s3_key": job["s3_key"], + "file_size": len(b"pdf"), } - assert uploaded_calls == [ - (str(downloaded_path), s3_key, settings.S3_BUCKET_NAME), - ] - assert os.path.exists(downloaded_path) is False - progress = redis_service.hgetall(f"task:{job_id}:progress") + uploaded_source = contract.verify_uploaded_source_object(job["s3_key"]) + assert uploaded_source["exists"] is True + assert uploaded_source["size"] == len(b"pdf") + assert contract.read_uploaded_source_object(job["s3_key"]) == b"pdf" + + progress = contract.get_task_progress(job["job_id"]) assert progress["progress"] == 100 assert progress["message"] == "URL file upload complete, waiting for processing..." assert progress["timestamp"] - with engine.begin() as connection: - job_row = ( - connection.execute( - text( - """ - SELECT status, source_type, s3_key - FROM jobs - WHERE job_id = :job_id - """ - ), - {"job_id": job_id}, - ) - .mappings() - .one() - ) - + job_row = contract.observe_job(job["job_id"]) assert job_row["status"] == "waiting-file" assert job_row["source_type"] == "url" - assert job_row["s3_key"] == s3_key + assert job_row["s3_key"] == job["s3_key"] diff --git a/apps/worker/tests/support/worker_parse_contract.py b/apps/worker/tests/support/worker_parse_contract.py new file mode 100644 index 000000000..f88ea0542 --- /dev/null +++ b/apps/worker/tests/support/worker_parse_contract.py @@ -0,0 +1,435 @@ +from __future__ import annotations + +import json +import sys +import zipfile +from dataclasses import dataclass +from pathlib import Path +from typing import Any +from uuid import uuid4 + +from pytest import MonkeyPatch +from sqlalchemy import text +from sqlalchemy.engine import Engine + +from support.contract_database import insert_contract_job + +_DOCUMENT_PARSE_JOB_TYPE = "document_ingestion" +_DOCUMENT_PARSE_TASK_NAME = "app.core.tasks.document_ingestion_tasks.parse_task" + + +@dataclass(frozen=True) +class WorkerParseContract: + engine: Engine + settings: Any + storage: Any + sync_redis_service_factory: Any + + @classmethod + def create(cls) -> WorkerParseContract: + from shared.core.config import settings + from shared.core.database_sync import get_sync_engine + from shared.services.redis.redis_sync_service import ( + SyncRedisServiceFactory, + ) + from shared.services.storage.job_file_storage import JobFileStorage + + return cls( + engine=get_sync_engine(), + settings=settings, + storage=JobFileStorage(), + sync_redis_service_factory=SyncRedisServiceFactory, + ) + + def use_workspace_root( + self, + monkeypatch: MonkeyPatch, + workspace_root: Path, + ) -> None: + monkeypatch.setattr(self.settings, "TMP_PATH", str(workspace_root)) + + def use_billing(self, monkeypatch: MonkeyPatch, is_enabled: bool) -> None: + monkeypatch.setenv("BILLING_ENABLED", "true" if is_enabled else "false") + monkeypatch.setattr(self.settings, "BILLING_ENABLED", is_enabled) + for loaded_module in list(sys.modules.values()): + module_settings = getattr(loaded_module, "settings", None) + if hasattr(module_settings, "BILLING_ENABLED"): + monkeypatch.setattr( + module_settings, + "BILLING_ENABLED", + is_enabled, + raising=False, + ) + def use_pdf_page_limit(self, monkeypatch: MonkeyPatch, page_limit: int) -> None: + monkeypatch.setenv("MAX_PDF_PAGE_LIMIT", str(page_limit)) + monkeypatch.setattr(self.settings, "MAX_PDF_PAGE_LIMIT", page_limit) + for loaded_module in list(sys.modules.values()): + module_settings = getattr(loaded_module, "settings", None) + if hasattr(module_settings, "MAX_PDF_PAGE_LIMIT"): + monkeypatch.setattr( + module_settings, + "MAX_PDF_PAGE_LIMIT", + page_limit, + raising=False, + ) + + def create_file_job( + self, + *, + source_file_name: str, + user_id: str | None = None, + status: str = "pending", + billing_status: str = "pending", + job_id_prefix: str = "job_parse", + ) -> dict[str, Any]: + effective_user_id = user_id or f"worker-contract-user-{uuid4().hex[:12]}" + job_id = f"{job_id_prefix}_{uuid4().hex[:12]}" + s3_key = self.storage.build_upload_key( + job_id=job_id, + file_extension=Path(source_file_name).suffix, + ) + with self.engine.begin() as connection: + self.ensure_user(connection, user_id=effective_user_id) + insert_contract_job( + connection, + job_id=job_id, + user_id=effective_user_id, + status=status, + source_type="file", + file_path=source_file_name, + s3_key=s3_key, + job_metadata={ + "namespace": "worker-contract", + "source_type": "file", + "source_file_name": source_file_name, + "parsing_params": { + "doc_type": "auto", + "smart_title_parse": False, + "summary_image": False, + "summary_table": False, + "summary_txt": False, + }, + }, + billing_status=billing_status, + ) + + return { + "job_id": job_id, + "user_id": effective_user_id, + "source_file_name": source_file_name, + "s3_key": s3_key, + } + + @staticmethod + def ensure_user(connection: Any, *, user_id: str) -> None: + connection.execute( + text( + """ + INSERT INTO "user" (id, name, email) + VALUES (:user_id, :name, :email) + ON CONFLICT (id) DO NOTHING + """ + ), + { + "user_id": user_id, + "name": f"Worker Contract User {user_id}", + "email": f"{user_id}@worker-contract.knowhere.local", + }, + ) + + def upload_source_file(self, *, local_file_path: Path, s3_key: str) -> None: + upload_result = self.storage.upload_source_file(str(local_file_path), s3_key) + assert upload_result["status"] == "success" + assert self.storage.verify_upload_exists(s3_key)["exists"] is True + + def enqueue_parse_task(self, *, job_id: str, user_id: str) -> Any: + from shared.core.celery_app import get_celery_app + from shared.core.celery_router import task_router + + queue_name = task_router.get_queue_for_job(_DOCUMENT_PARSE_JOB_TYPE, user_id) + task_signature = get_celery_app().signature( + _DOCUMENT_PARSE_TASK_NAME, + args=[job_id], + kwargs={ + "user_id": user_id, + "job_type": _DOCUMENT_PARSE_JOB_TYPE, + }, + ).set(queue=queue_name) + return task_signature.apply_async() + + def get_task_status(self, job_id: str) -> Any: + redis_service = self.sync_redis_service_factory.get_service() + return redis_service.get(f"task:{job_id}:status") + + def get_task_progress(self, job_id: str) -> dict[str, Any]: + redis_service = self.sync_redis_service_factory.get_service() + return redis_service.hgetall(f"task:{job_id}:progress") + + def observe_successful_job(self, job_id: str) -> dict[str, Any]: + with self.engine.connect() as connection: + job_row = ( + connection.execute( + text( + """ + SELECT + status, + billing_status, + page_count, + credits_charged, + error_code, + error_message + FROM jobs + WHERE job_id = :job_id + """ + ), + {"job_id": job_id}, + ) + .mappings() + .one() + ) + result_row = ( + connection.execute( + text( + """ + SELECT + id, + document_id, + result_s3_key, + result_size, + document_metadata + FROM job_results + WHERE job_id = :job_id + """ + ), + {"job_id": job_id}, + ) + .mappings() + .one() + ) + job_chunks = ( + connection.execute( + text( + """ + SELECT chunk_id, chunk_type, text, path, chunk_metadata + FROM job_chunks + WHERE job_result_id = :job_result_id + ORDER BY sort_order + """ + ), + {"job_result_id": result_row["id"]}, + ) + .mappings() + .all() + ) + document_chunks = ( + connection.execute( + text( + """ + SELECT chunk_id, chunk_type, content, source_chunk_path + FROM document_chunks + WHERE job_result_id = :job_result_id + ORDER BY sort_order + """ + ), + {"job_result_id": result_row["id"]}, + ) + .mappings() + .all() + ) + document_sections_count = connection.execute( + text( + """ + SELECT COUNT(*) + FROM document_sections + WHERE job_result_id = :job_result_id + """ + ), + {"job_result_id": result_row["id"]}, + ).scalar_one() + + return { + "job": job_row, + "result": result_row, + "job_chunks": job_chunks, + "document_chunks": document_chunks, + "document_sections_count": document_sections_count, + } + + def observe_job_status(self, job_id: str) -> dict[str, Any]: + with self.engine.connect() as connection: + return ( + connection.execute( + text( + """ + SELECT + status, + billing_status, + page_count, + credits_charged, + error_code, + error_message + FROM jobs + WHERE job_id = :job_id + """ + ), + {"job_id": job_id}, + ) + .mappings() + .one() + ) + + def get_job_metadata(self, job_id: str) -> dict[str, Any]: + from shared.services.redis.redis_sync_service import SyncJobMetadataService + + redis_service = self.sync_redis_service_factory.get_service() + metadata = SyncJobMetadataService(redis_service).get_metadata(job_id) + if isinstance(metadata, dict) and metadata: + return dict(metadata) + + with self.engine.connect() as connection: + stored_metadata = connection.execute( + text("SELECT job_metadata FROM jobs WHERE job_id = :job_id"), + {"job_id": job_id}, + ).scalar_one_or_none() + return dict(stored_metadata or {}) + + def observe_user_billing(self, user_id: str) -> dict[str, Any]: + with self.engine.connect() as connection: + balance = connection.execute( + text( + """ + SELECT credits_balance + FROM user_balances + WHERE user_id = :user_id + """ + ), + {"user_id": user_id}, + ).scalar_one_or_none() + transaction_types = list( + connection.execute( + text( + """ + SELECT transaction_type + FROM credits_transactions + WHERE user_id = :user_id + ORDER BY created_at ASC, id ASC + """ + ), + {"user_id": user_id}, + ) + .scalars() + .all() + ) + transaction_count_rows = ( + connection.execute( + text( + """ + SELECT transaction_type, COUNT(*) AS count + FROM credits_transactions + WHERE user_id = :user_id + GROUP BY transaction_type + ORDER BY transaction_type + """ + ), + {"user_id": user_id}, + ) + .mappings() + .all() + ) + system_grant_payment_count = int( + connection.execute( + text( + """ + SELECT COUNT(*) + FROM payment_records + WHERE user_id = :user_id + AND payment_type = 'system_grant' + """ + ), + {"user_id": user_id}, + ).scalar_one() + ) + + return { + "balance": int(balance) if balance is not None else None, + "transaction_types": transaction_types, + "transaction_counts": { + str(row["transaction_type"]): int(row["count"]) + for row in transaction_count_rows + }, + "system_grant_payment_count": system_grant_payment_count, + } + + def observe_job_state_transitions(self, job_id: str) -> list[tuple[str, str]]: + with self.engine.connect() as connection: + rows = ( + connection.execute( + text( + """ + SELECT transition_reason, to_state + FROM job_state_audit_logs + WHERE job_id = :job_id + ORDER BY id ASC + """ + ), + {"job_id": job_id}, + ) + .mappings() + .all() + ) + return [ + (str(row["transition_reason"]), str(row["to_state"])) for row in rows + ] + + def count_job_results(self, job_id: str) -> int: + with self.engine.connect() as connection: + return int( + connection.execute( + text("SELECT COUNT(*) FROM job_results WHERE job_id = :job_id"), + {"job_id": job_id}, + ).scalar_one() + ) + + def verify_result_zip_object(self, result_s3_key: str) -> dict[str, Any]: + return self.storage.verify_exists( + result_s3_key, + bucket=self.settings.S3_RESULTS_BUCKET, + ) + + def read_result_zip( + self, + *, + result_s3_key: str, + tmp_path: Path, + ) -> dict[str, Any]: + downloaded_zip_path = tmp_path / "result.zip" + self.storage.download_to_path( + result_s3_key, + str(downloaded_zip_path), + bucket=self.settings.S3_RESULTS_BUCKET, + ) + + with zipfile.ZipFile(downloaded_zip_path) as archive: + members = set(archive.namelist()) + with archive.open("chunks.json") as chunks_file: + chunks_payload = json.loads(chunks_file.read().decode("utf-8")) + with archive.open("manifest.json") as manifest_file: + manifest_payload = json.loads(manifest_file.read().decode("utf-8")) + + return { + "path": downloaded_zip_path, + "members": members, + "chunks": chunks_payload, + "manifest": manifest_payload, + } + + def find_task_workspaces(self, workspace_root: Path, job_id: str) -> list[Path]: + return sorted( + path + for path in workspace_root.iterdir() + if path.is_dir() + and path.name.startswith(f"document_ingestion_task_{job_id}_") + ) + + +__all__ = ["WorkerParseContract"] diff --git a/apps/worker/tests/support/worker_url_upload_contract.py b/apps/worker/tests/support/worker_url_upload_contract.py new file mode 100644 index 000000000..c5b249b4e --- /dev/null +++ b/apps/worker/tests/support/worker_url_upload_contract.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Any +from uuid import uuid4 + +from pytest import MonkeyPatch +from sqlalchemy import text +from sqlalchemy.engine import Engine + +from support.contract_database import insert_contract_job, insert_contract_user + +_DOCUMENT_INGESTION_JOB_TYPE = "document_ingestion" +_URL_UPLOAD_TASK_NAME = "app.core.tasks.document_ingestion_tasks.upload_url_file_task" + + +@dataclass(frozen=True) +class WorkerUrlUploadContract: + engine: Engine + settings: Any + storage: Any + sync_redis_service_factory: Any + + @classmethod + def create(cls) -> WorkerUrlUploadContract: + from shared.core.config import settings + from shared.core.database_sync import get_sync_engine + from shared.services.redis.redis_sync_service import SyncRedisServiceFactory + from shared.services.storage.job_file_storage import JobFileStorage + + return cls( + engine=get_sync_engine(), + settings=settings, + storage=JobFileStorage(), + sync_redis_service_factory=SyncRedisServiceFactory, + ) + + def use_workspace_root( + self, + monkeypatch: MonkeyPatch, + workspace_root: Path, + ) -> None: + monkeypatch.setattr(self.settings, "TMP_PATH", str(workspace_root)) + + def allow_private_url_sources(self, monkeypatch: MonkeyPatch) -> None: + monkeypatch.setattr(self.settings, "ENVIRONMENT", "development") + + def create_url_job(self, *, source_url: str) -> dict[str, str]: + user_id = f"worker-url-contract-user-{uuid4().hex[:12]}" + job_id = f"job_url_upload_{uuid4().hex[:12]}" + s3_key = self.storage.build_upload_key(job_id=job_id, file_extension=".pdf") + + with self.engine.begin() as connection: + insert_contract_user(connection, user_id=user_id) + insert_contract_job( + connection, + job_id=job_id, + user_id=user_id, + status="waiting-file", + source_type="url", + s3_key=s3_key, + job_metadata={ + "namespace": "worker-contract", + "source_type": "url", + "source_url": source_url, + "source_file_name": "contract-source.pdf", + }, + ) + + return { + "job_id": job_id, + "user_id": user_id, + "source_url": source_url, + "s3_key": s3_key, + } + + def enqueue_upload_url_task( + self, + *, + job_id: str, + source_url: str, + user_id: str, + ) -> Any: + from shared.core.celery_app import get_celery_app + + task_signature = get_celery_app().signature(_URL_UPLOAD_TASK_NAME) + return task_signature.apply_async( + args=[job_id, source_url, user_id], + kwargs={"job_type": _DOCUMENT_INGESTION_JOB_TYPE}, + ) + + def get_task_progress(self, job_id: str) -> dict[str, Any]: + redis_service = self.sync_redis_service_factory.get_service() + return redis_service.hgetall(f"task:{job_id}:progress") + + def observe_job(self, job_id: str) -> dict[str, Any]: + with self.engine.connect() as connection: + return ( + connection.execute( + text( + """ + SELECT status, source_type, s3_key + FROM jobs + WHERE job_id = :job_id + """ + ), + {"job_id": job_id}, + ) + .mappings() + .one() + ) + + def verify_uploaded_source_object(self, s3_key: str) -> dict[str, Any]: + return self.storage.verify_upload_exists(s3_key) + + def read_uploaded_source_object(self, s3_key: str) -> bytes: + return self.storage.storage_adapter.download_fileobj( + s3_key, + bucket=self.settings.S3_BUCKET_NAME, + ) + + +__all__ = ["WorkerUrlUploadContract"] diff --git a/packages/shared-python/shared/core/config/storage.py b/packages/shared-python/shared/core/config/storage.py index 556879f7f..3b3afbe61 100644 --- a/packages/shared-python/shared/core/config/storage.py +++ b/packages/shared-python/shared/core/config/storage.py @@ -23,7 +23,10 @@ class StorageConfig(BaseModel): """Storage configuration.""" # Storage backend selection. - S3_TYPE: str = Field(default="s3", description="Storage backend: s3, oss, or minio") + S3_TYPE: str = Field( + default="s3", + description="Storage backend: s3, oss, minio, or filesystem", + ) # Shared S3-style configuration used by S3, OSS, and MinIO. S3_BUCKET_NAME: str = Field(..., description="Bucket name") @@ -34,6 +37,10 @@ class StorageConfig(BaseModel): ) S3_PRIVATE_DOMAIN: str = Field(default="", description="Private asset domain") S3_TEMP_PATH: str = Field(..., description="Temporary path") + OBJECT_STORAGE_LOCAL_ROOT: str = Field( + default="", + description="Local root directory used when S3_TYPE=filesystem", + ) # Advanced S3 client configuration. S3_REGION: str = Field( @@ -146,6 +153,15 @@ def get_storage_adapter(self): """ storage_type = os.getenv("S3_TYPE", self.S3_TYPE).lower() + if storage_type == "filesystem": + from shared.services.storage.adapters import FileSystemStorageAdapter + + local_root = self.OBJECT_STORAGE_LOCAL_ROOT or os.path.join( + self.S3_TEMP_PATH, + "object-storage", + ) + return FileSystemStorageAdapter(local_root, self.S3_BUCKET_NAME) + if storage_type == "oss": # OSS storage adapter (imported lazily). from shared.services.storage.adapters.oss_adapter import OSSStorageAdapter diff --git a/packages/shared-python/shared/services/billing/work_billing_service.py b/packages/shared-python/shared/services/billing/work_billing_service.py index 28c82d1fe..11bcdaaca 100644 --- a/packages/shared-python/shared/services/billing/work_billing_service.py +++ b/packages/shared-python/shared/services/billing/work_billing_service.py @@ -5,7 +5,6 @@ from sqlalchemy.orm import Session from shared.core.billing import BillingCalculator -from shared.core.config import settings from shared.services.billing.credits_sync_service import SyncCreditsService _SKIPPED_BILLING_STATUS: str = "skipped" @@ -51,6 +50,8 @@ def charge_for_pages( filename: str, ) -> WorkBillingResult: """Authorize and charge per-page work if billing is enabled.""" + from shared.core.config import settings + if not settings.BILLING_ENABLED: return WorkBillingResult.skipped() diff --git a/packages/shared-python/shared/services/jobs/lifecycle/service.py b/packages/shared-python/shared/services/jobs/lifecycle/service.py index 3a611d1f8..3683fca30 100644 --- a/packages/shared-python/shared/services/jobs/lifecycle/service.py +++ b/packages/shared-python/shared/services/jobs/lifecycle/service.py @@ -126,20 +126,21 @@ def update_progress( job_id: str, progress: int, message: str = "", + redis_service: Any | None = None, ) -> bool: """Write job progress directly to Redis (replaces publish_progress_update). Best-effort — failures are logged but do not raise. """ try: - redis_service = SyncRedisServiceFactory.get_service() - task_ttl = redis_key_builder.get_key_ttl(RedisKeyType.TASK) - progress_key = redis_service._build_key( - redis_key_builder.task_progress(job_id) + active_redis_service = ( + redis_service + if redis_service is not None + else SyncRedisServiceFactory.get_service() ) - - pipe = redis_service.pipeline() - pipe.hset( + task_ttl = redis_key_builder.get_key_ttl(RedisKeyType.TASK) + progress_key = redis_key_builder.task_progress(job_id) + active_redis_service.hset( progress_key, mapping={ "progress": str(progress), @@ -147,8 +148,7 @@ def update_progress( "timestamp": str(int(time.time())), }, ) - pipe.expire(progress_key, task_ttl) - pipe.execute() + active_redis_service.expire(progress_key, task_ttl) return True except Exception as exc: logger.warning(f"Failed to update progress for job {job_id}: {exc}") diff --git a/packages/shared-python/shared/services/redis/redis_sync_service.py b/packages/shared-python/shared/services/redis/redis_sync_service.py index 90e363aa1..812b838f6 100644 --- a/packages/shared-python/shared/services/redis/redis_sync_service.py +++ b/packages/shared-python/shared/services/redis/redis_sync_service.py @@ -359,11 +359,9 @@ def get_metadata(self, job_id: str) -> Optional[Dict[str, Any]]: def update_metadata(self, job_id: str, updates: Dict[str, Any]) -> bool: try: - metadata = self.get_metadata(job_id) - if metadata: - metadata.update(updates) - return self.save_metadata(job_id, metadata) - return False + metadata = self.get_metadata(job_id) or {} + metadata.update(updates) + return self.save_metadata(job_id, metadata) except Exception as e: logger.error(f"Failed to update metadata: {e}") return False diff --git a/packages/shared-python/shared/services/storage/adapters/__init__.py b/packages/shared-python/shared/services/storage/adapters/__init__.py index 59a663584..f2c2b6ba4 100644 --- a/packages/shared-python/shared/services/storage/adapters/__init__.py +++ b/packages/shared-python/shared/services/storage/adapters/__init__.py @@ -1,10 +1,11 @@ """Storage adapter exports.""" +from .filesystem_adapter import FileSystemStorageAdapter from .s3_adapter import S3StorageAdapter # Import OSSStorageAdapter lazily so environments without oss2 still import safely. -__all__ = ["S3StorageAdapter"] +__all__ = ["FileSystemStorageAdapter", "S3StorageAdapter"] def get_oss_adapter(): diff --git a/packages/shared-python/shared/services/storage/adapters/filesystem_adapter.py b/packages/shared-python/shared/services/storage/adapters/filesystem_adapter.py new file mode 100644 index 000000000..efd4389fb --- /dev/null +++ b/packages/shared-python/shared/services/storage/adapters/filesystem_adapter.py @@ -0,0 +1,160 @@ +"""Filesystem-backed object storage adapter for contract runtimes.""" + +import shutil +from pathlib import Path, PurePosixPath +from typing import Any, BinaryIO, Iterator, Optional +from urllib.parse import quote + +from shared.core.exceptions.domain_exceptions import StorageServiceException +from shared.services.storage.storage_adapter import StorageAdapter + + +class FileSystemStorageAdapter(StorageAdapter): + """Store object-storage buckets and keys under a local directory.""" + + def __init__(self, root_path: str, default_bucket: str) -> None: + self.root_path = Path(root_path).resolve() + self.default_bucket = default_bucket + self.root_path.mkdir(parents=True, exist_ok=True) + + def upload_file( + self, local_path: str, key: str, bucket: Optional[str] = None + ) -> dict[str, Any]: + bucket_name = self._get_bucket(bucket) + object_path = self._resolve_object_path(key, bucket_name) + object_path.parent.mkdir(parents=True, exist_ok=True) + shutil.copyfile(local_path, object_path) + return {"bucket": bucket_name, "key": key, "status": "success"} + + def upload_fileobj( + self, + file_obj: BinaryIO, + key: str, + bucket: Optional[str] = None, + content_type: Optional[str] = None, + ) -> dict[str, Any]: + bucket_name = self._get_bucket(bucket) + object_path = self._resolve_object_path(key, bucket_name) + object_path.parent.mkdir(parents=True, exist_ok=True) + with object_path.open("wb") as output_file: + output_file.write(file_obj.read()) + return { + "bucket": bucket_name, + "key": key, + "status": "success", + "content_type": content_type, + } + + def download_file( + self, key: str, local_path: str, bucket: Optional[str] = None + ) -> str: + source_path = self._resolve_object_path(key, self._get_bucket(bucket)) + if not source_path.is_file(): + raise StorageServiceException( + internal_message=f"Filesystem object not found: {key}", + operation="download_file", + ) + + destination_path = Path(local_path) + destination_path.parent.mkdir(parents=True, exist_ok=True) + shutil.copyfile(source_path, destination_path) + return local_path + + def download_fileobj(self, key: str, bucket: Optional[str] = None) -> bytes: + source_path = self._resolve_object_path(key, self._get_bucket(bucket)) + if not source_path.is_file(): + raise StorageServiceException( + internal_message=f"Filesystem object not found: {key}", + operation="download_fileobj", + ) + return source_path.read_bytes() + + def delete_object(self, key: str, bucket: Optional[str] = None) -> bool: + object_path = self._resolve_object_path(key, self._get_bucket(bucket)) + if not object_path.exists(): + return False + object_path.unlink() + return True + + def list_objects( + self, prefix: str = "", bucket: Optional[str] = None + ) -> Iterator[str]: + bucket_path = self._resolve_bucket_path(self._get_bucket(bucket)) + if not bucket_path.is_dir(): + return iter(()) + + object_keys = ( + file_path.relative_to(bucket_path).as_posix() + for file_path in bucket_path.rglob("*") + if file_path.is_file() + ) + return iter(sorted(key for key in object_keys if key.startswith(prefix))) + + def generate_presigned_url( + self, + key: str, + expiration: int = 3600, + bucket: Optional[str] = None, + method: str = "GET", + headers: Optional[dict[str, str]] = None, + ) -> str: + bucket_name = self._get_bucket(bucket) + encoded_key = quote(key, safe="/") + return ( + f"filesystem://{bucket_name}/{encoded_key}" + f"?method={method.upper()}&expires_in={expiration}" + ) + + def exists(self, key: str, bucket: Optional[str] = None) -> bool: + return self._resolve_object_path(key, self._get_bucket(bucket)).is_file() + + def get_object_size(self, key: str, bucket: Optional[str] = None) -> Optional[int]: + object_path = self._resolve_object_path(key, self._get_bucket(bucket)) + if not object_path.is_file(): + return None + return object_path.stat().st_size + + def _get_bucket(self, bucket: Optional[str] = None) -> str: + bucket_name = bucket or self.default_bucket + if not bucket_name: + raise StorageServiceException( + internal_message="Filesystem bucket name cannot be empty", + operation="resolve_bucket", + ) + if any(part in {"", ".", ".."} for part in bucket_name.split("/")): + raise StorageServiceException( + internal_message=f"Invalid filesystem bucket name: {bucket_name}", + operation="resolve_bucket", + ) + return bucket_name + + def _resolve_bucket_path(self, bucket: str) -> Path: + bucket_path = (self.root_path / bucket).resolve() + if not bucket_path.is_relative_to(self.root_path): + raise StorageServiceException( + internal_message=f"Bucket escapes object-storage root: {bucket}", + operation="resolve_bucket", + ) + return bucket_path + + def _resolve_object_path(self, key: str, bucket: str) -> Path: + key_parts = PurePosixPath(key.strip("/")).parts + if not key_parts: + raise StorageServiceException( + internal_message="Filesystem object key cannot be empty", + operation="resolve_object_key", + ) + if any(part in {"", ".", ".."} for part in key_parts): + raise StorageServiceException( + internal_message=f"Invalid filesystem object key: {key}", + operation="resolve_object_key", + ) + + bucket_path = self._resolve_bucket_path(bucket) + object_path = bucket_path.joinpath(*key_parts).resolve() + if not object_path.is_relative_to(bucket_path): + raise StorageServiceException( + internal_message=f"Object key escapes bucket root: {key}", + operation="resolve_object_key", + ) + return object_path diff --git a/packages/shared-python/shared/testing/contract_runtime.py b/packages/shared-python/shared/testing/contract_runtime.py index ce029a8c6..27a3c7de1 100644 --- a/packages/shared-python/shared/testing/contract_runtime.py +++ b/packages/shared-python/shared/testing/contract_runtime.py @@ -41,6 +41,9 @@ _API_ROOT: Path = _REPO_ROOT / "apps" / "api" _SHARED_ROOT: Path = _REPO_ROOT / "packages" / "shared-python" _TEST_TMP_ROOT: Path = Path("/tmp/knowhere-api-tests") +_TEST_OBJECT_STORAGE_ROOT: Path = _TEST_TMP_ROOT / "object-storage" +_CONTRACT_UPLOADS_BUCKET: str = "knowhere-test-bucket" +_CONTRACT_RESULTS_BUCKET: str = "knowhere-test-results" _STATIC_TABLES_TO_PRESERVE: frozenset[str] = frozenset( { "alembic_version", @@ -328,13 +331,15 @@ def configure_contract_environment( f"redis://{CONTRACT_REDIS_HOST}:{CONTRACT_REDIS_PORT}/{CONTRACT_REDIS_DATABASE}" ), "TMP_PATH": str(_TEST_TMP_ROOT), - "S3_BUCKET_NAME": "knowhere-test-bucket", + "S3_TYPE": "filesystem", + "S3_BUCKET_NAME": _CONTRACT_UPLOADS_BUCKET, "S3_ACCESS_KEY_ID": "test-access-key", "S3_SECRET_ACCESS_KEY": "test-secret-key", "S3_TEMP_PATH": str(_TEST_TMP_ROOT), - "S3_ENDPOINT_URL": "http://127.0.0.1:4566", - "S3_PRIVATE_DOMAIN": "http://127.0.0.1:4566", - "S3_RESULTS_BUCKET": "knowhere-test-results", + "S3_ENDPOINT_URL": "", + "S3_PRIVATE_DOMAIN": "", + "S3_RESULTS_BUCKET": _CONTRACT_RESULTS_BUCKET, + "OBJECT_STORAGE_LOCAL_ROOT": str(_TEST_OBJECT_STORAGE_ROOT), "S3_REGION": "us-west-1", "S3_USE_SSL": "false", "S3_ADDRESSING_STYLE": "path", @@ -377,12 +382,28 @@ def clear_application_modules() -> None: sys.modules.pop(module_name, None) continue + if module_name == "shared.services.billing" or module_name.startswith( + "shared.services.billing." + ): + sys.modules.pop(module_name, None) + continue + if module_name == "shared.services.jobs" or module_name.startswith( "shared.services.jobs." ): sys.modules.pop(module_name, None) continue + if module_name == "app.services.document_ingestion" or module_name.startswith( + "app.services.document_ingestion." + ): + sys.modules.pop(module_name, None) + continue + + if module_name == "app.core.tasks.document_ingestion_tasks": + sys.modules.pop(module_name, None) + continue + if module_name == "shared.services.webhook" or module_name.startswith( "shared.services.webhook." ): @@ -702,6 +723,7 @@ async def prepare_contract_storage() -> None: await reset_contract_database() await reset_contract_redis() + reset_contract_object_storage() async def seed_contract_developer() -> dict[str, str | int]: @@ -779,3 +801,15 @@ async def reset_contract_redis() -> None: await redis_client.flushdb() finally: await redis_client.aclose() + + +def reset_contract_object_storage() -> None: + shutil.rmtree(_TEST_OBJECT_STORAGE_ROOT, ignore_errors=True) + (_TEST_OBJECT_STORAGE_ROOT / _CONTRACT_UPLOADS_BUCKET).mkdir( + parents=True, + exist_ok=True, + ) + (_TEST_OBJECT_STORAGE_ROOT / _CONTRACT_RESULTS_BUCKET).mkdir( + parents=True, + exist_ok=True, + )