Merge pull request #217 from AAEE86/1233

perf: 优化请求鉴权链路并批量化统计/调度查询
This commit is contained in:
fawney19
2026-03-10 23:31:28 +08:00
committed by GitHub
44 changed files with 2745 additions and 571 deletions

View File

@@ -0,0 +1,143 @@
from __future__ import annotations
from datetime import datetime, timezone
from decimal import Decimal
from types import SimpleNamespace
from typing import Any
import pytest
from src.api.admin.usage.routes import AdminUsageRecordsAdapter
class _FakeQuery:
def __init__(
self,
*,
scalar_result: int | None = None,
all_result: list[Any] | None = None,
) -> None:
self.scalar_result = scalar_result
self.all_result = all_result or []
self.options_args: tuple[Any, ...] = ()
def outerjoin(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def join(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def options(self, *args: Any) -> _FakeQuery:
self.options_args = args
return self
def order_by(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def offset(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def limit(self, *args: Any, **kwargs: Any) -> _FakeQuery:
return self
def scalar(self) -> int | None:
return self.scalar_result
def all(self) -> list[Any]:
return self.all_result
class _FakeDb:
def __init__(self, queries: list[_FakeQuery]) -> None:
self._queries = queries
self.query_calls: list[tuple[Any, ...]] = []
def query(self, *args: Any) -> _FakeQuery:
self.query_calls.append(args)
return self._queries[len(self.query_calls) - 1]
@pytest.mark.asyncio
async def test_admin_usage_records_returns_model_version_without_request_metadata(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr("src.utils.cache_decorator.get_redis_client_sync", lambda: None)
usage = SimpleNamespace(
id="usage-1",
request_id=None,
user_id="user-1",
api_key_id=None,
provider_name="google",
provider_id=None,
provider_endpoint_id=None,
provider_api_key_id=None,
model="gemini-2.5-pro",
target_model=None,
input_tokens=120,
output_tokens=80,
cache_creation_input_tokens=0,
cache_read_input_tokens=0,
total_tokens=200,
total_cost_usd=Decimal("1.25"),
actual_total_cost_usd=Decimal("1.25"),
rate_multiplier=Decimal("1.0"),
response_time_ms=850,
first_byte_time_ms=230,
created_at=datetime(2026, 3, 9, 8, 30, tzinfo=timezone.utc),
is_stream=False,
status_code=200,
error_message=None,
status="completed",
api_format="gemini:chat",
endpoint_api_format=None,
has_format_conversion=False,
input_price_per_1m=Decimal("0.10"),
output_price_per_1m=Decimal("0.30"),
cache_creation_price_per_1m=None,
cache_read_price_per_1m=None,
)
user = SimpleNamespace(id="user-1", email="user@example.com", username="tester")
count_query = _FakeQuery(scalar_result=1)
data_query = _FakeQuery(
all_result=[
(usage, user, None, None, None, "gemini-2.5-pro-001"),
]
)
db = _FakeDb([count_query, data_query])
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id="admin-1"),
add_audit_metadata=lambda **_: None,
)
adapter = AdminUsageRecordsAdapter(
time_range=None,
search=None,
user_id=None,
username=None,
model=None,
provider=None,
api_format=None,
status=None,
limit=100,
offset=0,
)
result = await adapter.handle(context)
assert len(db.query_calls) == 2
assert len(db.query_calls[1]) == 6
assert getattr(db.query_calls[1][-1], "name", None) == "model_version"
record = result["records"][0]
assert record["model_version"] == "gemini-2.5-pro-001"
assert "request_metadata" not in record
usage_load_only = data_query.options_args[0]
usage_paths = {str(option.path) for option in usage_load_only.context}
assert "ORM Path[Mapper[Usage(usage)] -> Usage.request_metadata]" not in usage_paths

View File

@@ -0,0 +1,93 @@
from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Any
from unittest.mock import MagicMock
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from src.api.admin.users.routes import router as admin_users_router
from src.database import get_db
def _build_admin_users_app(db: MagicMock, monkeypatch: pytest.MonkeyPatch) -> TestClient:
app = FastAPI()
app.include_router(admin_users_router)
app.dependency_overrides[get_db] = lambda: db
async def _fake_pipeline_run(
*, adapter: Any, http_request: object, db: MagicMock, mode: object
) -> Any:
_ = http_request, mode
context = SimpleNamespace(
db=db,
request=SimpleNamespace(state=SimpleNamespace()),
user=SimpleNamespace(id="admin-1"),
ensure_json_body=lambda: {},
add_audit_metadata=lambda **_: None,
)
return await adapter.handle(context)
monkeypatch.setattr("src.api.admin.users.routes.pipeline.run", _fake_pipeline_run)
return TestClient(app)
def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
client = _build_admin_users_app(db, monkeypatch)
now = datetime.now(timezone.utc)
users = [
SimpleNamespace(
id="user-1",
email="u1@example.com",
username="user1",
role=SimpleNamespace(value="user"),
allowed_providers=None,
allowed_api_formats=None,
allowed_models=None,
is_active=True,
created_at=now,
updated_at=now,
last_login_at=None,
),
SimpleNamespace(
id="user-2",
email="u2@example.com",
username="user2",
role=SimpleNamespace(value="admin"),
allowed_providers=None,
allowed_api_formats=None,
allowed_models=None,
is_active=True,
created_at=now,
updated_at=None,
last_login_at=None,
),
]
wallets_by_user_id = {
"user-1": SimpleNamespace(limit_mode="unlimited"),
}
batch_getter = MagicMock(return_value=wallets_by_user_id)
monkeypatch.setattr(
"src.api.admin.users.routes.UserService.list_users", lambda *_a, **_k: users
)
monkeypatch.setattr(
"src.api.admin.users.routes.WalletService.get_wallets_by_user_ids",
batch_getter,
)
monkeypatch.setattr(
"src.api.admin.users.routes.WalletService.get_wallet",
lambda *_a, **_k: (_ for _ in ()).throw(AssertionError("不应回退到逐个钱包查询")),
)
response = client.get("/api/admin/users")
assert response.status_code == 200
assert response.json()[0]["unlimited"] is True
assert response.json()[1]["unlimited"] is False
batch_getter.assert_called_once()
assert batch_getter.call_args.args[1] == ["user-1", "user-2"]

View File

@@ -8,54 +8,139 @@ API Pipeline 测试
"""
from datetime import datetime, timezone
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from src.api.base.adapter import ApiMode
from src.api.base.pipeline import ApiRequestPipeline
from src.core.enums import UserRole
from src.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS
class TestPipelineBalanceCalculation:
"""测试 Pipeline 余额计算"""
"""Balance calculation tests for Pipeline."""
@pytest.fixture
def pipeline(self) -> ApiRequestPipeline:
return ApiRequestPipeline()
def test_calculate_balance_remaining_with_balance(self, pipeline: ApiRequestPipeline) -> None:
"""测试有限制钱包时计算剩余余额"""
@pytest.mark.asyncio
async def test_calculate_balance_remaining_with_balance(
self, pipeline: ApiRequestPipeline
) -> None:
"""Returns remaining balance for limited wallets."""
mock_user = MagicMock()
mock_db = MagicMock()
mock_user.id = "user-123"
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=70.0,
):
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
thread_db = MagicMock()
db_user = MagicMock()
thread_db.query.return_value.filter.return_value.first.return_value = db_user
with patch("src.api.base.pipeline.create_session", return_value=thread_db):
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=70.0,
):
remaining = await pipeline._calculate_balance_remaining_async(mock_user)
assert remaining == 70.0
thread_db.close.assert_called_once()
def test_calculate_balance_remaining_unlimited(self, pipeline: ApiRequestPipeline) -> None:
"""测试无限制钱包时返回 None"""
@pytest.mark.asyncio
async def test_calculate_balance_remaining_unlimited(
self, pipeline: ApiRequestPipeline
) -> None:
"""Returns None for unlimited wallets."""
mock_user = MagicMock()
mock_db = MagicMock()
mock_user.id = "user-123"
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=None,
thread_db = MagicMock()
db_user = MagicMock()
thread_db.query.return_value.filter.return_value.first.return_value = db_user
with patch("src.api.base.pipeline.create_session", return_value=thread_db):
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=None,
):
remaining = await pipeline._calculate_balance_remaining_async(mock_user)
assert remaining is None
thread_db.close.assert_called_once()
@pytest.mark.asyncio
async def test_calculate_balance_remaining_none_user(
self, pipeline: ApiRequestPipeline
) -> None:
"""Returns None when user is missing."""
remaining = await pipeline._calculate_balance_remaining_async(None)
assert remaining is None
class TestPipelineRunModes:
"""Returns remaining balance for limited wallets."""
@pytest.fixture
def pipeline(self) -> ApiRequestPipeline:
return ApiRequestPipeline()
@pytest.mark.asyncio
async def test_run_management_mode_skips_balance_calculation(
self, pipeline: ApiRequestPipeline
) -> None:
"""Management mode skips balance calculation."""
mock_request = MagicMock()
mock_request.method = "GET"
mock_request.url.path = "/api/admin/tokens"
mock_request.state = MagicMock()
mock_db = MagicMock()
mock_user = MagicMock()
mock_user.id = "admin-123"
mock_token = MagicMock()
mock_token.id = "mt-123"
mock_adapter = MagicMock()
mock_adapter.name = "test-adapter"
mock_adapter.authorize = MagicMock(return_value=None)
mock_response = MagicMock()
mock_response.status_code = 200
mock_adapter.handle = AsyncMock(return_value=mock_response)
mock_context = MagicMock()
mock_context.db = mock_db
mock_context.request = mock_request
with patch.object(
pipeline,
"_authenticate_management",
new_callable=AsyncMock,
return_value=(mock_user, mock_token),
):
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
with patch(
"src.api.base.pipeline.ApiRequestContext.build",
return_value=mock_context,
):
with patch.object(
pipeline,
"_calculate_balance_remaining_async",
new_callable=AsyncMock,
) as mock_balance:
with patch.object(pipeline, "_record_audit_event"):
response = await pipeline.run(
mock_adapter,
mock_request,
mock_db,
mode=ApiMode.MANAGEMENT,
)
assert remaining is None
def test_calculate_balance_remaining_none_user(self, pipeline: ApiRequestPipeline) -> None:
"""测试用户为 None 时返回 None"""
mock_db = MagicMock()
remaining = pipeline._calculate_balance_remaining(mock_db, None)
assert remaining is None
assert response == mock_response
assert mock_context.management_token == mock_token
mock_balance.assert_not_called()
class TestPipelineAuditLogging:
@@ -213,7 +298,8 @@ class TestPipelineAuthentication:
def pipeline(self) -> ApiRequestPipeline:
return ApiRequestPipeline()
def test_authenticate_client_missing_key(self, pipeline: ApiRequestPipeline) -> None:
@pytest.mark.asyncio
async def test_authenticate_client_missing_key(self, pipeline: ApiRequestPipeline) -> None:
"""测试缺少 API Key 时抛出异常"""
mock_request = MagicMock()
mock_request.headers = {}
@@ -226,12 +312,13 @@ class TestPipelineAuthentication:
mock_adapter.extract_api_key = MagicMock(return_value=None)
with pytest.raises(HTTPException) as exc_info:
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
assert exc_info.value.status_code == 401
assert "API密钥" in exc_info.value.detail
def test_authenticate_client_invalid_key(self, pipeline: ApiRequestPipeline) -> None:
@pytest.mark.asyncio
async def test_authenticate_client_invalid_key(self, pipeline: ApiRequestPipeline) -> None:
"""测试无效的 API Key"""
mock_request = MagicMock()
mock_request.headers = {"Authorization": "Bearer sk-invalid"}
@@ -245,15 +332,17 @@ class TestPipelineAuthentication:
with patch.object(
pipeline.auth_service,
"authenticate_api_key",
"authenticate_api_key_threadsafe",
new_callable=AsyncMock,
return_value=None,
):
with pytest.raises(HTTPException) as exc_info:
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
assert exc_info.value.status_code == 401
def test_authenticate_client_balance_exceeded(self, pipeline: ApiRequestPipeline) -> None:
@pytest.mark.asyncio
async def test_authenticate_client_balance_exceeded(self, pipeline: ApiRequestPipeline) -> None:
"""测试余额不足时抛出异常"""
mock_user = MagicMock()
mock_user.id = "user-123"
@@ -268,28 +357,198 @@ class TestPipelineAuthentication:
mock_request.state = MagicMock()
mock_db = MagicMock()
db_user = MagicMock()
db_user.id = "user-123"
db_user.is_active = True
db_user.is_deleted = False
db_api_key = MagicMock()
db_api_key.id = "key-123"
db_api_key.user_id = "user-123"
db_api_key.is_active = True
db_api_key.is_locked = False
db_api_key.is_standalone = False
db_api_key.expires_at = None
user_query = MagicMock()
user_query.filter.return_value.first.return_value = db_user
api_key_query = MagicMock()
api_key_query.filter.return_value.first.return_value = db_api_key
mock_db.query.side_effect = [user_query, api_key_query]
mock_adapter = MagicMock()
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
with patch.object(
pipeline.auth_service,
"authenticate_api_key",
return_value=(mock_user, mock_api_key),
"authenticate_api_key_threadsafe",
new_callable=AsyncMock,
return_value=MagicMock(
user=mock_user,
api_key=mock_api_key,
access_allowed=False,
balance_remaining=0.0,
),
):
with patch.object(
pipeline.usage_service,
"check_request_balance",
return_value=(False, "余额不足"),
):
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=0.0,
):
from src.core.exceptions import BalanceInsufficientException
from src.core.exceptions import BalanceInsufficientException
with pytest.raises(BalanceInsufficientException):
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
with pytest.raises(BalanceInsufficientException):
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
@pytest.mark.asyncio
async def test_authenticate_client_requery_detects_inactive_user(
self, pipeline: ApiRequestPipeline
) -> None:
mock_user = MagicMock()
mock_user.id = "user-123"
mock_api_key = MagicMock()
mock_api_key.id = "key-123"
mock_request = MagicMock()
mock_request.headers = {"Authorization": "Bearer sk-test"}
mock_request.url.path = "/v1/messages"
mock_request.state = MagicMock()
mock_adapter = MagicMock()
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
db_user = MagicMock()
db_user.id = "user-123"
db_user.is_active = False
db_user.is_deleted = False
db_api_key = MagicMock()
db_api_key.id = "key-123"
db_api_key.user_id = "user-123"
db_api_key.is_active = True
db_api_key.is_locked = False
db_api_key.is_standalone = False
db_api_key.expires_at = None
mock_db = MagicMock()
user_query = MagicMock()
user_query.filter.return_value.first.return_value = db_user
api_key_query = MagicMock()
api_key_query.filter.return_value.first.return_value = db_api_key
mock_db.query.side_effect = [user_query, api_key_query]
with patch.object(
pipeline.auth_service,
"authenticate_api_key_threadsafe",
new_callable=AsyncMock,
return_value=MagicMock(
user=mock_user,
api_key=mock_api_key,
access_allowed=True,
balance_remaining=10.0,
),
):
with pytest.raises(HTTPException) as exc_info:
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
assert exc_info.value.status_code == 401
@pytest.mark.asyncio
async def test_authenticate_client_requery_detects_locked_key(
self, pipeline: ApiRequestPipeline
) -> None:
mock_user = MagicMock()
mock_user.id = "user-123"
mock_api_key = MagicMock()
mock_api_key.id = "key-123"
mock_request = MagicMock()
mock_request.headers = {"Authorization": "Bearer sk-test"}
mock_request.url.path = "/v1/messages"
mock_request.state = MagicMock()
mock_adapter = MagicMock()
mock_adapter.extract_api_key = MagicMock(return_value="sk-test")
db_user = MagicMock()
db_user.id = "user-123"
db_user.is_active = True
db_user.is_deleted = False
db_api_key = MagicMock()
db_api_key.id = "key-123"
db_api_key.user_id = "user-123"
db_api_key.is_active = True
db_api_key.is_locked = True
db_api_key.is_standalone = False
db_api_key.expires_at = None
mock_db = MagicMock()
user_query = MagicMock()
user_query.filter.return_value.first.return_value = db_user
api_key_query = MagicMock()
api_key_query.filter.return_value.first.return_value = db_api_key
mock_db.query.side_effect = [user_query, api_key_query]
with patch.object(
pipeline.auth_service,
"authenticate_api_key_threadsafe",
new_callable=AsyncMock,
return_value=MagicMock(
user=mock_user,
api_key=mock_api_key,
access_allowed=True,
balance_remaining=10.0,
),
):
with pytest.raises(HTTPException) as exc_info:
await pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
assert exc_info.value.status_code == 403
assert "锁定" in str(exc_info.value.detail)
class TestPipelineTokenPrefixAuth:
"""Tests token-prefix auth isolation."""
@pytest.fixture
def pipeline(self) -> ApiRequestPipeline:
return ApiRequestPipeline()
@pytest.mark.asyncio
async def test_try_token_prefix_auth_uses_isolated_session(
self, pipeline: ApiRequestPipeline
) -> None:
mock_request = MagicMock()
mock_request.headers = {}
mock_request.client = MagicMock(host="127.0.0.1")
route_db = MagicMock()
auth_db = MagicMock()
mock_user = MagicMock()
mock_token = MagicMock()
async def authenticate(db: Any, token: str, client_ip: str) -> tuple[Any, Any]:
assert db is auth_db
assert token == "ae_test"
assert client_ip == "127.0.0.1"
return mock_user, mock_token
with patch("src.api.base.pipeline.create_session", return_value=auth_db):
with patch("src.utils.request_utils.get_client_ip", return_value="127.0.0.1"):
with patch("src.core.modules.hooks.get_hook_dispatcher") as mock_get_dispatcher:
dispatcher = MagicMock()
dispatcher.dispatch = AsyncMock(
return_value=[
{
"prefix": "ae_",
"module": "management_tokens",
"authenticate": authenticate,
}
]
)
mock_get_dispatcher.return_value = dispatcher
result = await pipeline._try_token_prefix_auth(
"ae_test", mock_request, route_db
)
assert result == (mock_user, mock_token)
dispatcher.dispatch.assert_awaited_once_with(AUTH_TOKEN_PREFIX_AUTHENTICATORS)
auth_db.expunge.assert_any_call(mock_user)
auth_db.expunge.assert_any_call(mock_token)
auth_db.close.assert_called_once()
class TestPipelineAdminAuth:

View File

@@ -0,0 +1,175 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from src.api.user_me.routes import GetUsageAdapter
from src.core.enums import UserRole
@pytest.mark.asyncio
async def test_get_usage_adapter_uses_coarse_summary_grouping(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
query = MagicMock()
count_query = MagicMock()
count_query.scalar.return_value = 0
query.outerjoin.return_value = query
query.filter.return_value = query
query.with_entities.return_value = count_query
query.options.return_value = query
query.order_by.return_value = query
query.offset.return_value = query
query.limit.return_value = query
query.all.return_value = []
db.query.return_value = query
summary_getter = MagicMock(
return_value=[
{
"provider": "provider-a",
"model": "gpt-4o",
"requests": 2,
"input_tokens": 10,
"output_tokens": 5,
"total_tokens": 15,
"total_cost_usd": 1.5,
"actual_total_cost_usd": 1.2,
"success_count": 2,
"success_response_time_sum_ms": 1000.0,
"success_response_time_count": 2,
},
{
"provider": "pending",
"model": "gpt-4o",
"requests": 99,
"input_tokens": 999,
"output_tokens": 999,
"total_tokens": 1998,
"total_cost_usd": 9.9,
"actual_total_cost_usd": 9.9,
"success_count": 0,
"success_response_time_sum_ms": 0.0,
"success_response_time_count": 0,
},
]
)
monkeypatch.setattr("src.api.user_me.routes.UsageService.get_usage_summary", summary_getter)
monkeypatch.setattr("src.api.user_me.routes.WalletService.get_wallet", lambda *_a, **_k: None)
monkeypatch.setattr(
"src.api.user_me.routes.WalletService.serialize_wallet_summary",
lambda _wallet: {"limit_mode": "finite"},
)
adapter = GetUsageAdapter(time_range=None, limit=20, offset=0)
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id="user-1", role=UserRole.USER),
request=SimpleNamespace(state=SimpleNamespace()),
)
result = await adapter.handle(context)
assert result["total_requests"] == 2
assert result["total_tokens"] == 15
assert result["summary_by_model"] == [
{
"model": "gpt-4o",
"requests": 2,
"input_tokens": 10,
"output_tokens": 5,
"total_tokens": 15,
"total_cost_usd": 1.5,
}
]
assert "total_actual_cost" not in result
assert result["summary_by_provider"] == [
{
"provider": "provider-a",
"requests": 2,
"total_tokens": 15,
"total_cost_usd": 1.5,
"success_rate": 100.0,
"avg_response_time_ms": 500.0,
}
]
assert summary_getter.call_args.kwargs["group_by"] is None
@pytest.mark.asyncio
async def test_get_usage_adapter_provider_success_rate_uses_success_count(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
query = MagicMock()
count_query = MagicMock()
count_query.scalar.return_value = 0
query.outerjoin.return_value = query
query.filter.return_value = query
query.with_entities.return_value = count_query
query.options.return_value = query
query.order_by.return_value = query
query.offset.return_value = query
query.limit.return_value = query
query.all.return_value = []
db.query.return_value = query
summary_getter = MagicMock(
return_value=[
{
"provider": "provider-a",
"model": "gpt-4o",
"requests": 3,
"input_tokens": 30,
"output_tokens": 15,
"total_tokens": 45,
"total_cost_usd": 4.5,
"actual_total_cost_usd": 4.5,
"success_count": 2,
"success_response_time_sum_ms": 600.0,
"success_response_time_count": 2,
},
{
"provider": "provider-a",
"model": "gpt-4.1",
"requests": 1,
"input_tokens": 10,
"output_tokens": 5,
"total_tokens": 15,
"total_cost_usd": 1.5,
"actual_total_cost_usd": 1.5,
"success_count": 0,
"success_response_time_sum_ms": 0.0,
"success_response_time_count": 0,
},
]
)
monkeypatch.setattr("src.api.user_me.routes.UsageService.get_usage_summary", summary_getter)
monkeypatch.setattr("src.api.user_me.routes.WalletService.get_wallet", lambda *_a, **_k: None)
monkeypatch.setattr(
"src.api.user_me.routes.WalletService.serialize_wallet_summary",
lambda _wallet: {"limit_mode": "finite"},
)
adapter = GetUsageAdapter(time_range=None, limit=20, offset=0)
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id="user-1", role=UserRole.USER),
request=SimpleNamespace(state=SimpleNamespace()),
)
result = await adapter.handle(context)
assert result["summary_by_provider"] == [
{
"provider": "provider-a",
"requests": 4,
"total_tokens": 60,
"total_cost_usd": 6.0,
"success_rate": 50.0,
"avg_response_time_ms": 300.0,
}
]

View File

@@ -69,26 +69,33 @@ async def test_list_all_candidates_returns_provider_batch_count_even_when_candid
global_model = _make_global_model(gid="gm1", name="gpt-4o")
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
with patch.object(scheduler._candidate_builder, "_query_providers", return_value=providers):
with patch(
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
new=AsyncMock(return_value=global_model),
):
with patch.object(
scheduler._candidate_builder,
"_query_provider_refs",
return_value=[("p1", "p1"), ("p2", "p2")],
):
with patch.object(scheduler._candidate_builder, "_query_providers") as query_providers:
with patch(
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
return_value=True,
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
new=AsyncMock(return_value=global_model),
):
candidates, global_model_id, provider_batch_count = (
await scheduler.list_all_candidates(
db=db,
api_format="openai:chat",
model_name="gpt-4o",
affinity_key=None,
user_api_key=user_api_key, # type: ignore[arg-type]
provider_offset=0,
provider_limit=20,
with patch(
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
return_value=True,
):
candidates, global_model_id, provider_batch_count = (
await scheduler.list_all_candidates(
db=db,
api_format="openai:chat",
model_name="gpt-4o",
affinity_key=None,
user_api_key=user_api_key, # type: ignore[arg-type]
provider_offset=0,
provider_limit=20,
)
)
)
query_providers.assert_not_called()
assert candidates == []
assert global_model_id == "gm1"

View File

@@ -0,0 +1,93 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
def _make_db() -> MagicMock:
db = MagicMock()
db.new = []
db.dirty = []
db.deleted = []
db.in_transaction.return_value = False
return db
def _make_global_model() -> SimpleNamespace:
return SimpleNamespace(
id="gm1",
name="gpt-4o",
is_active=True,
config={},
supported_capabilities=[],
)
@pytest.mark.asyncio
async def test_list_all_candidates_prefilters_provider_graph_by_allowed_providers() -> None:
scheduler = CacheAwareScheduler()
scheduler.scheduling_mode = CacheAwareScheduler.SCHEDULING_MODE_FIXED_ORDER
db = _make_db()
global_model = _make_global_model()
user_api_key = SimpleNamespace(
id="ak1",
allowed_providers=["provider-b"],
allowed_models=None,
allowed_api_formats=None,
user=None,
)
filtered_provider = SimpleNamespace(
id="provider-b",
name="provider-b",
endpoints=[],
models=[],
provider_priority=2,
)
with patch.object(scheduler, "_ensure_initialized", new=AsyncMock(return_value=None)):
with patch.object(
scheduler._candidate_builder,
"_query_provider_refs",
return_value=[("provider-a", "provider-a"), ("provider-b", "provider-b")],
) as refs_mock:
with patch.object(
scheduler._candidate_builder,
"_query_providers",
return_value=[filtered_provider],
) as providers_mock:
with patch(
"src.services.scheduling.aware_scheduler.ModelCacheService.get_global_model_by_name",
new=AsyncMock(return_value=global_model),
):
with patch(
"src.services.scheduling.aware_scheduler.SystemConfigService.is_format_conversion_enabled",
return_value=True,
):
with patch.object(
scheduler._candidate_builder,
"_build_candidates",
new=AsyncMock(return_value=[]),
):
candidates, global_model_id, provider_batch_count = (
await scheduler.list_all_candidates(
db=db,
api_format="openai:chat",
model_name="gpt-4o",
affinity_key=None,
user_api_key=user_api_key, # type: ignore[arg-type]
provider_offset=0,
provider_limit=20,
)
)
assert candidates == []
assert global_model_id == "gm1"
assert provider_batch_count == 2
refs_mock.assert_called_once()
providers_mock.assert_called_once_with(db=db, provider_ids=["provider-b"])

View File

@@ -8,18 +8,20 @@
"""
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from unittest.mock import AsyncMock, MagicMock, patch
import jwt
import pytest
from src.core.exceptions import ForbiddenException
from src.core.enums import AuthSource
from src.core.exceptions import ForbiddenException
from src.models.database import UserRole
from src.services.auth.service import (
JWT_ALGORITHM,
JWT_EXPIRATION_HOURS,
JWT_SECRET_KEY,
AuthenticatedUserSnapshot,
AuthService,
)
@@ -235,6 +237,115 @@ class TestUserAuthentication:
assert result is None
@pytest.mark.asyncio
async def test_authenticate_user_threadsafe_uses_isolated_session_for_local_login(self) -> None:
mock_user = MagicMock()
mock_user.id = "user-123"
mock_user.email = "test@example.com"
mock_user.username = "tester"
mock_user.created_at = datetime.now(timezone.utc)
mock_user.is_deleted = False
mock_user.is_active = True
mock_user.auth_source = AuthSource.LOCAL
mock_user.role = UserRole.USER
mock_user.verify_password.return_value = True
thread_db = MagicMock()
thread_db.query.return_value.filter.return_value.first.return_value = mock_user
route_db = MagicMock()
with patch("src.services.auth.service.create_session", return_value=thread_db):
with patch(
"src.services.auth.service.UserCacheService.invalidate_user_cache",
new_callable=AsyncMock,
) as invalidate_cache:
result = await AuthService.authenticate_user_threadsafe(
route_db,
"test@example.com",
"password123",
)
assert isinstance(result, AuthenticatedUserSnapshot)
assert result.user_id == "user-123"
assert result.username == "tester"
thread_db.commit.assert_called_once()
thread_db.close.assert_called_once()
route_db.commit.assert_not_called()
invalidate_cache.assert_awaited_once_with("user-123", "test@example.com")
@pytest.mark.asyncio
async def test_load_user_for_pipeline_threadsafe_prefetches_balance(self) -> None:
mock_user = MagicMock()
mock_user.id = "user-123"
mock_user.is_active = True
mock_user.is_deleted = False
thread_db = MagicMock()
thread_db.query.return_value.filter.return_value.first.return_value = mock_user
with patch("src.services.auth.service.create_session", return_value=thread_db):
with patch(
"src.services.wallet.service.WalletService.get_balance_snapshot",
return_value=Decimal("7.5"),
):
result = await AuthService.load_user_for_pipeline_threadsafe(
"user-123",
include_balance=True,
)
assert result is not None
assert result.user == mock_user
assert result.balance_remaining == 7.5
thread_db.expunge.assert_called_with(mock_user)
thread_db.close.assert_called_once()
@pytest.mark.asyncio
async def test_authenticate_api_key_threadsafe_returns_balance_and_access_result(self) -> None:
mock_user = MagicMock()
mock_user.id = "user-123"
mock_api_key = MagicMock()
mock_api_key.id = "key-123"
thread_db = MagicMock()
with patch("src.services.auth.service.create_session", return_value=thread_db):
with patch.object(
AuthService,
"authenticate_api_key",
return_value=(mock_user, mock_api_key),
):
with patch(
"src.services.usage.service.UsageService.check_request_balance_details",
return_value=MagicMock(allowed=False, message="????", remaining=0.0),
) as mock_balance_details:
with patch(
"src.services.wallet.service.WalletService.get_balance_snapshot"
) as mock_balance_snapshot:
result = await AuthService.authenticate_api_key_threadsafe("sk-test")
assert result is not None
assert result.user == mock_user
assert result.api_key == mock_api_key
assert result.access_ok is False
assert result.balance_remaining == 0.0
assert result.access_message == "????"
mock_balance_details.assert_called_once()
mock_balance_snapshot.assert_not_called()
thread_db.expunge.assert_any_call(mock_user)
thread_db.expunge.assert_any_call(mock_api_key)
thread_db.close.assert_called_once()
def test_detach_instance_logs_debug_when_expunge_fails(self) -> None:
mock_db = MagicMock()
mock_db.expunge.side_effect = RuntimeError("expunge boom")
mock_instance = MagicMock()
with patch("src.services.auth.service.logger.debug") as mock_debug:
AuthService._detach_instance(mock_db, mock_instance)
mock_debug.assert_called_once()
assert "expunge failed" in mock_debug.call_args[0][0]
class TestAPIKeyAuthentication:
"""测试 API Key 认证"""

View File

@@ -1,36 +1,45 @@
from types import SimpleNamespace
from typing import Any, Callable, cast
import pytest
from src.core.cache_service import CacheService
from src.models.database import GlobalModel, Model
from src.services.cache.model_cache import ModelCacheService
from src.core.cache_service import CacheService
class _FakeQuery:
def __init__(self, *, first_result=None, all_result=None, on_all=None):
def __init__(
self,
*,
first_result: Any = None,
all_result: list[Any] | None = None,
on_all: Callable[[], None] | None = None,
) -> None:
self._first_result = first_result
self._all_result = all_result if all_result is not None else []
self._on_all = on_all
def join(self, *_args, **_kwargs):
def join(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def filter(self, *_args, **_kwargs):
def filter(self, *_args: object, **_kwargs: object) -> "_FakeQuery":
return self
def first(self):
def first(self) -> Any:
return self._first_result
def all(self):
def all(self) -> list[Any]:
if self._on_all:
self._on_all()
return self._all_result
class _FakeSession:
def __init__(self, *, direct_match: GlobalModel):
def __init__(self, *, direct_match: GlobalModel) -> None:
self._direct_match = direct_match
def query(self, *entities):
def query(self, *entities: object) -> "_FakeQuery":
if entities == (GlobalModel,):
return _FakeQuery(first_result=self._direct_match)
@@ -43,12 +52,43 @@ class _FakeSession:
raise AssertionError(f"Unexpected query entities: {entities}")
class _MappingIndexSession:
def __init__(
self,
*,
provider_mapping_rows: list[tuple[object, GlobalModel]],
) -> None:
self._provider_mapping_rows = provider_mapping_rows
self.provider_mapping_scan_count = 0
self.model_global_query_count = 0
def query(self, *entities: object) -> "_FakeQuery":
if entities == (GlobalModel,):
return _FakeQuery(first_result=None, all_result=[])
if entities == (Model, GlobalModel):
self.model_global_query_count += 1
if self.model_global_query_count in {1, 3}:
return _FakeQuery(all_result=[])
if self.model_global_query_count == 2:
return _FakeQuery(
all_result=self._provider_mapping_rows,
on_all=self._record_provider_mapping_scan,
)
raise AssertionError("provider_model_mappings 全量扫描被重复触发")
raise AssertionError(f"Unexpected query entities: {entities}")
def _record_provider_mapping_scan(self) -> None:
self.provider_mapping_scan_count += 1
@pytest.mark.asyncio
async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
async def _fake_get(_key: str):
async def test_resolve_global_model_prefers_direct_match(monkeypatch: pytest.MonkeyPatch) -> None:
async def _fake_get(_key: str) -> None:
return None
async def _fake_set(_key: str, _value, ttl_seconds: int = 60): # noqa: ARG001
async def _fake_set(_key: str, _value: object, ttl_seconds: int = 60) -> bool: # noqa: ARG001
return True
monkeypatch.setattr(CacheService, "get", staticmethod(_fake_get))
@@ -67,6 +107,98 @@ async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
db = _FakeSession(direct_match=global_model)
resolved = await ModelCacheService.resolve_global_model_by_name_or_mapping(
db, global_model.name
cast(Any, db),
cast(str, global_model.name),
)
assert resolved is global_model
@pytest.mark.asyncio
async def test_resolve_global_model_reuses_provider_mapping_index_cache(
monkeypatch: pytest.MonkeyPatch,
) -> None:
cache_store: dict[str, object] = {}
async def _fake_get(key: str) -> object | None:
return cache_store.get(key)
async def _fake_set(key: str, value: object, ttl_seconds: int = 60) -> bool: # noqa: ARG001
cache_store[key] = value
return True
monkeypatch.setattr(CacheService, "get", staticmethod(_fake_get))
monkeypatch.setattr(CacheService, "set", staticmethod(_fake_set))
global_model_one = GlobalModel(
id="gm-1",
name="gpt-4o",
display_name="GPT-4o",
supported_capabilities=[],
config={},
default_tiered_pricing=None,
default_price_per_request=None,
is_active=True,
)
global_model_two = GlobalModel(
id="gm-2",
name="claude-3-7-sonnet",
display_name="Claude 3.7 Sonnet",
supported_capabilities=[],
config={},
default_tiered_pricing=None,
default_price_per_request=None,
is_active=True,
)
model_one = SimpleNamespace(
id="m-1",
provider_model_mappings=[{"name": "mapped-one"}],
)
model_two = SimpleNamespace(
id="m-2",
provider_model_mappings=[{"name": "mapped-two"}],
)
db = _MappingIndexSession(
provider_mapping_rows=[
(model_one, global_model_one),
(model_two, global_model_two),
]
)
resolved_one = await ModelCacheService.resolve_global_model_by_name_or_mapping(
cast(Any, db), "mapped-one"
)
resolved_two = await ModelCacheService.resolve_global_model_by_name_or_mapping(
cast(Any, db), "mapped-two"
)
assert resolved_one is not None
assert resolved_one.name == "gpt-4o"
assert resolved_two is not None
assert resolved_two.name == "claude-3-7-sonnet"
assert db.provider_mapping_scan_count == 1
@pytest.mark.asyncio
async def test_invalidate_model_cache_clears_provider_mapping_index(
monkeypatch: pytest.MonkeyPatch,
) -> None:
deleted_keys: list[str] = []
async def _fake_delete(key: str) -> bool:
deleted_keys.append(key)
return True
monkeypatch.setattr(CacheService, "delete", staticmethod(_fake_delete))
await ModelCacheService.invalidate_model_cache(
model_id="model-1",
provider_model_name="provider-model",
provider_model_mappings=[{"name": "alias-model"}],
)
assert "model:id:model-1" in deleted_keys
assert "global_model:resolve:provider-model" in deleted_keys
assert "global_model:resolve:alias-model" in deleted_keys
assert ModelCacheService.PROVIDER_MAPPING_INDEX_CACHE_KEY in deleted_keys

View File

@@ -0,0 +1,182 @@
from __future__ import annotations
from datetime import date, datetime, time, timedelta, timezone
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.models.database import StatsDaily, StatsUserDaily
from src.services.system.stats_aggregator import (
AggregatedStats,
StatsAggregatorService,
query_stats_hybrid,
)
from src.services.system.time_range import TimeRangeParams
class _FakeQuery:
def __init__(self, *, all_result: list[Any] | None = None) -> None:
self._all_result = all_result if all_result is not None else []
def filter(self, *_args: object, **_kwargs: object) -> _FakeQuery:
return self
def group_by(self, *_args: object, **_kwargs: object) -> _FakeQuery:
return self
def all(self) -> list[Any]:
return self._all_result
class _HybridQuerySession:
def __init__(self, stats_daily_rows: list[SimpleNamespace]) -> None:
self._stats_daily_rows = stats_daily_rows
self.stats_daily_query_count = 0
def query(self, entity: object) -> _FakeQuery:
if entity is StatsDaily:
self.stats_daily_query_count += 1
return _FakeQuery(all_result=self._stats_daily_rows)
raise AssertionError(f"Unexpected query entity: {entity}")
class _BatchUserStatsSession:
def __init__(
self, existing_rows: list[StatsUserDaily], aggregated_rows: list[SimpleNamespace]
) -> None:
self._responses: list[list[Any]] = [list(existing_rows), list(aggregated_rows)]
self.added: list[StatsUserDaily] = []
self.commit_count = 0
def query(self, *_entities: object) -> _FakeQuery:
if not self._responses:
raise AssertionError("Unexpected extra query")
return _FakeQuery(all_result=self._responses.pop(0))
def add(self, row: StatsUserDaily) -> None:
self.added.append(row)
def commit(self) -> None:
self.commit_count += 1
def test_query_stats_hybrid_batches_statsdaily_lookup_and_merges_realtime_ranges(
monkeypatch: pytest.MonkeyPatch,
) -> None:
today = datetime.now(timezone.utc).date()
historical_cached_day = today - timedelta(days=4)
historical_missing_day = today - timedelta(days=3)
realtime_day = today
cached_row = SimpleNamespace(
date=datetime.combine(historical_cached_day, time.min, tzinfo=timezone.utc),
total_requests=10,
success_requests=9,
error_requests=1,
input_tokens=100,
output_tokens=50,
cache_creation_tokens=5,
cache_read_tokens=3,
cache_creation_cost=1.2,
cache_read_cost=0.8,
total_cost=3.5,
actual_total_cost=3.0,
avg_response_time_ms=200.0,
)
db = _HybridQuerySession(stats_daily_rows=[cached_row])
calls: list[tuple[datetime, datetime]] = []
def _fake_aggregate_usage_range(
_db: object,
start_utc: datetime,
end_utc: datetime,
filters: object | None = None, # noqa: ARG001
) -> AggregatedStats:
calls.append((start_utc, end_utc))
return AggregatedStats(total_requests=1, success_requests=1)
class _FakeParams:
def get_complete_utc_dates(self) -> tuple[list[date], None, None]:
return [historical_cached_day, historical_missing_day, realtime_day], None, None
monkeypatch.setattr(
"src.services.system.stats_aggregator.aggregate_usage_range",
_fake_aggregate_usage_range,
)
result = query_stats_hybrid(cast(Any, db), cast(Any, _FakeParams()))
assert db.stats_daily_query_count == 1
assert calls == [
(
datetime.combine(historical_missing_day, time.min, tzinfo=timezone.utc),
datetime.combine(
historical_missing_day + timedelta(days=1), time.min, tzinfo=timezone.utc
),
),
(
datetime.combine(realtime_day, time.min, tzinfo=timezone.utc),
datetime.combine(realtime_day + timedelta(days=1), time.min, tzinfo=timezone.utc),
),
]
assert result.total_requests == 12
assert result.success_requests == 11
def test_aggregate_user_daily_stats_batch_updates_all_users_in_two_queries() -> None:
target_day = datetime(2026, 3, 1, tzinfo=timezone.utc)
aggregated_rows = [
SimpleNamespace(
user_id="user-1",
username="alice",
total_requests=4,
error_requests=1,
input_tokens=20,
output_tokens=8,
cache_creation_tokens=2,
cache_read_tokens=1,
total_cost=1.5,
)
]
db = _BatchUserStatsSession(existing_rows=[], aggregated_rows=aggregated_rows)
result = StatsAggregatorService.aggregate_user_daily_stats_batch(
cast(Any, db),
target_day,
["user-1", "user-2"],
commit=True,
)
assert len(result) == 2
assert db.commit_count == 1
assert len(db.added) == 2
user_one = next(row for row in result if row.user_id == "user-1")
user_two = next(row for row in result if row.user_id == "user-2")
assert user_one.username == "alice"
assert user_one.total_requests == 4
assert user_one.success_requests == 3
assert user_one.total_cost == 1.5
assert user_two.total_requests == 0
assert user_two.success_requests == 0
assert user_two.error_requests == 0
assert user_two.total_cost == 0.0
def test_compute_percentiles_by_local_day_returns_sqlite_fallback_without_queries() -> None:
db = SimpleNamespace(bind=SimpleNamespace(dialect=SimpleNamespace(name="sqlite")))
time_range = TimeRangeParams(
start_date=date(2026, 3, 1),
end_date=date(2026, 3, 3),
timezone="Asia/Singapore",
)
result = StatsAggregatorService.compute_percentiles_by_local_day(cast(Any, db), time_range)
assert [row["date"] for row in result] == ["2026-03-01", "2026-03-02", "2026-03-03"]
assert all(row["p50_response_time_ms"] is None for row in result)
assert all(row["p50_first_byte_time_ms"] is None for row in result)

View File

@@ -8,7 +8,7 @@ UsageService 测试
"""
from decimal import Decimal
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import MagicMock, patch
import pytest
@@ -146,6 +146,58 @@ class TestBalanceCheck:
assert is_ok is True
def test_check_request_balance_details_returns_remaining(self) -> None:
"""Balance detail helper returns remaining."""
mock_user = MagicMock()
mock_user.role = MagicMock()
mock_user.role.value = "user"
mock_api_key = MagicMock()
mock_api_key.is_standalone = False
mock_db = MagicMock()
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(
False, Decimal("12.5"), "\u94b1\u5305\u4f59\u989d\u4e0d\u8db3"
),
):
result = UsageService.check_request_balance_details(
mock_db, mock_user, api_key=mock_api_key
)
assert result.allowed is False
assert result.remaining == 12.5
assert "\u4f59\u989d\u4e0d\u8db3" in result.message
def test_check_request_balance_details_maps_overdue_message(self) -> None:
"""欠费状态应映射为对外统一文案。"""
mock_user = MagicMock()
mock_api_key = MagicMock()
mock_api_key.is_standalone = False
mock_db = MagicMock()
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(False, Decimal("-1"), "钱包欠费,请先充值"),
):
normal_result = UsageService.check_request_balance_details(
mock_db, mock_user, api_key=mock_api_key
)
mock_api_key.is_standalone = True
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(False, Decimal("-1"), "钱包欠费,请先充值"),
):
standalone_result = UsageService.check_request_balance_details(
mock_db, mock_user, api_key=mock_api_key
)
assert normal_result.message == "账户欠费,请先充值"
assert standalone_result.message == "Key欠费请先调账或充值"
def test_check_request_balance_exceeded(self) -> None:
"""测试余额耗尽时拦截新请求"""
mock_user = MagicMock()

View File

@@ -1,5 +1,6 @@
from decimal import Decimal
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock, patch
import pytest
@@ -52,7 +53,9 @@ def test_get_or_create_wallet_prefers_user_owner_for_non_standalone_key() -> Non
user = SimpleNamespace(id="user-1")
api_key = SimpleNamespace(id="key-1", is_standalone=False)
wallet = WalletService.get_or_create_wallet(db, user=user, api_key=api_key)
wallet = WalletService.get_or_create_wallet(
db, user=cast(Any, user), api_key=cast(Any, api_key)
)
assert wallet is not None
assert wallet.user_id == "user-1"
@@ -68,13 +71,35 @@ def test_get_or_create_wallet_uses_api_key_owner_for_standalone_key() -> None:
user = SimpleNamespace(id="user-1")
api_key = SimpleNamespace(id="key-1", is_standalone=True)
wallet = WalletService.get_or_create_wallet(db, user=user, api_key=api_key)
wallet = WalletService.get_or_create_wallet(
db, user=cast(Any, user), api_key=cast(Any, api_key)
)
assert wallet is not None
assert wallet.user_id is None
assert wallet.api_key_id == "key-1"
def test_get_wallets_by_user_ids_returns_mapping() -> None:
db = MagicMock()
wallet_1 = SimpleNamespace(user_id="user-1")
wallet_2 = SimpleNamespace(user_id="user-2")
db.query.return_value.filter.return_value.all.return_value = [wallet_1, wallet_2]
result = WalletService.get_wallets_by_user_ids(db, ["user-1", "user-2"])
assert result == {"user-1": wallet_1, "user-2": wallet_2}
def test_get_wallets_by_user_ids_skips_query_for_empty_ids() -> None:
db = MagicMock()
result = WalletService.get_wallets_by_user_ids(db, [])
assert result == {}
db.query.assert_not_called()
def test_check_request_allowed_denies_when_recharge_negative_even_total_positive() -> None:
wallet = _build_wallet(recharge="-1", gift="10", limit_mode="finite")
db = MagicMock()
@@ -104,7 +129,7 @@ def test_admin_adjust_balance_negative_from_gift_spills_to_recharge() -> None:
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
wallet=cast(Any, wallet),
amount_usd=Decimal("-10"),
balance_type="gift",
operator_id="admin-1",
@@ -127,7 +152,7 @@ def test_admin_adjust_balance_negative_from_recharge_then_gift() -> None:
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
wallet=cast(Any, wallet),
amount_usd=Decimal("-4"),
balance_type="recharge",
operator_id="admin-1",
@@ -148,7 +173,7 @@ def test_admin_adjust_balance_positive_adds_to_selected_bucket_without_offset()
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
wallet=cast(Any, wallet),
amount_usd=Decimal("1"),
balance_type="gift",
operator_id="admin-1",
@@ -181,7 +206,9 @@ def test_apply_usage_charge_prefers_gift_then_recharge() -> None:
db = _build_locked_db(wallet)
with patch.object(WalletService, "_resolve_wallet_for_usage", return_value=wallet):
before, after = WalletService.apply_usage_charge(db, usage=usage, amount_usd=Decimal("6"))
before, after = WalletService.apply_usage_charge(
db, usage=cast(Any, usage), amount_usd=Decimal("6")
)
assert before == Decimal("8.00000000")
assert after == Decimal("2.00000000")
@@ -212,7 +239,9 @@ def test_apply_usage_charge_unlimited_wallet_keeps_balances() -> None:
db = _build_locked_db(wallet)
with patch.object(WalletService, "_resolve_wallet_for_usage", return_value=wallet):
before, after = WalletService.apply_usage_charge(db, usage=usage, amount_usd=Decimal("4"))
before, after = WalletService.apply_usage_charge(
db, usage=cast(Any, usage), amount_usd=Decimal("4")
)
assert before == Decimal("8.00000000")
assert after == Decimal("8.00000000")
@@ -237,7 +266,7 @@ def test_complete_refund_requires_processing_status() -> None:
db.query.return_value = query
with pytest.raises(ValueError, match="processing"):
WalletService.complete_refund(db, refund=refund)
WalletService.complete_refund(db, refund=cast(Any, refund))
def test_get_or_create_wallet_reuses_existing_after_integrity_error() -> None:
@@ -252,7 +281,7 @@ def test_get_or_create_wallet_reuses_existing_after_integrity_error() -> None:
with patch.object(WalletService, "get_wallet", side_effect=[None, existing_wallet]):
wallet = WalletService.get_or_create_wallet(
db,
user=SimpleNamespace(id="user-1"),
user=cast(Any, SimpleNamespace(id="user-1")),
api_key=None,
)
@@ -287,18 +316,20 @@ def test_create_refund_request_rejects_uncredited_payment_order() -> None:
db.query.side_effect = _query
with patch.object(WalletService, "_get_pending_refund_reserved_amount", return_value=Decimal("0")):
with patch.object(
WalletService, "_get_pending_refund_reserved_amount", return_value=Decimal("0")
):
with pytest.raises(ValueError, match="payment order is not refundable"):
WalletService.create_refund_request(
db,
wallet=wallet,
wallet=cast(Any, wallet),
user_id="user-1",
amount_usd=Decimal("2"),
refund_no="rf-1",
source_type="payment_order",
source_id="order-1",
refund_mode="original_channel",
payment_order=payment_order,
payment_order=cast(Any, payment_order),
)
@@ -326,7 +357,7 @@ def test_create_refund_request_reserves_pending_wallet_amount() -> None:
with pytest.raises(ValueError, match="available refundable recharge balance"):
WalletService.create_refund_request(
db,
wallet=wallet,
wallet=cast(Any, wallet),
user_id="user-1",
amount_usd=Decimal("2"),
refund_no="rf-2",
@@ -372,14 +403,14 @@ def test_create_refund_request_reserves_pending_order_amount() -> None:
with pytest.raises(ValueError, match="available refundable amount"):
WalletService.create_refund_request(
db,
wallet=wallet,
wallet=cast(Any, wallet),
user_id="user-1",
amount_usd=Decimal("2"),
refund_no="rf-3",
source_type="payment_order",
source_id="order-1",
refund_mode="original_channel",
payment_order=payment_order,
payment_order=cast(Any, payment_order),
)
@@ -430,9 +461,13 @@ def test_move_refund_to_processing_rejects_double_transition() -> None:
tx = SimpleNamespace(id="tx-1")
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
first_tx = WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
first_tx = WalletService.move_refund_to_processing(
db, refund=cast(Any, refund), operator_id="admin-1"
)
with pytest.raises(ValueError, match="not approvable"):
WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
WalletService.move_refund_to_processing(
db, refund=cast(Any, refund), operator_id="admin-1"
)
assert first_tx is tx
assert create_tx.call_count == 1
@@ -488,7 +523,9 @@ def test_move_refund_to_processing_rechecks_payment_order_refundable_amount() ->
with patch.object(WalletService, "create_wallet_transaction") as create_tx:
with pytest.raises(ValueError, match="refund amount exceeds refundable amount"):
WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
WalletService.move_refund_to_processing(
db, refund=cast(Any, refund), operator_id="admin-1"
)
create_tx.assert_not_called()
assert refund.status == "pending_approval"
@@ -531,14 +568,14 @@ def test_fail_refund_rejects_invalid_status_after_first_failure() -> None:
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
first_tx = WalletService.fail_refund(
db,
refund=refund,
refund=cast(Any, refund),
reason="first-failure",
operator_id="admin-1",
)
with pytest.raises(ValueError, match="cannot fail refund in status: failed"):
WalletService.fail_refund(
db,
refund=refund,
refund=cast(Any, refund),
reason="retry-failure",
operator_id="admin-1",
)
@@ -577,7 +614,7 @@ def test_fail_refund_rejects_succeeded_status() -> None:
with pytest.raises(ValueError, match="cannot fail refund in status: succeeded"):
WalletService.fail_refund(
db,
refund=refund,
refund=cast(Any, refund),
reason="should-not-override",
operator_id="admin-1",
)

View File

@@ -11,6 +11,13 @@ from src.api.base.context import ApiRequestContext
def _build_request(headers: dict[str, str] | None = None) -> Request:
return _build_request_with_body(b"", headers=headers)
def _build_request_with_body(
body: bytes,
headers: dict[str, str] | None = None,
) -> Request:
header_items = [
(str(key).encode("latin-1"), str(value).encode("latin-1"))
for key, value in (headers or {}).items()
@@ -28,8 +35,14 @@ def _build_request(headers: dict[str, str] | None = None) -> Request:
"server": ("testserver", 80),
}
received = False
async def receive() -> dict[str, object]:
return {"type": "http.request", "body": b"", "more_body": False}
nonlocal received
if received:
return {"type": "http.request", "body": b"", "more_body": False}
received = True
return {"type": "http.request", "body": body, "more_body": False}
request = Request(scope, receive)
request.state.perf_metrics = {}
@@ -89,3 +102,22 @@ class TestApiRequestContextEnsureJsonBody:
assert context.client_content_encoding == "gzip"
assert context.client_accept_encoding == "gzip, deflate"
@pytest.mark.asyncio
async def test_ensure_json_body_async_loads_body_lazily(self) -> None:
payload = {"message": "hello", "count": 2}
request = _build_request_with_body(json.dumps(payload).encode("utf-8"))
context = ApiRequestContext.build(
request=request,
db=None, # type: ignore[arg-type]
user=None,
api_key=None,
raw_body=None,
)
assert context.raw_body is None
result = await context.ensure_json_body_async()
assert result == payload
assert context.raw_body == json.dumps(payload).encode("utf-8")

View File

@@ -0,0 +1,40 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
from src.services.request.candidate import RequestCandidateService
def _build_db_with_candidate(candidate: SimpleNamespace) -> MagicMock:
query = MagicMock()
query.filter.return_value.first.return_value = candidate
db = MagicMock()
db.query.return_value = query
db.info = {"managed_by_middleware": True}
return db
def test_mark_candidate_started_flushes_without_immediate_commit() -> None:
candidate = SimpleNamespace(status="available", started_at=None)
db = _build_db_with_candidate(candidate)
RequestCandidateService.mark_candidate_started(db, "candidate-1")
assert candidate.status == "pending"
assert candidate.started_at is not None
db.flush.assert_called_once()
db.commit.assert_not_called()
def test_mark_candidate_streaming_flushes_without_immediate_commit() -> None:
candidate = SimpleNamespace(status="pending", concurrent_requests=None)
db = _build_db_with_candidate(candidate)
RequestCandidateService.mark_candidate_streaming(db, "candidate-1", concurrent_requests=3)
assert candidate.status == "streaming"
assert candidate.concurrent_requests == 3
db.flush.assert_called_once()
db.commit.assert_not_called()