修复语音聊天和实时聊天token统计

This commit is contained in:
taiyi 2026-05-24 15:44:51 +08:00
parent 43f17dcdae
commit 21bd6d0a57
2 changed files with 177 additions and 9 deletions

View File

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

View File

@ -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