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()