feat(wallet): 钱包系统替代配额系统,新增支付与退款机制

- 新增钱包余额管理、充值、扣费、退款完整流程
- 新增支付网关抽象层(支持手动/支付宝/微信)
- 用量计费从配额系统迁移到钱包余额扣费
- 新增管理员钱包管理与支付订单管理页面
- 新增用户钱包中心页面
- 移除独立 Key 锁定机制,统一由钱包余额控制
- 新增相关 API 路由、序列化器与数据库迁移
- 新增钱包、支付、退款相关测试
This commit is contained in:
LewisPen
2026-03-08 00:05:48 +08:00
committed by fawney19
parent 9cdcce1b5f
commit 783f654953
108 changed files with 13152 additions and 3372 deletions

View File

@@ -0,0 +1,289 @@
from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
from src.api.admin.api_keys.routes import (
AdminGetFullKeyAdapter,
AdminToggleApiKeyAdapter,
router as admin_api_keys_router,
)
from src.api.admin.users.routes import (
AdminGetUserKeyFullKeyAdapter,
AdminToggleUserKeyLockAdapter,
router as admin_users_router,
)
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.database import get_db
def _build_context(db: MagicMock) -> SimpleNamespace:
return SimpleNamespace(
db=db,
add_audit_metadata=lambda **_: None,
)
def _mock_query_first(db: MagicMock, value: object | None) -> None:
db.query.return_value.filter.return_value.first.return_value = value
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: object, http_request: object, db: MagicMock, mode: object) -> object:
_ = http_request, mode
context = SimpleNamespace(
db=db,
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 _build_admin_api_keys_app(db: MagicMock, monkeypatch: pytest.MonkeyPatch) -> TestClient:
app = FastAPI()
app.include_router(admin_api_keys_router)
app.dependency_overrides[get_db] = lambda: db
async def _fake_pipeline_run(*, adapter: object, http_request: object, db: MagicMock, mode: object) -> object:
_ = http_request, mode
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id="admin-1"),
ensure_json_body=lambda: {},
add_audit_metadata=lambda **_: None,
)
return await adapter.handle(context)
monkeypatch.setattr("src.api.admin.api_keys.routes.pipeline.run", _fake_pipeline_run)
return TestClient(app)
@pytest.mark.asyncio
async def test_toggle_user_key_lock_adapter_success() -> None:
db = MagicMock()
api_key = SimpleNamespace(id="key-1", user_id="user-1", is_standalone=False, is_locked=False)
_mock_query_first(db, api_key)
adapter = AdminToggleUserKeyLockAdapter(user_id="user-1", key_id="key-1")
result = await adapter.handle(_build_context(db))
assert result["id"] == "key-1"
assert result["is_locked"] is True
assert "锁定" in result["message"]
db.commit.assert_called_once()
db.refresh.assert_called_once_with(api_key)
@pytest.mark.asyncio
async def test_toggle_user_key_lock_adapter_not_found_for_standalone_or_wrong_owner() -> None:
db = MagicMock()
_mock_query_first(db, None)
adapter = AdminToggleUserKeyLockAdapter(user_id="user-1", key_id="key-standalone")
with pytest.raises(NotFoundException):
await adapter.handle(_build_context(db))
db.commit.assert_not_called()
@pytest.mark.asyncio
async def test_get_user_key_full_key_adapter_success(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
api_key = SimpleNamespace(
id="key-2",
user_id="user-1",
is_standalone=False,
key_encrypted="encrypted-value",
)
_mock_query_first(db, api_key)
monkeypatch.setattr("src.core.crypto.crypto_service.decrypt", lambda _v: "sk-user-full-key")
adapter = AdminGetUserKeyFullKeyAdapter(user_id="user-1", key_id="key-2")
result = await adapter.handle(_build_context(db))
assert result == {"key": "sk-user-full-key"}
@pytest.mark.asyncio
async def test_get_user_key_full_key_adapter_requires_encrypted_key() -> None:
db = MagicMock()
api_key = SimpleNamespace(
id="key-3",
user_id="user-1",
is_standalone=False,
key_encrypted=None,
)
_mock_query_first(db, api_key)
adapter = AdminGetUserKeyFullKeyAdapter(user_id="user-1", key_id="key-3")
with pytest.raises(InvalidRequestException):
await adapter.handle(_build_context(db))
@pytest.mark.asyncio
async def test_get_user_key_full_key_adapter_returns_500_on_decrypt_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
api_key = SimpleNamespace(
id="key-4",
user_id="user-1",
is_standalone=False,
key_encrypted="encrypted-value",
)
_mock_query_first(db, api_key)
def _raise(_: str) -> str:
raise ValueError("decrypt failed")
monkeypatch.setattr("src.core.crypto.crypto_service.decrypt", _raise)
adapter = AdminGetUserKeyFullKeyAdapter(user_id="user-1", key_id="key-4")
with pytest.raises(HTTPException) as exc_info:
await adapter.handle(_build_context(db))
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
async def test_standalone_toggle_adapters_reject_normal_user_key() -> None:
db = MagicMock()
normal_key = SimpleNamespace(
id="key-user",
user_id="user-1",
is_standalone=False,
is_active=True,
is_locked=False,
key_encrypted="encrypted-value",
updated_at=datetime.now(timezone.utc),
)
_mock_query_first(db, normal_key)
context = _build_context(db)
with pytest.raises(InvalidRequestException):
await AdminToggleApiKeyAdapter(key_id="key-user").handle(context)
with pytest.raises(InvalidRequestException):
await AdminGetFullKeyAdapter(key_id="key-user").handle(context)
def test_user_key_lock_route_path_smoke(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
api_key = SimpleNamespace(id="key-5", user_id="user-2", is_standalone=False, is_locked=False)
_mock_query_first(db, api_key)
client = _build_admin_users_app(db, monkeypatch)
response = client.patch("/api/admin/users/user-2/api-keys/key-5/lock")
assert response.status_code == 200
assert response.json()["id"] == "key-5"
assert response.json()["is_locked"] is True
def test_user_key_full_key_route_path_smoke(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
api_key = SimpleNamespace(
id="key-6",
user_id="user-2",
is_standalone=False,
key_encrypted="enc",
)
_mock_query_first(db, api_key)
monkeypatch.setattr("src.core.crypto.crypto_service.decrypt", lambda _v: "sk-user-route-key")
client = _build_admin_users_app(db, monkeypatch)
response = client.get("/api/admin/users/user-2/api-keys/key-6/full-key")
assert response.status_code == 200
assert response.json() == {"key": "sk-user-route-key"}
def test_standalone_lock_route_removed(monkeypatch: pytest.MonkeyPatch) -> None:
client = _build_admin_api_keys_app(MagicMock(), monkeypatch)
response = client.patch("/api/admin/api-keys/key-1/lock")
assert response.status_code == 404
def test_standalone_list_route_does_not_expose_is_locked(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
api_key = SimpleNamespace(
id="sa-key-1",
user_id="admin-1",
name="Standalone Key",
get_display_key=lambda: "sk-stand...1234",
is_active=True,
is_standalone=True,
total_requests=0,
total_cost_usd=0,
rate_limit=None,
allowed_providers=None,
allowed_api_formats=None,
allowed_models=None,
last_used_at=None,
expires_at=None,
created_at=datetime.now(timezone.utc),
updated_at=None,
auto_delete_on_expiry=False,
)
query = db.query.return_value.filter.return_value
query.count.return_value = 1
query.order_by.return_value.offset.return_value.limit.return_value.all.return_value = [api_key]
monkeypatch.setattr(
"src.api.admin.api_keys.routes.WalletService.get_wallet",
lambda _db, user_id=None, api_key_id=None, user=None, api_key=None: SimpleNamespace(id="w-1"),
)
client = _build_admin_api_keys_app(db, monkeypatch)
response = client.get("/api/admin/api-keys")
assert response.status_code == 200
payload = response.json()
assert len(payload["api_keys"]) == 1
assert "is_locked" not in payload["api_keys"][0]
def test_standalone_detail_route_does_not_expose_is_locked(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
api_key = SimpleNamespace(
id="sa-key-2",
user_id="admin-1",
name="Standalone Key 2",
get_display_key=lambda: "sk-stand...5678",
is_active=True,
is_standalone=True,
total_requests=0,
total_cost_usd=0,
rate_limit=None,
allowed_providers=[],
allowed_api_formats=[],
allowed_models=[],
last_used_at=None,
expires_at=None,
created_at=datetime.now(timezone.utc),
updated_at=None,
)
_mock_query_first(db, api_key)
monkeypatch.setattr(
"src.api.admin.api_keys.routes.WalletService.get_wallet",
lambda _db, user_id=None, api_key_id=None, user=None, api_key=None: None,
)
client = _build_admin_api_keys_app(db, monkeypatch)
response = client.get("/api/admin/api-keys/sa-key-2")
assert response.status_code == 200
payload = response.json()
assert payload["id"] == "sa-key-2"
assert "is_locked" not in payload

View File

@@ -0,0 +1,232 @@
from __future__ import annotations
from decimal import Decimal
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from src.api.admin.payments.routes import AdminPaymentOrderCreditAdapter
from src.api.payment.routes import router as payment_router
from src.config import config
from src.database import get_db
from src.models.database import PaymentOrder
from src.services.payment.gateway import get_payment_gateway
CALLBACK_SECRET = "test-callback-secret"
def _build_payment_app(db: MagicMock) -> TestClient:
app = FastAPI()
app.include_router(payment_router)
app.dependency_overrides[get_db] = lambda: db
return TestClient(app)
def _sign_payload(payload: dict[str, object]) -> str:
gateway = get_payment_gateway("alipay")
signature = gateway.build_callback_signature(payload=payload, callback_secret=CALLBACK_SECRET)
assert signature is not None
return signature
def test_specific_wechat_callback_route_is_not_shadowed(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
client = _build_payment_app(db)
monkeypatch.setattr(config, "payment_callback_secret", CALLBACK_SECRET)
captured_kwargs: dict[str, object] = {}
def _fake_handle_callback(*args: object, **kwargs: object) -> dict[str, object]:
captured_kwargs.update(kwargs)
return {
"ok": True,
"credited": True,
"duplicate": False,
"payment_method_seen": kwargs["payment_method"],
}
monkeypatch.setattr("src.api.payment.routes.PaymentService.handle_callback", _fake_handle_callback)
callback_payload = {"callback_key": "cb-wechat", "amount_usd": 1.0}
response = client.post(
"/api/payment/callback/wechat",
json=callback_payload,
headers={
"x-payment-callback-token": CALLBACK_SECRET,
"x-payment-callback-signature": _sign_payload(callback_payload),
},
)
assert response.status_code == 200
payload = response.json()
assert payload["payment_method"] == "wechat"
assert payload["payment_method_seen"] == "wechat"
assert payload["request_path"] == "/api/payment/callback/wechat"
assert captured_kwargs["callback_signature"] == _sign_payload(callback_payload)
assert captured_kwargs["callback_secret"] == CALLBACK_SECRET
assert "signature_valid" not in captured_kwargs
db.commit.assert_called_once()
def test_generic_payment_callback_route_still_handles_custom_methods(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
client = _build_payment_app(db)
monkeypatch.setattr(config, "payment_callback_secret", CALLBACK_SECRET)
captured_kwargs: dict[str, object] = {}
def _fake_handle_callback(*args: object, **kwargs: object) -> dict[str, object]:
captured_kwargs.update(kwargs)
return {
"ok": True,
"credited": False,
"duplicate": False,
"payment_method_seen": kwargs["payment_method"],
}
monkeypatch.setattr("src.api.payment.routes.PaymentService.handle_callback", _fake_handle_callback)
callback_payload = {"callback_key": "cb-generic", "amount_usd": 1.0}
response = client.post(
"/api/payment/callback/mockpay",
json=callback_payload,
headers={
"x-payment-callback-token": CALLBACK_SECRET,
"x-payment-callback-signature": _sign_payload(callback_payload),
},
)
assert response.status_code == 200
payload = response.json()
assert payload["payment_method"] == "mockpay"
assert payload["payment_method_seen"] == "mockpay"
assert captured_kwargs["callback_signature"] == _sign_payload(callback_payload)
assert captured_kwargs["callback_secret"] == CALLBACK_SECRET
assert "signature_valid" not in captured_kwargs
def test_callback_requires_shared_token(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
client = _build_payment_app(db)
monkeypatch.setattr(config, "payment_callback_secret", CALLBACK_SECRET)
response = client.post(
"/api/payment/callback/alipay",
json={"callback_key": "cb-missing-token", "amount_usd": 1.0},
)
assert response.status_code == 401
db.commit.assert_not_called()
def test_callback_rejects_invalid_shared_token(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
client = _build_payment_app(db)
monkeypatch.setattr(config, "payment_callback_secret", CALLBACK_SECRET)
callback_payload = {"callback_key": "cb-invalid-token", "amount_usd": 1.0}
response = client.post(
"/api/payment/callback/alipay",
json=callback_payload,
headers={
"x-payment-callback-token": "wrong-token",
"x-payment-callback-signature": _sign_payload(callback_payload),
},
)
assert response.status_code == 401
db.commit.assert_not_called()
def test_callback_rejects_missing_signature(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
client = _build_payment_app(db)
monkeypatch.setattr(config, "payment_callback_secret", CALLBACK_SECRET)
response = client.post(
"/api/payment/callback/alipay",
json={"callback_key": "cb-missing-signature", "amount_usd": 1.0},
headers={"x-payment-callback-token": CALLBACK_SECRET},
)
assert response.status_code == 401
db.commit.assert_not_called()
def test_callback_disabled_when_secret_not_configured(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
client = _build_payment_app(db)
monkeypatch.setattr(config, "payment_callback_secret", "")
response = client.post(
"/api/payment/callback/alipay",
json={"callback_key": "cb-secret-missing", "amount_usd": 1.0},
)
assert response.status_code == 503
db.commit.assert_not_called()
@pytest.mark.asyncio
async def test_admin_payment_credit_adapter_marks_manual_credit(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
order = PaymentOrder(
id="po-credit",
order_no="order-credit",
wallet_id="w1",
user_id="u1",
amount_usd=Decimal("8.00000000"),
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("8.00000000"),
payment_method="alipay",
status="pending",
gateway_response={"existing": True},
)
adapter = AdminPaymentOrderCreditAdapter(order_id=order.id)
context = SimpleNamespace(
db=db,
raw_body=b"{}",
ensure_json_body=lambda: {
"pay_amount": 58.0,
"pay_currency": "CNY",
"exchange_rate": 7.25,
},
user=SimpleNamespace(id="admin-1"),
)
monkeypatch.setattr(
"src.api.admin.payments.routes.PaymentService.get_order",
lambda _db, order_id: order if order_id == "po-credit" else None,
)
captured: dict[str, object] = {}
def _fake_credit_order(_db: MagicMock, **kwargs: object) -> tuple[PaymentOrder, bool]:
captured.update(kwargs)
return order, True
monkeypatch.setattr(
"src.api.admin.payments.routes.PaymentService.credit_order",
_fake_credit_order,
)
result = await adapter.handle(context)
assert result["credited"] is True
assert result["order"]["id"] == "po-credit"
gateway_response = captured["gateway_response"]
assert isinstance(gateway_response, dict)
assert gateway_response["existing"] is True
assert gateway_response["manual_credit"] is True
assert gateway_response["credited_by"] == "admin-1"
db.commit.assert_called_once()

View File

@@ -3,7 +3,7 @@ API Pipeline 测试
测试 ApiRequestPipeline 的核心功能:
- 认证流程API Key、JWT Token
- 额计算
- 额计算
- 审计日志记录
"""
@@ -17,56 +17,43 @@ from src.api.base.pipeline import ApiRequestPipeline
from src.core.enums import UserRole
class TestPipelineQuotaCalculation:
"""测试 Pipeline 额计算"""
class TestPipelineBalanceCalculation:
"""测试 Pipeline 额计算"""
@pytest.fixture
def pipeline(self) -> ApiRequestPipeline:
return ApiRequestPipeline()
def test_calculate_quota_remaining_with_quota(self, pipeline: ApiRequestPipeline) -> None:
"""测试有配额限制时计算剩余"""
def test_calculate_balance_remaining_with_balance(self, pipeline: ApiRequestPipeline) -> None:
"""测试有限制钱包时计算剩余"""
mock_user = MagicMock()
mock_user.quota_usd = 100.0
mock_user.used_usd = 30.0
mock_db = MagicMock()
remaining = pipeline._calculate_quota_remaining(mock_user)
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=70.0,
):
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
assert remaining == 70.0
def test_calculate_quota_remaining_no_quota(self, pipeline: ApiRequestPipeline) -> None:
"""测试无配额限制时返回 None"""
def test_calculate_balance_remaining_unlimited(self, pipeline: ApiRequestPipeline) -> None:
"""测试无限制钱包时返回 None"""
mock_user = MagicMock()
mock_user.quota_usd = None
mock_user.used_usd = 30.0
mock_db = MagicMock()
remaining = pipeline._calculate_quota_remaining(mock_user)
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=None,
):
remaining = pipeline._calculate_balance_remaining(mock_db, mock_user)
assert remaining is None
def test_calculate_quota_remaining_negative_quota(self, pipeline: ApiRequestPipeline) -> None:
"""测试负配额时返回 None"""
mock_user = MagicMock()
mock_user.quota_usd = -1
mock_user.used_usd = 0.0
remaining = pipeline._calculate_quota_remaining(mock_user)
assert remaining is None
def test_calculate_quota_remaining_exceeded(self, pipeline: ApiRequestPipeline) -> None:
"""测试配额已超时返回 0"""
mock_user = MagicMock()
mock_user.quota_usd = 100.0
mock_user.used_usd = 150.0
remaining = pipeline._calculate_quota_remaining(mock_user)
assert remaining == 0.0
def test_calculate_quota_remaining_none_user(self, pipeline: ApiRequestPipeline) -> None:
def test_calculate_balance_remaining_none_user(self, pipeline: ApiRequestPipeline) -> None:
"""测试用户为 None 时返回 None"""
remaining = pipeline._calculate_quota_remaining(None)
mock_db = MagicMock()
remaining = pipeline._calculate_balance_remaining(mock_db, None)
assert remaining is None
@@ -266,12 +253,10 @@ class TestPipelineAuthentication:
assert exc_info.value.status_code == 401
def test_authenticate_client_quota_exceeded(self, pipeline: ApiRequestPipeline) -> None:
"""测试配额超限时抛出异常"""
def test_authenticate_client_balance_exceeded(self, pipeline: ApiRequestPipeline) -> None:
"""测试余额不足时抛出异常"""
mock_user = MagicMock()
mock_user.id = "user-123"
mock_user.quota_usd = 100.0
mock_user.used_usd = 100.0
mock_api_key = MagicMock()
mock_api_key.id = "key-123"
@@ -294,13 +279,17 @@ class TestPipelineAuthentication:
):
with patch.object(
pipeline.usage_service,
"check_user_quota",
return_value=(False, "额不足"),
"check_request_balance",
return_value=(False, "额不足"),
):
from src.core.exceptions import QuotaExceededException
with patch(
"src.api.base.pipeline.WalletService.get_balance_snapshot",
return_value=0.0,
):
from src.core.exceptions import BalanceInsufficientException
with pytest.raises(QuotaExceededException):
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
with pytest.raises(BalanceInsufficientException):
pipeline._authenticate_client(mock_request, mock_db, mock_adapter)
class TestPipelineAdminAuth:

View File

@@ -0,0 +1,202 @@
from __future__ import annotations
from decimal import Decimal
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from src.api.wallet.routes import router as wallet_router
from src.database import get_db
def _build_wallet_app(
db: MagicMock,
monkeypatch: pytest.MonkeyPatch,
*,
payload: dict[str, object],
user_id: str = "user-1",
) -> TestClient:
app = FastAPI()
app.include_router(wallet_router)
app.dependency_overrides[get_db] = lambda: db
async def _fake_pipeline_run(*, adapter: object, http_request: object, db: MagicMock, mode: object) -> object:
_ = http_request, mode
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id=user_id),
ensure_json_body=lambda: payload,
add_audit_metadata=lambda **_: None,
)
return await adapter.handle(context)
monkeypatch.setattr("src.api.wallet.routes.pipeline.run", _fake_pipeline_run)
return TestClient(app)
def test_create_refund_route_maps_uncredited_order_to_400(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
payload = {"amount_usd": 2.0, "payment_order_id": "order-1"}
client = _build_wallet_app(db, monkeypatch, payload=payload)
wallet = SimpleNamespace(id="wallet-1")
payment_order = SimpleNamespace(id="order-1", wallet_id="wallet-1", payment_method="alipay")
monkeypatch.setattr(
"src.api.wallet.routes.WalletService.get_or_create_wallet",
lambda _db, user: wallet,
)
db.query.return_value.filter.return_value.first.return_value = payment_order
def _raise(*args: object, **kwargs: object) -> object:
raise ValueError("payment order is not refundable")
monkeypatch.setattr("src.api.wallet.routes.WalletService.create_refund_request", _raise)
response = client.post("/api/wallet/refunds", json=payload)
assert response.status_code == 400
assert "not refundable" in response.json()["detail"]
db.rollback.assert_called_once()
db.commit.assert_not_called()
def test_create_refund_route_maps_reserved_wallet_amount_to_400(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
payload = {"amount_usd": 2.0}
client = _build_wallet_app(db, monkeypatch, payload=payload)
wallet = SimpleNamespace(id="wallet-1")
monkeypatch.setattr(
"src.api.wallet.routes.WalletService.get_or_create_wallet",
lambda _db, user: wallet,
)
def _raise(*args: object, **kwargs: object) -> object:
raise ValueError("refund amount exceeds available refundable recharge balance")
monkeypatch.setattr("src.api.wallet.routes.WalletService.create_refund_request", _raise)
response = client.post("/api/wallet/refunds", json=payload)
assert response.status_code == 400
assert "available refundable recharge balance" in response.json()["detail"]
db.rollback.assert_called_once()
db.commit.assert_not_called()
def test_create_refund_route_passes_default_order_refund_mode_and_commits(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
payload = {"amount_usd": 2.0, "payment_order_id": "order-1", "reason": "test"}
client = _build_wallet_app(db, monkeypatch, payload=payload)
wallet = SimpleNamespace(id="wallet-1")
payment_order = SimpleNamespace(id="order-1", wallet_id="wallet-1", payment_method="alipay")
refund = SimpleNamespace(
id="refund-1",
refund_no="rf-1",
payment_order_id="order-1",
source_type="payment_order",
source_id="order-1",
refund_mode="original_channel",
amount_usd=Decimal("2.00000000"),
status="pending_approval",
reason="test",
failure_reason=None,
gateway_refund_id=None,
payout_method=None,
payout_reference=None,
payout_proof=None,
created_at="2026-03-07T00:00:00Z",
updated_at="2026-03-07T00:00:00Z",
processed_at=None,
completed_at=None,
)
monkeypatch.setattr(
"src.api.wallet.routes.WalletService.get_or_create_wallet",
lambda _db, user: wallet,
)
db.query.return_value.filter.return_value.first.return_value = payment_order
captured: dict[str, object] = {}
def _create_refund_request(_db: MagicMock, **kwargs: object) -> object:
captured.update(kwargs)
return refund
monkeypatch.setattr(
"src.api.wallet.routes.WalletService.create_refund_request",
_create_refund_request,
)
response = client.post("/api/wallet/refunds", json=payload)
assert response.status_code == 200
body = response.json()
assert body["id"] == "refund-1"
assert body["status"] == "pending_approval"
assert captured["refund_mode"] == "original_channel"
assert captured["source_type"] == "payment_order"
assert captured["source_id"] == "order-1"
assert captured["payment_order"] is payment_order
db.commit.assert_called_once()
db.refresh.assert_called_once_with(refund)
db.rollback.assert_not_called()
def test_create_refund_route_uses_offline_payout_for_manual_recharge(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
payload = {"amount_usd": 2.0, "payment_order_id": "order-2"}
client = _build_wallet_app(db, monkeypatch, payload=payload)
wallet = SimpleNamespace(id="wallet-1")
payment_order = SimpleNamespace(id="order-2", wallet_id="wallet-1", payment_method="admin_manual")
refund = SimpleNamespace(
id="refund-2",
refund_no="rf-2",
payment_order_id="order-2",
source_type="payment_order",
source_id="order-2",
refund_mode="offline_payout",
amount_usd=Decimal("2.00000000"),
status="pending_approval",
reason=None,
failure_reason=None,
gateway_refund_id=None,
payout_method=None,
payout_reference=None,
payout_proof=None,
created_at="2026-03-07T00:00:00Z",
updated_at="2026-03-07T00:00:00Z",
processed_at=None,
completed_at=None,
)
monkeypatch.setattr(
"src.api.wallet.routes.WalletService.get_or_create_wallet",
lambda _db, user: wallet,
)
db.query.return_value.filter.return_value.first.return_value = payment_order
captured: dict[str, object] = {}
def _create_refund_request(_db: MagicMock, **kwargs: object) -> object:
captured.update(kwargs)
return refund
monkeypatch.setattr(
"src.api.wallet.routes.WalletService.create_refund_request",
_create_refund_request,
)
response = client.post("/api/wallet/refunds", json=payload)
assert response.status_code == 200
assert captured["refund_mode"] == "offline_payout"

View File

@@ -38,15 +38,16 @@ async def test_video_cancel_openai_route_end_to_end(monkeypatch: pytest.MonkeyPa
"""
pipeline = ApiRequestPipeline()
# Pipeline auth/quota/audit shortcuts
user = SimpleNamespace(id="u1", username="u1", role="user", quota_usd=None, used_usd=0.0)
# Pipeline auth/balance/audit shortcuts
user = SimpleNamespace(id="u1", username="u1", role="user")
api_key = SimpleNamespace(id="ak1", user_id="u1", is_standalone=False)
monkeypatch.setattr(
pipeline.auth_service, "authenticate_api_key", lambda _db, _k: (user, api_key)
)
monkeypatch.setattr(
pipeline.usage_service, "check_user_quota", lambda *_args, **_kwargs: (True, "ok")
pipeline.usage_service, "check_request_balance", lambda *_args, **_kwargs: (True, "ok")
)
monkeypatch.setattr(pipeline, "_calculate_balance_remaining", lambda *_args, **_kwargs: None)
monkeypatch.setattr(pipeline.audit_service, "log_event", MagicMock())
# DB stubs used by TaskService.cancel
@@ -152,15 +153,16 @@ async def test_video_cancel_gemini_route_end_to_end(monkeypatch: pytest.MonkeyPa
"""
pipeline = ApiRequestPipeline()
# Pipeline auth/quota/audit shortcuts
user = SimpleNamespace(id="u1", username="u1", role="user", quota_usd=None, used_usd=0.0)
# Pipeline auth/balance/audit shortcuts
user = SimpleNamespace(id="u1", username="u1", role="user")
api_key = SimpleNamespace(id="ak1", user_id="u1", is_standalone=False)
monkeypatch.setattr(
pipeline.auth_service, "authenticate_api_key", lambda _db, _k: (user, api_key)
)
monkeypatch.setattr(
pipeline.usage_service, "check_user_quota", lambda *_args, **_kwargs: (True, "ok")
pipeline.usage_service, "check_request_balance", lambda *_args, **_kwargs: (True, "ok")
)
monkeypatch.setattr(pipeline, "_calculate_balance_remaining", lambda *_args, **_kwargs: None)
monkeypatch.setattr(pipeline.audit_service, "log_event", MagicMock())
# DB stubs used by TaskService.cancel

View File

@@ -13,6 +13,7 @@ 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.models.database import UserRole
from src.services.auth.service import (
@@ -251,7 +252,6 @@ class TestAPIKeyAuthentication:
mock_api_key.is_locked = False
mock_api_key.expires_at = None
mock_api_key.user = mock_user
mock_api_key.balance_used_usd = 0.0
mock_db = MagicMock()
mock_db.query.return_value.options.return_value.filter.return_value.first.return_value = (
@@ -259,11 +259,7 @@ class TestAPIKeyAuthentication:
)
with patch("src.services.auth.service.ApiKey.hash_key", return_value="hashed_key"):
with patch(
"src.services.auth.service.ApiKeyService.check_balance",
return_value=(True, 100.0),
):
result = AuthService.authenticate_api_key(mock_db, "sk-test-key")
result = AuthService.authenticate_api_key(mock_db, "sk-test-key")
assert result is not None
assert result[0] == mock_user
@@ -283,7 +279,6 @@ class TestAPIKeyAuthentication:
mock_api_key.is_locked = False
mock_api_key.expires_at = None
mock_api_key.user = mock_user
mock_api_key.balance_used_usd = 0.0
mock_db = MagicMock()
mock_db.expire_on_commit = True
@@ -298,11 +293,7 @@ class TestAPIKeyAuthentication:
with patch("src.services.auth.service._should_update_last_used", return_value=True):
with patch("src.services.auth.service.ApiKey.hash_key", return_value="hashed_key"):
with patch(
"src.services.auth.service.ApiKeyService.check_balance",
return_value=(True, 100.0),
):
result = AuthService.authenticate_api_key(mock_db, "sk-test-key")
result = AuthService.authenticate_api_key(mock_db, "sk-test-key")
assert result is not None
assert mock_db.expire_on_commit is True
@@ -336,6 +327,51 @@ class TestAPIKeyAuthentication:
assert result is None
def test_authenticate_api_key_locked_non_standalone_raises_forbidden(self) -> None:
"""测试普通用户 API Key 被锁定会拒绝认证"""
mock_api_key = MagicMock()
mock_api_key.is_active = True
mock_api_key.is_locked = True
mock_api_key.is_standalone = False
mock_db = MagicMock()
mock_db.query.return_value.options.return_value.filter.return_value.first.return_value = (
mock_api_key
)
with patch("src.services.auth.service.ApiKey.hash_key", return_value="hashed_key"):
with pytest.raises(ForbiddenException):
AuthService.authenticate_api_key(mock_db, "sk-locked-key")
def test_authenticate_api_key_locked_standalone_can_pass(self) -> None:
"""测试独立 Key 即使历史上被锁定也不因锁定字段拒绝认证"""
mock_user = MagicMock()
mock_user.id = "user-standalone"
mock_user.email = "standalone@example.com"
mock_user.is_active = True
mock_user.is_deleted = False
mock_api_key = MagicMock()
mock_api_key.id = "key-standalone"
mock_api_key.is_active = True
mock_api_key.is_locked = True
mock_api_key.is_standalone = True
mock_api_key.expires_at = None
mock_api_key.user = mock_user
mock_db = MagicMock()
mock_db.query.return_value.options.return_value.filter.return_value.first.return_value = (
mock_api_key
)
with patch("src.services.auth.service._should_update_last_used", return_value=False):
with patch("src.services.auth.service.ApiKey.hash_key", return_value="hashed_key"):
result = AuthService.authenticate_api_key(mock_db, "sk-standalone-key")
assert result is not None
assert result[0] == mock_user
assert result[1] == mock_api_key
def test_authenticate_api_key_expired(self) -> None:
"""测试 API Key 已过期"""
mock_api_key = MagicMock()

View File

@@ -1,166 +0,0 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock
import pytest
from src.services.system.maintenance_scheduler import MaintenanceScheduler
@pytest.mark.asyncio
async def test_user_quota_reset_disabled(monkeypatch):
scheduler = MaintenanceScheduler()
mock_db = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.create_session",
lambda: mock_db,
)
def fake_get_config(cls, db, key, default=None):
if key == "enable_user_quota_reset":
return False
return default
mock_set_config = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.get_config",
classmethod(fake_get_config),
)
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.set_config",
mock_set_config,
)
await scheduler._perform_user_quota_reset()
assert not mock_db.query.called
assert not mock_db.commit.called
assert not mock_set_config.called
@pytest.mark.asyncio
async def test_user_quota_reset_not_due_skips(monkeypatch):
scheduler = MaintenanceScheduler()
mock_db = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.create_session",
lambda: mock_db,
)
last_reset_at = (datetime.now(timezone.utc) - timedelta(days=1)).isoformat()
def fake_get_config(cls, db, key, default=None):
if key == "enable_user_quota_reset":
return True
if key == "user_quota_reset_interval_days":
return 2
if key == "user_quota_last_reset_at":
return last_reset_at
return default
mock_set_config = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.get_config",
classmethod(fake_get_config),
)
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.set_config",
mock_set_config,
)
await scheduler._perform_user_quota_reset()
assert not mock_db.query.called
assert not mock_db.commit.called
assert not mock_set_config.called
@pytest.mark.asyncio
async def test_user_quota_reset_due_runs(monkeypatch):
scheduler = MaintenanceScheduler()
mock_db = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.create_session",
lambda: mock_db,
)
last_reset_at = (datetime.now(timezone.utc) - timedelta(days=2)).isoformat()
def fake_get_config(cls, db, key, default=None):
if key == "enable_user_quota_reset":
return True
if key == "user_quota_reset_interval_days":
return 2
if key == "user_quota_last_reset_at":
return last_reset_at
return default
mock_set_config = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.get_config",
classmethod(fake_get_config),
)
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.set_config",
mock_set_config,
)
mock_query = MagicMock()
mock_filter = MagicMock()
mock_filter.update.return_value = 7
mock_query.filter.return_value = mock_filter
mock_db.query.return_value = mock_query
await scheduler._perform_user_quota_reset()
mock_db.query.assert_called_once()
mock_filter.update.assert_called_once()
_, update_kwargs = mock_filter.update.call_args
assert update_kwargs["synchronize_session"] is False
mock_db.commit.assert_called_once()
mock_set_config.assert_called_once()
@pytest.mark.asyncio
async def test_user_quota_reset_invalid_interval_defaults_to_1(monkeypatch):
scheduler = MaintenanceScheduler()
mock_db = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.create_session",
lambda: mock_db,
)
def fake_get_config(cls, db, key, default=None):
if key == "enable_user_quota_reset":
return True
if key == "user_quota_reset_interval_days":
return "abc"
if key == "user_quota_last_reset_at":
return None
return default
mock_set_config = MagicMock()
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.get_config",
classmethod(fake_get_config),
)
monkeypatch.setattr(
"src.services.system.maintenance_scheduler.SystemConfigService.set_config",
mock_set_config,
)
mock_query = MagicMock()
mock_filter = MagicMock()
mock_filter.update.return_value = 1
mock_query.filter.return_value = mock_filter
mock_db.query.return_value = mock_query
await scheduler._perform_user_quota_reset()
mock_db.commit.assert_called_once()
mock_set_config.assert_called_once()

View File

@@ -0,0 +1,328 @@
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from src.models.database import PaymentOrder, Wallet
from src.services.payment.gateway import get_payment_gateway
from src.services.payment import PaymentService
CALLBACK_SECRET = "test-callback-secret"
def _sign_payload(payload: dict[str, object]) -> str:
gateway = get_payment_gateway("alipay")
signature = gateway.build_callback_signature(payload=payload, callback_secret=CALLBACK_SECRET)
assert signature is not None
return signature
def test_create_recharge_order_creates_pending_order(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
user = SimpleNamespace(id="u1")
wallet = SimpleNamespace(id="w1", status="active")
monkeypatch.setattr(
"src.services.payment.service.WalletService.get_or_create_wallet",
lambda _db, user: wallet,
)
order = PaymentService.create_recharge_order(
db,
user=user,
amount_usd="12.5",
payment_method="alipay",
pay_amount="88.00",
pay_currency="CNY",
exchange_rate="7.04",
)
db.add.assert_called_once()
db.flush.assert_called_once()
assert order.wallet_id == "w1"
assert order.user_id == "u1"
assert order.status == "pending"
assert order.payment_method == "alipay"
assert order.gateway_order_id == f"ali_{order.order_no}"
assert isinstance(order.gateway_response, dict)
assert order.gateway_response["gateway"] == "alipay"
assert Decimal(order.amount_usd) == Decimal("12.50000000")
assert Decimal(order.refundable_amount_usd) == Decimal("0")
def test_refresh_order_status_marks_expired_pending_order() -> None:
order = PaymentOrder(
id="po-expired",
order_no="order-expired",
wallet_id="w1",
user_id="u1",
amount_usd=Decimal("2.00000000"),
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("2.00000000"),
payment_method="wechat",
status="pending",
)
order.expires_at = datetime(2000, 1, 1, tzinfo=timezone.utc)
changed = PaymentService.refresh_order_status(order)
assert changed is True
assert order.status == "expired"
def test_handle_callback_is_idempotent_for_processed_callback(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
callback = SimpleNamespace(status="processed", payment_order_id="po1")
monkeypatch.setattr(
"src.services.payment.service.PaymentService.log_callback",
lambda *args, **kwargs: (callback, False),
)
result = PaymentService.handle_callback(
db,
payment_method="alipay",
callback_key="cb1",
payload={"foo": "bar"},
callback_signature=_sign_payload({"foo": "bar"}),
callback_secret=CALLBACK_SECRET,
order_no="order-1",
)
assert result["ok"] is True
assert result["duplicate"] is True
assert result["credited"] is False
assert result["order_id"] == "po1"
def test_handle_callback_fails_on_amount_mismatch(monkeypatch: pytest.MonkeyPatch) -> None:
db = MagicMock()
callback = SimpleNamespace(
status="received",
payment_order_id=None,
order_no=None,
gateway_order_id=None,
error_message=None,
processed_at=None,
)
order = SimpleNamespace(
id="po1",
order_no="order-1",
gateway_order_id=None,
amount_usd=Decimal("10.00000000"),
)
monkeypatch.setattr(
"src.services.payment.service.PaymentService.log_callback",
lambda *args, **kwargs: (callback, True),
)
monkeypatch.setattr(
"src.services.payment.service.PaymentService.get_order",
lambda *args, **kwargs: order,
)
result = PaymentService.handle_callback(
db,
payment_method="wechat",
callback_key="cb2",
payload={"foo": "bar"},
callback_signature=_sign_payload({"foo": "bar"}),
callback_secret=CALLBACK_SECRET,
order_no="order-1",
amount_usd="9.99",
)
assert result["ok"] is False
assert "mismatch" in result["error"]
assert callback.status == "failed"
def test_handle_callback_rejects_invalid_signature() -> None:
db = MagicMock()
result = PaymentService.handle_callback(
db,
payment_method="alipay",
callback_key="cb-invalid-signature",
payload={"foo": "bar"},
callback_signature="invalid-signature",
callback_secret=CALLBACK_SECRET,
order_no="order-1",
amount_usd="9.99",
)
assert result["ok"] is False
assert "signature" in result["error"]
def test_credit_order_applies_wallet_recharge_once(monkeypatch: pytest.MonkeyPatch) -> None:
order = PaymentOrder(
id="po1",
order_no="order-1",
wallet_id="w1",
user_id="u1",
amount_usd=Decimal("5.00000000"),
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("0"),
payment_method="alipay",
status="pending",
)
wallet = Wallet(
id="w1",
user_id="u1",
balance=Decimal("1.00000000"),
status="active",
total_recharged=Decimal("1.00000000"),
total_consumed=Decimal("0"),
total_refunded=Decimal("0"),
total_adjusted=Decimal("0"),
currency="USD",
)
class DummyOrderQuery:
def filter(self, *args: object, **kwargs: object) -> "DummyOrderQuery":
return self
def with_for_update(self) -> "DummyOrderQuery":
return self
def one_or_none(self) -> PaymentOrder:
return order
class DummyWalletQuery:
def filter(self, *args: object, **kwargs: object) -> "DummyWalletQuery":
return self
def first(self) -> Wallet:
return wallet
db = MagicMock()
db.query.side_effect = [DummyOrderQuery(), DummyWalletQuery()]
create_tx = MagicMock()
monkeypatch.setattr(
"src.services.payment.service.WalletService.create_wallet_transaction",
create_tx,
)
updated, credited = PaymentService.credit_order(
db,
order=order,
gateway_order_id="gw-1",
pay_amount="36.00",
pay_currency="CNY",
exchange_rate="7.20",
)
assert credited is True
assert updated.status == "credited"
assert updated.gateway_order_id == "gw-1"
assert updated.paid_at is not None
assert updated.credited_at is not None
assert Decimal(updated.refundable_amount_usd) == Decimal("5.00000000")
create_tx.assert_called_once()
def test_expire_order_marks_pending_order_expired() -> None:
order = PaymentOrder(
id="po2",
order_no="order-2",
wallet_id="w1",
user_id="u1",
amount_usd=Decimal("3.00000000"),
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("3.00000000"),
payment_method="wechat",
status="pending",
)
order.expires_at = datetime.now(timezone.utc) + timedelta(minutes=30)
class DummyOrderQuery:
def filter(self, *args: object, **kwargs: object) -> "DummyOrderQuery":
return self
def with_for_update(self) -> "DummyOrderQuery":
return self
def one_or_none(self) -> PaymentOrder:
return order
db = MagicMock()
db.query.return_value = DummyOrderQuery()
updated, changed = PaymentService.expire_order(db, order=order, reason="ops_close")
assert changed is True
assert updated.status == "expired"
assert updated.gateway_response["expire_reason"] == "ops_close"
assert "expired_at" in updated.gateway_response
def test_expire_order_is_idempotent_for_existing_expired_order() -> None:
order = PaymentOrder(
id="po3",
order_no="order-3",
wallet_id="w1",
user_id="u1",
amount_usd=Decimal("1.00000000"),
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("1.00000000"),
payment_method="alipay",
status="expired",
)
class DummyOrderQuery:
def filter(self, *args: object, **kwargs: object) -> "DummyOrderQuery":
return self
def with_for_update(self) -> "DummyOrderQuery":
return self
def one_or_none(self) -> PaymentOrder:
return order
db = MagicMock()
db.query.return_value = DummyOrderQuery()
updated, changed = PaymentService.expire_order(db, order=order)
assert changed is False
assert updated.status == "expired"
def test_list_orders_expires_overdue_pending_before_filtered_count() -> None:
expiry_query = MagicMock()
expiry_query.filter.return_value = expiry_query
expiry_query.update.return_value = 2
list_query = MagicMock()
list_query.filter.return_value = list_query
list_query.count.return_value = 1
ordered_query = MagicMock()
list_query.order_by.return_value = ordered_query
ordered_query.offset.return_value = ordered_query
ordered_query.limit.return_value = ordered_query
ordered_query.all.return_value = ["order-1"]
db = MagicMock()
db.query.side_effect = [expiry_query, list_query]
items, total, changed = PaymentService.list_orders(
db,
status="pending",
payment_method="alipay",
limit=5,
offset=10,
)
assert items == ["order-1"]
assert total == 1
assert changed is True
expiry_query.update.assert_called_once()
ordered_query.offset.assert_called_once_with(10)
ordered_query.limit.assert_called_once_with(5)

View File

@@ -1304,6 +1304,9 @@ async def test_record_usage_batch_updates_when_status_completed_billing_pending(
def filter(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def with_for_update(self) -> "DummyQuery":
return self
def all(self) -> list[Any]:
return self._all_result
@@ -1362,10 +1365,10 @@ async def test_record_usage_batch_updates_when_status_completed_billing_pending(
@pytest.mark.asyncio
async def test_record_usage_batch_uses_bulk_insert_mappings_for_new_records(
async def test_record_usage_batch_uses_orm_insert_for_new_records(
monkeypatch: Any,
) -> None:
"""确保批量新建 Usage 走 bulk_insert_mappings"""
"""确保批量新建 Usage 走 ORM add 路径,以便复用统一结算逻辑"""
from src.models.database import Usage
from src.services.usage.service import UsageService
@@ -1380,27 +1383,18 @@ async def test_record_usage_batch_uses_bulk_insert_mappings_for_new_records(
def filter(self, *args: Any, **kwargs: Any) -> "DummyQuery":
return self
def with_for_update(self) -> "DummyQuery":
return self
def all(self) -> list[Any]:
return self._all_result
inserted = Usage(
request_id="req-usage-batch-new",
provider_name="openai",
model="gpt-4",
status="completed",
billing_status="settled",
)
usage_query_calls = {"count": 0}
def _query_side_effect(model: Any) -> Any:
if model is Usage:
usage_query_calls["count"] += 1
if usage_query_calls["count"] == 1:
# existing_records
return DummyQuery([])
# inserted_records
return DummyQuery([inserted])
return DummyQuery([])
return DummyQuery([])
db = MagicMock()
@@ -1431,13 +1425,10 @@ async def test_record_usage_batch_uses_bulk_insert_mappings_for_new_records(
],
)
db.bulk_insert_mappings.assert_called_once()
args, _kwargs = db.bulk_insert_mappings.call_args
assert args[0] is Usage
mappings = args[1]
assert isinstance(mappings, list) and len(mappings) == 1
assert mappings[0]["request_id"] == "req-usage-batch-new"
assert mappings[0]["billing_status"] == "settled"
assert mappings[0].get("finalized_at") is not None
db.add.assert_called_once()
db.bulk_insert_mappings.assert_not_called()
assert result and result[0] is inserted
assert result and isinstance(result[0], Usage)
assert result[0].request_id == "req-usage-batch-new"
assert result[0].billing_status == "settled"
assert result[0].finalized_at is not None

View File

@@ -3,15 +3,17 @@ UsageService 测试
测试用量统计服务的核心功能:
- 成本计算
- 配额检查
- 钱包准入检查
- 用量统计查询
"""
from decimal import Decimal
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from src.services.usage.service import UsageService
from src.services.wallet import WalletAccessResult
class TestCostCalculation:
@@ -120,14 +122,12 @@ class TestCostCalculation:
assert total_cost == 0
class TestQuotaCheck:
"""测试配额检查"""
class TestBalanceCheck:
"""测试钱包准入检查"""
def test_check_user_quota_sufficient(self) -> None:
"""测试额充足"""
def test_check_request_balance_sufficient(self) -> None:
"""测试额充足"""
mock_user = MagicMock()
mock_user.quota_usd = 100.0
mock_user.used_usd = 30.0
mock_user.role = MagicMock()
mock_user.role.value = "user"
@@ -136,15 +136,19 @@ class TestQuotaCheck:
mock_db = MagicMock()
is_ok, message = UsageService.check_user_quota(mock_db, mock_user, api_key=mock_api_key)
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(True, Decimal("70"), "OK"),
):
is_ok, message = UsageService.check_request_balance(
mock_db, mock_user, api_key=mock_api_key
)
assert is_ok is True
def test_check_user_quota_exceeded(self) -> None:
"""测试配额超限(当有预估成本时)"""
def test_check_request_balance_exceeded(self) -> None:
"""测试余额耗尽时拦截新请求"""
mock_user = MagicMock()
mock_user.quota_usd = 100.0
mock_user.used_usd = 99.0 # 接近配额上限
mock_user.role = MagicMock()
mock_user.role.value = "user"
@@ -153,19 +157,20 @@ class TestQuotaCheck:
mock_db = MagicMock()
# 当预估成本超过剩余配额时应该返回 False
is_ok, message = UsageService.check_user_quota(
mock_db, mock_user, estimated_cost=5.0, api_key=mock_api_key
)
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(False, Decimal("0"), "钱包余额不足"),
):
is_ok, message = UsageService.check_request_balance(
mock_db, mock_user, estimated_cost=5.0, api_key=mock_api_key
)
assert is_ok is False
assert "" in message
assert "" in message
def test_check_user_quota_no_limit(self) -> None:
def test_check_request_balance_no_limit(self) -> None:
"""测试无配额限制None"""
mock_user = MagicMock()
mock_user.quota_usd = None
mock_user.used_usd = 1000.0
mock_user.role = MagicMock()
mock_user.role.value = "user"
@@ -174,17 +179,21 @@ class TestQuotaCheck:
mock_db = MagicMock()
is_ok, message = UsageService.check_user_quota(mock_db, mock_user, api_key=mock_api_key)
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(True, None, "OK"),
):
is_ok, message = UsageService.check_request_balance(
mock_db, mock_user, api_key=mock_api_key
)
assert is_ok is True
def test_check_user_quota_admin_bypass(self) -> None:
"""测试管理员绕过额检查"""
def test_check_request_balance_admin_bypass(self) -> None:
"""测试管理员绕过额检查"""
from src.models.database import UserRole
mock_user = MagicMock()
mock_user.quota_usd = 0.0
mock_user.used_usd = 1000.0
mock_user.role = UserRole.ADMIN
mock_api_key = MagicMock()
@@ -192,55 +201,58 @@ class TestQuotaCheck:
mock_db = MagicMock()
is_ok, message = UsageService.check_user_quota(mock_db, mock_user, api_key=mock_api_key)
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(True, None, "OK"),
):
is_ok, message = UsageService.check_request_balance(
mock_db, mock_user, api_key=mock_api_key
)
assert is_ok is True
def test_check_standalone_api_key_balance(self) -> None:
"""测试独立 API Key 余额检查"""
"""测试独立 API Key 余额充足"""
mock_user = MagicMock()
mock_user.quota_usd = 0.0
mock_user.used_usd = 0.0
mock_user.role = MagicMock()
mock_user.role.value = "user"
mock_api_key = MagicMock()
mock_api_key.is_standalone = True
mock_api_key.current_balance_usd = 50.0
mock_api_key.balance_used_usd = 10.0
mock_db = MagicMock()
is_ok, message = UsageService.check_user_quota(mock_db, mock_user, api_key=mock_api_key)
with patch(
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(True, Decimal("40"), "OK"),
):
is_ok, message = UsageService.check_request_balance(
mock_db, mock_user, api_key=mock_api_key
)
assert is_ok is True
def test_check_standalone_api_key_insufficient_balance(self) -> None:
"""测试独立 API Key 余额不足"""
"""测试独立 API Key 余额耗尽时拦截"""
mock_user = MagicMock()
mock_user.quota_usd = 100.0
mock_user.used_usd = 0.0
mock_user.role = MagicMock()
mock_user.role.value = "user"
mock_api_key = MagicMock()
mock_api_key.is_standalone = True
mock_api_key.current_balance_usd = 10.0
mock_api_key.balance_used_usd = 9.0 # 剩余 $1
mock_db = MagicMock()
# 需要 mock ApiKeyService.get_remaining_balance
with patch(
"src.services.user.apikey.ApiKeyService.get_remaining_balance",
return_value=1.0,
"src.services.wallet.WalletService.check_request_allowed",
return_value=WalletAccessResult(False, Decimal("0"), "钱包余额不足"),
):
# 预估成本 $5 超过剩余余额 $1
is_ok, message = UsageService.check_user_quota(
is_ok, message = UsageService.check_request_balance(
mock_db, mock_user, estimated_cost=5.0, api_key=mock_api_key
)
assert is_ok is False
assert "Key余额不足" in message
class TestUsageStatistics:

View File

@@ -0,0 +1,87 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from src.models.database import PaymentOrder, RefundRequest, User, Wallet
from src.services.user.service import UserService
def test_delete_user_blocks_when_unfinished_refund_exists() -> None:
user = SimpleNamespace(id="user-1", email="u1@example.com")
user_query = MagicMock()
user_query.filter.return_value = user_query
user_query.first.return_value = user
wallet_ids_query = MagicMock()
wallet_ids_query.outerjoin.return_value = wallet_ids_query
wallet_ids_query.filter.return_value = wallet_ids_query
wallet_ids_query.all.return_value = [("wallet-1",)]
refund_query = MagicMock()
refund_query.filter.return_value = refund_query
refund_query.count.return_value = 1
db = MagicMock()
db.in_transaction.return_value = True
def _query(model: object) -> MagicMock:
if model is User:
return user_query
if model is Wallet.id:
return wallet_ids_query
if model is RefundRequest:
return refund_query
raise AssertionError(f"unexpected query target: {model}")
db.query.side_effect = _query
with pytest.raises(ValueError, match="未完结退款"):
UserService.delete_user.__wrapped__(db, "user-1")
db.delete.assert_not_called()
def test_delete_user_blocks_when_unfinished_payment_order_exists() -> None:
user = SimpleNamespace(id="user-2", email="u2@example.com")
user_query = MagicMock()
user_query.filter.return_value = user_query
user_query.first.return_value = user
wallet_ids_query = MagicMock()
wallet_ids_query.outerjoin.return_value = wallet_ids_query
wallet_ids_query.filter.return_value = wallet_ids_query
wallet_ids_query.all.return_value = [("wallet-2",)]
refund_query = MagicMock()
refund_query.filter.return_value = refund_query
refund_query.count.return_value = 0
order_query = MagicMock()
order_query.filter.return_value = order_query
order_query.count.return_value = 2
db = MagicMock()
db.in_transaction.return_value = True
def _query(model: object) -> MagicMock:
if model is User:
return user_query
if model is Wallet.id:
return wallet_ids_query
if model is RefundRequest:
return refund_query
if model is PaymentOrder:
return order_query
raise AssertionError(f"unexpected query target: {model}")
db.query.side_effect = _query
with pytest.raises(ValueError, match="未完结充值订单"):
UserService.delete_user.__wrapped__(db, "user-2")
db.delete.assert_not_called()

View File

@@ -0,0 +1,36 @@
from __future__ import annotations
from unittest.mock import MagicMock
from src.core.enums import UserRole
from src.services.user.service import UserService
def test_list_users_orders_by_created_at_desc_then_id_desc() -> None:
ordered_query = MagicMock()
ordered_query.offset.return_value = ordered_query
ordered_query.limit.return_value = ordered_query
ordered_query.all.return_value = ["user-1"]
query = MagicMock()
query.filter.return_value = query
query.order_by.return_value = ordered_query
db = MagicMock()
db.query.return_value = query
result = UserService.list_users(
db,
skip=5,
limit=10,
role=UserRole.ADMIN,
is_active=True,
)
assert result == ["user-1"]
order_args = query.order_by.call_args.args
assert len(order_args) == 2
assert str(order_args[0]) == "users.created_at DESC"
assert str(order_args[1]) == "users.id DESC"
ordered_query.offset.assert_called_once_with(5)
ordered_query.limit.assert_called_once_with(10)

View File

@@ -0,0 +1,586 @@
from decimal import Decimal
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy.exc import IntegrityError
from src.services.wallet.service import WalletService
def _build_wallet(
*,
wallet_id: str = "wallet-1",
recharge: str = "0",
gift: str = "0",
limit_mode: str = "finite",
status: str = "active",
api_key_id: str | None = None,
) -> SimpleNamespace:
return SimpleNamespace(
id=wallet_id,
balance=Decimal(recharge),
gift_balance=Decimal(gift),
total_recharged=Decimal("0"),
total_consumed=Decimal("0"),
total_refunded=Decimal("0"),
total_adjusted=Decimal("0"),
limit_mode=limit_mode,
currency="USD",
status=status,
api_key_id=api_key_id,
updated_at=None,
)
def _build_locked_db(wallet: SimpleNamespace) -> MagicMock:
db = MagicMock()
query = MagicMock()
query.filter.return_value = query
query.with_for_update.return_value = query
query.one_or_none.return_value = wallet
db.query.return_value = query
return db
def test_get_or_create_wallet_prefers_user_owner_for_non_standalone_key() -> None:
db = MagicMock()
db.query.return_value.filter.return_value.first.return_value = None
db.begin_nested.return_value.__enter__.return_value = None
db.begin_nested.return_value.__exit__.return_value = None
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)
assert wallet is not None
assert wallet.user_id == "user-1"
assert wallet.api_key_id is None
def test_get_or_create_wallet_uses_api_key_owner_for_standalone_key() -> None:
db = MagicMock()
db.query.return_value.filter.return_value.first.return_value = None
db.begin_nested.return_value.__enter__.return_value = None
db.begin_nested.return_value.__exit__.return_value = 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)
assert wallet is not None
assert wallet.user_id is None
assert wallet.api_key_id == "key-1"
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()
api_key = MagicMock()
with patch.object(WalletService, "get_or_create_wallet", return_value=wallet):
result = WalletService.check_request_allowed(db, user=None, api_key=api_key)
assert result.allowed is False
assert result.message == "钱包欠费,请先充值"
assert result.remaining == Decimal("-1.00000000")
def test_get_balance_snapshot_returns_negative_recharge_for_unlimited_wallet() -> None:
wallet = _build_wallet(recharge="-2", gift="100", limit_mode="unlimited")
db = MagicMock()
with patch.object(WalletService, "get_or_create_wallet", return_value=wallet):
snapshot = WalletService.get_balance_snapshot(db, user=None, api_key=MagicMock())
assert snapshot == Decimal("-2.00000000")
def test_admin_adjust_balance_negative_from_gift_spills_to_recharge() -> None:
wallet = _build_wallet(recharge="2", gift="3")
db = _build_locked_db(wallet)
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
amount_usd=Decimal("-10"),
balance_type="gift",
operator_id="admin-1",
description="test",
)
assert wallet.gift_balance == Decimal("0E-8")
assert wallet.balance == Decimal("-5.00000000")
assert wallet.total_adjusted == Decimal("-10.00000000")
assert tx.recharge_balance_after == Decimal("-5.00000000")
assert tx.gift_balance_after == Decimal("0E-8")
assert tx.amount == Decimal("-10.00000000")
assert tx.category == "adjust"
assert tx.reason_code == "adjust_admin"
def test_admin_adjust_balance_negative_from_recharge_then_gift() -> None:
wallet = _build_wallet(recharge="2", gift="3")
db = _build_locked_db(wallet)
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
amount_usd=Decimal("-4"),
balance_type="recharge",
operator_id="admin-1",
description="test",
)
assert wallet.balance == Decimal("0E-8")
assert wallet.gift_balance == Decimal("1.00000000")
assert wallet.total_adjusted == Decimal("-4.00000000")
assert tx.recharge_balance_after == Decimal("0E-8")
assert tx.gift_balance_after == Decimal("1.00000000")
assert tx.amount == Decimal("-4.00000000")
def test_admin_adjust_balance_positive_adds_to_selected_bucket_without_offset() -> None:
wallet = _build_wallet(recharge="-2", gift="5")
db = _build_locked_db(wallet)
tx = WalletService.admin_adjust_balance(
db,
wallet=wallet,
amount_usd=Decimal("1"),
balance_type="gift",
operator_id="admin-1",
description="test",
)
assert wallet.balance == Decimal("-2.00000000")
assert wallet.gift_balance == Decimal("6.00000000")
assert wallet.total_adjusted == Decimal("1.00000000")
assert tx.recharge_balance_after == Decimal("-2.00000000")
assert tx.gift_balance_after == Decimal("6.00000000")
assert tx.amount == Decimal("1.00000000")
assert tx.category == "adjust"
assert tx.reason_code == "adjust_admin"
def test_apply_usage_charge_prefers_gift_then_recharge() -> None:
wallet = _build_wallet(recharge="5", gift="3", limit_mode="finite")
usage = SimpleNamespace(
wallet_id=None,
api_key_id=None,
user_id="user-1",
wallet_balance_before=None,
wallet_balance_after=None,
wallet_recharge_balance_before=None,
wallet_recharge_balance_after=None,
wallet_gift_balance_before=None,
wallet_gift_balance_after=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"))
assert before == Decimal("8.00000000")
assert after == Decimal("2.00000000")
assert wallet.gift_balance == Decimal("0E-8")
assert wallet.balance == Decimal("2.00000000")
assert wallet.total_consumed == Decimal("6.00000000")
assert usage.wallet_balance_before == Decimal("8.00000000")
assert usage.wallet_balance_after == Decimal("2.00000000")
assert usage.wallet_recharge_balance_before == Decimal("5.00000000")
assert usage.wallet_recharge_balance_after == Decimal("2.00000000")
assert usage.wallet_gift_balance_before == Decimal("3.00000000")
assert usage.wallet_gift_balance_after == Decimal("0E-8")
def test_apply_usage_charge_unlimited_wallet_keeps_balances() -> None:
wallet = _build_wallet(recharge="5", gift="3", limit_mode="unlimited")
usage = SimpleNamespace(
wallet_id=None,
api_key_id=None,
user_id="user-1",
wallet_balance_before=None,
wallet_balance_after=None,
wallet_recharge_balance_before=None,
wallet_recharge_balance_after=None,
wallet_gift_balance_before=None,
wallet_gift_balance_after=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"))
assert before == Decimal("8.00000000")
assert after == Decimal("8.00000000")
assert wallet.balance == Decimal("5")
assert wallet.gift_balance == Decimal("3")
assert wallet.total_consumed == Decimal("4.00000000")
assert usage.wallet_balance_before == Decimal("8.00000000")
assert usage.wallet_balance_after == Decimal("8.00000000")
assert usage.wallet_recharge_balance_before == Decimal("5")
assert usage.wallet_recharge_balance_after == Decimal("5")
assert usage.wallet_gift_balance_before == Decimal("3")
assert usage.wallet_gift_balance_after == Decimal("3")
def test_complete_refund_requires_processing_status() -> None:
refund = SimpleNamespace(id="refund-1", status="failed")
db = MagicMock()
query = MagicMock()
query.filter.return_value = query
query.with_for_update.return_value = query
query.one_or_none.return_value = refund
db.query.return_value = query
with pytest.raises(ValueError, match="processing"):
WalletService.complete_refund(db, refund=refund)
def test_get_or_create_wallet_reuses_existing_after_integrity_error() -> None:
existing_wallet = _build_wallet(wallet_id="wallet-existing", recharge="1", gift="0")
db = MagicMock()
nested = MagicMock()
nested.__enter__.return_value = None
nested.__exit__.return_value = None
db.begin_nested.return_value = nested
db.flush.side_effect = IntegrityError("insert", {}, Exception("duplicate key"))
with patch.object(WalletService, "get_wallet", side_effect=[None, existing_wallet]):
wallet = WalletService.get_or_create_wallet(
db,
user=SimpleNamespace(id="user-1"),
api_key=None,
)
assert wallet is existing_wallet
def test_create_refund_request_rejects_uncredited_payment_order() -> None:
wallet = _build_wallet(wallet_id="wallet-1", recharge="10", gift="0")
payment_order = SimpleNamespace(
id="order-1",
wallet_id="wallet-1",
status="pending",
refundable_amount_usd=Decimal("10"),
)
wallet_query = MagicMock()
wallet_query.filter.return_value = wallet_query
wallet_query.with_for_update.return_value = wallet_query
wallet_query.one_or_none.return_value = wallet
order_query = MagicMock()
order_query.filter.return_value = order_query
order_query.with_for_update.return_value = order_query
order_query.one_or_none.return_value = payment_order
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "Wallet":
return wallet_query
if name == "PaymentOrder":
return order_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
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,
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,
)
def test_create_refund_request_reserves_pending_wallet_amount() -> None:
wallet = _build_wallet(wallet_id="wallet-1", recharge="10", gift="0")
wallet_query = MagicMock()
wallet_query.filter.return_value = wallet_query
wallet_query.with_for_update.return_value = wallet_query
wallet_query.one_or_none.return_value = wallet
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "Wallet":
return wallet_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
with patch.object(
WalletService,
"_get_pending_refund_reserved_amount",
return_value=Decimal("9"),
):
with pytest.raises(ValueError, match="available refundable recharge balance"):
WalletService.create_refund_request(
db,
wallet=wallet,
user_id="user-1",
amount_usd=Decimal("2"),
refund_no="rf-2",
source_type="wallet_balance",
source_id=None,
refund_mode="offline_payout",
)
def test_create_refund_request_reserves_pending_order_amount() -> None:
wallet = _build_wallet(wallet_id="wallet-1", recharge="10", gift="0")
payment_order = SimpleNamespace(
id="order-1",
wallet_id="wallet-1",
status="credited",
refundable_amount_usd=Decimal("5"),
)
wallet_query = MagicMock()
wallet_query.filter.return_value = wallet_query
wallet_query.with_for_update.return_value = wallet_query
wallet_query.one_or_none.return_value = wallet
order_query = MagicMock()
order_query.filter.return_value = order_query
order_query.with_for_update.return_value = order_query
order_query.one_or_none.return_value = payment_order
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "Wallet":
return wallet_query
if name == "PaymentOrder":
return order_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
with patch.object(
WalletService,
"_get_pending_refund_reserved_amount",
side_effect=[Decimal("0"), Decimal("4")],
):
with pytest.raises(ValueError, match="available refundable amount"):
WalletService.create_refund_request(
db,
wallet=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,
)
def test_move_refund_to_processing_rejects_double_transition() -> None:
refund = SimpleNamespace(
id="refund-1",
wallet_id="wallet-1",
payment_order_id="order-1",
amount_usd=Decimal("2"),
status="pending_approval",
approved_by=None,
processed_by=None,
processed_at=None,
updated_at=None,
)
wallet = _build_wallet(wallet_id="wallet-1", recharge="10", gift="0")
payment_order = SimpleNamespace(
id="order-1",
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("10"),
)
refund_query = MagicMock()
refund_query.filter.return_value = refund_query
refund_query.with_for_update.return_value = refund_query
refund_query.one_or_none.return_value = refund
wallet_query = MagicMock()
wallet_query.filter.return_value = wallet_query
wallet_query.with_for_update.return_value = wallet_query
wallet_query.one_or_none.return_value = wallet
order_query = MagicMock()
order_query.filter.return_value = order_query
order_query.with_for_update.return_value = order_query
order_query.one_or_none.return_value = payment_order
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "RefundRequest":
return refund_query
if name == "Wallet":
return wallet_query
if name == "PaymentOrder":
return order_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
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")
with pytest.raises(ValueError, match="not approvable"):
WalletService.move_refund_to_processing(db, refund=refund, operator_id="admin-1")
assert first_tx is tx
assert create_tx.call_count == 1
assert refund.status == "processing"
assert payment_order.refunded_amount_usd == Decimal("2.00000000")
assert payment_order.refundable_amount_usd == Decimal("8.00000000")
def test_move_refund_to_processing_rechecks_payment_order_refundable_amount() -> None:
refund = SimpleNamespace(
id="refund-1",
wallet_id="wallet-1",
payment_order_id="order-1",
amount_usd=Decimal("2"),
status="pending_approval",
approved_by=None,
processed_by=None,
processed_at=None,
updated_at=None,
)
wallet = _build_wallet(wallet_id="wallet-1", recharge="10", gift="0")
payment_order = SimpleNamespace(
id="order-1",
refunded_amount_usd=Decimal("0"),
refundable_amount_usd=Decimal("1"),
)
refund_query = MagicMock()
refund_query.filter.return_value = refund_query
refund_query.with_for_update.return_value = refund_query
refund_query.one_or_none.return_value = refund
wallet_query = MagicMock()
wallet_query.filter.return_value = wallet_query
wallet_query.with_for_update.return_value = wallet_query
wallet_query.one_or_none.return_value = wallet
order_query = MagicMock()
order_query.filter.return_value = order_query
order_query.with_for_update.return_value = order_query
order_query.one_or_none.return_value = payment_order
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "RefundRequest":
return refund_query
if name == "Wallet":
return wallet_query
if name == "PaymentOrder":
return order_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
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")
create_tx.assert_not_called()
assert refund.status == "pending_approval"
assert payment_order.refunded_amount_usd == Decimal("0")
assert payment_order.refundable_amount_usd == Decimal("1")
def test_fail_refund_rejects_invalid_status_after_first_failure() -> None:
refund = SimpleNamespace(
id="refund-2",
wallet_id="wallet-1",
payment_order_id=None,
amount_usd=Decimal("1"),
status="processing",
failure_reason=None,
updated_at=None,
)
wallet = _build_wallet(wallet_id="wallet-1", recharge="10", gift="0")
refund_query = MagicMock()
refund_query.filter.return_value = refund_query
refund_query.with_for_update.return_value = refund_query
refund_query.one_or_none.return_value = refund
wallet_query = MagicMock()
wallet_query.filter.return_value = wallet_query
wallet_query.first.return_value = wallet
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "RefundRequest":
return refund_query
if name == "Wallet":
return wallet_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
tx = SimpleNamespace(id="tx-revert")
with patch.object(WalletService, "create_wallet_transaction", return_value=tx) as create_tx:
first_tx = WalletService.fail_refund(
db,
refund=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,
reason="retry-failure",
operator_id="admin-1",
)
assert first_tx is tx
assert create_tx.call_count == 1
assert refund.status == "failed"
assert refund.failure_reason == "first-failure"
def test_fail_refund_rejects_succeeded_status() -> None:
refund = SimpleNamespace(
id="refund-succeeded",
wallet_id="wallet-1",
payment_order_id=None,
amount_usd=Decimal("1"),
status="succeeded",
failure_reason=None,
updated_at=None,
)
refund_query = MagicMock()
refund_query.filter.return_value = refund_query
refund_query.with_for_update.return_value = refund_query
refund_query.one_or_none.return_value = refund
db = MagicMock()
def _query(model: object) -> MagicMock:
name = getattr(model, "__name__", "")
if name == "RefundRequest":
return refund_query
raise AssertionError(f"unexpected model query: {name}")
db.query.side_effect = _query
with pytest.raises(ValueError, match="cannot fail refund in status: succeeded"):
WalletService.fail_refund(
db,
refund=refund,
reason="should-not-override",
operator_id="admin-1",
)
assert refund.status == "succeeded"
assert refund.failure_reason is None

View File

@@ -0,0 +1,79 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from src.plugins.auth.api_key import ApiKeyAuthPlugin
def _build_request() -> SimpleNamespace:
return SimpleNamespace(
headers={"x-api-key": "test-key"},
client=SimpleNamespace(host="127.0.0.1"),
)
@pytest.mark.asyncio
async def test_api_key_auth_plugin_uses_standalone_wallet_for_billing() -> None:
plugin = ApiKeyAuthPlugin()
request = _build_request()
db = MagicMock()
user = SimpleNamespace(id="user-1", username="alice", is_admin=False)
api_key = SimpleNamespace(id="key-1", name="standalone-key", is_standalone=True)
wallet = SimpleNamespace(id="wallet-key")
with (
patch(
"src.plugins.auth.api_key.AuthService.authenticate_api_key",
return_value=(user, api_key),
),
patch(
"src.plugins.auth.api_key.UsageService.check_request_balance",
return_value=(True, "OK"),
),
patch("src.plugins.auth.api_key.WalletService.get_wallet", return_value=wallet) as get_wallet,
patch(
"src.plugins.auth.api_key.WalletService.serialize_wallet_summary",
return_value={"id": "wallet-key"},
),
):
context = await plugin.authenticate(request, db)
assert context is not None
assert context.billing_info["billing"]["id"] == "wallet-key"
get_wallet.assert_called_once_with(db, api_key_id="key-1")
@pytest.mark.asyncio
async def test_api_key_auth_plugin_uses_user_wallet_for_normal_key_billing() -> None:
plugin = ApiKeyAuthPlugin()
request = _build_request()
db = MagicMock()
user = SimpleNamespace(id="user-2", username="bob", is_admin=False)
api_key = SimpleNamespace(id="key-2", name="normal-key", is_standalone=False)
wallet = SimpleNamespace(id="wallet-user")
with (
patch(
"src.plugins.auth.api_key.AuthService.authenticate_api_key",
return_value=(user, api_key),
),
patch(
"src.plugins.auth.api_key.UsageService.check_request_balance",
return_value=(True, "OK"),
),
patch("src.plugins.auth.api_key.WalletService.get_wallet", return_value=wallet) as get_wallet,
patch(
"src.plugins.auth.api_key.WalletService.serialize_wallet_summary",
return_value={"id": "wallet-user"},
),
):
context = await plugin.authenticate(request, db)
assert context is not None
assert context.billing_info["billing"]["id"] == "wallet-user"
get_wallet.assert_called_once_with(db, user_id="user-2")

View File

@@ -0,0 +1,68 @@
from __future__ import annotations
from decimal import Decimal
from types import SimpleNamespace
from src.api.serializers.wallet_payment import serialize_admin_wallet, serialize_wallet_transaction
def test_serialize_wallet_transaction_emits_split_balance_numbers() -> None:
tx = SimpleNamespace(
id="tx-1",
category="adjust",
reason_code="adjust_admin",
amount=Decimal("1.25"),
balance_before=Decimal("5.00"),
balance_after=Decimal("6.25"),
recharge_balance_before=Decimal("2.00"),
recharge_balance_after=Decimal("3.25"),
gift_balance_before=Decimal("3.00"),
gift_balance_after=Decimal("3.00"),
link_type="admin_action",
link_id="wallet-1",
operator_id="admin-1",
description="test",
created_at="2026-03-07T00:00:00Z",
)
payload = serialize_wallet_transaction(tx)
assert payload["recharge_balance_before"] == 2.0
assert payload["recharge_balance_after"] == 3.25
assert payload["gift_balance_before"] == 3.0
assert payload["gift_balance_after"] == 3.0
def test_serialize_admin_wallet_omits_version(monkeypatch) -> None:
wallet = SimpleNamespace(
id="wallet-1",
user_id="user-1",
api_key_id=None,
user=SimpleNamespace(username="alice"),
api_key=None,
created_at="2026-03-07T00:00:00Z",
)
monkeypatch.setattr(
"src.api.serializers.wallet_payment.WalletService.serialize_wallet_summary",
lambda _wallet: {
"balance": 10.0,
"recharge_balance": 7.0,
"gift_balance": 3.0,
"refundable_balance": 7.0,
"currency": "USD",
"status": "active",
"limit_mode": "finite",
"unlimited": False,
"total_recharged": 10.0,
"total_consumed": 0.0,
"total_refunded": 0.0,
"total_adjusted": 0.0,
"updated_at": "2026-03-07T00:00:00Z",
},
)
payload = serialize_admin_wallet(wallet)
assert "version" not in payload
assert payload["owner_name"] == "alice"