431 lines
18 KiB
Python
431 lines
18 KiB
Python
import uuid
|
|
import os
|
|
import hashlib
|
|
import hmac
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from threading import Barrier
|
|
from unittest.mock import Mock, patch
|
|
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
|
|
from database import SessionLocal
|
|
from main import app, seed
|
|
from models import Avatar, TokenAccount, TokenPaymentOrder, TokenPlan, TokenUsage, User
|
|
from routers.chat import _resolve_reply, _stream_reply
|
|
from services.token_billing import (
|
|
DEFAULT_TOKEN_GRANT,
|
|
InsufficientTokensError,
|
|
get_or_create_account,
|
|
release_reservation,
|
|
reserve_avatar_tokens,
|
|
settle_reservation,
|
|
)
|
|
|
|
client = TestClient(app)
|
|
|
|
|
|
def _enable_huihui_payment_login(context):
|
|
db = SessionLocal()
|
|
try:
|
|
user = db.query(User).filter(User.id == context["owner"].id).one()
|
|
user.huihui_token = f"huihui-payment-{context['suffix']}"
|
|
db.commit()
|
|
finally:
|
|
db.close()
|
|
|
|
|
|
def test_balance_is_user_scoped_and_defaults_to_one_million(authorization_context):
|
|
context = authorization_context
|
|
owner = client.get("/api/token/balance", headers=context["owner_headers"])
|
|
other = client.get("/api/token/balance", headers=context["other_headers"])
|
|
|
|
assert owner.status_code == 200
|
|
assert owner.json()["data"] == {
|
|
"balance": DEFAULT_TOKEN_GRANT,
|
|
"totalGranted": DEFAULT_TOKEN_GRANT,
|
|
"totalConsumed": 0,
|
|
}
|
|
assert other.json()["data"]["balance"] == DEFAULT_TOKEN_GRANT
|
|
assert client.get("/api/token/balance").status_code == 401
|
|
|
|
|
|
def test_seed_synchronizes_requested_recharge_plans():
|
|
seed()
|
|
db = SessionLocal()
|
|
try:
|
|
plans = db.query(TokenPlan).order_by(TokenPlan.price.asc()).all()
|
|
assert [(plan.price, plan.amount) for plan in plans] == [
|
|
(10, 2_000_000),
|
|
(100, 20_000_000),
|
|
(1000, 250_000_000),
|
|
(10000, 2_500_000_000),
|
|
]
|
|
finally:
|
|
db.close()
|
|
|
|
|
|
def test_charge_creates_huihui_order_without_early_points(authorization_context):
|
|
context = authorization_context
|
|
_enable_huihui_payment_login(context)
|
|
payment_client = Mock()
|
|
payment_client.create_payment.return_value = {
|
|
"orderId": "huihui-payment-id",
|
|
"orderNo": "huihui-payment-no",
|
|
"payMessage": {"mock": "payment-params"},
|
|
"payType": "WECHAT",
|
|
"paySubType": "APP",
|
|
"status": "pending",
|
|
}
|
|
env = {
|
|
"HUIHUI_PAYMENT_CALLBACK_BASE_URL": "https://digital.example",
|
|
"HUIHUI_PAYMENT_CALLBACK_SECRET": "test-callback-secret-123456",
|
|
}
|
|
with patch.dict(os.environ, env), patch("routers.tokens._payment_client", return_value=payment_client):
|
|
response = client.post(
|
|
"/api/token/charge",
|
|
headers=context["owner_headers"],
|
|
json={"planId": "1", "paymentMethod": "wechat", "payScene": "APP"},
|
|
)
|
|
|
|
assert response.status_code == 200
|
|
result = response.json()["data"]
|
|
assert result["status"] == "pending"
|
|
assert result["payType"] == "WECHAT"
|
|
assert result["payWay"] == "APP"
|
|
assert result["balance"] == DEFAULT_TOKEN_GRANT
|
|
assert payment_client.create_payment.call_args.kwargs["amount"] == "10.00"
|
|
callback_url = payment_client.create_payment.call_args.kwargs["callback_url"]
|
|
assert callback_url.startswith("https://digital.example/api/token/payment/callback/AV")
|
|
assert "test-callback-secret-123456" not in callback_url
|
|
|
|
|
|
def test_success_callback_credits_once_and_status_is_user_scoped(authorization_context):
|
|
context = authorization_context
|
|
_enable_huihui_payment_login(context)
|
|
payment_client = Mock()
|
|
payment_client.create_payment.return_value = {
|
|
"orderId": "huihui-payment-id",
|
|
"orderNo": "huihui-payment-no",
|
|
"payMessage": "payment-message",
|
|
"status": "pending",
|
|
}
|
|
secret = "test-callback-secret-123456"
|
|
env = {
|
|
"HUIHUI_PAYMENT_CALLBACK_BASE_URL": "https://digital.example",
|
|
"HUIHUI_PAYMENT_CALLBACK_SECRET": secret,
|
|
}
|
|
with patch.dict(os.environ, env), patch("routers.tokens._payment_client", return_value=payment_client):
|
|
created = client.post(
|
|
"/api/token/charge",
|
|
headers=context["owner_headers"],
|
|
json={"planId": "1", "paymentMethod": "alipay", "payScene": "APP"},
|
|
).json()["data"]
|
|
callback_body = {
|
|
"data": {
|
|
"masterOrderNo": created["orderNo"],
|
|
"status": "succeeded",
|
|
"payAmt": "10.00",
|
|
}
|
|
}
|
|
signature = hmac.new(
|
|
secret.encode(), created["orderNo"].encode(), hashlib.sha256
|
|
).hexdigest()
|
|
callback_path = f"/api/token/payment/callback/{created['orderNo']}/{signature}"
|
|
first = client.post(callback_path, json=callback_body)
|
|
second = client.post(callback_path, json=callback_body)
|
|
|
|
assert first.json()["data"] == {"received": True, "paid": True}
|
|
assert second.json()["data"] == {"received": True, "duplicate": True}
|
|
status = client.get(
|
|
f"/api/token/payment/{created['id']}", headers=context["owner_headers"]
|
|
).json()["data"]
|
|
assert status["status"] == "paid"
|
|
assert status["balance"] == DEFAULT_TOKEN_GRANT + 2_000_000
|
|
assert client.get(
|
|
f"/api/token/payment/{created['id']}", headers=context["other_headers"]
|
|
).json()["code"] == 404
|
|
|
|
|
|
def test_callback_amount_mismatch_never_credits_points(authorization_context):
|
|
context = authorization_context
|
|
_enable_huihui_payment_login(context)
|
|
payment_client = Mock()
|
|
payment_client.create_payment.return_value = {"status": "pending", "payMessage": "mock"}
|
|
secret = "test-callback-secret-123456"
|
|
env = {
|
|
"HUIHUI_PAYMENT_CALLBACK_BASE_URL": "https://digital.example",
|
|
"HUIHUI_PAYMENT_CALLBACK_SECRET": secret,
|
|
}
|
|
with patch.dict(os.environ, env), patch("routers.tokens._payment_client", return_value=payment_client):
|
|
created = client.post(
|
|
"/api/token/charge",
|
|
headers=context["owner_headers"],
|
|
json={"planId": "1", "paymentMethod": "wechat", "payScene": "APP"},
|
|
).json()["data"]
|
|
signature = hmac.new(
|
|
secret.encode(), created["orderNo"].encode(), hashlib.sha256
|
|
).hexdigest()
|
|
callback = client.post(
|
|
f"/api/token/payment/callback/{created['orderNo']}/{signature}",
|
|
json={
|
|
"masterOrderNo": created["orderNo"],
|
|
"status": "success",
|
|
"actualAmt": "9.99",
|
|
},
|
|
)
|
|
|
|
assert callback.json()["code"] == 422
|
|
db = SessionLocal()
|
|
try:
|
|
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.id == created["id"]).one()
|
|
account = get_or_create_account(db, context["owner"].id)
|
|
assert order.status == "pending"
|
|
assert account.balance == DEFAULT_TOKEN_GRANT
|
|
finally:
|
|
db.close()
|
|
|
|
|
|
def test_payment_callback_creates_missing_account_in_same_settlement(authorization_context):
|
|
context = authorization_context
|
|
_enable_huihui_payment_login(context)
|
|
payment_client = Mock()
|
|
payment_client.create_payment.return_value = {"status": "pending", "payMessage": "mock"}
|
|
secret = "test-callback-secret-123456"
|
|
env = {
|
|
"HUIHUI_PAYMENT_CALLBACK_BASE_URL": "https://digital.example",
|
|
"HUIHUI_PAYMENT_CALLBACK_SECRET": secret,
|
|
}
|
|
with patch.dict(os.environ, env), patch("routers.tokens._payment_client", return_value=payment_client):
|
|
created = client.post(
|
|
"/api/token/charge",
|
|
headers=context["owner_headers"],
|
|
json={"planId": "1", "paymentMethod": "alipay", "payScene": "APP"},
|
|
).json()["data"]
|
|
db = SessionLocal()
|
|
try:
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == context["owner"].id).delete()
|
|
db.commit()
|
|
finally:
|
|
db.close()
|
|
signature = hmac.new(
|
|
secret.encode(), created["orderNo"].encode(), hashlib.sha256
|
|
).hexdigest()
|
|
callback = client.post(
|
|
f"/api/token/payment/callback/{created['orderNo']}/{signature}",
|
|
json={
|
|
"masterOrderNo": created["orderNo"],
|
|
"status": "success",
|
|
"payAmt": "10.00",
|
|
},
|
|
)
|
|
|
|
assert callback.json()["data"] == {"received": True, "paid": True}
|
|
db = SessionLocal()
|
|
try:
|
|
account = db.query(TokenAccount).filter(TokenAccount.user_id == context["owner"].id).one()
|
|
assert account.balance == DEFAULT_TOKEN_GRANT + 2_000_000
|
|
assert account.total_granted == DEFAULT_TOKEN_GRANT + 2_000_000
|
|
finally:
|
|
db.close()
|
|
|
|
|
|
def test_multiple_avatars_share_owner_balance_and_usage_is_itemized():
|
|
suffix = uuid.uuid4().hex
|
|
db = SessionLocal()
|
|
user = User(id=f"token-user-{suffix}", huihui_user_id=f"token-huihui-{suffix}")
|
|
first = Avatar(id=f"token-avatar-a-{suffix}", owner_id=user.huihui_user_id, name="甲")
|
|
second = Avatar(id=f"token-avatar-b-{suffix}", owner_id=user.huihui_user_id, name="乙")
|
|
db.add_all([user, first, second])
|
|
db.commit()
|
|
try:
|
|
first_reservation = reserve_avatar_tokens(db, first, "chat", "qwen-test", [{"content": "问题一"}], 128)
|
|
settle_reservation(
|
|
db,
|
|
first_reservation,
|
|
{"prompt_tokens": 60, "completion_tokens": 40, "total_tokens": 100},
|
|
fallback_total=999,
|
|
)
|
|
second_reservation = reserve_avatar_tokens(db, second, "takeover", "qwen-test", [{"content": "问题二"}], 128)
|
|
settle_reservation(
|
|
db,
|
|
second_reservation,
|
|
{"prompt_tokens": 120, "completion_tokens": 80, "total_tokens": 200},
|
|
fallback_total=999,
|
|
)
|
|
|
|
account = get_or_create_account(db, user.id)
|
|
assert account.balance == DEFAULT_TOKEN_GRANT - 300
|
|
assert account.total_consumed == 300
|
|
usages = db.query(TokenUsage).filter(TokenUsage.user_id == user.id).order_by(TokenUsage.total_tokens).all()
|
|
assert [(row.avatar_id, row.source, row.total_tokens) for row in usages] == [
|
|
(first.id, "chat", 100),
|
|
(second.id, "takeover", 200),
|
|
]
|
|
finally:
|
|
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(Avatar).filter(Avatar.id.in_([first.id, second.id])).delete(synchronize_session=False)
|
|
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
|
|
db.commit()
|
|
db.close()
|
|
|
|
|
|
def test_concurrent_settlements_do_not_overwrite_each_other():
|
|
suffix = uuid.uuid4().hex
|
|
db = SessionLocal()
|
|
user = User(id=f"concurrent-user-{suffix}", huihui_user_id=f"concurrent-huihui-{suffix}")
|
|
avatar = Avatar(id=f"concurrent-avatar-{suffix}", owner_id=user.huihui_user_id, name="并发测试")
|
|
db.add_all([user, avatar])
|
|
db.commit()
|
|
first = reserve_avatar_tokens(db, avatar, "takeover", "qwen-test", [{"content": "甲"}], 128)
|
|
second = reserve_avatar_tokens(db, avatar, "takeover", "qwen-test", [{"content": "乙"}], 128)
|
|
db.close()
|
|
barrier = Barrier(2, timeout=3)
|
|
|
|
def settle(reservation, total):
|
|
thread_db = SessionLocal()
|
|
try:
|
|
barrier.wait()
|
|
settle_reservation(
|
|
thread_db,
|
|
reservation,
|
|
{"prompt_tokens": total - 20, "completion_tokens": 20, "total_tokens": total},
|
|
fallback_total=999,
|
|
)
|
|
finally:
|
|
thread_db.close()
|
|
|
|
with ThreadPoolExecutor(max_workers=2) as pool:
|
|
list(pool.map(lambda args: settle(*args), [(first, 100), (second, 200)]))
|
|
|
|
db = SessionLocal()
|
|
try:
|
|
account = get_or_create_account(db, user.id)
|
|
assert account.balance == DEFAULT_TOKEN_GRANT - 300
|
|
assert account.total_consumed == 300
|
|
finally:
|
|
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
|
|
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
|
|
db.commit()
|
|
db.close()
|
|
|
|
|
|
def test_failed_model_request_returns_the_full_reservation():
|
|
suffix = uuid.uuid4().hex
|
|
db = SessionLocal()
|
|
user = User(id=f"refund-user-{suffix}", huihui_user_id=f"refund-huihui-{suffix}")
|
|
avatar = Avatar(id=f"refund-avatar-{suffix}", owner_id=user.huihui_user_id, name="退款测试")
|
|
db.add_all([user, avatar])
|
|
db.commit()
|
|
try:
|
|
reservation = reserve_avatar_tokens(db, avatar, "chat", "qwen-test", [{"content": "问题"}], 128)
|
|
release_reservation(db, reservation, "provider error")
|
|
account = get_or_create_account(db, user.id)
|
|
usage = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).one()
|
|
assert account.balance == DEFAULT_TOKEN_GRANT
|
|
assert account.total_consumed == 0
|
|
assert usage.status == "failed"
|
|
finally:
|
|
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
|
|
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
|
|
db.commit()
|
|
db.close()
|
|
|
|
|
|
def test_insufficient_balance_rejects_before_model_usage_is_created():
|
|
suffix = uuid.uuid4().hex
|
|
db = SessionLocal()
|
|
user = User(id=f"empty-user-{suffix}", huihui_user_id=f"empty-huihui-{suffix}")
|
|
avatar = Avatar(id=f"empty-avatar-{suffix}", owner_id=user.huihui_user_id, name="余额不足")
|
|
db.add_all([user, avatar])
|
|
db.commit()
|
|
try:
|
|
account = get_or_create_account(db, user.id)
|
|
account.balance = 1
|
|
db.commit()
|
|
with pytest.raises(InsufficientTokensError):
|
|
reserve_avatar_tokens(db, avatar, "chat", "qwen-test", [{"content": "问题"}], 128)
|
|
db.refresh(account)
|
|
assert account.balance == 1
|
|
assert db.query(TokenUsage).filter(TokenUsage.user_id == user.id).count() == 0
|
|
finally:
|
|
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
|
|
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
|
|
db.commit()
|
|
db.close()
|
|
|
|
|
|
def test_chat_settles_from_provider_usage_not_fallback_estimate():
|
|
suffix = uuid.uuid4().hex
|
|
db = SessionLocal()
|
|
user = User(id=f"chat-user-{suffix}", huihui_user_id=f"chat-huihui-{suffix}")
|
|
avatar = Avatar(id=f"chat-avatar-{suffix}", owner_id=user.huihui_user_id, name="聊天测试", config={})
|
|
db.add_all([user, avatar])
|
|
db.commit()
|
|
try:
|
|
with patch(
|
|
"routers.chat._call_qwen",
|
|
return_value={
|
|
"answer": "测试回答",
|
|
"usage": {"prompt_tokens": 80, "completion_tokens": 20, "total_tokens": 100},
|
|
},
|
|
):
|
|
result = _resolve_reply(
|
|
db,
|
|
avatar,
|
|
"测试问题",
|
|
[],
|
|
qa_pairs=[],
|
|
search_fn=lambda *_args: [],
|
|
)
|
|
assert result["tokenUsage"]["totalTokens"] == 100
|
|
assert result["tokenUsage"]["balance"] == DEFAULT_TOKEN_GRANT - 100
|
|
finally:
|
|
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
|
|
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
|
|
db.commit()
|
|
db.close()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_streaming_chat_settles_final_provider_usage():
|
|
suffix = uuid.uuid4().hex
|
|
db = SessionLocal()
|
|
user = User(id=f"stream-user-{suffix}", huihui_user_id=f"stream-huihui-{suffix}")
|
|
avatar = Avatar(id=f"stream-avatar-{suffix}", owner_id=user.huihui_user_id, name="流式测试", config={})
|
|
db.add_all([user, avatar])
|
|
db.commit()
|
|
try:
|
|
chunks = iter([
|
|
{"content": "流式"},
|
|
{"content": "回答"},
|
|
{"usage": {"prompt_tokens": 90, "completion_tokens": 10, "total_tokens": 100}},
|
|
])
|
|
with patch("routers.chat._iter_qwen_stream", return_value=chunks):
|
|
response = _stream_reply(db, avatar, "测试问题", [])
|
|
body = []
|
|
async for chunk in response.body_iterator:
|
|
body.append(chunk.decode() if isinstance(chunk, bytes) else chunk)
|
|
assert "流式" in "".join(body)
|
|
account = get_or_create_account(db, user.id)
|
|
usage = db.query(TokenUsage).filter(TokenUsage.user_id == user.id).one()
|
|
assert account.balance == DEFAULT_TOKEN_GRANT - 100
|
|
assert usage.source == "chat_stream"
|
|
assert usage.total_tokens == 100
|
|
finally:
|
|
db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False)
|
|
db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False)
|
|
db.query(User).filter(User.id == user.id).delete(synchronize_session=False)
|
|
db.commit()
|
|
db.close()
|