mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: Provider 异步删除、可配置密码策略、Hub 超时优化及多项改进
- 新增 Provider 异步删除任务系统,后台分阶段删除子资源并清理残留引用 - 新增可配置密码策略等级(weak/medium/strong),支持系统设置面板调整 - aether-hub 升级至 0.1.4,idle timeout 支持禁用(设为 0),worker 默认超时调整为 120s - OAuth 手动续期增加 Redis 分布式锁,防止并发刷新冲突 - ProxyNode 心跳检测改为 asyncio.to_thread,避免阻塞事件循环 - 删除 ModelMultiSelect 和 useInvalidModels,MultiSelect 组件通用化 - 明确 allowed_providers/allowed_api_formats 的 NULL 与空数组语义 - 前端 StandaloneKeyFormDialog、UserFormDialog 等多处 UI 优化 - 新增 Alembic 迁移脚本清理 Provider 删除后的残留引用 - 补充相关测试用例
This commit is contained in:
@@ -9,8 +9,10 @@ from fastapi import FastAPI, HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.api.admin.api_keys.routes import (
|
||||
AdminCreateStandaloneKeyAdapter,
|
||||
AdminGetFullKeyAdapter,
|
||||
AdminToggleApiKeyAdapter,
|
||||
AdminUpdateApiKeyAdapter,
|
||||
)
|
||||
from src.api.admin.api_keys.routes import router as admin_api_keys_router
|
||||
from src.api.admin.users.routes import (
|
||||
@@ -20,6 +22,7 @@ from src.api.admin.users.routes import (
|
||||
from src.api.admin.users.routes import router as admin_users_router
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.database import get_db
|
||||
from src.models.api import CreateApiKeyRequest
|
||||
|
||||
|
||||
def _build_context(db: MagicMock) -> SimpleNamespace:
|
||||
@@ -296,3 +299,114 @@ def test_standalone_detail_route_does_not_expose_is_locked(monkeypatch: pytest.M
|
||||
payload = response.json()
|
||||
assert payload["id"] == "sa-key-2"
|
||||
assert "is_locked" not in payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_standalone_key_adapter_preserves_empty_restriction_lists(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
captured: dict[str, object] = {}
|
||||
created_key = SimpleNamespace(
|
||||
id="sa-key-3",
|
||||
name="Standalone Key 3",
|
||||
get_display_key=lambda: "sk-stand...9012",
|
||||
is_active=True,
|
||||
rate_limit=None,
|
||||
expires_at=None,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
allowed_providers=[],
|
||||
allowed_api_formats=[],
|
||||
allowed_models=[],
|
||||
)
|
||||
|
||||
def _create_api_key(**kwargs: object) -> tuple[SimpleNamespace, str]:
|
||||
captured.update(kwargs)
|
||||
return created_key, "sk-created"
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.api_keys.routes.ApiKeyService.create_api_key", _create_api_key
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.api_keys.routes.WalletService.initialize_api_key_wallet",
|
||||
lambda *_a, **_k: SimpleNamespace(id="wallet-1"),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.api_keys.routes.WalletService.serialize_wallet_summary",
|
||||
lambda _wallet: {"id": "wallet-1"},
|
||||
)
|
||||
|
||||
adapter = AdminCreateStandaloneKeyAdapter(
|
||||
CreateApiKeyRequest(
|
||||
name="Standalone Key 3",
|
||||
initial_balance_usd=10,
|
||||
allowed_providers=[],
|
||||
allowed_api_formats=[],
|
||||
allowed_models=[],
|
||||
)
|
||||
)
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
user=SimpleNamespace(id="admin-1"),
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result["id"] == "sa-key-3"
|
||||
assert captured["allowed_providers"] == []
|
||||
assert captured["allowed_api_formats"] == []
|
||||
assert captured["allowed_models"] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_standalone_key_adapter_preserves_empty_restriction_lists(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
existing_key = SimpleNamespace(id="sa-key-4", is_standalone=True)
|
||||
_mock_query_first(db, existing_key)
|
||||
|
||||
updated_key = SimpleNamespace(
|
||||
id="sa-key-4",
|
||||
name="Standalone Key 4",
|
||||
get_display_key=lambda: "sk-stand...3456",
|
||||
is_active=True,
|
||||
rate_limit=None,
|
||||
expires_at=None,
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _update_api_key(_db: MagicMock, _key_id: str, **kwargs: object) -> SimpleNamespace:
|
||||
captured.update(kwargs)
|
||||
return updated_key
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.api_keys.routes.ApiKeyService.update_api_key", _update_api_key
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.api_keys.routes._ensure_standalone_wallet",
|
||||
lambda *_a, **_k: SimpleNamespace(id="wallet-2"),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.api_keys.routes.WalletService.serialize_wallet_summary",
|
||||
lambda _wallet: {"id": "wallet-2"},
|
||||
)
|
||||
|
||||
adapter = AdminUpdateApiKeyAdapter(
|
||||
key_id="sa-key-4",
|
||||
key_data=CreateApiKeyRequest(
|
||||
allowed_providers=[],
|
||||
allowed_api_formats=[],
|
||||
allowed_models=[],
|
||||
),
|
||||
)
|
||||
|
||||
result = await adapter.handle(_build_context(db))
|
||||
|
||||
assert result["id"] == "sa-key-4"
|
||||
assert captured["allowed_providers"] == []
|
||||
assert captured["allowed_api_formats"] == []
|
||||
assert captured["allowed_models"] == []
|
||||
|
||||
128
tests/api/test_admin_provider_routes.py
Normal file
128
tests/api/test_admin_provider_routes.py
Normal file
@@ -0,0 +1,128 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.api.admin.providers.routes import (
|
||||
AdminDeleteProviderAdapter,
|
||||
AdminProviderDeleteTaskStatusAdapter,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_provider_adapter_submits_async_task_and_deactivates_provider(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
provider = SimpleNamespace(id="provider-1", name="Provider 1", is_active=True)
|
||||
db.query.return_value.filter.return_value.first.return_value = provider
|
||||
|
||||
submit_task_mock = AsyncMock(return_value="task-1")
|
||||
invalidate_models_mock = AsyncMock()
|
||||
invalidate_resolve_mock = AsyncMock()
|
||||
invalidate_provider_cache_mock = AsyncMock()
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.providers.routes.submit_provider_delete",
|
||||
submit_task_mock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.providers.routes.invalidate_models_list_cache",
|
||||
invalidate_models_mock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.providers.routes.ModelCacheService.invalidate_all_resolve_cache",
|
||||
invalidate_resolve_mock,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.providers.routes.ProviderCacheService.invalidate_provider_cache",
|
||||
invalidate_provider_cache_mock,
|
||||
)
|
||||
|
||||
audit_calls: list[dict[str, object]] = []
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
add_audit_metadata=lambda **kwargs: audit_calls.append(kwargs),
|
||||
)
|
||||
|
||||
adapter = AdminDeleteProviderAdapter(provider_id="provider-1")
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result == {
|
||||
"task_id": "task-1",
|
||||
"status": "pending",
|
||||
"message": "删除任务已提交,提供商已进入后台删除队列",
|
||||
}
|
||||
submit_task_mock.assert_awaited_once_with("provider-1")
|
||||
assert provider.is_active is False
|
||||
db.commit.assert_called_once()
|
||||
invalidate_models_mock.assert_awaited_once()
|
||||
invalidate_resolve_mock.assert_awaited_once()
|
||||
invalidate_provider_cache_mock.assert_awaited_once_with("provider-1")
|
||||
assert audit_calls[0]["action"] == "delete_provider"
|
||||
assert audit_calls[1]["task_id"] == "task-1"
|
||||
assert audit_calls[1]["provider_deactivated"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_provider_adapter_reuses_task_without_extra_commit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
provider = SimpleNamespace(id="provider-1", name="Provider 1", is_active=False)
|
||||
db.query.return_value.filter.return_value.first.return_value = provider
|
||||
|
||||
submit_task_mock = AsyncMock(return_value="task-1")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.providers.routes.submit_provider_delete",
|
||||
submit_task_mock,
|
||||
)
|
||||
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
add_audit_metadata=lambda **kwargs: None,
|
||||
)
|
||||
|
||||
adapter = AdminDeleteProviderAdapter(provider_id="provider-1")
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result["task_id"] == "task-1"
|
||||
db.commit.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_provider_task_status_adapter_returns_task_payload(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
task = SimpleNamespace(
|
||||
task_id="task-1",
|
||||
provider_id="provider-1",
|
||||
status="running",
|
||||
stage="deleting_keys",
|
||||
total_keys=100,
|
||||
deleted_keys=25,
|
||||
total_endpoints=8,
|
||||
deleted_endpoints=2,
|
||||
message="deleted key batch 1/2",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.providers.routes.get_provider_delete_task",
|
||||
AsyncMock(return_value=task),
|
||||
)
|
||||
|
||||
context = SimpleNamespace(db=MagicMock(), request=SimpleNamespace(state=SimpleNamespace()))
|
||||
adapter = AdminProviderDeleteTaskStatusAdapter(provider_id="provider-1", task_id="task-1")
|
||||
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result.task_id == "task-1"
|
||||
assert result.status == "running"
|
||||
assert result.stage == "deleting_keys"
|
||||
assert result.deleted_keys == 25
|
||||
assert result.deleted_endpoints == 2
|
||||
@@ -9,6 +9,7 @@ import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.api.admin.users.routes import AdminCreateUserAdapter
|
||||
from src.api.admin.users.routes import router as admin_users_router
|
||||
from src.database import get_db
|
||||
|
||||
@@ -91,3 +92,53 @@ def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) ->
|
||||
assert response.json()[1]["unlimited"] is False
|
||||
batch_getter.assert_called_once()
|
||||
assert batch_getter.call_args.args[1] == ["user-1", "user-2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_user_adapter_preserves_empty_restriction_lists(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _create_user(**kwargs: Any) -> SimpleNamespace:
|
||||
captured.update(kwargs)
|
||||
return SimpleNamespace(
|
||||
id="user-3",
|
||||
email="u3@example.com",
|
||||
username="user3",
|
||||
role=SimpleNamespace(value="user"),
|
||||
is_active=True,
|
||||
allowed_providers=[],
|
||||
allowed_api_formats=[],
|
||||
allowed_models=[],
|
||||
)
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes.UserService.create_user", _create_user)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes._serialize_user",
|
||||
lambda _db, user: {"id": user.id},
|
||||
)
|
||||
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
ensure_json_body=lambda: {
|
||||
"username": "user3",
|
||||
"password": "Abcd12",
|
||||
"email": "u3@example.com",
|
||||
"role": "user",
|
||||
"initial_gift_usd": 10,
|
||||
"allowed_providers": [],
|
||||
"allowed_api_formats": [],
|
||||
"allowed_models": [],
|
||||
},
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
|
||||
result = await AdminCreateUserAdapter().handle(context)
|
||||
|
||||
assert result == {"id": "user-3"}
|
||||
assert captured["allowed_providers"] == []
|
||||
assert captured["allowed_api_formats"] == []
|
||||
assert captured["allowed_models"] == []
|
||||
|
||||
@@ -140,8 +140,14 @@ def test_sync_delete_reports_progress_after_each_batch(
|
||||
|
||||
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: cleanup_batches.append(list(batch)),
|
||||
)
|
||||
monkeypatch.setattr("src.database.create_session", lambda: session)
|
||||
monkeypatch.setattr(
|
||||
"src.models.database.ProviderAPIKey",
|
||||
@@ -157,5 +163,6 @@ def test_sync_delete_reports_progress_after_each_batch(
|
||||
|
||||
assert affected == 3
|
||||
assert progress_updates == [2, 3]
|
||||
assert cleanup_batches == [["key-1", "key-2"], ["key-3"]]
|
||||
assert session.commits == 2
|
||||
assert session.closed is True
|
||||
|
||||
109
tests/services/test_provider_auth_detached.py
Normal file
109
tests/services/test_provider_auth_detached.py
Normal file
@@ -0,0 +1,109 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import types
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.provider import auth as module
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def __init__(self, row: Any | None) -> None:
|
||||
self._row = row
|
||||
|
||||
def filter(self, *_args: Any, **_kwargs: Any) -> "_FakeQuery":
|
||||
return self
|
||||
|
||||
def first(self) -> Any | None:
|
||||
return self._row
|
||||
|
||||
|
||||
class _FakeDB:
|
||||
def __init__(self, row: Any | None) -> None:
|
||||
self.row = row
|
||||
self.committed = False
|
||||
|
||||
def query(self, _model: Any) -> _FakeQuery:
|
||||
return _FakeQuery(self.row)
|
||||
|
||||
def commit(self) -> None:
|
||||
self.committed = True
|
||||
|
||||
|
||||
class _FakeSessionCtx:
|
||||
def __init__(self, db: _FakeDB) -> None:
|
||||
self.db = db
|
||||
|
||||
def __enter__(self) -> _FakeDB:
|
||||
return self.db
|
||||
|
||||
def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> bool:
|
||||
_ = exc_type, exc, tb
|
||||
return False
|
||||
|
||||
|
||||
def _install_module(monkeypatch: pytest.MonkeyPatch, name: str, attrs: dict[str, Any]) -> None:
|
||||
fake_module = types.ModuleType(name)
|
||||
for key, value in attrs.items():
|
||||
setattr(fake_module, key, value)
|
||||
monkeypatch.setitem(sys.modules, name, fake_module)
|
||||
|
||||
|
||||
def test_persist_refreshed_token_detached_key_does_not_raise(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
key = SimpleNamespace(
|
||||
id="key-1",
|
||||
api_key="old-api",
|
||||
auth_config="old-config",
|
||||
oauth_invalid_at=datetime.now(timezone.utc),
|
||||
oauth_invalid_reason="[REFRESH_FAILED] stale",
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
module, "object_session", lambda _key: (_ for _ in ()).throw(RuntimeError())
|
||||
)
|
||||
monkeypatch.setattr(module.crypto_service, "encrypt", lambda value: f"enc:{value}")
|
||||
|
||||
module._persist_refreshed_token(key, "new-token", {"refresh_token": "rt-2"})
|
||||
|
||||
assert key.api_key == "enc:new-token"
|
||||
assert key.auth_config == 'enc:{"refresh_token": "rt-2"}'
|
||||
assert key.oauth_invalid_at is None
|
||||
assert key.oauth_invalid_reason is None
|
||||
|
||||
|
||||
def test_mark_refresh_token_invalid_persists_detached_key(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
key = SimpleNamespace(id="key-1")
|
||||
row = SimpleNamespace(id="key-1", oauth_invalid_at=None, oauth_invalid_reason=None)
|
||||
fake_db = _FakeDB(row)
|
||||
|
||||
monkeypatch.setattr(
|
||||
module, "object_session", lambda _key: (_ for _ in ()).throw(RuntimeError())
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch, "src.database", {"create_session": lambda: _FakeSessionCtx(fake_db)}
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.models.database",
|
||||
{"ProviderAPIKey": type("ProviderAPIKey", (), {"id": "id"})},
|
||||
)
|
||||
|
||||
module._mark_refresh_token_invalid(
|
||||
key,
|
||||
401,
|
||||
'{"error": {"code": "refresh_token_reused", "message": "used"}}',
|
||||
)
|
||||
|
||||
assert fake_db.committed is True
|
||||
assert key.oauth_invalid_at is not None
|
||||
assert row.oauth_invalid_at is not None
|
||||
assert str(key.oauth_invalid_reason).startswith("[REFRESH_FAILED] Token 续期失败 (401)")
|
||||
assert "refresh_token_reused" in str(row.oauth_invalid_reason)
|
||||
159
tests/services/test_provider_delete_cleanup.py
Normal file
159
tests/services/test_provider_delete_cleanup.py
Normal file
@@ -0,0 +1,159 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.provider import delete_cleanup as cleanup_module
|
||||
from src.services.provider.delete_cleanup import (
|
||||
delete_provider_tree,
|
||||
prune_allowed_provider_list,
|
||||
prune_allowed_provider_refs,
|
||||
)
|
||||
|
||||
|
||||
def _query_mock(
|
||||
*,
|
||||
rows: list[object] | None = None,
|
||||
update_count: int | None = None,
|
||||
delete_count: int | None = None,
|
||||
) -> MagicMock:
|
||||
query = MagicMock()
|
||||
filtered = query.filter.return_value
|
||||
if rows is not None:
|
||||
filtered.all.return_value = rows
|
||||
if update_count is not None:
|
||||
filtered.update.return_value = update_count
|
||||
if delete_count is not None:
|
||||
filtered.delete.return_value = delete_count
|
||||
return query
|
||||
|
||||
|
||||
def test_prune_allowed_provider_list_removes_target_and_normalizes_empty() -> None:
|
||||
next_allowed, changed = prune_allowed_provider_list(["provider-a", "provider-b"], "provider-a")
|
||||
|
||||
assert changed is True
|
||||
assert next_allowed == ["provider-b"]
|
||||
|
||||
next_allowed, changed = prune_allowed_provider_list(["provider-a"], "provider-a")
|
||||
|
||||
assert changed is True
|
||||
assert next_allowed == []
|
||||
|
||||
|
||||
def test_prune_allowed_provider_refs_updates_matching_records_only() -> None:
|
||||
records = [
|
||||
SimpleNamespace(allowed_providers=["provider-a", "provider-b"]),
|
||||
SimpleNamespace(allowed_providers=["provider-b"]),
|
||||
SimpleNamespace(allowed_providers=None),
|
||||
]
|
||||
|
||||
updated = prune_allowed_provider_refs(records, "provider-a")
|
||||
|
||||
assert updated == 1
|
||||
assert records[0].allowed_providers == ["provider-b"]
|
||||
assert records[1].allowed_providers == ["provider-b"]
|
||||
assert records[2].allowed_providers is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleanup_deleted_provider_references_cleans_large_fanout_tables(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
|
||||
user_records = [SimpleNamespace(allowed_providers=["provider-a", "provider-b"])]
|
||||
api_key_records = [SimpleNamespace(allowed_providers=["provider-a"])]
|
||||
key_cleanup_calls: list[list[str]] = []
|
||||
|
||||
monkeypatch.setattr(
|
||||
cleanup_module,
|
||||
"cleanup_key_references",
|
||||
lambda _db, key_ids: key_cleanup_calls.append(list(key_ids)),
|
||||
)
|
||||
|
||||
db.query.side_effect = [
|
||||
_query_mock(rows=[("endpoint-1",), ("endpoint-2",)]),
|
||||
_query_mock(rows=[("key-1",), ("key-2",)]),
|
||||
_query_mock(rows=user_records),
|
||||
_query_mock(rows=api_key_records),
|
||||
_query_mock(update_count=2),
|
||||
_query_mock(update_count=3),
|
||||
_query_mock(update_count=4),
|
||||
_query_mock(update_count=5),
|
||||
_query_mock(update_count=6),
|
||||
_query_mock(delete_count=7),
|
||||
_query_mock(delete_count=8),
|
||||
]
|
||||
|
||||
stats = cleanup_module.cleanup_deleted_provider_references(db, "provider-a")
|
||||
|
||||
assert stats == {
|
||||
"users": 1,
|
||||
"api_keys": 1,
|
||||
"user_preferences": 2,
|
||||
"usage_provider": 3,
|
||||
"usage_endpoint": 5,
|
||||
"video_tasks_provider": 4,
|
||||
"video_tasks_endpoint": 6,
|
||||
"request_candidates_provider": 8,
|
||||
"request_candidates_endpoint": 7,
|
||||
}
|
||||
assert user_records[0].allowed_providers == ["provider-b"]
|
||||
assert api_key_records[0].allowed_providers == []
|
||||
assert key_cleanup_calls == [["key-1", "key-2"]]
|
||||
|
||||
|
||||
def test_delete_provider_tree_deletes_children_before_provider(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
|
||||
cleanup_mock = MagicMock(
|
||||
return_value={
|
||||
"users": 1,
|
||||
"api_keys": 2,
|
||||
"user_preferences": 3,
|
||||
"usage_provider": 4,
|
||||
"usage_endpoint": 5,
|
||||
"video_tasks_provider": 6,
|
||||
"video_tasks_endpoint": 7,
|
||||
"request_candidates_provider": 8,
|
||||
"request_candidates_endpoint": 9,
|
||||
}
|
||||
)
|
||||
monkeypatch.setattr(cleanup_module, "cleanup_deleted_provider_references", cleanup_mock)
|
||||
|
||||
db.query.side_effect = [
|
||||
_query_mock(rows=[("endpoint-1",), ("endpoint-2",)]),
|
||||
_query_mock(rows=[("key-1",), ("key-2",), ("key-3",)]),
|
||||
_query_mock(delete_count=10),
|
||||
_query_mock(delete_count=11),
|
||||
_query_mock(delete_count=12),
|
||||
_query_mock(delete_count=13),
|
||||
_query_mock(delete_count=14),
|
||||
_query_mock(delete_count=1),
|
||||
]
|
||||
|
||||
stats = delete_provider_tree(db, "provider-a")
|
||||
|
||||
cleanup_mock.assert_called_once_with(
|
||||
db,
|
||||
"provider-a",
|
||||
endpoint_ids=["endpoint-1", "endpoint-2"],
|
||||
key_ids=["key-1", "key-2", "key-3"],
|
||||
)
|
||||
assert stats == {
|
||||
"cleanup": cleanup_mock.return_value,
|
||||
"deleted": {
|
||||
"api_key_mappings": 10,
|
||||
"usage_tracking": 11,
|
||||
"models": 12,
|
||||
"api_keys": 13,
|
||||
"endpoints": 14,
|
||||
"providers": 1,
|
||||
},
|
||||
"key_count": 3,
|
||||
"endpoint_count": 2,
|
||||
}
|
||||
151
tests/services/test_provider_delete_task.py
Normal file
151
tests/services/test_provider_delete_task.py
Normal file
@@ -0,0 +1,151 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Coroutine
|
||||
from concurrent.futures import Future
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.provider import delete_task as taskmod
|
||||
|
||||
|
||||
class _FakeRedis:
|
||||
def __init__(self) -> None:
|
||||
self.store: dict[str, str] = {}
|
||||
|
||||
async def get(self, key: str) -> str | None:
|
||||
return self.store.get(key)
|
||||
|
||||
async def setex(self, key: str, _ttl: int, value: str) -> None:
|
||||
self.store[key] = value
|
||||
|
||||
|
||||
async def _fake_get_redis_client(*, require_redis: bool = False) -> _FakeRedis:
|
||||
_ = require_redis
|
||||
return _FakeRedis()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_provider_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(
|
||||
provider_id: str,
|
||||
progress_callback: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
_ = provider_id
|
||||
assert callable(progress_callback)
|
||||
progress_callback(
|
||||
{
|
||||
"stage": "deleting_keys",
|
||||
"total_keys": 4,
|
||||
"deleted_keys": 2,
|
||||
"message": "deleted key batch 1/2",
|
||||
}
|
||||
)
|
||||
return {
|
||||
"total_keys": 4,
|
||||
"deleted_keys": 4,
|
||||
"total_endpoints": 2,
|
||||
"deleted_endpoints": 2,
|
||||
"elapsed_seconds": 1.2,
|
||||
}
|
||||
|
||||
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, "_update_task_field", fake_update_task_field)
|
||||
monkeypatch.setattr(taskmod, "_sync_delete_provider", fake_sync_delete_provider)
|
||||
monkeypatch.setattr(taskmod.asyncio, "run_coroutine_threadsafe", fake_run_coroutine_threadsafe)
|
||||
monkeypatch.setattr(taskmod, "get_redis_client", _fake_get_redis_client)
|
||||
monkeypatch.setattr(taskmod, "invalidate_models_list_cache", AsyncMock())
|
||||
monkeypatch.setattr(taskmod.ModelCacheService, "invalidate_all_resolve_cache", AsyncMock())
|
||||
monkeypatch.setattr(taskmod.ProviderCacheService, "invalidate_provider_cache", AsyncMock())
|
||||
|
||||
task = asyncio.create_task(taskmod._run_provider_delete("task-1", "provider-1"))
|
||||
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
assert not task.done()
|
||||
assert updates == [
|
||||
{"status": taskmod.STATUS_RUNNING, "stage": "queued", "message": "delete task started"}
|
||||
]
|
||||
|
||||
release_progress.set()
|
||||
await task
|
||||
|
||||
assert updates == [
|
||||
{"status": taskmod.STATUS_RUNNING, "stage": "queued", "message": "delete task started"},
|
||||
{
|
||||
"stage": "deleting_keys",
|
||||
"total_keys": 4,
|
||||
"deleted_keys": 2,
|
||||
"message": "deleted key batch 1/2",
|
||||
},
|
||||
{
|
||||
"status": taskmod.STATUS_COMPLETED,
|
||||
"stage": "completed",
|
||||
"total_keys": 4,
|
||||
"deleted_keys": 4,
|
||||
"total_endpoints": 2,
|
||||
"deleted_endpoints": 2,
|
||||
"message": "provider deleted: keys=4, endpoints=2",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_submit_provider_delete_reuses_running_task(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
redis = _FakeRedis()
|
||||
started = asyncio.Event()
|
||||
finish = asyncio.Event()
|
||||
|
||||
async def fake_get_redis(*, require_redis: bool = False) -> _FakeRedis:
|
||||
_ = require_redis
|
||||
return redis
|
||||
|
||||
async def fake_run_provider_delete(task_id: str, provider_id: str) -> None:
|
||||
_ = task_id, provider_id
|
||||
started.set()
|
||||
await finish.wait()
|
||||
|
||||
monkeypatch.setattr(taskmod, "get_redis_client", fake_get_redis)
|
||||
monkeypatch.setattr(taskmod, "_run_provider_delete", fake_run_provider_delete)
|
||||
|
||||
task_id_1 = await taskmod.submit_provider_delete("provider-1")
|
||||
await started.wait()
|
||||
task_id_2 = await taskmod.submit_provider_delete("provider-1")
|
||||
|
||||
assert task_id_2 == task_id_1
|
||||
|
||||
finish.set()
|
||||
await asyncio.sleep(0)
|
||||
@@ -57,6 +57,18 @@ class _FakeClearOAuthDB:
|
||||
self.commit_count += 1
|
||||
|
||||
|
||||
class _FakeQueryAll:
|
||||
def __init__(self, rows: list[Any]) -> None:
|
||||
self._rows = rows
|
||||
|
||||
def filter(self, *args: Any, **kwargs: Any) -> _FakeQueryAll:
|
||||
_ = args, kwargs
|
||||
return self
|
||||
|
||||
def all(self) -> list[Any]:
|
||||
return self._rows
|
||||
|
||||
|
||||
def _build_key(**overrides: Any) -> SimpleNamespace:
|
||||
base: dict[str, Any] = {
|
||||
"auto_fetch_models": False,
|
||||
@@ -228,3 +240,121 @@ async def test_run_delete_key_side_effects_skip_disassociate(
|
||||
|
||||
assert captured["provider_id"] == "provider-1"
|
||||
assert captured["skip_disassociate"] is True
|
||||
|
||||
|
||||
def test_cleanup_key_references_preserves_usage_and_video_tasks(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class _FakeDeleteStatement:
|
||||
def __init__(self, model: Any) -> None:
|
||||
self.kind = "delete"
|
||||
self.model = model
|
||||
self.conditions: list[Any] = []
|
||||
|
||||
def where(self, *conditions: Any) -> _FakeDeleteStatement:
|
||||
self.conditions.extend(conditions)
|
||||
return self
|
||||
|
||||
class _FakeUpdateStatement:
|
||||
def __init__(self, model: Any) -> None:
|
||||
self.kind = "update"
|
||||
self.model = model
|
||||
self.conditions: list[Any] = []
|
||||
self.values_dict: dict[str, Any] = {}
|
||||
|
||||
def where(self, *conditions: Any) -> _FakeUpdateStatement:
|
||||
self.conditions.extend(conditions)
|
||||
return self
|
||||
|
||||
def values(self, **values: Any) -> _FakeUpdateStatement:
|
||||
self.values_dict.update(values)
|
||||
return self
|
||||
|
||||
class _FakeDB:
|
||||
def __init__(self) -> None:
|
||||
self.statements: list[Any] = []
|
||||
|
||||
def execute(self, statement: Any) -> None:
|
||||
self.statements.append(statement)
|
||||
|
||||
def get_bind(self) -> Any:
|
||||
return SimpleNamespace(dialect=SimpleNamespace(name="postgresql"))
|
||||
|
||||
monkeypatch.setattr(side_effects_module, "sa_delete", lambda model: _FakeDeleteStatement(model))
|
||||
monkeypatch.setattr(side_effects_module, "sa_update", lambda model: _FakeUpdateStatement(model))
|
||||
|
||||
db = _FakeDB()
|
||||
side_effects_module.cleanup_key_references(cast(Any, db), ["key-1", "key-2"])
|
||||
|
||||
assert [(stmt.kind, stmt.model.__name__) for stmt in db.statements] == [
|
||||
("delete", "RequestCandidate"),
|
||||
("delete", "GeminiFileMapping"),
|
||||
("update", "Usage"),
|
||||
("update", "VideoTask"),
|
||||
]
|
||||
assert db.statements[2].values_dict == {"provider_api_key_id": None}
|
||||
assert db.statements[3].values_dict == {"key_id": None}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_delete_endpoint_keys_response_cleans_related_references(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
cleanup_calls: list[list[str]] = []
|
||||
side_effect_calls: list[str | None] = []
|
||||
|
||||
class _FakeBatchDeleteDB:
|
||||
def __init__(self, keys: list[Any]) -> None:
|
||||
self._keys = keys
|
||||
self.commit_count = 0
|
||||
self.executed: list[Any] = []
|
||||
|
||||
def query(self, model: Any) -> _FakeQueryAll:
|
||||
_ = model
|
||||
return _FakeQueryAll(self._keys)
|
||||
|
||||
def execute(self, statement: Any) -> None:
|
||||
self.executed.append(statement)
|
||||
|
||||
def commit(self) -> None:
|
||||
self.commit_count += 1
|
||||
|
||||
def rollback(self) -> None:
|
||||
raise AssertionError("rollback should not be called")
|
||||
|
||||
async def _fake_run_delete_key_side_effects(
|
||||
db: Any,
|
||||
provider_id: str | None,
|
||||
deleted_key_allowed_models: list[str] | None,
|
||||
) -> None:
|
||||
_ = db, deleted_key_allowed_models
|
||||
side_effect_calls.append(provider_id)
|
||||
|
||||
monkeypatch.setattr(
|
||||
command_module,
|
||||
"cleanup_key_references",
|
||||
lambda _db, key_ids: cleanup_calls.append(list(key_ids)),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
command_module,
|
||||
"run_delete_key_side_effects",
|
||||
_fake_run_delete_key_side_effects,
|
||||
)
|
||||
|
||||
keys = [
|
||||
SimpleNamespace(id="key-1", provider_id="provider-1"),
|
||||
SimpleNamespace(id="key-2", provider_id="provider-1"),
|
||||
]
|
||||
db = _FakeBatchDeleteDB(keys)
|
||||
|
||||
result = await command_module.batch_delete_endpoint_keys_response(
|
||||
cast(Any, db),
|
||||
["key-1", "key-2"],
|
||||
)
|
||||
|
||||
assert result["success_count"] == 2
|
||||
assert result["failed_count"] == 0
|
||||
assert cleanup_calls == [["key-1", "key-2"]] or cleanup_calls == [["key-2", "key-1"]]
|
||||
assert side_effect_calls == ["provider-1"]
|
||||
assert db.commit_count == 1
|
||||
assert len(db.executed) == 1
|
||||
|
||||
54
tests/unit/test_password_policy.py
Normal file
54
tests/unit/test_password_policy.py
Normal file
@@ -0,0 +1,54 @@
|
||||
import pytest
|
||||
|
||||
from src.core.validators import PasswordPolicyLevel, PasswordValidator
|
||||
|
||||
|
||||
class TestPasswordPolicy:
|
||||
def test_default_policy_is_weak(self) -> None:
|
||||
valid, error = PasswordValidator.validate("123456")
|
||||
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("password", "expected_error"),
|
||||
[
|
||||
("1234567", "密码长度至少为8个字符"),
|
||||
("12345678", "密码必须包含至少一个字母"),
|
||||
("abcdefgh", "密码必须包含至少一个数字"),
|
||||
],
|
||||
)
|
||||
def test_medium_policy_rejects_weak_patterns(self, password: str, expected_error: str) -> None:
|
||||
valid, error = PasswordValidator.validate(password, policy=PasswordPolicyLevel.MEDIUM)
|
||||
|
||||
assert valid is False
|
||||
assert error == expected_error
|
||||
|
||||
def test_medium_policy_accepts_letters_and_digits(self) -> None:
|
||||
valid, error = PasswordValidator.validate("abc12345", policy=PasswordPolicyLevel.MEDIUM)
|
||||
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("password", "expected_error"),
|
||||
[
|
||||
("abc12345", "密码必须包含至少一个大写字母"),
|
||||
("ABC12345", "密码必须包含至少一个小写字母"),
|
||||
("Abcdefgh", "密码必须包含至少一个数字"),
|
||||
("Abcd1234", "密码必须包含至少一个特殊字符"),
|
||||
],
|
||||
)
|
||||
def test_strong_policy_requires_upper_lower_and_digit(
|
||||
self, password: str, expected_error: str
|
||||
) -> None:
|
||||
valid, error = PasswordValidator.validate(password, policy=PasswordPolicyLevel.STRONG)
|
||||
|
||||
assert valid is False
|
||||
assert error == expected_error
|
||||
|
||||
def test_strong_policy_accepts_mixed_password(self) -> None:
|
||||
valid, error = PasswordValidator.validate("Abcd1234!", policy="strong")
|
||||
|
||||
assert valid is True
|
||||
assert error is None
|
||||
Reference in New Issue
Block a user