diff --git a/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py b/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py index 3890ef58680..46817f4c6b3 100644 --- a/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py +++ b/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py @@ -68,7 +68,6 @@ from repositories.api_workflow_run_repository import APIWorkflowRunRepository from repositories.sqlalchemy_workflow_trigger_log_repository import SQLAlchemyWorkflowTriggerLogRepository from services.billing_service import BillingService from services.retention.workflow_run.archive_bundle_index import ( - ArchiveBundleManifest, ArchiveBundleTableManifestEntry, decode_archive_bundle_manifest, upsert_archive_bundle_index_from_manifest, @@ -110,6 +109,7 @@ class ArchiveManifestDict(TypedDict): max_run_id: str archived_at: str campaign_id: str + # None means the archive scan has no inclusive lower time bound. archive_window_start: str | None archive_window_end: str run_shard: str @@ -873,7 +873,7 @@ class WorkflowRunArchiver: identity: ArchiveBundleIdentity, runs: Sequence[WorkflowRun], table_stats: list[TableStats], - ) -> ArchiveBundleManifest: + ) -> ArchiveManifestDict: """Generate a manifest for the archived workflow run bundle.""" tables: dict[str, ArchiveBundleTableManifestEntry] = { stat.table_name: { diff --git a/api/tests/integration_tests/services/retention/test_workflow_run_archiver.py b/api/tests/unit_tests/services/retention/workflow_run/test_workflow_run_archiver_sqlite.py similarity index 75% rename from api/tests/integration_tests/services/retention/test_workflow_run_archiver.py rename to api/tests/unit_tests/services/retention/workflow_run/test_workflow_run_archiver_sqlite.py index 098069776f2..fa4de531b7f 100644 --- a/api/tests/integration_tests/services/retention/test_workflow_run_archiver.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_workflow_run_archiver_sqlite.py @@ -1,37 +1,52 @@ +"""SQLite and in-process coverage for workflow-run bundle archiving.""" + import datetime import json import uuid +from types import SimpleNamespace +from typing import override from unittest.mock import ANY, MagicMock, patch import pyarrow as pa import pyarrow.parquet as pq import pytest +from sqlalchemy import select +from sqlalchemy.engine import Engine from sqlalchemy.exc import OperationalError +from sqlalchemy.orm import Session, sessionmaker from enums import DeploymentEdition -from models.workflow import WorkflowRunArchiveBundle +from graphon.enums import WorkflowExecutionStatus +from libs.archive_storage import ArchiveStorage +from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom +from models.workflow import WorkflowRun, WorkflowRunArchiveBundle, WorkflowType from services.retention.workflow_run.archive_paid_plan_workflow_run import ( ArchiveResult, ArchiveSummary, + TableStats, WorkflowRunArchiver, ) from services.retention.workflow_run.constants import ARCHIVE_BUNDLE_FORMAT, ARCHIVE_BUNDLE_SCHEMA_VERSION -class FakeArchiveStorage: - def __init__(self, objects: dict[str, bytes] | None = None): - self.objects = objects or {} +class FakeArchiveStorage(ArchiveStorage): + def __init__(self, objects: dict[str, bytes] | None = None) -> None: + self.objects = {} if objects is None else objects + @override def object_exists(self, key: str) -> bool: return key in self.objects + @override def get_object(self, key: str) -> bytes: return self.objects[key] + @override def put_object(self, key: str, data: bytes) -> str: self.objects[key] = data return "checksum" + @override def list_objects(self, prefix: str) -> list[str]: return sorted(key for key in self.objects if key.startswith(prefix)) @@ -45,63 +60,73 @@ def _db_disconnect_error() -> OperationalError: ) -def _run(run_id: str = "run-1"): - run = MagicMock() - run.id = run_id - run.tenant_id = "tenant-1" - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) - return run - - -def _session_context(session): - context = MagicMock() - context.__enter__.return_value = session - context.__exit__.return_value = False - return context +def _run(run_id: str = "run-1") -> WorkflowRun: + """Build the mapped entity consumed by the archiver without persisting it.""" + return WorkflowRun( + id=run_id, + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + type=WorkflowType.WORKFLOW, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + version="1", + graph="{}", + inputs="{}", + status=WorkflowExecutionStatus.SUCCEEDED, + outputs="{}", + error=None, + elapsed_time=0, + total_tokens=0, + total_steps=0, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + created_at=datetime.datetime(2025, 3, 15, 10, 0, 0), + finished_at=datetime.datetime(2025, 3, 15, 10, 1, 0), + ) class TestWorkflowRunArchiverInit: - def test_start_from_without_end_before_raises(self): + def test_start_from_without_end_before_raises(self) -> None: with pytest.raises(ValueError, match="start_from and end_before must be provided together"): WorkflowRunArchiver(start_from=datetime.datetime(2025, 1, 1)) - def test_end_before_without_start_from_raises(self): + def test_end_before_without_start_from_raises(self) -> None: with pytest.raises(ValueError, match="start_from and end_before must be provided together"): WorkflowRunArchiver(end_before=datetime.datetime(2025, 1, 1)) - def test_start_equals_end_raises(self): + def test_start_equals_end_raises(self) -> None: ts = datetime.datetime(2025, 1, 1) with pytest.raises(ValueError, match="start_from must be earlier than end_before"): WorkflowRunArchiver(start_from=ts, end_before=ts) - def test_start_after_end_raises(self): + def test_start_after_end_raises(self) -> None: with pytest.raises(ValueError, match="start_from must be earlier than end_before"): WorkflowRunArchiver( start_from=datetime.datetime(2025, 6, 1), end_before=datetime.datetime(2025, 1, 1), ) - def test_workers_zero_raises(self): + def test_workers_zero_raises(self) -> None: with pytest.raises(ValueError, match="workers must be at least 1"): WorkflowRunArchiver(workers=0) - def test_run_shard_index_without_total_raises(self): + def test_run_shard_index_without_total_raises(self) -> None: with pytest.raises(ValueError, match="run_shard_index and run_shard_total must be provided together"): WorkflowRunArchiver(run_shard_index=0) - def test_run_shard_total_without_index_raises(self): + def test_run_shard_total_without_index_raises(self) -> None: with pytest.raises(ValueError, match="run_shard_index and run_shard_total must be provided together"): WorkflowRunArchiver(run_shard_total=4) - def test_run_shard_total_above_supported_range_raises(self): + def test_run_shard_total_above_supported_range_raises(self) -> None: with pytest.raises(ValueError, match="run_shard_total must be between 1 and 16"): WorkflowRunArchiver(run_shard_index=0, run_shard_total=17) - def test_run_shard_index_must_be_less_than_total(self): + def test_run_shard_index_must_be_less_than_total(self) -> None: with pytest.raises(ValueError, match="run_shard_index must be between 0 and run_shard_total - 1"): WorkflowRunArchiver(run_shard_index=4, run_shard_total=4) - def test_valid_init_defaults(self): + def test_valid_init_defaults(self) -> None: archiver = WorkflowRunArchiver(days=30, batch_size=50) assert archiver.days == 30 assert archiver.batch_size == 50 @@ -109,7 +134,7 @@ class TestWorkflowRunArchiverInit: assert archiver.delete_after_archive is False assert archiver.start_from is None - def test_valid_init_with_time_range(self): + def test_valid_init_with_time_range(self) -> None: start = datetime.datetime(2025, 1, 1) end = datetime.datetime(2025, 6, 1) archiver = WorkflowRunArchiver(start_from=start, end_before=end, workers=2) @@ -117,13 +142,14 @@ class TestWorkflowRunArchiverInit: assert archiver.end_before is not None assert archiver.workers == 2 - def test_delete_after_archive_is_not_supported_for_bundle_archive(self): + def test_delete_after_archive_is_not_supported_for_bundle_archive(self) -> None: with pytest.raises(ValueError, match="delete_after_archive is not supported by bundle archive"): WorkflowRunArchiver(delete_after_archive=True) - def test_get_runs_batch_passes_shard_options(self): + def test_get_runs_batch_passes_shard_options(self) -> None: repo = MagicMock() - repo.get_runs_batch_by_time_range.return_value = [] + empty_runs: list[WorkflowRun] = [] + repo.get_runs_batch_by_time_range.return_value = empty_runs archiver = WorkflowRunArchiver( tenant_prefixes=["0", "a"], run_shard_index=1, @@ -138,9 +164,10 @@ class TestWorkflowRunArchiverInit: assert repo.get_runs_batch_by_time_range.call_args.kwargs["run_shard_index"] == 1 assert repo.get_runs_batch_by_time_range.call_args.kwargs["run_shard_total"] == 4 - def test_get_runs_batch_prefers_planned_tenant_ids_over_prefix_filter(self): + def test_get_runs_batch_prefers_planned_tenant_ids_over_prefix_filter(self) -> None: repo = MagicMock() - repo.get_runs_batch_by_time_range.return_value = [] + empty_runs: list[WorkflowRun] = [] + repo.get_runs_batch_by_time_range.return_value = empty_runs archiver = WorkflowRunArchiver( tenant_ids=["0tenant"], tenant_prefixes=["0"], @@ -154,9 +181,10 @@ class TestWorkflowRunArchiverInit: assert repo.get_runs_batch_by_time_range.call_args.kwargs["tenant_ids"] == ["0tenant"] assert repo.get_runs_batch_by_time_range.call_args.kwargs["tenant_prefixes"] is None - def test_get_runs_batch_uses_current_tenant_scan_scope(self): + def test_get_runs_batch_uses_current_tenant_scan_scope(self) -> None: repo = MagicMock() - repo.get_runs_batch_by_time_range.return_value = [] + empty_runs: list[WorkflowRun] = [] + repo.get_runs_batch_by_time_range.return_value = empty_runs archiver = WorkflowRunArchiver( tenant_ids=["tenant-a", "tenant-b"], workflow_run_repo=repo, @@ -167,9 +195,10 @@ class TestWorkflowRunArchiverInit: repo.get_runs_batch_by_time_range.assert_called_once() assert repo.get_runs_batch_by_time_range.call_args.kwargs["tenant_ids"] == ["tenant-b"] - def test_get_runs_batch_retries_retryable_db_disconnect(self): + def test_get_runs_batch_retries_retryable_db_disconnect(self) -> None: repo = MagicMock() - repo.get_runs_batch_by_time_range.side_effect = [_db_disconnect_error(), []] + empty_runs: list[WorkflowRun] = [] + repo.get_runs_batch_by_time_range.side_effect = [_db_disconnect_error(), empty_runs] archiver = WorkflowRunArchiver(workflow_run_repo=repo) with patch("services.retention.workflow_run.db_retry.time.sleep") as sleep: @@ -179,7 +208,7 @@ class TestWorkflowRunArchiverInit: assert repo.get_runs_batch_by_time_range.call_count == 2 sleep.assert_called_once_with(1.0) - def test_get_runs_batch_does_not_retry_non_db_broken_pipe_error(self): + def test_get_runs_batch_does_not_retry_non_db_broken_pipe_error(self) -> None: repo = MagicMock() repo.get_runs_batch_by_time_range.side_effect = RuntimeError("broken pipe") archiver = WorkflowRunArchiver(workflow_run_repo=repo) @@ -193,7 +222,7 @@ class TestWorkflowRunArchiverInit: repo.get_runs_batch_by_time_range.assert_called_once() sleep.assert_not_called() - def test_start_message_includes_shard(self): + def test_start_message_includes_shard(self) -> None: archiver = WorkflowRunArchiver(tenant_prefixes=["0"], run_shard_index=1, run_shard_total=4) message = archiver._build_start_message() @@ -201,7 +230,7 @@ class TestWorkflowRunArchiverInit: assert "tenant_prefixes=0" in message assert "run_shard=1/4" in message - def test_start_message_summarizes_large_planned_tenant_list(self): + def test_start_message_summarizes_large_planned_tenant_list(self) -> None: tenant_ids = [f"tenant-{index}" for index in range(11)] archiver = WorkflowRunArchiver(tenant_ids=tenant_ids, tenant_prefixes=["0"]) @@ -212,12 +241,10 @@ class TestWorkflowRunArchiverInit: class TestBuildArchiveBundle: - def test_bundle_contains_manifest_and_all_table_objects(self): + def test_bundle_contains_manifest_and_all_table_objects(self) -> None: archiver = WorkflowRunArchiver(days=90) - run = MagicMock() - run.id = str(uuid.uuid4()) + run = _run(str(uuid.uuid4())) run.tenant_id = str(uuid.uuid4()) - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) identity = archiver._build_bundle_identity([run]) table_data = {"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}]} @@ -233,16 +260,12 @@ class TestBuildArchiveBundle: class TestGenerateManifest: - def test_manifest_structure(self): + def test_manifest_structure(self) -> None: start = datetime.datetime(2025, 1, 1, tzinfo=datetime.UTC) end = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC) archiver = WorkflowRunArchiver(start_from=start, end_before=end, run_shard_index=1, run_shard_total=4) - from services.retention.workflow_run.archive_paid_plan_workflow_run import TableStats - - run = MagicMock() - run.id = str(uuid.uuid4()) + run = _run(str(uuid.uuid4())) run.tenant_id = str(uuid.uuid4()) - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) identity = archiver._build_bundle_identity([run]) stats = [ @@ -282,7 +305,7 @@ class TestGenerateManifest: class TestFilterPaidTenants: - def test_all_tenants_paid_in_community_edition(self): + def test_all_tenants_paid_in_community_edition(self) -> None: archiver = WorkflowRunArchiver(days=90) tenant_ids = {"t1", "t2", "t3"} @@ -292,7 +315,7 @@ class TestFilterPaidTenants: assert result == tenant_ids - def test_empty_tenants_returns_empty(self): + def test_empty_tenants_returns_empty(self) -> None: archiver = WorkflowRunArchiver(days=90) with patch("services.retention.workflow_run.archive_paid_plan_workflow_run.dify_config") as cfg: @@ -301,7 +324,7 @@ class TestFilterPaidTenants: assert result == set() - def test_only_paid_plans_returned(self): + def test_only_paid_plans_returned(self) -> None: archiver = WorkflowRunArchiver(days=90) mock_bulk = { @@ -322,7 +345,7 @@ class TestFilterPaidTenants: assert "t3" in result assert "t2" not in result - def test_billing_api_failure_returns_empty(self): + def test_billing_api_failure_returns_empty(self) -> None: archiver = WorkflowRunArchiver(days=90) with ( @@ -335,7 +358,7 @@ class TestFilterPaidTenants: assert result == set() - def test_planned_paid_tenants_skip_billing_lookup(self): + def test_planned_paid_tenants_skip_billing_lookup(self) -> None: archiver = WorkflowRunArchiver(days=90, paid_tenant_ids=["t1", "t3"]) with ( @@ -351,31 +374,32 @@ class TestFilterPaidTenants: class TestDryRunArchive: @patch("services.retention.workflow_run.archive_paid_plan_workflow_run.get_archive_storage") - def test_dry_run_does_not_call_storage(self, mock_get_storage, flask_req_ctx): + def test_dry_run_does_not_call_storage(self, mock_get_storage: MagicMock, sqlite_engine: Engine) -> None: archiver = WorkflowRunArchiver(days=90, dry_run=True) - with patch.object(archiver, "_get_runs_batch", return_value=[]): + with ( + patch( + "services.retention.workflow_run.archive_paid_plan_workflow_run.db", + SimpleNamespace(engine=sqlite_engine), + ), + patch.object(archiver, "_get_runs_batch", return_value=[]), + ): summary = archiver.run() mock_get_storage.assert_not_called() assert isinstance(summary, ArchiveSummary) assert summary.runs_failed == 0 - def test_dry_run_estimates_table_and_object_sizes(self): + def test_dry_run_estimates_table_and_object_sizes(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90, dry_run=True) - run = MagicMock() - run.id = "run-1" - run.tenant_id = "tenant-1" - run.app_id = "app-1" - run.workflow_id = "workflow-1" - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) + run = _run() table_data = { "workflow_runs": [{"id": "run-1", "tenant_id": "tenant-1"}], "workflow_app_logs": [{"id": "log-1", "workflow_run_id": "run-1"}], } with patch.object(archiver, "_extract_bundle_data", return_value=table_data): - result = archiver._archive_bundle(MagicMock(), None, [run]) + result = archiver._archive_bundle(sqlite_session, None, [run]) stats_by_table = {stat.table_name: stat for stat in result.tables} assert result.success is True @@ -387,14 +411,20 @@ class TestDryRunArchive: assert stats_by_table["workflow_node_executions"].row_count == 0 assert stats_by_table["workflow_node_executions"].size_bytes > 0 - def test_summary_merges_dry_run_estimates(self): + def test_summary_merges_dry_run_estimates(self) -> None: summary = ArchiveSummary() - result = MagicMock() - result.object_size_bytes = 128 - result.tables = [ - MagicMock(table_name="workflow_runs", row_count=1, size_bytes=64), - MagicMock(table_name="workflow_app_logs", row_count=2, size_bytes=32), - ] + result = ArchiveResult( + bundle_id="bundle-1", + tenant_id="tenant-1", + object_prefix="prefix", + run_count=1, + success=True, + object_size_bytes=128, + tables=[ + TableStats(table_name="workflow_runs", row_count=1, checksum="", size_bytes=64), + TableStats(table_name="workflow_app_logs", row_count=2, checksum="", size_bytes=32), + ], + ) WorkflowRunArchiver._merge_result_stats(summary, result) @@ -406,15 +436,12 @@ class TestDryRunArchive: class TestArchiveDbRetry: - def test_archive_bundle_groups_retries_with_fresh_session(self): + def test_archive_bundle_groups_retries_with_fresh_session( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: archiver = WorkflowRunArchiver(days=90) run = _run() - session_maker = MagicMock( - side_effect=[ - _session_context(MagicMock(name="session-1")), - _session_context(MagicMock(name="session-2")), - ] - ) + session_maker = MagicMock(wraps=sqlite_session_factory) success = ArchiveResult( bundle_id=archiver._build_bundle_identity([run]).bundle_id, tenant_id=run.tenant_id, @@ -434,16 +461,12 @@ class TestArchiveDbRetry: assert session_maker.call_count == 2 sleep.assert_called_once_with(1.0) - def test_archive_bundle_groups_returns_failed_result_after_retry_exhaustion(self): + def test_archive_bundle_groups_returns_failed_result_after_retry_exhaustion( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: archiver = WorkflowRunArchiver(days=90) run = _run() - session_maker = MagicMock( - side_effect=[ - _session_context(MagicMock(name="session-1")), - _session_context(MagicMock(name="session-2")), - _session_context(MagicMock(name="session-3")), - ] - ) + session_maker = MagicMock(wraps=sqlite_session_factory) with ( patch.object(archiver, "_archive_bundle", side_effect=[_db_disconnect_error()] * 3) as archive_bundle, @@ -458,21 +481,22 @@ class TestArchiveDbRetry: assert session_maker.call_count == archiver.DB_RETRY_ATTEMPTS assert sleep.call_count == archiver.DB_RETRY_ATTEMPTS - 1 - def test_archive_bundle_uses_safe_rollback_when_failure_rolls_back_badly(self): + def test_archive_bundle_uses_safe_rollback_when_failure_rolls_back_badly(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90, dry_run=True) - session = MagicMock() - session.rollback.side_effect = RuntimeError("rollback failed") - with patch.object(archiver, "_extract_bundle_data", side_effect=RuntimeError("extract failed")): - result = archiver._archive_bundle(session, None, [_run()]) + with ( + patch.object(sqlite_session, "rollback", side_effect=RuntimeError("rollback failed")) as rollback, + patch.object(archiver, "_extract_bundle_data", side_effect=RuntimeError("extract failed")), + ): + result = archiver._archive_bundle(sqlite_session, None, [_run()]) assert result.success is False assert result.error == "extract failed" - session.rollback.assert_called_once() + rollback.assert_called_once() class TestArchiveRunIdempotency: - def _index_payload(self, archiver: WorkflowRunArchiver, run_ids: list[str], run) -> tuple[str, bytes]: + def _index_payload(self, archiver: WorkflowRunArchiver, run_ids: list[str], run: WorkflowRun) -> tuple[str, bytes]: identity = archiver._build_bundle_identity([run]) index_key = archiver._get_index_object_key(identity) payload = json.dumps( @@ -487,56 +511,52 @@ class TestArchiveRunIdempotency: ).encode() return index_key, payload - def test_locked_bundle_is_skipped(self): + def test_locked_bundle_is_skipped(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) - run = MagicMock() - run.id = "run-1" - run.tenant_id = "tenant-1" - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) + run = _run() with ( patch.object(archiver, "_lock_runs_for_archive", return_value=[]), ): storage = MagicMock() storage.object_exists.return_value = False - result = archiver._archive_bundle(MagicMock(), storage, [run]) + result = archiver._archive_bundle(sqlite_session, storage, [run]) assert result.success is True assert result.skipped is True assert result.error == "one or more runs locked or deleted by another archiver" - def test_already_archived_bundle_is_skipped(self): + def test_already_archived_bundle_is_skipped(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) - run = MagicMock() - run.id = "run-1" - run.tenant_id = "tenant-1" - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) + run = _run() storage = MagicMock() storage.object_exists.return_value = True with patch.object(archiver, "_sync_existing_bundle_index") as sync_existing_bundle_index: - result = archiver._archive_bundle(MagicMock(), storage, [run]) + result = archiver._archive_bundle(sqlite_session, storage, [run]) assert result.success is True assert result.skipped is True assert result.error == "bundle already archived" sync_existing_bundle_index.assert_called_once() - def test_existing_bundle_catalog_publication_failure_is_not_success(self): + def test_existing_bundle_catalog_publication_failure_is_not_success(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) run = _run() - session = MagicMock() storage = MagicMock() storage.object_exists.return_value = True - with patch.object(archiver, "_sync_existing_bundle_index", side_effect=RuntimeError("catalog unavailable")): - result = archiver._archive_bundle(session, storage, [run]) + with ( + patch.object(sqlite_session, "rollback", wraps=sqlite_session.rollback) as rollback, + patch.object(archiver, "_sync_existing_bundle_index", side_effect=RuntimeError("catalog unavailable")), + ): + result = archiver._archive_bundle(sqlite_session, storage, [run]) assert result.success is False assert result.error == "catalog unavailable" - session.rollback.assert_called_once() + rollback.assert_called_once() - def test_retry_repairs_index_after_catalog_commit_then_index_write_failure(self): + def test_retry_repairs_index_after_catalog_commit_then_index_write_failure(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) run = _run() identity = archiver._build_bundle_identity([run]) @@ -555,15 +575,13 @@ class TestArchiveRunIdempotency: return original_put_object(key, data) storage.put_object = MagicMock(side_effect=put_object) - first_session = MagicMock() - first_session.scalar.return_value = None table_data = {"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}]} with ( patch.object(archiver, "_lock_runs_for_archive", return_value=[run]), patch.object(archiver, "_extract_bundle_data", return_value=table_data), ): - first_result = archiver._archive_bundle(first_session, storage, [run]) + first_result = archiver._archive_bundle(sqlite_session, storage, [run]) assert first_result.success is False assert first_result.error == "index write failed" @@ -571,10 +589,7 @@ class TestArchiveRunIdempotency: assert json.loads(storage.objects[index_key])["run_ids"] == [] storage.list_objects = MagicMock(wraps=storage.list_objects) - retry_session = MagicMock() - retry_session.scalar.return_value = None - - retry_result = archiver._archive_bundle(retry_session, storage, [run]) + retry_result = archiver._archive_bundle(sqlite_session, storage, [run]) assert retry_result.success is True assert retry_result.skipped is True @@ -582,7 +597,7 @@ class TestArchiveRunIdempotency: assert json.loads(storage.objects[index_key])["run_ids"] == [run.id] storage.list_objects.assert_not_called() - def test_existing_manifest_with_missing_index_fails_without_partial_rebuild(self): + def test_existing_manifest_with_missing_index_fails_without_partial_rebuild(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) run = _run() identity = archiver._build_bundle_identity([run]) @@ -596,21 +611,17 @@ class TestArchiveRunIdempotency: storage = FakeArchiveStorage({manifest_key: manifest_data}) storage.list_objects = MagicMock(wraps=storage.list_objects) - result = archiver._archive_bundle(MagicMock(), storage, [run]) + result = archiver._archive_bundle(sqlite_session, storage, [run]) assert result.success is False assert "archive shard index missing" in (result.error or "") assert index_key not in storage.objects storage.list_objects.assert_not_called() - def test_successful_bundle_persists_archive_index(self): + def test_successful_bundle_persists_archive_index(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) - run = MagicMock() - run.id = str(uuid.uuid4()) + run = _run(str(uuid.uuid4())) run.tenant_id = str(uuid.uuid4()) - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) - session = MagicMock() - session.scalar.return_value = None storage = MagicMock() storage.object_exists.return_value = False table_data = { @@ -622,9 +633,11 @@ class TestArchiveRunIdempotency: patch.object(archiver, "_lock_runs_for_archive", return_value=[run]), patch.object(archiver, "_extract_bundle_data", return_value=table_data), ): - result = archiver._archive_bundle(session, storage, [run]) + result = archiver._archive_bundle(sqlite_session, storage, [run]) - archived_bundle = session.add.call_args.args[0] + archived_bundle = sqlite_session.scalar( + select(WorkflowRunArchiveBundle).where(WorkflowRunArchiveBundle.tenant_id == run.tenant_id) + ) assert result.success is True assert isinstance(archived_bundle, WorkflowRunArchiveBundle) assert archived_bundle.tenant_id == run.tenant_id @@ -632,40 +645,36 @@ class TestArchiveRunIdempotency: assert archived_bundle.month == 3 assert archived_bundle.workflow_run_count == 1 assert archived_bundle.row_count == 2 - session.commit.assert_called_once() - def test_new_bundle_catalog_commit_failure_is_not_success(self): + def test_new_bundle_catalog_commit_failure_is_not_success(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) run = _run(str(uuid.uuid4())) run.tenant_id = str(uuid.uuid4()) - session = MagicMock() - session.scalar.return_value = None - session.commit.side_effect = RuntimeError("catalog commit failed") storage = MagicMock() storage.object_exists.return_value = False - storage.list_objects.return_value = [] + empty_keys: list[str] = [] + storage.list_objects.return_value = empty_keys table_data = {"workflow_runs": [{"id": run.id, "tenant_id": run.tenant_id}]} with ( + patch.object(sqlite_session, "commit", side_effect=RuntimeError("catalog commit failed")), + patch.object(sqlite_session, "rollback", wraps=sqlite_session.rollback) as rollback, patch.object(archiver, "_lock_runs_for_archive", return_value=[run]), patch.object(archiver, "_extract_bundle_data", return_value=table_data), ): - result = archiver._archive_bundle(session, storage, [run]) + result = archiver._archive_bundle(sqlite_session, storage, [run]) assert result.success is False assert result.error == "catalog commit failed" - session.rollback.assert_called_once() + rollback.assert_called_once() - def test_index_skips_all_already_archived_runs(self): + def test_index_skips_all_already_archived_runs(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) - run = MagicMock() - run.id = "run-1" - run.tenant_id = "tenant-1" - run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) + run = _run() index_key, index_payload = self._index_payload(archiver, ["run-1"], run) storage = FakeArchiveStorage({index_key: index_payload}) - result = archiver._archive_bundle(MagicMock(), storage, [run]) + result = archiver._archive_bundle(sqlite_session, storage, [run]) assert result.success is True assert result.skipped is True @@ -673,15 +682,10 @@ class TestArchiveRunIdempotency: assert result.skipped_run_count == 1 assert result.error == "all runs already archived in shard index" - def test_index_filters_duplicate_runs_before_archive(self): + def test_index_filters_duplicate_runs_before_archive(self, sqlite_session: Session) -> None: archiver = WorkflowRunArchiver(days=90) - archived_run = MagicMock() - archived_run.id = "run-1" - archived_run.tenant_id = "tenant-1" - archived_run.created_at = datetime.datetime(2025, 3, 15, 10, 0, 0) - new_run = MagicMock() - new_run.id = "run-2" - new_run.tenant_id = "tenant-1" + archived_run = _run() + new_run = _run("run-2") new_run.created_at = datetime.datetime(2025, 3, 15, 11, 0, 0) index_key, index_payload = self._index_payload(archiver, ["run-1"], archived_run) storage = FakeArchiveStorage({index_key: index_payload}) @@ -690,7 +694,7 @@ class TestArchiveRunIdempotency: patch.object(archiver, "_lock_runs_for_archive", return_value=[new_run]) as lock_runs, patch.object(archiver, "_extract_bundle_data", return_value={"workflow_runs": [{"id": "run-2"}]}), ): - result = archiver._archive_bundle(MagicMock(), storage, [archived_run, new_run]) + result = archiver._archive_bundle(sqlite_session, storage, [archived_run, new_run]) assert result.success is True assert result.skipped is False