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

@@ -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