Files
Aether/tests/services/test_batch_delete_task.py
fawney19 776dd2f8ea feat(cleanup): 解耦 request_candidates 与 provider_api_keys 生命周期
- 移除 request_candidates.key_id 对 provider_api_keys 的外键约束(含迁移脚本)
- 删除 Key 时不再级联删除候选记录,改为独立按保留天数定时清理
- 新增 request_candidates_retention_days / request_candidates_cleanup_batch_size 配置项
- batch_delete_task 增加 lock_timeout 及超时自动降批重试机制
- cleanup_key_references 提取阶段化清理流程,移除 RequestCandidate 联动删除
- 前端 CleanupPolicySection 新增候选记录保留天数和清理批次配置

Closes #227

Co-authored-by: Entropy-Xu <entropy.xu@cloudhabitatsh.com>
2026-03-14 01:33:08 +08:00

252 lines
7.7 KiB
Python

from __future__ import annotations
import asyncio
from collections.abc import Coroutine, Generator
from concurrent.futures import Future
from contextlib import contextmanager
from types import SimpleNamespace
import pytest
from src.services.provider_keys import batch_delete_task as taskmod
@contextmanager
def _fake_db_context() -> Generator[object, None, None]:
yield object()
async def _fake_get_redis_client(*, require_redis: bool = False) -> object:
_ = require_redis
return object()
async def _noop_delete_side_effects(**_kwargs: object) -> None:
return None
@pytest.mark.asyncio
async def test_run_batch_delete_waits_for_progress_updates_before_completion(
monkeypatch: pytest.MonkeyPatch,
) -> None:
release_progress = asyncio.Event()
updates: list[dict[str, object]] = []
async def fake_update_task_field(
task_id: str,
r: object | None = None,
**fields: object,
) -> None:
_ = task_id, r
updates.append(dict(fields))
def fake_sync_delete(
provider_id: str,
key_ids: list[str],
progress_callback: object | None = None,
) -> int:
_ = provider_id, key_ids
assert callable(progress_callback)
progress_callback(1)
return 1
def fake_run_coroutine_threadsafe(
coro: Coroutine[object, object, object],
loop: asyncio.AbstractEventLoop,
) -> Future[object]:
future: Future[object] = Future()
async def runner() -> None:
await release_progress.wait()
try:
result = await coro
except Exception as exc:
future.set_exception(exc)
else:
future.set_result(result)
loop.call_soon_threadsafe(lambda: asyncio.create_task(runner()))
return future
monkeypatch.setattr(taskmod, "get_redis_client", _fake_get_redis_client)
monkeypatch.setattr(taskmod, "_update_task_field", fake_update_task_field)
monkeypatch.setattr(taskmod, "_sync_delete", fake_sync_delete)
monkeypatch.setattr(taskmod.asyncio, "run_coroutine_threadsafe", fake_run_coroutine_threadsafe)
monkeypatch.setattr("src.database.get_db_context", _fake_db_context)
monkeypatch.setattr(
"src.services.provider_keys.key_side_effects.run_delete_key_side_effects",
_noop_delete_side_effects,
)
task = asyncio.create_task(taskmod._run_batch_delete("task-1", "provider-1", ["key-1"]))
await asyncio.sleep(0.05)
assert not task.done()
assert updates == [{"status": taskmod.STATUS_RUNNING}]
release_progress.set()
await task
assert updates == [
{"status": taskmod.STATUS_RUNNING},
{"deleted": 1},
{
"status": taskmod.STATUS_COMPLETED,
"deleted": 1,
"message": "1 keys deleted",
},
]
def test_sync_delete_reports_progress_after_each_batch(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class _FakeColumn:
def __eq__(self, other: object) -> tuple[str, object]: # type: ignore[override]
return ("eq", other)
def in_(self, values: list[str]) -> tuple[str, tuple[str, ...]]:
return ("in", tuple(values))
class _FakeProviderAPIKey:
provider_id = _FakeColumn()
id = _FakeColumn()
class _FakeDeleteStatement:
def where(self, *_conditions: object) -> "_FakeDeleteStatement":
return self
class _FakeSession:
def __init__(self) -> None:
self.rowcounts = [2, 1]
self.commits = 0
self.closed = False
self.text_statements: list[str] = []
def execute(self, _statement: object) -> SimpleNamespace:
# SET LOCAL statement_timeout / lock_timeout 不消耗 rowcount
if hasattr(_statement, "text"):
self.text_statements.append(_statement.text)
return SimpleNamespace(rowcount=0)
return SimpleNamespace(rowcount=self.rowcounts.pop(0))
def commit(self) -> None:
self.commits += 1
def rollback(self) -> None:
raise AssertionError("rollback should not be called")
def close(self) -> None:
self.closed = True
session = _FakeSession()
progress_updates: list[int] = []
cleanup_batches: list[list[str]] = []
monkeypatch.setattr(taskmod, "_CLEANUP_BATCH_SIZE", 2)
monkeypatch.setattr(
taskmod,
"cleanup_key_references",
lambda _db, batch, **_kwargs: cleanup_batches.append(list(batch)),
)
monkeypatch.setattr("src.database.create_session", lambda: session)
monkeypatch.setattr(
"src.models.database.ProviderAPIKey",
_FakeProviderAPIKey,
)
monkeypatch.setattr(taskmod, "sa_delete", lambda _model: _FakeDeleteStatement())
affected = taskmod._sync_delete(
"provider-1",
["key-1", "key-2", "key-3"],
progress_updates.append,
)
assert affected == 3
assert progress_updates == [2, 3]
assert cleanup_batches == [["key-1", "key-2"], ["key-3"]]
assert session.commits == 2
assert session.text_statements == [
"SET LOCAL statement_timeout = '30000'",
"SET LOCAL lock_timeout = '5000'",
"SET LOCAL statement_timeout = '30000'",
"SET LOCAL lock_timeout = '5000'",
]
assert session.closed is True
def test_sync_delete_retries_with_smaller_batches_after_timeout(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class _FakeColumn:
def __eq__(self, other: object) -> tuple[str, object]: # type: ignore[override]
return ("eq", other)
def in_(self, values: list[str]) -> tuple[str, tuple[str, ...]]:
return ("in", tuple(values))
class _FakeProviderAPIKey:
provider_id = _FakeColumn()
id = _FakeColumn()
class _FakeDeleteStatement:
def where(self, *_conditions: object) -> "_FakeDeleteStatement":
return self
class _FakeSession:
def __init__(self) -> None:
self.rowcounts = [2, 2]
self.commits = 0
self.rollbacks = 0
self.closed = False
def execute(self, _statement: object) -> SimpleNamespace:
if hasattr(_statement, "text"):
return SimpleNamespace(rowcount=0)
return SimpleNamespace(rowcount=self.rowcounts.pop(0))
def commit(self) -> None:
self.commits += 1
def rollback(self) -> None:
self.rollbacks += 1
def close(self) -> None:
self.closed = True
session = _FakeSession()
progress_updates: list[int] = []
cleanup_batches: list[list[str]] = []
def fake_cleanup(_db: object, batch: list[str], **_kwargs: object) -> None:
cleanup_batches.append(list(batch))
if len(batch) >= 4:
raise RuntimeError(
"(psycopg2.errors.QueryCanceled) canceling statement due to statement timeout"
)
monkeypatch.setattr(taskmod, "_CLEANUP_BATCH_SIZE", 4)
monkeypatch.setattr(taskmod, "_MIN_RETRY_BATCH_SIZE", 1)
monkeypatch.setattr(taskmod, "cleanup_key_references", fake_cleanup)
monkeypatch.setattr("src.database.create_session", lambda: session)
monkeypatch.setattr("src.models.database.ProviderAPIKey", _FakeProviderAPIKey)
monkeypatch.setattr(taskmod, "sa_delete", lambda _model: _FakeDeleteStatement())
affected = taskmod._sync_delete(
"provider-1",
["key-1", "key-2", "key-3", "key-4"],
progress_updates.append,
)
assert affected == 4
assert progress_updates == [4]
assert cleanup_batches == [
["key-1", "key-2", "key-3", "key-4"],
["key-1", "key-2"],
["key-3", "key-4"],
]
assert session.commits == 2
assert session.rollbacks == 1
assert session.closed is True