mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
feat(wallet): 钱包系统替代配额系统,新增支付与退款机制
- 新增钱包余额管理、充值、扣费、退款完整流程 - 新增支付网关抽象层(支持手动/支付宝/微信) - 用量计费从配额系统迁移到钱包余额扣费 - 新增管理员钱包管理与支付订单管理页面 - 新增用户钱包中心页面 - 移除独立 Key 锁定机制,统一由钱包余额控制 - 新增相关 API 路由、序列化器与数据库迁移 - 新增钱包、支付、退款相关测试
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
328
tests/services/test_payment_service.py
Normal file
328
tests/services/test_payment_service.py
Normal 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)
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
87
tests/services/test_user_service_delete_user.py
Normal file
87
tests/services/test_user_service_delete_user.py
Normal 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()
|
||||
36
tests/services/test_user_service_list_users.py
Normal file
36
tests/services/test_user_service_list_users.py
Normal 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)
|
||||
586
tests/services/test_wallet_service_rules.py
Normal file
586
tests/services/test_wallet_service_rules.py
Normal 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
|
||||
Reference in New Issue
Block a user