""" 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 返回结果中提取 标签内容 Args: text: LLM 返回的完整文本 Returns: 提取的结果,如果没有标签则返回原文 """ match = re.search(r'(.*?)', text, re.DOTALL | re.IGNORECASE) return match.group(1).strip() if match else text.strip() class AuditorAgent(BaseAgent): """ 审计 Agent:判断当前发言是否完成 功能: - 分析累积的文本缓冲 - 判断说话人是否已经完成一次完整的发言 - 判断依据:语义完整性、停顿暗示等 """ # 系统提示词:定义 Agent 的角色和输出格式 SYSTEM_PROMPT = "你是一个对话状态监测专家。请判断发言是否结束,放入 true/false 中。" 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} 中选择发言人,放入 名称 中。" # 增强版提示词:结合长期记忆 SYSTEM_PROMPT_WITH_MEMORY = f"""你是一个会议身份识别专家。请根据以下信息识别发言人: 候选说话人列表:{CANDIDATE_SPEAKERS} === 会议历史记录(按时间顺序)=== {{history}} === 说话人特征库(长期记忆)=== {{speaker_profiles}} === 当前待识别的发言 === {{current_text}} === 识别指南 === 请仔细分析并综合判断: 1. **对话流程分析**: - 查看历史记录中的对话顺序和上下文 - 判断当前发言是对之前谁的话题的回应、延续或补充 - 注意对话的逻辑关系和互动模式 2. **说话风格匹配**: - 每个角色的用词习惯、句式结构、语气的特点 - 专有名词、技术术语的使用习惯 - 提问方式、陈述方式、指令方式的差异 3. **角色职责判断**: - 根据各角色的职责范围判断发言内容的合理性 - 例如:局长通常做总结和决策,主持人负责控场,专员负责具体业务汇报 4. **话题连贯性**: - 当前发言与该角色之前发言的话题是否连贯 - 是否体现了该角色应有的专业领域关注点 将识别结果放入 说话人名称 中。""" 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