test: migrate retention and task sessions to SQLite (#40090)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Asuka Minato
2026-08-11 16:05:36 +09:00
committed by GitHub
parent 500e37e2fd
commit dd3cf53044
11 changed files with 689 additions and 499 deletions
@@ -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
@@ -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")
@@ -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,
@@ -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:
@@ -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=[],
)
@@ -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
@@ -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()
@@ -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)
@@ -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"]
@@ -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",
@@ -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",