253 lines
11 KiB
Python
253 lines
11 KiB
Python
import uuid
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from threading import Barrier
|
|
from unittest.mock import patch
|
|
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
|
|
from database import SessionLocal
|
|
from main import app, seed
|
|
from models import Avatar, TokenAccount, 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 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_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()
|