修复语音聊天和实时聊天token统计
This commit is contained in:
parent
43f17dcdae
commit
21bd6d0a57
@ -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))
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user