mirror of
https://github.com/langgenius/dify.git
synced 2026-09-01 15:09:21 +08:00
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:
@@ -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
|
||||
|
||||
+36
-13
@@ -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,
|
||||
|
||||
+172
-148
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user