770 lines
26 KiB
Python
770 lines
26 KiB
Python
"""
|
||
WebSocket 智能语义分析服务
|
||
|
||
功能说明:
|
||
1. 作为 WebSocket 代理,接收前端音频数据并转发到 ASR 引擎
|
||
2. 实时处理 ASR 识别结果,应用多 Agent 协作(A2A)模式
|
||
3. Agent 功能:
|
||
- AuditorAgent: 判断说话人是否完成发言
|
||
- ProfilerAgent: 识别当前说话人的身份
|
||
4. 使用 Redis 存储会话缓冲和历史记录
|
||
5. 将处理后的带说话人身份的结果推送回前端
|
||
|
||
架构模式:Agent-to-Agent (A2A) 编排
|
||
"""
|
||
|
||
import asyncio
|
||
import json
|
||
import os
|
||
import websockets
|
||
import re
|
||
import logging
|
||
from pathlib import Path
|
||
from typing import List, Optional, Dict, Set
|
||
from dataclasses import dataclass
|
||
from dotenv import load_dotenv
|
||
|
||
try:
|
||
import redis.asyncio as redis
|
||
except ModuleNotFoundError:
|
||
redis = None
|
||
|
||
try:
|
||
from openai import AsyncOpenAI
|
||
except ModuleNotFoundError:
|
||
AsyncOpenAI = None
|
||
# ========================
|
||
# 1. 日志配置
|
||
# ========================
|
||
ENV_PATH = Path(__file__).resolve().parent.parent / ".env"
|
||
load_dotenv(ENV_PATH)
|
||
|
||
|
||
def get_env(name: str, default: str) -> str:
|
||
return os.getenv(name, default).strip()
|
||
|
||
|
||
def get_env_bool(name: str, default: bool = False) -> bool:
|
||
value = os.getenv(name)
|
||
if value is None:
|
||
return default
|
||
return value.strip().lower() in {"1", "true", "yes", "on"}
|
||
|
||
|
||
def setup_logging():
|
||
"""配置日志系统"""
|
||
# 创建日志目录
|
||
log_dir = "logs"
|
||
os.makedirs(log_dir, exist_ok=True)
|
||
|
||
# 配置日志格式
|
||
log_format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||
date_format = "%Y-%m-%d %H:%M:%S"
|
||
|
||
# 配置根日志记录器
|
||
logger = logging.getLogger()
|
||
logger.setLevel(logging.INFO)
|
||
|
||
# 清除现有处理器
|
||
logger.handlers.clear()
|
||
|
||
# 控制台处理器
|
||
console_handler = logging.StreamHandler()
|
||
console_handler.setLevel(logging.INFO)
|
||
console_formatter = logging.Formatter(log_format, datefmt=date_format)
|
||
console_handler.setFormatter(console_formatter)
|
||
logger.addHandler(console_handler)
|
||
|
||
# 文件处理器(所有日志)
|
||
file_handler = logging.FileHandler(
|
||
f"{log_dir}/a2a_wss.log",
|
||
encoding="utf-8"
|
||
)
|
||
file_handler.setLevel(logging.DEBUG)
|
||
file_formatter = logging.Formatter(log_format, datefmt=date_format)
|
||
file_handler.setFormatter(file_formatter)
|
||
logger.addHandler(file_handler)
|
||
|
||
# 错误日志文件处理器
|
||
error_handler = logging.FileHandler(
|
||
f"{log_dir}/a2a_wss_error.log",
|
||
encoding="utf-8"
|
||
)
|
||
error_handler.setLevel(logging.ERROR)
|
||
error_formatter = logging.Formatter(log_format, datefmt=date_format)
|
||
error_handler.setFormatter(error_formatter)
|
||
logger.addHandler(error_handler)
|
||
|
||
return logger
|
||
|
||
# 初始化日志系统
|
||
logger = setup_logging()
|
||
|
||
# ========================
|
||
# 2. 基础配置
|
||
# ========================
|
||
|
||
# 目标 ASR 服务的 WebSocket 地址(负责语音识别)
|
||
# 注意:移除末尾的斜杠,避免路径匹配问题
|
||
TARGET_WS_URL = get_env("TARGET_WS_URL", "ws://localhost:59805/ws/asr")
|
||
|
||
# 本地服务监听地址和端口(接收前端连接)
|
||
LOCAL_HOST = get_env("LOCAL_HOST", "0.0.0.0")
|
||
LOCAL_PORT = int(get_env("LOCAL_PORT", "10095"))
|
||
|
||
# Redis 连接配置(用于存储会话状态和历史记录)
|
||
REDIS_URL = get_env("REDIS_URL", "redis://localhost:6379/0")
|
||
|
||
# LLM(大语言模型)配置(用于 Agent 推理)
|
||
LLM_API_KEY = get_env("LLM_API_KEY", "none")
|
||
LLM_BASE_URL = get_env("LLM_BASE_URL", "http://192.168.0.46:59800/v1")
|
||
LLM_MODEL = get_env("LLM_MODEL", "46-qwen3.5-35B")
|
||
USE_LLM_ROLE_IDENTIFICATION = get_env_bool("USE_LLM_ROLE_IDENTIFICATION", False)
|
||
# 候选说话人列表(用于身份识别 Agent)
|
||
CANDIDATE_SPEAKERS = ["局长", "主持人", "商务专员", "市场专员", "控制要素席", "空中侦察席"]
|
||
|
||
# ========================
|
||
# 3. Redis 管理器
|
||
# ========================
|
||
|
||
class RedisManager:
|
||
"""
|
||
Redis 会话管理器
|
||
|
||
职责:
|
||
- 管理与 Redis 的连接池
|
||
- 存储每个会话的文本缓冲(累积不完整的识别片段)
|
||
- 存储历史对话记录(用于上下文理解)
|
||
"""
|
||
|
||
def __init__(self, url):
|
||
"""
|
||
初始化 Redis 连接池
|
||
|
||
Args:
|
||
url: Redis 连接 URL
|
||
"""
|
||
if redis is None:
|
||
raise RuntimeError("redis package is not installed")
|
||
# 使用连接池提高并发性能,decode_responses=True 自动将字节转为字符串
|
||
self.pool = redis.ConnectionPool.from_url(url, decode_responses=True)
|
||
# 键前缀,用于命名隔离,避免与其他业务冲突
|
||
self.prefix = "meeting:"
|
||
|
||
async def _get_conn(self):
|
||
"""从连接池获取一个 Redis 连接"""
|
||
return redis.Redis(connection_pool=self.pool)
|
||
|
||
async def get_buffer(self, sid):
|
||
"""
|
||
获取指定会话的文本缓冲
|
||
|
||
Args:
|
||
sid: 会话 ID(Session ID)
|
||
|
||
Returns:
|
||
该会话当前累积的文本片段
|
||
"""
|
||
r = await self._get_conn()
|
||
val = await r.get(f"{self.prefix}buffer:{sid}") or ""
|
||
return val
|
||
|
||
async def append_buffer(self, sid, text):
|
||
"""
|
||
向会话缓冲追加新的文本片段
|
||
|
||
Args:
|
||
sid: 会话 ID
|
||
text: 新识别的文本片段
|
||
"""
|
||
r = await self._get_conn()
|
||
# 追加文本并添加空格分隔
|
||
await r.append(f"{self.prefix}buffer:{sid}", text + " ")
|
||
|
||
async def clear_buffer(self, sid):
|
||
"""
|
||
清空指定会话的缓冲区
|
||
|
||
用途:当判断发言完成后,清空缓冲准备下一次发言
|
||
|
||
Args:
|
||
sid: 会话 ID
|
||
"""
|
||
r = await self._get_conn()
|
||
await r.delete(f"{self.prefix}buffer:{sid}")
|
||
|
||
async def push_history(self, sid, role, content):
|
||
"""
|
||
将完成的对话推入历史记录
|
||
|
||
Args:
|
||
sid: 会话 ID
|
||
role: 说话人角色
|
||
content: 对话内容
|
||
"""
|
||
r = await self._get_conn()
|
||
key = f"{self.prefix}history:{sid}"
|
||
# 将对话记录推入列表右端
|
||
await r.rpush(key, json.dumps({"role": role, "content": content}))
|
||
# 只保留最近 20 条记录,避免历史过长影响性能
|
||
await r.ltrim(key, -20, -1)
|
||
|
||
async def get_history(self, sid):
|
||
"""
|
||
获取会话的历史对话记录
|
||
|
||
Args:
|
||
sid: 会话 ID
|
||
|
||
Returns:
|
||
历史对话列表,每项包含 role 和 content
|
||
"""
|
||
r = await self._get_conn()
|
||
items = await r.lrange(f"{self.prefix}history:{sid}", 0, -1)
|
||
return [json.loads(i) for i in items]
|
||
|
||
async def save_speaker_profile(self, sid, speaker, text, features=None):
|
||
"""
|
||
保存说话人的特征到长期记忆
|
||
|
||
Args:
|
||
sid: 会话 ID
|
||
speaker: 说话人名称
|
||
text: 发言文本
|
||
features: 可选的特征字典(如关键词、语速等)
|
||
"""
|
||
r = await self._get_conn()
|
||
key = f"{self.prefix}speaker_profile:{sid}:{speaker}"
|
||
|
||
# 构造特征记录
|
||
profile_entry = {
|
||
"text": text[:200], # 只保留前200个字符作为示例
|
||
"timestamp": asyncio.get_event_loop().time(),
|
||
"features": features or {}
|
||
}
|
||
|
||
# 将记录推入该说话人的历史列表
|
||
await r.rpush(key, json.dumps(profile_entry))
|
||
# 只保留最近10条该说话人的记录
|
||
await r.ltrim(key, -10, -1)
|
||
|
||
async def get_speaker_profiles(self, sid):
|
||
"""
|
||
获取所有说话人的历史特征记录
|
||
|
||
Args:
|
||
sid: 会话 ID
|
||
|
||
Returns:
|
||
字典,key为说话人名称,value为该说话人的历史记录列表
|
||
"""
|
||
r = await self._get_conn()
|
||
# 扫描所有说话人的profile键
|
||
pattern = f"{self.prefix}speaker_profile:{sid}:*"
|
||
profiles = {}
|
||
|
||
async for key in r.scan_iter(match=pattern):
|
||
# 从键名中提取说话人名称
|
||
speaker = key.decode() if isinstance(key, bytes) else key
|
||
speaker = speaker.split(":")[-1]
|
||
|
||
# 获取该说话人的所有历史记录
|
||
items = await r.lrange(key, 0, -1)
|
||
profiles[speaker] = [json.loads(i) for i in items]
|
||
|
||
return profiles
|
||
|
||
async def close(self):
|
||
"""关闭连接池,释放资源(修复 DeprecationWarning)"""
|
||
await self.pool.disconnect()
|
||
|
||
# ========================
|
||
# 3. Agent 实现
|
||
# ========================
|
||
|
||
class BaseAgent:
|
||
"""
|
||
Agent 基类
|
||
|
||
提供所有 Agent 的通用能力:
|
||
- 调用 LLM 进行推理
|
||
- 提取 LLM 返回的 XML 标签结果
|
||
"""
|
||
|
||
def __init__(self, client: AsyncOpenAI):
|
||
"""
|
||
初始化 Agent
|
||
|
||
Args:
|
||
client: OpenAI 异步客户端
|
||
"""
|
||
self.client = client
|
||
|
||
async def _call_llm(self, system_prompt: str, user_content: str) -> str:
|
||
"""
|
||
调用大语言模型进行推理
|
||
|
||
Args:
|
||
system_prompt: 系统提示词(定义 Agent 的角色和行为)
|
||
user_content: 用户输入内容
|
||
|
||
Returns:
|
||
LLM 返回的文本结果,失败时返回空字符串
|
||
"""
|
||
try:
|
||
response = await self.client.chat.completions.create(
|
||
model=LLM_MODEL,
|
||
messages=[
|
||
{"role": "system", "content": system_prompt},
|
||
{"role": "user", "content": user_content}
|
||
],
|
||
temperature=0.1, # 低温度使输出更确定、一致
|
||
)
|
||
return response.choices[0].message.content
|
||
except Exception as e:
|
||
logger.error(f"LLM调用失败: {e}", exc_info=True)
|
||
return ""
|
||
|
||
def _extract_result(self, text: str) -> str:
|
||
"""
|
||
从 LLM 返回结果中提取 <result> 标签内容
|
||
|
||
Args:
|
||
text: LLM 返回的完整文本
|
||
|
||
Returns:
|
||
提取的结果,如果没有标签则返回原文
|
||
"""
|
||
match = re.search(r'<result>(.*?)</result>', text, re.DOTALL | re.IGNORECASE)
|
||
return match.group(1).strip() if match else text.strip()
|
||
|
||
|
||
class AuditorAgent(BaseAgent):
|
||
"""
|
||
审计 Agent:判断当前发言是否完成
|
||
|
||
功能:
|
||
- 分析累积的文本缓冲
|
||
- 判断说话人是否已经完成一次完整的发言
|
||
- 判断依据:语义完整性、停顿暗示等
|
||
"""
|
||
|
||
# 系统提示词:定义 Agent 的角色和输出格式
|
||
SYSTEM_PROMPT = "你是一个对话状态监测专家。请判断发言是否结束,放入 <result>true/false</result> 中。"
|
||
|
||
async def check(self, text: str) -> bool:
|
||
"""
|
||
判断给定文本是否代表一次完整的发言
|
||
|
||
Args:
|
||
text: 当前累积的文本片段
|
||
|
||
Returns:
|
||
True 表示发言完成,False 表示发言未完成(继续累积)
|
||
"""
|
||
res = await self._call_llm(self.SYSTEM_PROMPT, f"内容:{text}")
|
||
return "true" in self._extract_result(res).lower()
|
||
|
||
|
||
class ProfilerAgent(BaseAgent):
|
||
"""
|
||
分析 Agent:识别说话人身份
|
||
|
||
功能:
|
||
- 基于文本内容和说话风格识别说话人
|
||
- 利用历史对话上下文提高识别准确率
|
||
- 利用长期记忆(说话人特征库)增强识别
|
||
- 从预定义的候选说话人列表中选择
|
||
"""
|
||
|
||
# 系统提示词:定义识别任务和候选说话人列表
|
||
SYSTEM_PROMPT = f"你是一个会议身份识别专家。请从 {CANDIDATE_SPEAKERS} 中选择发言人,放入 <result>名称</result> 中。"
|
||
|
||
# 增强版提示词:结合长期记忆
|
||
SYSTEM_PROMPT_WITH_MEMORY = f"""你是一个会议身份识别专家。请根据以下信息识别发言人:
|
||
|
||
候选说话人列表:{CANDIDATE_SPEAKERS}
|
||
|
||
=== 会议历史记录(按时间顺序)===
|
||
{{history}}
|
||
|
||
=== 说话人特征库(长期记忆)===
|
||
{{speaker_profiles}}
|
||
|
||
=== 当前待识别的发言 ===
|
||
{{current_text}}
|
||
|
||
=== 识别指南 ===
|
||
请仔细分析并综合判断:
|
||
|
||
1. **对话流程分析**:
|
||
- 查看历史记录中的对话顺序和上下文
|
||
- 判断当前发言是对之前谁的话题的回应、延续或补充
|
||
- 注意对话的逻辑关系和互动模式
|
||
|
||
2. **说话风格匹配**:
|
||
- 每个角色的用词习惯、句式结构、语气的特点
|
||
- 专有名词、技术术语的使用习惯
|
||
- 提问方式、陈述方式、指令方式的差异
|
||
|
||
3. **角色职责判断**:
|
||
- 根据各角色的职责范围判断发言内容的合理性
|
||
- 例如:局长通常做总结和决策,主持人负责控场,专员负责具体业务汇报
|
||
|
||
4. **话题连贯性**:
|
||
- 当前发言与该角色之前发言的话题是否连贯
|
||
- 是否体现了该角色应有的专业领域关注点
|
||
|
||
将识别结果放入 <result>说话人名称</result> 中。"""
|
||
|
||
async def identify(self, text: str, history: List[Dict], speaker_profiles: Dict = None) -> str:
|
||
"""
|
||
识别当前文本的说话人身份
|
||
|
||
Args:
|
||
text: 当前发言文本
|
||
history: 历史对话记录(用于上下文理解)
|
||
speaker_profiles: 说话人特征库(长期记忆),key为说话人名称,value为历史记录列表
|
||
|
||
Returns:
|
||
识别出的说话人名称
|
||
"""
|
||
# 如果有长期记忆,使用增强版提示词
|
||
if speaker_profiles:
|
||
# 使用更多历史记录(最多15条),保持完整的会议上下文
|
||
h_str = self._format_history(history[-15:])
|
||
|
||
# 构建说话人特征库摘要
|
||
profiles_summary = self._format_profiles(speaker_profiles)
|
||
|
||
prompt = self.SYSTEM_PROMPT_WITH_MEMORY.format(
|
||
history=h_str or "暂无历史记录",
|
||
speaker_profiles=profiles_summary,
|
||
current_text=text
|
||
)
|
||
|
||
res = await self._call_llm(prompt, "")
|
||
return self._extract_result(res)
|
||
else:
|
||
# 原有的简单识别逻辑
|
||
h_str = "\n".join([f"{h['role']}: {h['content']}" for h in history[-5:]])
|
||
res = await self._call_llm(self.SYSTEM_PROMPT, f"历史:{h_str}\n当前:{text}")
|
||
return self._extract_result(res)
|
||
|
||
def _format_history(self, history: List[Dict]) -> str:
|
||
"""
|
||
格式化历史对话记录,添加序号和时间顺序
|
||
|
||
Args:
|
||
history: 历史对话列表
|
||
|
||
Returns:
|
||
格式化后的对话文本
|
||
"""
|
||
if not history:
|
||
return "暂无历史记录"
|
||
|
||
lines = []
|
||
for idx, h in enumerate(history, 1):
|
||
content_preview = h['content'][:150] # 只显示前150字符
|
||
if len(h['content']) > 150:
|
||
content_preview += "..."
|
||
lines.append(f"[{idx}] {h['role']}: {content_preview}")
|
||
|
||
return "\n".join(lines)
|
||
|
||
def _format_profiles(self, speaker_profiles: Dict) -> str:
|
||
"""
|
||
格式化说话人特征库为可读文本
|
||
|
||
Args:
|
||
speaker_profiles: 说话人特征字典
|
||
|
||
Returns:
|
||
格式化后的文本摘要
|
||
"""
|
||
if not speaker_profiles:
|
||
return "暂无说话人特征记录"
|
||
|
||
summary_lines = []
|
||
for speaker, profiles in speaker_profiles.items():
|
||
# 提取该说话人最近5条发言的文本片段
|
||
recent_texts = [p['text'] for p in profiles[-5:]]
|
||
texts_str = "; ".join(recent_texts)
|
||
|
||
# 统计该说话人的发言次数
|
||
speaking_count = len(profiles)
|
||
|
||
summary_lines.append(f"- 【{speaker}】(发言{speaking_count}次): {texts_str}")
|
||
|
||
return "\n".join(summary_lines)
|
||
|
||
# ========================
|
||
# 4. A2A 编排器
|
||
# ========================
|
||
|
||
class A2AOrchestrator:
|
||
"""
|
||
Agent-to-Agent 编排器
|
||
|
||
职责:
|
||
- 协调多个 Agent 的协作流程
|
||
- 处理 ASR 识别结果,依次调用 Auditor 和 Profiler
|
||
- 管理会话状态(缓冲和历史记录)
|
||
- 将最终结果推送给前端
|
||
"""
|
||
|
||
def __init__(self, client: AsyncOpenAI, redis_mgr: RedisManager):
|
||
"""
|
||
初始化编排器
|
||
|
||
Args:
|
||
client: OpenAI 异步客户端
|
||
redis_mgr: Redis 管理器实例
|
||
"""
|
||
self.redis = redis_mgr
|
||
self.auditor = AuditorAgent(client) # 审计 Agent
|
||
self.profiler = ProfilerAgent(client) # 分析 Agent
|
||
|
||
async def process_asr_fragment(self, session_id: str, fragment: str, websocket):
|
||
"""
|
||
处理 ASR 识别的文本片段(核心流程)
|
||
|
||
流程:
|
||
1. 将新片段追加到缓冲区
|
||
2. 让 Auditor 判断发言是否完成
|
||
3. 如果完成,让 Profiler 使用长期记忆识别说话人
|
||
4. 将识别结果保存到长期记忆
|
||
5. 将结果存入历史,清空缓冲
|
||
6. 推送结果到前端
|
||
|
||
Args:
|
||
session_id: 会话 ID
|
||
fragment: ASR 识别的新文本片段
|
||
websocket: 前端 WebSocket 连接(用于推送结果)
|
||
"""
|
||
# 1. 追加新片段到缓冲区
|
||
await self.redis.append_buffer(session_id, fragment)
|
||
buf = await self.redis.get_buffer(session_id)
|
||
|
||
# 2. 检查发言是否完成
|
||
if await self.auditor.check(buf):
|
||
logger.info(f"[Auditor] 发言完成检测通过 - 会话ID: {session_id}")
|
||
|
||
# 3. 获取长期记忆(说话人特征库)
|
||
speaker_profiles = await self.redis.get_speaker_profiles(session_id)
|
||
logger.debug(f"获取说话人特征库 - 会话ID: {session_id}, 说话人数量: {len(speaker_profiles)}")
|
||
|
||
# 4. 识别说话人身份(使用长期记忆)
|
||
history = await self.redis.get_history(session_id)
|
||
speaker = await self.profiler.identify(buf, history, speaker_profiles)
|
||
logger.info(f"[Profiler] 识别说话人: {speaker} - 会话ID: {session_id}")
|
||
|
||
# 5. 保存到长期记忆(更新说话人特征库)
|
||
await self.redis.save_speaker_profile(session_id, speaker, buf)
|
||
|
||
# 6. 保存到历史记录并清空缓冲
|
||
await self.redis.push_history(session_id, speaker, buf)
|
||
await self.redis.clear_buffer(session_id)
|
||
|
||
# 7. 构造结果消息
|
||
segments = [{
|
||
"id": 0,
|
||
"speaker": speaker,
|
||
"embedding": None, # 占位(预留字段)
|
||
"seek": 0,
|
||
"full_text": buf.strip()
|
||
}]
|
||
|
||
msg = {
|
||
"mode": "offline-speaker",
|
||
"type": "asr_with_speaker",
|
||
"segments": segments,
|
||
"speaker": speaker
|
||
}
|
||
|
||
# 推送结果到前端(失败时静默处理)
|
||
try:
|
||
await websocket.send(json.dumps(msg, ensure_ascii=False))
|
||
except:
|
||
pass
|
||
|
||
|
||
def build_acoustic_speaker_message(data: Dict) -> Optional[Dict]:
|
||
text = (data.get("text") or "").strip()
|
||
speaker = (data.get("spk_name") or "").strip()
|
||
if not text or not speaker or speaker.lower() == "unknown":
|
||
return None
|
||
|
||
segment = {
|
||
"id": 0,
|
||
"speaker": speaker,
|
||
"embedding": None,
|
||
"seek": 0,
|
||
"full_text": text,
|
||
"start": None,
|
||
"end": None,
|
||
}
|
||
msg = {
|
||
"mode": "offline-speaker",
|
||
"type": "asr_with_speaker",
|
||
"source": "acoustic",
|
||
"speaker": speaker,
|
||
"spk_score": data.get("spk_score"),
|
||
"text": text,
|
||
"wav_name": data.get("wav_name"),
|
||
"is_final": data.get("is_final", True),
|
||
"segments": [segment],
|
||
}
|
||
if "timestamp" in data:
|
||
msg["timestamp"] = data["timestamp"]
|
||
if "sentence_info" in data:
|
||
msg["sentence_info"] = data["sentence_info"]
|
||
return msg
|
||
|
||
|
||
# ========================
|
||
# 5. WebSocket 处理与任务追踪
|
||
# ========================
|
||
|
||
async def forward_audio(websocket):
|
||
"""
|
||
WebSocket 连接处理函数(核心入口)
|
||
|
||
职责:
|
||
1. 接受前端 WebSocket 连接
|
||
2. 建立到 ASR 服务的下游连接
|
||
3. 双向转发消息(前端 -> ASR,ASR -> 前端)
|
||
4. 拦截 ASR 结果并触 Agent 处理
|
||
5. 确保所有后台任务完成后才退出(防止数据丢失)
|
||
|
||
Args:
|
||
websocket: 前端 WebSocket 连接对象
|
||
"""
|
||
logger.info("客户端已连接")
|
||
|
||
client = None
|
||
redis_mgr = None
|
||
orchestrator = None
|
||
use_llm_role_identification = USE_LLM_ROLE_IDENTIFICATION
|
||
|
||
if use_llm_role_identification:
|
||
if AsyncOpenAI is None:
|
||
logger.error("openai package is not installed; LLM role identification disabled")
|
||
use_llm_role_identification = False
|
||
elif redis is None:
|
||
logger.error("redis package is not installed; LLM role identification disabled")
|
||
use_llm_role_identification = False
|
||
else:
|
||
# 初始化 LLM 客户端、Redis 管理器和编排器
|
||
client = AsyncOpenAI(api_key=LLM_API_KEY, base_url=LLM_BASE_URL)
|
||
redis_mgr = RedisManager(REDIS_URL)
|
||
orchestrator = A2AOrchestrator(client, redis_mgr)
|
||
|
||
# 用于追踪后台任务,防止连接关闭时任务被强制销毁
|
||
background_tasks: Set[asyncio.Task] = set()
|
||
|
||
try:
|
||
# 连接到目标 ASR 服务
|
||
# 添加连接参数以提高鲁棒性
|
||
async with websockets.connect(
|
||
TARGET_WS_URL,
|
||
close_timeout=10, # 关闭超时
|
||
ping_timeout=20, # ping 超时
|
||
max_queue=1024, # 消息队列大小
|
||
subprotocols=["binary"],
|
||
) as target_ws:
|
||
logger.info(f"已连接到 ASR 服务: {TARGET_WS_URL}")
|
||
|
||
# 发送 ASR 配置(双通道模式,PCM 音频格式)
|
||
config = {"mode": "2pass", "chunk_size": [10, 10, 10], "wav_format": "pcm"}
|
||
await target_ws.send(json.dumps(config))
|
||
logger.debug(f"发送 ASR 配置: {config}")
|
||
|
||
async def frontend_to_target():
|
||
"""
|
||
上行转发:前端音频数据 -> ASR 服务
|
||
"""
|
||
async for message in websocket:
|
||
await target_ws.send(message)
|
||
|
||
async def target_to_frontend():
|
||
"""
|
||
下行转发:ASR 识别结果 -> 前端
|
||
同时拦截结果触发 Agent 处理
|
||
"""
|
||
async for message in target_ws:
|
||
# 1. 透传 ASR 原始结果到前端(字幕显示)
|
||
await websocket.send(message)
|
||
|
||
# 2. 如果是文本消息,尝试触发 Agent 处理
|
||
if isinstance(message, str):
|
||
try:
|
||
data = json.loads(message)
|
||
# 只处理最终识别结果(2pass-offline 模式)
|
||
if data.get("mode") == "2pass-offline":
|
||
text = data.get("text", "").strip()
|
||
speaker_msg = build_acoustic_speaker_message(data)
|
||
if speaker_msg is not None:
|
||
await websocket.send(json.dumps(speaker_msg, ensure_ascii=False))
|
||
|
||
if text and use_llm_role_identification and orchestrator is not None:
|
||
logger.debug(f"收到 ASR 识别结果: {text[:50]}...")
|
||
# 创建后台任务处理 Agent 逻辑(不阻塞主循环)
|
||
t = asyncio.create_task(
|
||
orchestrator.process_asr_fragment("room_101", text, websocket)
|
||
)
|
||
background_tasks.add(t)
|
||
# 任务完成后自动从追踪集合中移除
|
||
t.add_done_callback(background_tasks.discard)
|
||
except Exception as e:
|
||
logger.error(f"处理 ASR 结果时出错: {e}", exc_info=True)
|
||
|
||
# 并发运行两个方向的数据流
|
||
await asyncio.gather(frontend_to_target(), target_to_frontend())
|
||
|
||
except websockets.exceptions.InvalidMessage as e:
|
||
logger.error(f"WebSocket 握手失败 - ASR 服务可能未正常运行或 URL 配置错误: {e}")
|
||
logger.error(f"请检查: 1) ASR 服务是否在 {TARGET_WS_URL} 正常运行")
|
||
logger.error(f" 2) ASR 服务的 URL 路径是否正确")
|
||
logger.error(f" 3) ASR 服务是否需要特定的认证或配置")
|
||
except ConnectionRefusedError:
|
||
logger.error(f"连接被拒绝 - 无法连接到 ASR 服务 {TARGET_WS_URL}")
|
||
logger.error(f"请检查 ASR 服务是否已启动")
|
||
except Exception as e:
|
||
logger.error(f"连接错误: {type(e).__name__}: {e}", exc_info=True)
|
||
finally:
|
||
# --- 优雅退出处理 ---
|
||
if background_tasks:
|
||
logger.info(f"等待 {len(background_tasks)} 个后台 Agent 任务完成...")
|
||
# 等待所有剩余任务跑完,最多等 5 秒(防止无限等待)
|
||
await asyncio.wait(background_tasks, timeout=5)
|
||
logger.debug("所有后台任务已完成或超时")
|
||
|
||
# 释放资源
|
||
if client is not None:
|
||
await client.close()
|
||
if redis_mgr is not None:
|
||
await redis_mgr.close()
|
||
logger.info("客户端已断开连接")
|
||
|
||
|
||
async def main():
|
||
"""
|
||
主函数:启动 WebSocket 服务器
|
||
"""
|
||
logger.info(f"启动 WebSocket 智能语义分析服务: ws://{LOCAL_HOST}:{LOCAL_PORT}")
|
||
logger.info(f"Agent 模型: {LLM_MODEL}")
|
||
logger.info(f"ASR 服务地址: {TARGET_WS_URL}")
|
||
logger.info(f"Redis 地址: {REDIS_URL}")
|
||
# 启动服务器并持续运行(asyncio.Future() 永不完成)
|
||
async with websockets.serve(forward_audio, LOCAL_HOST, LOCAL_PORT):
|
||
await asyncio.Future()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
try:
|
||
asyncio.run(main())
|
||
except KeyboardInterrupt:
|
||
pass
|