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