From dd3cf53044d95056bcbaf35be3a0f4e1f56fc579 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Tue, 11 Aug 2026 16:05:36 +0900 Subject: [PATCH] test: migrate retention and task sessions to SQLite (#40090) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../workflow_run/test_archive_bundle_index.py | 58 ++-- .../test_archive_download_preparation.py | 49 ++- .../workflow_run/test_archive_log_service.py | 160 ++++----- .../test_bundle_archive_maintenance.py | 320 ++++++++++-------- .../tasks/test_batch_clean_document_task.py | 73 ++-- .../tasks/test_community_telemetry_task.py | 39 ++- .../tasks/test_enable_segment_index_tasks.py | 221 ++++++------ .../test_retry_document_indexing_task.py | 65 +++- .../tasks/test_segment_index_cleanup_tasks.py | 148 +++++--- .../tasks/test_trigger_processing_tasks.py | 9 - .../tasks/test_workflow_execute_task.py | 46 +-- 11 files changed, 689 insertions(+), 499 deletions(-) diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py b/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py index df26e7a47af..6dd0098c863 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py @@ -3,6 +3,10 @@ import json from typing import cast from unittest.mock import MagicMock +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.archive_bundle_index import ( ARCHIVE_BUNDLE_ROOT_PREFIX, ArchiveBundleManifest, @@ -90,12 +94,11 @@ def test_decode_and_calculate_archive_bundle_index_values() -> None: assert values.archived_at == datetime.datetime(2026, 6, 25, 8, 0) -def test_upsert_archive_bundle_index_inserts_new_bundle() -> None: - session = MagicMock() - session.scalar.return_value = None +def test_upsert_archive_bundle_index_inserts_new_bundle(sqlite_session: Session) -> None: data = _manifest_bytes() - bundle = upsert_archive_bundle_index_from_manifest(session, decode_archive_bundle_manifest(data), len(data)) + bundle = upsert_archive_bundle_index_from_manifest(sqlite_session, decode_archive_bundle_manifest(data), len(data)) + sqlite_session.flush() assert bundle.tenant_id == TENANT_ID assert bundle.year == 2025 @@ -103,34 +106,42 @@ def test_upsert_archive_bundle_index_inserts_new_bundle() -> None: assert bundle.workflow_run_count == 2 assert bundle.row_count == 5 assert bundle.archive_bytes == len(data) + 300 - session.add.assert_called_once_with(bundle) + assert sqlite_session.get(WorkflowRunArchiveBundle, bundle.id) is bundle -def test_upsert_archive_bundle_index_updates_existing_bundle() -> None: - existing = MagicMock() - session = MagicMock() - session.scalar.return_value = existing +def test_upsert_archive_bundle_index_updates_existing_bundle(sqlite_session: Session) -> None: + existing = WorkflowRunArchiveBundle( + tenant_id=TENANT_ID, + year=2025, + month=3, + shard="00-of-01", + bundle_id=BUNDLE_ID, + workflow_run_count=1, + row_count=1, + archive_bytes=1, + archived_at=datetime.datetime(2025, 3, 1), + ) + sqlite_session.add(existing) + sqlite_session.flush() data = _manifest_bytes() - bundle = upsert_archive_bundle_index_from_manifest(session, decode_archive_bundle_manifest(data), len(data)) + bundle = upsert_archive_bundle_index_from_manifest(sqlite_session, decode_archive_bundle_manifest(data), len(data)) assert bundle is existing assert existing.workflow_run_count == 2 assert existing.row_count == 5 assert existing.archive_bytes == len(data) + 300 assert existing.archived_at == datetime.datetime(2026, 6, 25, 8, 0) - session.add.assert_not_called() + assert sqlite_session.query(WorkflowRunArchiveBundle).count() == 1 -def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index() -> None: +def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: storage = FakeArchiveStorage({MANIFEST_KEY: _manifest_bytes()}) - session = MagicMock() - session.scalar.return_value = None - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session backfill = WorkflowRunArchiveBundleIndexBackfill( storage=cast(MagicMock, storage), - session_factory=session_factory, + session_factory=sqlite_session_factory, ) summary = backfill.run(tenant_ids=[TENANT_ID], year=2025, month=3) @@ -142,11 +153,13 @@ def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index() -> None: assert summary.bundles_processed == 1 assert summary.bundles_upserted == 1 assert summary.bundles_failed == 0 - session.add.assert_called_once() - session.commit.assert_called_once() + sqlite_session.expire_all() + assert sqlite_session.scalar(select(WorkflowRunArchiveBundle)) is not None -def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None: +def test_backfill_dry_run_filters_by_year_month_without_database_write( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: other_month_prefix = OBJECT_PREFIX.replace("month=03", "month=04") storage = FakeArchiveStorage( { @@ -156,10 +169,9 @@ def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None ), } ) - session_factory = MagicMock() backfill = WorkflowRunArchiveBundleIndexBackfill( storage=cast(MagicMock, storage), - session_factory=session_factory, + session_factory=sqlite_session_factory, ) summary = backfill.run(tenant_prefixes=["1"], year=2025, month=3, dry_run=True) @@ -169,4 +181,4 @@ def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None assert summary.bundles_processed == 1 assert summary.bundles_upserted == 0 assert summary.archive_bytes > 0 - session_factory.assert_not_called() + assert sqlite_session.scalar(select(WorkflowRunArchiveBundle)) is None diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py b/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py index 2bf96dc9701..d83a7bc978e 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py @@ -4,7 +4,6 @@ import io import json import zipfile from contextlib import nullcontext -from types import SimpleNamespace from typing import cast from unittest.mock import MagicMock @@ -83,7 +82,17 @@ def _object_prefix(bundle_id: str = BUNDLE_ID) -> str: def _bundle(bundle_id: str = BUNDLE_ID) -> WorkflowRunArchiveBundle: - return cast(WorkflowRunArchiveBundle, SimpleNamespace(shard=SHARD, bundle_id=bundle_id)) + return WorkflowRunArchiveBundle( + tenant_id=TENANT_ID, + year=2025, + month=3, + shard=SHARD, + bundle_id=bundle_id, + workflow_run_count=1, + row_count=1, + archive_bytes=100, + archived_at=datetime.datetime(2026, 6, 25, 8), + ) def _task(bundle_refs: list[tuple[str, str]] | None = None) -> WorkflowRunArchiveDownloadTask: @@ -145,16 +154,16 @@ def _preparer( archive_storage: FakeArchiveStorage | None = None, download_storage: FakeArchiveStorage | None = None, cache: FakeTaskCache, + session_factory: sessionmaker[Session], bundles: list[WorkflowRunArchiveBundle] | None = None, ) -> WorkflowRunArchiveDownloadPreparer: source_storage = archive_storage or storage target_storage = download_storage or storage assert source_storage is not None assert target_storage is not None - session = MagicMock() - session.scalars.return_value = bundles or [_bundle()] - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session + with session_factory() as session: + session.add_all([_bundle()] if bundles is None else bundles) + session.commit() return WorkflowRunArchiveDownloadPreparer( archive_storage=cast(ArchiveStorage, source_storage), download_storage=cast(ArchiveStorage, target_storage), @@ -169,7 +178,9 @@ def _parquet_bytes(records: list[dict[str, object]]) -> bytes: return buffer.getvalue() -def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() -> None: +def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready( + sqlite_session_factory: sessionmaker[Session], +) -> None: bundle_refs = [(SHARD, "bundle-a"), (SHARD, "bundle-b")] task = _task(bundle_refs) first_bundle_payloads = { @@ -206,6 +217,7 @@ def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() archive_storage=archive_storage, download_storage=download_storage, cache=cache, + session_factory=sqlite_session_factory, bundles=[_bundle("bundle-a"), _bundle("bundle-b")], ) @@ -232,7 +244,9 @@ def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() assert '"run-b","failed","safe"' in workflow_runs_csv -def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch() -> None: +def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch( + sqlite_session_factory: sessionmaker[Session], +) -> None: task = _task() table_payloads = {"workflow_runs": _parquet_bytes([{"id": "run-a", "status": "succeeded"}])} manifest_data = json.loads(_manifest_bytes(table_payloads).decode("utf-8")) @@ -244,7 +258,7 @@ def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch } ) cache = FakeTaskCache(task) - preparer = _preparer(storage=storage, cache=cache) + preparer = _preparer(storage=storage, cache=cache, session_factory=sqlite_session_factory) result = preparer.prepare(tenant_id=TENANT_ID, download_id=task.download_id) @@ -254,11 +268,18 @@ def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch assert storage.put_objects == {} -def test_prepare_workflow_run_archive_download_skips_duplicate_worker() -> None: +def test_prepare_workflow_run_archive_download_skips_duplicate_worker( + sqlite_session_factory: sessionmaker[Session], +) -> None: task = _task().model_copy(update={"celery_task_id": "celery-task-1"}) storage = FakeArchiveStorage({}) cache = FakeTaskCache(task) - preparer = _preparer(storage=storage, cache=cache, bundles=[]) + preparer = _preparer( + storage=storage, + cache=cache, + session_factory=sqlite_session_factory, + bundles=[], + ) nested_results: list[WorkflowRunArchiveDownloadTask | None] = [] preparer._get_task_bundles = MagicMock(return_value=[]) @@ -277,13 +298,15 @@ def test_prepare_workflow_run_archive_download_skips_duplicate_worker() -> None: preparer._build_zip_payload.assert_called_once() -def test_failed_worker_cannot_overwrite_ready_task() -> None: +def test_failed_worker_cannot_overwrite_ready_task( + sqlite_session_factory: sessionmaker[Session], +) -> None: processing_task = _task().model_copy( update={"status": WorkflowRunArchiveDownloadStatus.PROCESSING, "celery_task_id": "celery-task-1"} ) ready_task = processing_task.model_copy(update={"status": WorkflowRunArchiveDownloadStatus.READY}) cache = FakeTaskCache(ready_task) - preparer = _preparer(storage=FakeArchiveStorage({}), cache=cache) + preparer = _preparer(storage=FakeArchiveStorage({}), cache=cache, session_factory=sqlite_session_factory) result = preparer._mark_failed(processing_task, error="late failure") diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py b/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py index 6fdf9d8358e..8817c71744f 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py @@ -1,10 +1,9 @@ import datetime from contextlib import nullcontext -from types import SimpleNamespace from typing import cast -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.archive_download_task_cache import ( @@ -63,18 +62,16 @@ def _bundle( row_count: int = 9, archived_at: datetime.datetime | None = None, ) -> WorkflowRunArchiveBundle: - return cast( - WorkflowRunArchiveBundle, - SimpleNamespace( - year=year, - month=month, - shard=shard, - bundle_id=bundle_id, - workflow_run_count=workflow_run_count, - row_count=row_count, - archive_bytes=archive_bytes, - archived_at=archived_at or datetime.datetime(2026, 6, 25, 8, 0), - ), + return WorkflowRunArchiveBundle( + tenant_id="tenant-1", + year=year, + month=month, + shard=shard, + bundle_id=bundle_id, + workflow_run_count=workflow_run_count, + row_count=row_count, + archive_bytes=archive_bytes, + archived_at=archived_at or datetime.datetime(2026, 6, 25, 8, 0), ) @@ -89,10 +86,9 @@ def _fake_dispatcher(dispatched_tasks: list[WorkflowRunArchiveDownloadTask]) -> return dispatch -def test_list_workflow_run_archives_aggregates_month_rows() -> None: +def test_list_workflow_run_archives_aggregates_month_rows(sqlite_session: Session) -> None: latest = datetime.datetime(2026, 6, 25, 8, 0) previous = datetime.datetime(2026, 6, 24, 8, 0) - session = MagicMock() march_download_id = build_archive_download_id( tenant_id="tenant-1", year=2025, @@ -117,40 +113,45 @@ def test_list_workflow_run_archives_aggregates_month_rows() -> None: } ) cache = FakeTaskCache(tasks_by_download_id={march_download_id: ready_task}) - session.scalars.return_value = [ - _bundle( - year=2025, - month=3, - shard="00-of-01", - bundle_id="bundle-a", - workflow_run_count=40, - row_count=360, - archive_bytes=1024, - archived_at=previous, - ), - _bundle( - year=2025, - month=3, - shard="00-of-01", - bundle_id="bundle-b", - workflow_run_count=60, - row_count=540, - archive_bytes=3072, - archived_at=latest, - ), - _bundle( - year=2025, - month=2, - shard="00-of-01", - bundle_id="bundle-c", - workflow_run_count=20, - row_count=180, - archive_bytes=1024, - archived_at=previous, - ), - ] + sqlite_session.add_all( + [ + _bundle( + year=2025, + month=3, + shard="00-of-01", + bundle_id="bundle-a", + workflow_run_count=40, + row_count=360, + archive_bytes=1024, + archived_at=previous, + ), + _bundle( + year=2025, + month=3, + shard="00-of-01", + bundle_id="bundle-b", + workflow_run_count=60, + row_count=540, + archive_bytes=3072, + archived_at=latest, + ), + _bundle( + year=2025, + month=2, + shard="00-of-01", + bundle_id="bundle-c", + workflow_run_count=20, + row_count=180, + archive_bytes=1024, + archived_at=previous, + ), + ] + ) + sqlite_session.flush() - result = list_workflow_run_archives(session, "tenant-1", cache=cast(WorkflowRunArchiveDownloadTaskCache, cache)) + result = list_workflow_run_archives( + sqlite_session, "tenant-1", cache=cast(WorkflowRunArchiveDownloadTaskCache, cache) + ) assert result.summary.archived_month_count == 2 assert result.summary.workflow_run_count == 120 @@ -165,17 +166,19 @@ def test_list_workflow_run_archives_aggregates_month_rows() -> None: assert result.months[1].download_task is None -def test_create_workflow_run_archive_download_task_creates_stable_pending_task() -> None: - session = MagicMock() - session.scalars.return_value = [ - _bundle(shard="01-of-02", bundle_id="bundle-b", archive_bytes=2048), - _bundle(shard="00-of-02", bundle_id="bundle-a", archive_bytes=1024), - ] +def test_create_workflow_run_archive_download_task_creates_stable_pending_task(sqlite_session: Session) -> None: + sqlite_session.add_all( + [ + _bundle(shard="01-of-02", bundle_id="bundle-b", archive_bytes=2048), + _bundle(shard="00-of-02", bundle_id="bundle-a", archive_bytes=1024), + ] + ) + sqlite_session.flush() cache = FakeTaskCache() dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = [] task = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -188,22 +191,24 @@ def test_create_workflow_run_archive_download_task_creates_stable_pending_task() tenant_id="tenant-1", year=2025, month=3, - bundle_refs=[("01-of-02", "bundle-b"), ("00-of-02", "bundle-a")], + bundle_refs=[("00-of-02", "bundle-a"), ("01-of-02", "bundle-b")], ) assert task.requested_by == "account-1" - assert task.bundle_ids == ["bundle-b", "bundle-a"] + assert task.bundle_ids == ["bundle-a", "bundle-b"] assert [(ref.shard, ref.bundle_id) for ref in task.bundle_refs] == [ - ("01-of-02", "bundle-b"), ("00-of-02", "bundle-a"), + ("01-of-02", "bundle-b"), ] assert task.archive_bytes == 3072 assert cache.saved_task == dispatched_tasks[0] assert task.celery_task_id is not None -def test_create_workflow_run_archive_download_task_returns_existing_task_when_cache_key_exists() -> None: - session = MagicMock() - session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)] +def test_create_workflow_run_archive_download_task_returns_existing_task_when_cache_key_exists( + sqlite_session: Session, +) -> None: + sqlite_session.add(_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)) + sqlite_session.flush() existing_task = build_pending_archive_download_task( tenant_id="tenant-1", requested_by="account-1", @@ -217,7 +222,7 @@ def test_create_workflow_run_archive_download_task_returns_existing_task_when_ca dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = [] task = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -231,9 +236,11 @@ def test_create_workflow_run_archive_download_task_returns_existing_task_when_ca assert dispatched_tasks == [] -def test_create_workflow_run_archive_download_task_retries_failed_cached_task() -> None: - session = MagicMock() - session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)] +def test_create_workflow_run_archive_download_task_retries_failed_cached_task( + sqlite_session: Session, +) -> None: + sqlite_session.add(_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)) + sqlite_session.flush() existing_task = build_pending_archive_download_task( tenant_id="tenant-1", requested_by="account-1", @@ -253,7 +260,7 @@ def test_create_workflow_run_archive_download_task_retries_failed_cached_task() dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = [] task = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -269,9 +276,11 @@ def test_create_workflow_run_archive_download_task_retries_failed_cached_task() @pytest.mark.parametrize("retry_failed_task", [False, True]) -def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_failed_task: bool) -> None: - session = MagicMock() - session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)] +def test_create_workflow_run_archive_download_task_claims_dispatch_once( + retry_failed_task: bool, sqlite_session: Session +) -> None: + sqlite_session.add(_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)) + sqlite_session.flush() download_id = build_archive_download_id( tenant_id="tenant-1", year=2025, @@ -301,7 +310,7 @@ def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_fa dispatched_tasks.append(task) concurrent_results.append( create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-2", year=2025, @@ -313,7 +322,7 @@ def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_fa return task result = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -326,13 +335,10 @@ def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_fa assert concurrent_results == [result] -def test_create_workflow_run_archive_download_task_rejects_missing_month() -> None: - session = MagicMock() - session.scalars.return_value = [] - +def test_create_workflow_run_archive_download_task_rejects_missing_month(sqlite_session: Session) -> None: with pytest.raises(WorkflowRunArchiveNotFoundError): create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py b/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py index 5fefb149f3a..82b616f3b5e 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py @@ -1,10 +1,14 @@ +import datetime import json from types import SimpleNamespace from typing import Any, cast from unittest.mock import MagicMock, call, patch import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker +from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.bundle_archive_maintenance import ( ARCHIVED_TABLES, ArchiveBundleCatalogEntry, @@ -35,14 +39,16 @@ def _table_records( return records -def _catalog_entry(*, catalog_id: str = CATALOG_ID, shard: str = "00-of-01") -> ArchiveBundleCatalogEntry: +def _catalog_entry( + *, catalog_id: str = CATALOG_ID, shard: str = "00-of-01", bundle_id: str = BUNDLE_ID +) -> ArchiveBundleCatalogEntry: return ArchiveBundleCatalogEntry( catalog_id=catalog_id, tenant_id=TENANT_ID, year=2025, month=3, shard=shard, - bundle_id=BUNDLE_ID, + bundle_id=bundle_id, workflow_run_count=0, row_count=0, archive_bytes=0, @@ -93,10 +99,25 @@ def _manifest( ).encode() -def _session_factory(session: MagicMock) -> MagicMock: - factory = MagicMock() - factory.return_value.__enter__.return_value = session - return factory +def _bundle_model(entry: ArchiveBundleCatalogEntry) -> WorkflowRunArchiveBundle: + bundle = WorkflowRunArchiveBundle( + tenant_id=entry.tenant_id, + year=entry.year, + month=entry.month, + shard=entry.shard, + bundle_id=entry.bundle_id, + workflow_run_count=entry.workflow_run_count, + row_count=entry.row_count, + archive_bytes=entry.archive_bytes, + archived_at=datetime.datetime(2026, 1, 1), + ) + bundle.id = entry.catalog_id + return bundle + + +def _persist_catalog(session: Session, entry: ArchiveBundleCatalogEntry) -> None: + session.add(_bundle_model(entry)) + session.flush() def _bundle_reference( @@ -141,26 +162,17 @@ def _sample_archive_records() -> dict[str, list[dict[str, Any]]]: ) -def test_catalog_discovery_is_ordered_and_limited_before_storage_io() -> None: - entry = _catalog_entry() - bundle = SimpleNamespace( - id=entry.catalog_id, - tenant_id=entry.tenant_id, - year=entry.year, - month=entry.month, - shard=entry.shard, - bundle_id=entry.bundle_id, - workflow_run_count=entry.workflow_run_count, - row_count=entry.row_count, - archive_bytes=entry.archive_bytes, - ) - session = MagicMock() - session.get.return_value = bundle - session.scalars.return_value = [bundle] +def test_catalog_discovery_is_ordered_and_limited_before_storage_io( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + cursor = _catalog_entry() + entry = _catalog_entry(catalog_id="019f63b7-5ca4-7681-9ce0-800283608f40", bundle_id="bundle-b") + sqlite_session.add_all([_bundle_model(cursor), _bundle_model(entry)]) + sqlite_session.commit() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) entries = maintenance._list_catalog_entries( @@ -171,36 +183,29 @@ def test_catalog_discovery_is_ordered_and_limited_before_storage_io() -> None: limit=2, ) - statement = session.scalars.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.year" in rendered - assert "workflow_run_archive_bundles.month" in rendered - assert "workflow_run_archive_bundles.id >" in rendered - assert "ORDER BY workflow_run_archive_bundles.id ASC" in rendered - assert "LIMIT" in rendered assert entries == [entry] storage.list_objects.assert_not_called() -def test_catalog_discovery_filters_and_validates_the_requested_shard() -> None: - entry = _catalog_entry(shard="03-of-16") - bundle = SimpleNamespace( - id=entry.catalog_id, - tenant_id=entry.tenant_id, - year=entry.year, - month=entry.month, - shard=entry.shard, - bundle_id=entry.bundle_id, - workflow_run_count=entry.workflow_run_count, - row_count=entry.row_count, - archive_bytes=entry.archive_bytes, +def test_catalog_discovery_filters_and_validates_the_requested_shard( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + cursor = _catalog_entry(shard="03-of-16") + entry = _catalog_entry( + catalog_id="019f63b7-5ca4-7681-9ce0-800283608f40", + shard="03-of-16", + bundle_id="bundle-b", ) - session = MagicMock() - session.get.return_value = bundle - session.scalars.return_value = [bundle] + wrong_cursor = _catalog_entry( + catalog_id="019f63b7-5ca4-7681-9ce0-800283608f41", + shard="04-of-16", + bundle_id="bundle-c", + ) + sqlite_session.add_all([_bundle_model(cursor), _bundle_model(entry), _bundle_model(wrong_cursor)]) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) entries = maintenance._list_catalog_entries( @@ -212,34 +217,27 @@ def test_catalog_discovery_filters_and_validates_the_requested_shard() -> None: shard="03-of-16", ) - statement = session.scalars.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.shard =" in rendered assert entries == [entry] - session.get.return_value = SimpleNamespace( - year=2025, - month=3, - tenant_id=TENANT_ID, - shard="04-of-16", - ) with pytest.raises(ValueError, match="requested archive shard"): maintenance._list_catalog_entries( tenant_ids=None, target_year=2025, target_month=3, - after_catalog_id=CATALOG_ID, + after_catalog_id=wrong_cursor.catalog_id, limit=2, shard="03-of-16", ) -def test_catalog_shard_preflight_rejects_mixed_layout_before_delete() -> None: - session = MagicMock() - session.scalars.return_value = ["00-of-01"] +def test_catalog_shard_preflight_rejects_mixed_layout_before_delete( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + sqlite_session.add(_bundle_model(_catalog_entry(shard="00-of-01"))) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) with pytest.raises(ValueError, match=r"unexpected shards.*00-of-01"): @@ -249,19 +247,15 @@ def test_catalog_shard_preflight_rejects_mixed_layout_before_delete() -> None: shard_total=16, ) - statement = session.scalars.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.year" in rendered - assert "workflow_run_archive_bundles.month" in rendered - assert "workflow_run_archive_bundles.shard NOT IN" in rendered - -def test_catalog_shard_preflight_accepts_an_expected_subset() -> None: - session = MagicMock() - session.scalars.return_value = [] +def test_catalog_shard_preflight_accepts_an_expected_subset( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + sqlite_session.add(_bundle_model(_catalog_entry(shard="03-of-16"))) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) maintenance.validate_catalog_shards( @@ -271,12 +265,17 @@ def test_catalog_shard_preflight_accepts_an_expected_subset() -> None: ) -def test_catalog_shard_preflight_uses_requested_tenant_scope() -> None: - session = MagicMock() - session.scalars.return_value = [] +def test_catalog_shard_preflight_uses_requested_tenant_scope( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + other_entry = _catalog_entry(shard="00-of-01") + other_bundle = _bundle_model(other_entry) + other_bundle.tenant_id = "other-tenant" + sqlite_session.add(other_bundle) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) maintenance.validate_catalog_shards( @@ -286,9 +285,6 @@ def test_catalog_shard_preflight_uses_requested_tenant_scope() -> None: tenant_ids=[TENANT_ID], ) - statement = session.scalars.call_args.args[0] - assert "workflow_run_archive_bundles.tenant_id IN" in str(statement) - @pytest.mark.parametrize( ("cursor_bundle", "tenant_ids", "error_message"), @@ -310,12 +306,20 @@ def test_catalog_discovery_rejects_cursor_outside_requested_scope( cursor_bundle: SimpleNamespace | None, tenant_ids: list[str] | None, error_message: str, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: - session = MagicMock() - session.get.return_value = cursor_bundle + if cursor_bundle is not None: + entry = _catalog_entry() + stored = _bundle_model(entry) + stored.year = cursor_bundle.year + stored.month = cursor_bundle.month + stored.tenant_id = cursor_bundle.tenant_id + sqlite_session.add(stored) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) with pytest.raises(ValueError, match=error_message): @@ -327,43 +331,38 @@ def test_catalog_discovery_rejects_cursor_outside_requested_scope( limit=1, ) - session.scalars.assert_not_called() - -def test_catalog_manifest_identity_mismatch_fails_closed() -> None: +def test_catalog_manifest_identity_mismatch_fails_closed( + unbound_session_factory: sessionmaker[Session], +) -> None: entry = _catalog_entry() storage = MagicMock() storage.get_object.return_value = _manifest(entry, bundle_id="other-bundle") maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(MagicMock())), + session_factory=unbound_session_factory, ) with pytest.raises(ValueError, match="identity does not match catalog"): maintenance._build_bundle_reference(cast(MagicMock, storage), entry) -def test_bundle_maintenance_locks_the_existing_catalog_row() -> None: +def test_bundle_maintenance_locks_the_existing_catalog_row(sqlite_session: Session) -> None: entry = _catalog_entry() - session = MagicMock() - session.scalar.return_value = entry.catalog_id + sqlite_session.add(_bundle_model(entry)) + sqlite_session.flush() - WorkflowRunBundleArchiveMaintenance._lock_catalog_entry(session, entry) - - statement = session.scalar.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.id" in rendered - assert "workflow_run_archive_bundles.tenant_id" in rendered - assert "FOR UPDATE" in rendered + WorkflowRunBundleArchiveMaintenance._lock_catalog_entry(sqlite_session, entry) -def test_failure_and_dry_run_do_not_return_a_persistable_cursor() -> None: +def test_failure_and_dry_run_do_not_return_a_persistable_cursor( + sqlite_session_factory: sessionmaker[Session], +) -> None: entry = _catalog_entry() - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) with ( @@ -385,7 +384,7 @@ def test_failure_and_dry_run_do_not_return_a_persistable_cursor() -> None: dry_run = WorkflowRunBundleArchiveMaintenance( dry_run=True, storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) bundle_ref = BundleReference( catalog=entry, @@ -468,12 +467,13 @@ def test_live_archive_subset_rejects_content_mismatch() -> None: ) -def test_live_bundle_scope_includes_archived_ids_and_indirect_children() -> None: +def test_live_bundle_scope_includes_archived_ids_and_indirect_children( + unbound_session: Session, unbound_session_factory: sessionmaker[Session] +) -> None: archive_records = _sample_archive_records() manifest = _bundle_reference(_catalog_entry(), table_records=archive_records).manifest - session = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=unbound_session_factory, ) def select_live_parent_ids(_session, model, _run_ids): @@ -488,7 +488,7 @@ def test_live_bundle_scope_includes_archived_ids_and_indirect_children() -> None patch.object(maintenance, "_load_records_by_column", return_value=[]) as load_records, ): maintenance._load_live_bundle_records( - session, + unbound_session, manifest, archive_records, lock=True, @@ -511,8 +511,11 @@ def test_live_bundle_scope_includes_archived_ids_and_indirect_children() -> None assert ("workflow_app_logs", "id", {"app-log-1"}, True) in queries -def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subset() -> None: +def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subset( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) archive_records = _sample_archive_records() bundle_ref = _bundle_reference(entry, table_records=archive_records) partial_records = _table_records( @@ -520,11 +523,13 @@ def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subse workflow_pause_reasons=archive_records["workflow_pause_reasons"], ) expected_deleted_counts = {table_name: len(partial_records[table_name]) for table_name in ARCHIVED_TABLES} - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + transaction_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit")) + event.listen(sqlite_session, "after_rollback", lambda _session: transaction_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=False), @@ -548,24 +553,28 @@ def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subse patch.object(maintenance, "_mark_deleted") as mark_deleted, patch.object(maintenance, "_delete_marker"), ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert result.success - delete_bundle_rows.assert_called_once_with(session, partial_records) - session.commit.assert_called_once_with() - session.rollback.assert_not_called() + delete_bundle_rows.assert_called_once_with(sqlite_session, partial_records) + assert transaction_events == ["commit"] mark_deleted.assert_called_once_with(storage, bundle_ref.object_prefix) -def test_delete_bundle_marks_an_already_absent_source_without_deleting_rows() -> None: +def test_delete_bundle_marks_an_already_absent_source_without_deleting_rows( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) archive_records = _sample_archive_records() bundle_ref = _bundle_reference(entry, table_records=archive_records) - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + transaction_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit")) + event.listen(sqlite_session, "after_rollback", lambda _session: transaction_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=False), @@ -580,28 +589,31 @@ def test_delete_bundle_marks_an_already_absent_source_without_deleting_rows() -> patch.object(maintenance, "_mark_deleted") as mark_deleted, patch.object(maintenance, "_delete_marker") as delete_marker, ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert result.success delete_bundle_rows.assert_not_called() - session.commit.assert_not_called() - session.rollback.assert_not_called() + assert transaction_events == [] mark_deleted.assert_called_once_with(storage, bundle_ref.object_prefix) assert delete_marker.call_count == 2 -def test_delete_bundle_with_deleted_marker_rejects_remaining_orphan_children() -> None: +def test_delete_bundle_with_deleted_marker_rejects_remaining_orphan_children( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) archive_records = _sample_archive_records() bundle_ref = _bundle_reference(entry, table_records=archive_records) live_records = _table_records( workflow_node_execution_offload=archive_records["workflow_node_execution_offload"], ) - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=False), @@ -614,42 +626,45 @@ def test_delete_bundle_with_deleted_marker_rejects_remaining_orphan_children() - patch.object(maintenance, "_load_live_bundle_records", return_value=live_records), patch.object(maintenance, "_delete_bundle_rows") as delete_bundle_rows, ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert not result.success assert "Live rows exist for bundle with deleted marker" in result.error delete_bundle_rows.assert_not_called() - session.commit.assert_not_called() - session.rollback.assert_called_once_with() + assert rollback_events == ["rollback"] -def test_delete_bundle_rejects_an_in_progress_restore() -> None: +def test_delete_bundle_rejects_an_in_progress_restore( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) bundle_ref = _bundle_reference(entry) - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=True), patch.object(maintenance, "_validate_archive_object") as validate_archive, ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert not result.success assert "reconcile restore before delete" in result.error validate_archive.assert_not_called() - session.commit.assert_not_called() - session.rollback.assert_called_once_with() + assert rollback_events == ["rollback"] -def test_delete_bundle_rows_use_only_verified_primary_keys() -> None: +def test_delete_bundle_rows_use_only_verified_primary_keys( + unbound_session: Session, unbound_session_factory: sessionmaker[Session] +) -> None: live_records = _sample_archive_records() - session = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=unbound_session_factory, ) with patch.object( @@ -657,7 +672,7 @@ def test_delete_bundle_rows_use_only_verified_primary_keys() -> None: "_delete_by_column", side_effect=lambda _session, _model, _column, values: len(values), ) as delete_by_column: - deleted_counts = maintenance._delete_bundle_rows(session, live_records) + deleted_counts = maintenance._delete_bundle_rows(unbound_session, live_records) expected_table_order = [ "workflow_pause_reasons", @@ -673,33 +688,39 @@ def test_delete_bundle_rows_use_only_verified_primary_keys() -> None: assert deleted_counts == {table_name: len(live_records[table_name]) for table_name in ARCHIVED_TABLES} -def test_restore_does_not_skip_an_interrupted_delete_without_deleted_marker() -> None: +def test_restore_does_not_skip_an_interrupted_delete_without_deleted_marker( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() - session = MagicMock() + _persist_catalog(sqlite_session, entry) maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + commit_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: commit_events.append("commit")) with ( patch.object(maintenance, "_is_deleted", return_value=False), patch.object(maintenance, "_is_delete_started", return_value=True), patch.object(maintenance, "_validate_live_counts") as validate_live_counts, ): - result = maintenance._restore_bundle(session, MagicMock(), _bundle_reference(entry)) + result = maintenance._restore_bundle(sqlite_session, MagicMock(), _bundle_reference(entry)) assert not result.success assert "reconcile delete first" in result.error validate_live_counts.assert_not_called() - session.commit.assert_not_called() + assert commit_events == [] -def test_restore_does_not_skip_missing_source_rows_without_deleted_marker() -> None: +def test_restore_does_not_skip_missing_source_rows_without_deleted_marker( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() - session = MagicMock() + _persist_catalog(sqlite_session, entry) maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) bundle_ref = _bundle_reference(entry) @@ -712,22 +733,25 @@ def test_restore_does_not_skip_missing_source_rows_without_deleted_marker() -> N side_effect=ValueError("source rows are missing"), ) as validate_live_counts, ): - result = maintenance._restore_bundle(session, MagicMock(), bundle_ref) + result = maintenance._restore_bundle(sqlite_session, MagicMock(), bundle_ref) assert not result.success assert "source rows are missing" in result.error - validate_live_counts.assert_called_once_with(session, bundle_ref.manifest, expected_present=True) - session.commit.assert_not_called() + validate_live_counts.assert_called_once_with(sqlite_session, bundle_ref.manifest, expected_present=True) -def test_restore_reconciles_a_started_marker_after_the_source_commit() -> None: +def test_restore_reconciles_a_started_marker_after_the_source_commit( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() - session = MagicMock() + _persist_catalog(sqlite_session, entry) storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + commit_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: commit_events.append("commit")) bundle_ref = _bundle_reference(entry) with ( @@ -737,12 +761,12 @@ def test_restore_reconciles_a_started_marker_after_the_source_commit() -> None: patch.object(maintenance, "_validate_live_counts") as validate_live_counts, patch.object(maintenance, "_mark_restored") as mark_restored, ): - result = maintenance._restore_bundle(session, storage, bundle_ref) + result = maintenance._restore_bundle(sqlite_session, storage, bundle_ref) assert result.success - validate_live_counts.assert_called_once_with(session, bundle_ref.manifest, expected_present=True) + validate_live_counts.assert_called_once_with(sqlite_session, bundle_ref.manifest, expected_present=True) mark_restored.assert_called_once_with(storage, bundle_ref.object_prefix) - session.commit.assert_not_called() + assert commit_events == [] def test_mark_restored_clears_stale_delete_marker_before_releasing_restore_fence() -> None: diff --git a/api/tests/unit_tests/tasks/test_batch_clean_document_task.py b/api/tests/unit_tests/tasks/test_batch_clean_document_task.py index 6386c72188e..59fab89ab9c 100644 --- a/api/tests/unit_tests/tasks/test_batch_clean_document_task.py +++ b/api/tests/unit_tests/tasks/test_batch_clean_document_task.py @@ -1,46 +1,73 @@ -from unittest.mock import MagicMock, patch +import uuid +from unittest.mock import patch +import pytest +from sqlalchemy.orm import Session + +import tasks.batch_clean_document_task as task_module +from models.dataset import Dataset, DocumentSegment +from models.enums import DataSourceType from tasks.batch_clean_document_task import batch_clean_document_task -def _setup_cleanup_dependencies(): - session = MagicMock() - segment = MagicMock(id="segment-1", index_node_id="node-1", content="content") - dataset = MagicMock(id="dataset-1", tenant_id="tenant-1") - session.scalars.return_value.all.return_value = [segment] - session.scalar.return_value = dataset - - context_manager = MagicMock() - context_manager.__enter__.return_value = session - context_manager.__exit__.return_value = None - return session, context_manager +@pytest.fixture +def cleanup_rows(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> tuple[str, str, str]: + tenant_id = str(uuid.uuid4()) + dataset_id = str(uuid.uuid4()) + document_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Batch cleanup dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + ) + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_id=document_id, + position=1, + content="content", + word_count=1, + tokens=1, + created_by=created_by, + index_node_id="node-1", + ) + sqlite_session.add_all([dataset, segment]) + sqlite_session.commit() + engine = sqlite_session.get_bind() + monkeypatch.setattr( + task_module.session_factory, + "create_session", + lambda: Session(engine, expire_on_commit=False), + ) + return dataset_id, document_id, tenant_id -def test_successful_vector_cleanup_schedules_billing_refresh(): - _, context_manager = _setup_cleanup_dependencies() +def test_successful_vector_cleanup_schedules_billing_refresh(cleanup_rows: tuple[str, str, str]): + dataset_id, document_id, tenant_id = cleanup_rows with ( - patch("tasks.batch_clean_document_task.session_factory.create_session", return_value=context_manager), patch("tasks.batch_clean_document_task.get_image_upload_file_ids", return_value=[]), patch("tasks.batch_clean_document_task.IndexProcessorFactory") as processor_factory, patch("tasks.batch_clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh, ): batch_clean_document_task( - document_ids=["document-1"], - dataset_id="dataset-1", + document_ids=[document_id], + dataset_id=dataset_id, doc_form="paragraph", file_ids=[], ) processor_factory.return_value.init_index_processor.return_value.clean.assert_called_once() - schedule_refresh.assert_called_once_with("tenant-1") + schedule_refresh.assert_called_once_with(tenant_id) -def test_failed_vector_cleanup_does_not_schedule_billing_refresh(): - _, context_manager = _setup_cleanup_dependencies() +def test_failed_vector_cleanup_does_not_schedule_billing_refresh(cleanup_rows: tuple[str, str, str]): + dataset_id, document_id, _tenant_id = cleanup_rows with ( - patch("tasks.batch_clean_document_task.session_factory.create_session", return_value=context_manager), patch("tasks.batch_clean_document_task.get_image_upload_file_ids", return_value=[]), patch("tasks.batch_clean_document_task.IndexProcessorFactory") as processor_factory, patch("tasks.batch_clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh, @@ -49,8 +76,8 @@ def test_failed_vector_cleanup_does_not_schedule_billing_refresh(): "vector cleanup failed" ) batch_clean_document_task( - document_ids=["document-1"], - dataset_id="dataset-1", + document_ids=[document_id], + dataset_id=dataset_id, doc_form="paragraph", file_ids=[], ) diff --git a/api/tests/unit_tests/tasks/test_community_telemetry_task.py b/api/tests/unit_tests/tasks/test_community_telemetry_task.py index c90ecd67109..00f97a5f31b 100644 --- a/api/tests/unit_tests/tasks/test_community_telemetry_task.py +++ b/api/tests/unit_tests/tasks/test_community_telemetry_task.py @@ -1,41 +1,50 @@ import logging from types import SimpleNamespace -from unittest.mock import MagicMock, Mock +from unittest.mock import Mock import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from tasks import community_telemetry_task -def _configure_task_session(monkeypatch: pytest.MonkeyPatch) -> Mock: - session = Mock() - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session - monkeypatch.setattr(community_telemetry_task, "db", SimpleNamespace(engine=object())) - monkeypatch.setattr(community_telemetry_task, "sessionmaker", Mock(return_value=session_factory)) - return session +def _bind_task_to_sqlite(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None: + """Bind the task-local sessionmaker to the isolated SQLite database.""" + monkeypatch.setattr(community_telemetry_task, "db", SimpleNamespace(engine=sqlite_engine)) -def test_send_community_telemetry_heartbeat_reports_with_a_database_session(monkeypatch: pytest.MonkeyPatch): - session = _configure_task_session(monkeypatch) - report_heartbeat = Mock() +def test_send_community_telemetry_heartbeat_reports_with_a_database_session( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine +) -> None: + _bind_task_to_sqlite(monkeypatch, sqlite_engine) + received_sessions: list[Session] = [] + + def report_heartbeat(*, session: Session) -> None: + received_sessions.append(session) + assert session.get_bind() is sqlite_engine + monkeypatch.setattr(community_telemetry_task.CommunityTelemetryService, "report_heartbeat", report_heartbeat) community_telemetry_task.send_community_telemetry_heartbeat.run() - report_heartbeat.assert_called_once_with(session=session) + assert len(received_sessions) == 1 + assert isinstance(received_sessions[0], Session) def test_send_community_telemetry_heartbeat_swallows_report_errors( - monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture -): - _configure_task_session(monkeypatch) + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + caplog: pytest.LogCaptureFixture, +) -> None: + _bind_task_to_sqlite(monkeypatch, sqlite_engine) monkeypatch.setattr( community_telemetry_task.CommunityTelemetryService, "report_heartbeat", Mock(side_effect=RuntimeError("telemetry unavailable")), ) caplog.set_level(logging.DEBUG, logger=community_telemetry_task.logger.name) + community_telemetry_task.send_community_telemetry_heartbeat.run() assert "Failed to process community telemetry heartbeat" in caplog.text diff --git a/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py b/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py index af7ecb8c570..fb64f5b5ba4 100644 --- a/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py +++ b/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py @@ -1,43 +1,97 @@ -from contextlib import nullcontext -from types import SimpleNamespace +import uuid +from collections.abc import Iterator +from contextlib import contextmanager from unittest.mock import MagicMock, patch +import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker + from core.rag.index_processor.constant.index_type import IndexStructureType -from models.enums import IndexingStatus, SegmentStatus +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus from tasks.enable_segment_to_index_task import enable_segment_to_index_task from tasks.enable_segments_to_index_task import enable_segments_to_index_task -def test_enable_segment_commits_index_rows_after_loading() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, +@pytest.fixture +def indexed_segment(sqlite_session: Session) -> tuple[Dataset, Document, DocumentSegment]: + """Persist the complete owner chain consumed by segment indexing tasks.""" + tenant_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Indexing dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + is_multimodal=False, + ) + document = Document( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="document.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, indexing_status=IndexingStatus.COMPLETED, doc_form=IndexStructureType.PARAGRAPH_INDEX, ) - segment = SimpleNamespace( - id="segment-1", - status=SegmentStatus.COMPLETED, + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset.id, + document_id=document.id, + position=1, content="content", + word_count=1, + tokens=1, + created_by=created_by, index_node_id="node-1", index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - get_dataset=MagicMock(return_value=dataset), - get_document=MagicMock(return_value=document), + status=SegmentStatus.COMPLETED, ) - session = MagicMock() - session.scalar.return_value = segment + sqlite_session.add_all([dataset, document, segment]) + sqlite_session.commit() + return dataset, document, segment + + +@contextmanager +def _record_transaction_events( + sqlite_session_factory: sessionmaker[Session], phase_events: list[str] +) -> Iterator[None]: + """Record real transaction boundaries from task-owned SQLite sessions.""" + session_type = sqlite_session_factory.class_ + + def after_commit(_session: Session) -> None: + phase_events.append("commit") + + def after_rollback(_session: Session) -> None: + phase_events.append("rollback") + + event.listen(session_type, "after_commit", after_commit) + event.listen(session_type, "after_rollback", after_rollback) + try: + yield + finally: + event.remove(session_type, "after_commit", after_commit) + event.remove(session_type, "after_rollback", after_rollback) + + +def test_enable_segment_commits_index_rows_after_loading( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, _document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") index_processor = MagicMock() index_processor.load.side_effect = lambda *_args, **_kwargs: phase_events.append("load") enable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch("tasks.enable_segment_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segment_to_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.enable_summaries_for_segments", @@ -49,53 +103,28 @@ def test_enable_segment_commits_index_rows_after_loading() -> None: enable_segment_to_index_task.run(segment.id) assert phase_events == ["load", "commit", "summary"] + enable_summaries.assert_called_once() + assert enable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert enable_summaries.call_args.kwargs["segment_ids"] == [segment.id] -def test_enable_segment_rolls_back_before_error_compensation() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status=IndexingStatus.COMPLETED, - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) - segment = SimpleNamespace( - id="segment-1", - status=SegmentStatus.COMPLETED, - content="content", - index_node_id="node-1", - index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - enabled=True, - disabled_at=None, - error=None, - get_dataset=MagicMock(return_value=dataset), - get_document=MagicMock(return_value=document), - ) +def test_enable_segment_rolls_back_before_error_compensation( + sqlite_session: Session, + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + _dataset, _document, segment = indexed_segment phase_events: list[str] = [] - session = MagicMock() - session.scalar.return_value = segment - session.rollback.side_effect = lambda: phase_events.append("rollback") - - def commit() -> None: - assert segment.enabled is False - assert segment.status == SegmentStatus.ERROR - assert segment.error == "load failed" - phase_events.append("commit") - - session.commit.side_effect = commit index_processor = MagicMock() - def fail_load(*_args, **_kwargs) -> None: + def fail_load(*_args: object, **_kwargs: object) -> None: phase_events.append("load") raise RuntimeError("load failed") index_processor.load.side_effect = fail_load with ( - patch("tasks.enable_segment_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segment_to_index_task.IndexProcessorFactory") as processor_factory, patch("services.summary_index_service.SummaryIndexService.enable_summaries_for_segments") as enable_summaries, patch("tasks.enable_segment_to_index_task.redis_client.delete"), @@ -103,38 +132,29 @@ def test_enable_segment_rolls_back_before_error_compensation() -> None: processor_factory.return_value.init_index_processor.return_value = index_processor enable_segment_to_index_task.run(segment.id) + sqlite_session.expire_all() + persisted_segment = sqlite_session.get(DocumentSegment, segment.id) + assert persisted_segment is not None + assert persisted_segment.enabled is False + assert persisted_segment.status == SegmentStatus.ERROR + assert persisted_segment.error == "load failed" + assert persisted_segment.disabled_at is not None assert phase_events == ["load", "rollback", "commit"] enable_summaries.assert_not_called() -def test_enable_segments_commits_index_rows_after_loading() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) - segment = SimpleNamespace( - id="segment-1", - content="content", - index_node_id="node-1", - index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - ) - session = MagicMock() - session.scalar.side_effect = [dataset, document] - session.scalars.return_value.all.return_value = [segment] +def test_enable_segments_commits_index_rows_after_loading( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") index_processor = MagicMock() index_processor.load.side_effect = lambda *_args, **_kwargs: phase_events.append("load") enable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch("tasks.enable_segments_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segments_to_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.enable_summaries_for_segments", @@ -146,42 +166,28 @@ def test_enable_segments_commits_index_rows_after_loading() -> None: enable_segments_to_index_task.run([segment.id], dataset.id, document.id) assert phase_events == ["load", "commit", "summary"] + enable_summaries.assert_called_once() + assert enable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert enable_summaries.call_args.kwargs["segment_ids"] == [segment.id] -def test_enable_segments_rolls_back_before_error_compensation() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) - segment = SimpleNamespace( - id="segment-1", - content="content", - index_node_id="node-1", - index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - ) +def test_enable_segments_rolls_back_before_error_compensation( + sqlite_session: Session, + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session = MagicMock() - session.scalar.side_effect = [dataset, document] - session.scalars.return_value.all.return_value = [segment] - session.rollback.side_effect = lambda: phase_events.append("rollback") - session.execute.side_effect = lambda *_args, **_kwargs: phase_events.append("compensate") - session.commit.side_effect = lambda: phase_events.append("commit") index_processor = MagicMock() - def fail_load(*_args, **_kwargs) -> None: + def fail_load(*_args: object, **_kwargs: object) -> None: phase_events.append("load") raise RuntimeError("load failed") index_processor.load.side_effect = fail_load with ( - patch("tasks.enable_segments_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segments_to_index_task.IndexProcessorFactory") as processor_factory, patch("services.summary_index_service.SummaryIndexService.enable_summaries_for_segments") as enable_summaries, patch("tasks.enable_segments_to_index_task.redis_client.delete"), @@ -189,5 +195,12 @@ def test_enable_segments_rolls_back_before_error_compensation() -> None: processor_factory.return_value.init_index_processor.return_value = index_processor enable_segments_to_index_task.run([segment.id], dataset.id, document.id) - assert phase_events == ["load", "rollback", "compensate", "commit"] + sqlite_session.expire_all() + persisted_segment = sqlite_session.get(DocumentSegment, segment.id) + assert persisted_segment is not None + assert persisted_segment.enabled is False + assert persisted_segment.status == SegmentStatus.ERROR + assert persisted_segment.error == "load failed" + assert persisted_segment.disabled_at is not None + assert phase_events == ["load", "rollback", "commit"] enable_summaries.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py b/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py index cd3afb4904b..e5e1961112d 100644 --- a/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py +++ b/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py @@ -1,28 +1,57 @@ from unittest.mock import MagicMock, patch +from uuid import uuid4 +from sqlalchemy.orm import Session + +from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType +from models import Account, Tenant, TenantAccountJoin +from models.account import TenantAccountRole +from models.dataset import Dataset, Document +from models.enums import DatasetRuntimeMode, DataSourceType, DocumentCreatedFrom, IndexingStatus from tasks.retry_document_indexing_task import retry_document_indexing_task -def test_retry_enforces_vector_space_admission() -> None: - session = MagicMock() - dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", runtime_mode="general") - user = MagicMock(id="user-1") - tenant = MagicMock(id="tenant-1") - document = MagicMock(id="document-1", dataset_id="dataset-1", doc_form="paragraph") - session.scalar.side_effect = [dataset, user, tenant, document] - empty_segments: list[MagicMock] = [] - session.scalars.return_value.all.return_value = empty_segments - - session_context = MagicMock() - session_context.__enter__.return_value = session +def test_retry_enforces_vector_space_admission(sqlite_session: Session) -> None: + tenant = Tenant(name="Retry tenant") + user = Account(name="Retry user", email=f"retry-{uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=user.id, + current=True, + role=TenantAccountRole.OWNER, + ) + dataset = Dataset( + id=str(uuid4()), + tenant_id=tenant.id, + name="Retry dataset", + created_by=user.id, + data_source_type=DataSourceType.UPLOAD_FILE, + indexing_technique=IndexTechniqueType.ECONOMY, + chunk_structure=IndexStructureType.PARAGRAPH_INDEX, + runtime_mode=DatasetRuntimeMode.GENERAL, + ) + document = Document( + id=str(uuid4()), + tenant_id=tenant.id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + data_source_info="{}", + batch="retry-batch", + name="Retry document", + created_from=DocumentCreatedFrom.WEB, + created_by=user.id, + indexing_status=IndexingStatus.COMPLETED, + enabled=True, + archived=False, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + sqlite_session.add_all([tenant, user, membership, dataset, document]) + sqlite_session.commit() features = MagicMock() features.billing.enabled = False with ( - patch( - "tasks.retry_document_indexing_task.session_factory.create_session", - return_value=session_context, - ), patch("tasks.retry_document_indexing_task.FeatureService.get_features", return_value=features), patch("tasks.retry_document_indexing_task.IndexProcessorFactory"), patch("tasks.retry_document_indexing_task.IndexingRunner") as indexing_runner, @@ -31,4 +60,6 @@ def test_retry_enforces_vector_space_admission() -> None: retry_document_indexing_task.run(dataset.id, [document.id], user.id) indexing_runner.assert_called_once_with(enforce_vector_space_admission=True) - indexing_runner.return_value.run.assert_called_once_with([document], session) + run_documents, run_session = indexing_runner.return_value.run.call_args.args + assert [item.id for item in run_documents] == [document.id] + assert isinstance(run_session, Session) diff --git a/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py b/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py index 3454035fe5f..8af7ac6a3b5 100644 --- a/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py +++ b/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py @@ -1,36 +1,94 @@ -from contextlib import nullcontext -from types import SimpleNamespace +import uuid +from collections.abc import Generator +from contextlib import contextmanager from unittest.mock import MagicMock, patch -from models.enums import SegmentStatus +import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker + +from core.rag.index_processor.constant.index_type import IndexStructureType +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus from tasks.delete_segment_from_index_task import delete_segment_from_index_task from tasks.disable_segment_from_index_task import disable_segment_from_index_task from tasks.disable_segments_from_index_task import disable_segments_from_index_task -def test_disable_segment_commits_index_cleanup() -> None: - dataset = SimpleNamespace(id="dataset-1") - document = SimpleNamespace(enabled=True, archived=False, indexing_status="completed", doc_form="text_model") - segment = SimpleNamespace( - id="segment-1", - status=SegmentStatus.COMPLETED, - index_node_id="node-1", - disabled_by="user-1", - get_dataset=MagicMock(return_value=dataset), - get_document=MagicMock(return_value=document), +@pytest.fixture +def indexed_segment(sqlite_session: Session) -> tuple[Dataset, Document, DocumentSegment]: + """Persist the complete owner chain consumed by segment cleanup tasks.""" + tenant_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Cleanup dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + is_multimodal=False, ) - session = MagicMock() - session.scalar.return_value = segment + document = Document( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="document.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + indexing_status=IndexingStatus.COMPLETED, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset.id, + document_id=document.id, + position=1, + content="content", + word_count=1, + tokens=1, + created_by=created_by, + index_node_id="node-1", + index_node_hash="hash-1", + disabled_by=created_by, + status=SegmentStatus.COMPLETED, + ) + sqlite_session.add_all([dataset, document, segment]) + sqlite_session.commit() + return dataset, document, segment + + +@contextmanager +def _record_transaction_events( + sqlite_session_factory: sessionmaker[Session], phase_events: list[str] +) -> Generator[None]: + """Record real commits made by the task-owned SQLite session.""" + session_type = sqlite_session_factory.class_ + + def after_commit(_session: Session) -> None: + phase_events.append("commit") + + event.listen(session_type, "after_commit", after_commit) + try: + yield + finally: + event.remove(session_type, "after_commit", after_commit) + + +def test_disable_segment_commits_index_cleanup( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, _document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") processor = MagicMock() processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") disable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch( - "tasks.disable_segment_from_index_task.session_factory.create_session", return_value=nullcontext(session) - ), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.disable_segment_from_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.disable_summaries_for_segments", @@ -42,31 +100,24 @@ def test_disable_segment_commits_index_cleanup() -> None: disable_segment_from_index_task.run(segment.id) assert phase_events == ["clean", "commit", "summary"] + disable_summaries.assert_called_once() + assert disable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert disable_summaries.call_args.kwargs["segment_ids"] == [segment.id] + assert disable_summaries.call_args.kwargs["disabled_by"] == segment.disabled_by -def test_disable_segments_commits_index_cleanup() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form="text_model", - ) - segment = SimpleNamespace(id="segment-1", index_node_id="node-1", disabled_by="user-1") - session = MagicMock() - session.scalar.side_effect = [dataset, document] - session.scalars.return_value.all.return_value = [segment] +def test_disable_segments_commits_index_cleanup( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") processor = MagicMock() processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") disable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch( - "tasks.disable_segments_from_index_task.session_factory.create_session", return_value=nullcontext(session) - ), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.disable_segments_from_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.disable_summaries_for_segments", @@ -78,29 +129,26 @@ def test_disable_segments_commits_index_cleanup() -> None: disable_segments_from_index_task.run([segment.id], dataset.id, document.id) assert phase_events == ["clean", "commit", "summary"] + disable_summaries.assert_called_once() + assert disable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert disable_summaries.call_args.kwargs["segment_ids"] == [segment.id] + assert disable_summaries.call_args.kwargs["disabled_by"] == segment.disabled_by -def test_delete_segment_commits_index_cleanup_without_attachments() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form="text_model", - ) - session = MagicMock() - session.scalar.side_effect = [dataset, document] +def test_delete_segment_commits_index_cleanup_without_attachments( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") processor = MagicMock() processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") with ( - patch("tasks.delete_segment_from_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.delete_segment_from_index_task.IndexProcessorFactory") as processor_factory, ): processor_factory.return_value.init_index_processor.return_value = processor - delete_segment_from_index_task.run(["node-1"], dataset.id, document.id, ["segment-1"]) + delete_segment_from_index_task.run(["node-1"], dataset.id, document.id, [segment.id]) assert phase_events == ["clean", "commit"] diff --git a/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py b/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py index cd5df4466e3..84f95d649c4 100644 --- a/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py +++ b/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py @@ -55,10 +55,6 @@ class TestDispatchTriggeredWorkflow: (``get_workflows``, ``reserve``, ``create_end_user_batch``, ...) to drive the path it targets. """ - session_cm = MagicMock() - session_cm.__enter__.return_value = MagicMock() - session_cm.__exit__.return_value = False - invoke_response = MagicMock() invoke_response.cancelled = False invoke_response.variables = {} @@ -105,11 +101,6 @@ class TestDispatchTriggeredWorkflow: "create_end_user_batch", return_value={}, ) as create_end_user_batch, - patch.object( - trigger_processing_tasks_module.session_factory, - "create_session", - return_value=session_cm, - ), patch.object( trigger_processing_tasks_module.QuotaService, "reserve", diff --git a/api/tests/unit_tests/tasks/test_workflow_execute_task.py b/api/tests/unit_tests/tasks/test_workflow_execute_task.py index f99d4a1d942..7c7f8f34b08 100644 --- a/api/tests/unit_tests/tasks/test_workflow_execute_task.py +++ b/api/tests/unit_tests/tasks/test_workflow_execute_task.py @@ -18,6 +18,7 @@ from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity from graphon.entities import WorkflowStartReason from graphon.enums import WorkflowExecutionStatus +from models.account import Account from models.base import TypeBase from models.enums import ConversationFromSource, CreatorUserRole, WorkflowRunTriggeredFrom from models.model import App, AppMode, Conversation, Message @@ -632,44 +633,50 @@ def test_app_runner_streaming_failure_publishes_started_then_failed_workflow_fin assert finished_payload["data"]["files"] == [] -def test_app_runner_resolves_account_without_switching_tenant(monkeypatch: pytest.MonkeyPatch): +def test_app_runner_resolves_account_without_switching_tenant( + sqlite_session_factory: sessionmaker[Session], +): + account_id = str(uuid.uuid4()) exec_params = AppExecutionParams( app_id="app-id", workflow_id="workflow-id", tenant_id="resource-tenant-id", app_mode=AppMode.WORKFLOW, - user={"TYPE": "account", "user_id": "user-id"}, + user={"TYPE": "account", "user_id": account_id}, args={"inputs": {}}, invoke_from=InvokeFrom.EXPLORE, streaming=True, workflow_run_id="workflow-run-id", ) - runner = _AppRunner(session_factory=MagicMock(), exec_params=exec_params) - account = MagicMock() - session = MagicMock() - session.get.return_value = account - monkeypatch.setattr(runner, "_session", lambda: nullcontext(session)) + with sqlite_session_factory() as session: + account = Account(name="Runner Account", email="runner@example.com") + account.id = account_id + session.add(account) + session.commit() + runner = _AppRunner(session_factory=sqlite_session_factory, exec_params=exec_params) resolved_user = runner._resolve_user() - assert resolved_user is account - account.set_tenant_id_with_session.assert_not_called() + assert resolved_user.id == account_id + assert resolved_user.current_tenant is None -def test_resolve_account_for_run_without_switching_tenant(): - account = MagicMock() - session = MagicMock() - session.get.return_value = account +def test_resolve_account_for_run_without_switching_tenant(sqlite_session: Session): + account_id = str(uuid.uuid4()) + account = Account(name="Run Account", email="run@example.com") + account.id = account_id + sqlite_session.add(account) + sqlite_session.commit() workflow_run = MagicMock( created_by_role=CreatorUserRole.ACCOUNT, - created_by="user-id", + created_by=account_id, tenant_id="resource-tenant-id", ) - resolved_user = workflow_execute_task_module._resolve_user_for_run(session, workflow_run) + resolved_user = workflow_execute_task_module._resolve_user_for_run(sqlite_session, workflow_run) assert resolved_user is account - account.set_tenant_id_with_session.assert_not_called() + assert account.current_tenant is None def test_app_runner_streaming_failure_keeps_existing_pre_runtime_helper_behavior( @@ -847,6 +854,7 @@ def test_resume_app_execution_returns_early_when_advanced_chat_missing_conversat def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session], ): generate_entity = _build_advanced_chat_generate_entity(conversation_id="conversation-id") @@ -873,8 +881,6 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs( "tasks.app_generate.workflow_execute_task.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", lambda **kwargs: MagicMock(), ) - session = MagicMock() - _resume_advanced_chat( app_model=SimpleNamespace(id="app-id", tenant_id="resource-tenant-id"), workflow=workflow, @@ -888,12 +894,12 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs( pause_state_config=MagicMock(), workflow_run_id="workflow-run-id", workflow_run=SimpleNamespace(triggered_from="app_run"), - session=session, + session=sqlite_session, ) resumed_entity = generator_instance.resume.call_args.kwargs["application_generate_entity"] assert resumed_entity.stream is True - assert generator_instance.resume.call_args.kwargs["session"] is session + assert generator_instance.resume.call_args.kwargs["session"] is sqlite_session publish_streaming_response.assert_called_once_with( response_stream, "workflow-run-id",