diff --git a/backend/api/voice_chat_router.py b/backend/api/voice_chat_router.py index d02fcbf..11ecb2a 100644 --- a/backend/api/voice_chat_router.py +++ b/backend/api/voice_chat_router.py @@ -39,6 +39,108 @@ logger = logging.getLogger(__name__) voice_chat_router = APIRouter(prefix="/voice", tags=["voice"]) +VOICE_TOKEN_FALLBACKS = { + "voice_chat": 200, + "voice_realtime_call": 200, +} + +VOICE_CHARGE_MODES = { + "voice_chat": "语音聊天", + "voice_realtime_call": "实时通话", +} + + +async def _charge_voice_tokens( + user_id: str, + mode: str, + session_state: Optional[Dict[str, Any]] = None, + tokens_used: Optional[Dict[str, Any]] = None, + fallback_amount: Optional[int] = None, + reason: str = "manual_end" +) -> Dict[str, Any]: + """统一语音扣费入口,补充清晰日志。""" + session_state = session_state or {} + if session_state.get("token_charged"): + logger.info( + "[VOICE_CHARGE_SKIP] user_id=%s mode=%s reason=%s conversation_id=%s trace_id=%s already_charged=True", + user_id, + mode, + reason, + session_state.get("conversation_id"), + session_state.get("trace_id"), + ) + return {"success": True, "skipped": True, "message": "already charged"} + + token_service = session_state.get("token_service") + if token_service is None: + async with AsyncSessionLocal() as db: + token_service = TokenService(db) + return await _charge_voice_tokens( + user_id=user_id, + mode=mode, + session_state={**session_state, "token_service": token_service}, + tokens_used=tokens_used, + fallback_amount=fallback_amount, + reason=reason, + ) + + charge_cfg_fallback = fallback_amount if fallback_amount is not None else VOICE_TOKEN_FALLBACKS.get(mode, 200) + result = await token_service.charge_tokens( + user_id=int(user_id), + mode=mode, + tokens_used=tokens_used, + fallback_amount=charge_cfg_fallback, + ) + + token_used_amount = int(result.get("amount") or 0) + logger.info( + "[VOICE_CHARGE_%s] user_id=%s mode=%s reason=%s conversation_id=%s trace_id=%s amount=%s available_tokens=%s success=%s message=%s", + "OK" if result.get("success") else "FAIL", + user_id, + mode, + reason, + session_state.get("conversation_id"), + session_state.get("trace_id"), + token_used_amount, + result.get("available_tokens"), + result.get("success"), + result.get("message"), + ) + if result.get("success"): + session_state["token_charged"] = True + session_state["tokens_used"] = tokens_used or session_state.get("tokens_used") + return result + + +async def _finalize_voice_session( + user_id: str, + mode: str, + session_state: Optional[Dict[str, Any]] = None, + reason: str = "manual_end", + fallback_amount: Optional[int] = None, +) -> Dict[str, Any]: + """结束语音会话并确保扣费。""" + session_state = session_state or {} + tokens_used = session_state.get("tokens_used") if isinstance(session_state.get("tokens_used"), dict) else None + result = await _charge_voice_tokens( + user_id=user_id, + mode=mode, + session_state=session_state, + tokens_used=tokens_used, + fallback_amount=fallback_amount, + reason=reason, + ) + logger.info( + "[VOICE_SESSION_FINALIZED] user_id=%s mode=%s reason=%s conversation_id=%s trace_id=%s charged=%s", + user_id, + mode, + reason, + session_state.get("conversation_id"), + session_state.get("trace_id"), + result.get("success"), + ) + return result + # ============ 请求/响应模型 ============ @@ -167,6 +269,7 @@ async def websocket_voice_chat(websocket: WebSocket, user_id: str): "user_msg": "", "ai_msg": "", "tokens_used": None, + "token_charged": False, "last_persisted_user_msg": "", "last_persisted_ai_len": 0, } @@ -212,8 +315,7 @@ async def websocket_voice_chat(websocket: WebSocket, user_id: str): session_state["user_msg"] = "" session_state["ai_msg"] = "" session_state["tokens_used"] = None - session_state["last_persisted_user_msg"] = "" - session_state["last_persisted_ai_len"] = 0 + session_state["token_charged"] = False session_state["last_persisted_user_msg"] = "" session_state["last_persisted_ai_len"] = 0 @@ -271,6 +373,13 @@ async def websocket_voice_chat(websocket: WebSocket, user_id: str): pass await client.finish_session() + await _ensure_voice_token_charge( + user_id=user_id, + mode="voice_chat", + session_state=session_state, + fallback_amount=get_token_charge_config().get("voice_chat_fallback_tokens", VOICE_TOKEN_FALLBACKS["voice_chat"]), + reason="stop_session" + ) await websocket.send_text(json.dumps({ "type": "session_stopped", "timestamp": datetime.now().isoformat() @@ -333,6 +442,19 @@ async def websocket_realtime_call(websocket: WebSocket, user_id: str): client: Optional[RealtimeVoiceClient] = None receive_task: Optional[asyncio.Task] = None audio_send_task: Optional[asyncio.Task] = None + session_state: Dict[str, Any] = { + "mode": "voice_realtime_call", + "conversation_id": None, + "trace_id": None, + "pet_id": None, + "bg_id": None, + "user_msg": "", + "ai_msg": "", + "tokens_used": None, + "token_charged": False, + "last_persisted_user_msg": "", + "last_persisted_ai_len": 0, + } try: # 发送连接成功消息 @@ -415,8 +537,13 @@ async def websocket_realtime_call(websocket: WebSocket, user_id: str): })) # 启动事件接收任务 + session_state["pet_id"] = pet_id + session_state["bg_id"] = message.get("background_id") + session_state["conversation_id"] = f"realtime-{user_id}-{uuid.uuid4().hex[:12]}" + session_state["trace_id"] = f"trace-{user_id}-{uuid.uuid4().hex[:12]}" + session_state["token_charged"] = False receive_task = asyncio.create_task( - receive_events_loop(websocket, client, user_id, "voice_realtime_call") + receive_events_loop(websocket, client, user_id, "voice_realtime_call", session_state) ) # 可选:发送打招呼 @@ -445,6 +572,13 @@ async def websocket_realtime_call(websocket: WebSocket, user_id: str): pass await client.close() + await _ensure_voice_token_charge( + user_id=user_id, + mode="voice_realtime_call", + session_state=session_state, + fallback_amount=get_token_charge_config().get("voice_realtime_call_fallback_tokens", VOICE_TOKEN_FALLBACKS["voice_realtime_call"]), + reason="end_call" + ) client = None await websocket.send_text(json.dumps({ @@ -650,6 +784,31 @@ async def _append_voice_log_chunk( logger.warning(f"Failed to append voice log chunk: user_id={user_id}, mode={mode}, error={e}") +async def _ensure_voice_token_charge( + user_id: str, + mode: str, + session_state: Optional[Dict[str, Any]] = None, + fallback_amount: int = 200, + reason: str = "ensure_end" +) -> Dict[str, Any]: + """确保语音会话最终一定会扣费;若已按 usage 扣费则跳过。""" + if _is_guest_user_id(user_id): + logger.info("[VOICE_CHARGE_SKIP] user_id=%s mode=%s reason=%s guest_user=True", user_id, mode, reason) + return { + "success": True, + "skipped": True, + "message": "Guest user skipped" + } + + return await _finalize_voice_session( + user_id=user_id, + mode=mode, + session_state=session_state, + reason=reason, + fallback_amount=fallback_amount, + ) + + async def _charge_user_tokens( user_id: str, mode: str, @@ -817,7 +976,8 @@ async def receive_events_loop( persist_kind='ai' ) elif event.event_type == 'usage' and isinstance(event.data, dict): - session_state['tokens_used'] = event.data.get('usage', {}) + usage_data = event.data.get('usage', {}) + session_state['tokens_used'] = usage_data await websocket.send_text(json.dumps(msg)) diff --git a/frontend/src/stores/chatStore.js b/frontend/src/stores/chatStore.js index b316236..11ef259 100644 --- a/frontend/src/stores/chatStore.js +++ b/frontend/src/stores/chatStore.js @@ -580,10 +580,14 @@ export const useChatStore = defineStore('chat', () => { */ function disconnectVoiceWebSocket() { if (voiceSocket.value) { - // 发送停止会话消息 + // 发送停止会话消息,给后端一个正常结算机会 sendVoiceMessage({ type: 'stop_session' }) - voiceSocket.value.close() - voiceSocket.value = null + setTimeout(() => { + if (voiceSocket.value) { + voiceSocket.value.close() + voiceSocket.value = null + } + }, 300) } isVoiceConnected.value = false } @@ -828,8 +832,12 @@ export const useChatStore = defineStore('chat', () => { function disconnectRealtimeWebSocket() { if (realtimeSocket) { sendRealtimeMessage({ type: 'end_call' }) - realtimeSocket.close() - realtimeSocket = null + setTimeout(() => { + if (realtimeSocket) { + realtimeSocket.close() + realtimeSocket = null + } + }, 300) } isRealtimeConnected.value = false isInCall.value = false