This commit is contained in:
田岩
2025-12-03 20:54:23 +08:00
parent f250b21b38
commit 11ce05edcf
972 changed files with 121839 additions and 824 deletions
+4 -3
View File
@@ -6,9 +6,10 @@ LOG_LEVEL=INFO
LOG_FILE=logs/app.log
# ASR服务配置
ASR_SERVICE_URL=http://localhost:5000/asr
ASR_TIMEOUT=30 # 超时时间(秒)
ASR_RETRY_TIMES=2 # 重试次数
ASR_HOST=10.10.10.202
ASR_PORT=10096
ASR_TIMEOUT=30 # 超时时间(秒)
ASR_RETRY_TIMES=2 # 重试次数
# LLM服务配置
LLM_SERVICE_URL=http://localhost:6000/chat
+7
View File
@@ -0,0 +1,7 @@
<component name="ProjectDictionaryState">
<dictionary name="project">
<words>
<w>tymas</w>
</words>
</dictionary>
</component>
+56 -16
View File
@@ -1,20 +1,17 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from pydantic import Field
from pathlib import Path
from typing import List
ROOT_DIR = Path(__file__).parent.parent.parent
class Settings(BaseSettings):
# 应用配置
APP_PORT: int = Field(default=8000, description="服务端口")
LOG_LEVEL: str = Field(default="INFO", description="日志级别")
LOG_FILE: Path = Field(default=ROOT_DIR / "logs/app.log", description="日志文件路径")
# ASR服务配置
ASR_SERVICE_URL: str = Field(..., description="ASR服务地址")
ASR_TIMEOUT: int = Field(default=30, description="ASR超时时间(秒)")
ASR_RETRY_TIMES: int = Field(default=2, description="ASR重试次数")
# LLM服务配置
LLM_SERVICE_URL: str = Field(..., description="LLM服务地址")
LLM_TIMEOUT: int = Field(default=60, description="LLM超时时间(秒)")
@@ -38,20 +35,62 @@ class Settings(BaseSettings):
# -------------------------- 新增:服务版本配置 --------------------------
# ASR当前使用版本(对应ASR_REGISTRY中的key
ASR_CURRENT_VERSION: str = Field(default="local_v1", description="ASR服务当前版本")
# 本地ASR专属配置(仅local_v1版本使用)
LOCAL_ASR_MODEL_PATH: Path = Field(default=ROOT_DIR / "models/asr/local_model", description="本地ASR模型路径")
# 百度云ASR专属配置(仅baidu_v2版本使用)
BAIDU_ASR_API_KEY: str = Field(default="", description="百度云ASR API Key")
BAIDU_ASR_SECRET_KEY: str = Field(default="", description="百度云ASR Secret Key")
ASR_CURRENT_VERSION: str = Field(default="FunASR", description="ASR服务当前版本")
# ASR 公共配置
ASR_TIMEOUT: int = Field(default=5, description="ASR连接超时时间(秒)")
ASR_RETRY_TIMES: int = Field(default=3, description="ASR公共重试次数")
# ASR服务配置
# FunASR 专属配置
ASR_HOST: str = Field(..., description="FunASR服务地址")
ASR_PORT: int = Field(..., description="FunASR服务端口")
ASR_MODE: str = Field(default="2pass", description="FunASR识别模式")
ASR_CHUNK_SIZE: List[int] = Field(default=[5, 10, 5], description="FunASR分片大小配置")
ASR_CHUNK_INTERVAL: int = Field(default=10, description="FunASR分片间隔(毫秒)")
ASR_USE_ITN: int = Field(default=1, description="FunASR是否启用数字转换(1=启用,0=禁用)")
ASR_HOTWORDS: str = Field(default="", description="FunASR热词列表(逗号分隔)")
ASR_RECONNECT_MAX_TIMES: int = Field(default=3, description="FunASR连接重连最大次数")
ASR_POOL_SIZE: int = Field(default=2, description="FunASR连接池大小")
ASR_AUDIO_QUEUE_SIZE: int = Field(default=10000, description="FunASR音频队列最大长度")
# 音频参数(FunASR要求)
ASR_SAMPLE_RATE: int = Field(default=16000, description="音频采样率(Hz")
ASR_CHANNELS: int = Field(default=1, description="音频声道数(1=单声道)")
ASR_SAMPLE_WIDTH: int = Field(default=2, description="音频采样宽度(字节)")
ASR_FRAME_SIZE: int = Field(default=1024, description="音频帧大小")
# LLM当前使用版本(对应LLM_REGISTRY中的key
LLM_CURRENT_VERSION: str = Field(default="local", description="LLM服务当前版本")
# OpenAI LLM专属配置(仅openai版本使用)
OPENAI_API_KEY: str = Field(default="", description="OpenAI API Key")
OPENAI_BASE_URL: str = Field(default="https://api.openai.com/v1", description="OpenAI接口地址")
# 本地LLM专属配置(仅local版本使用)
LOCAL_LLM_MODEL_PATH: Path = Field(default=ROOT_DIR / "models/llm/local_model", description="本地LLM模型路径")
# -------------------------- Dify API 配置(专属) --------------------------
DIFY_BASE_URL: str = Field(
default="http://10.10.10.202:8088/v1",
description="Dify平台API基础URL(如http://xxx:8088/v1"
)
DIFY_API_KEY: str = Field(
default="app-m7HZNV1aGiheh3wr6wNVHFxX",
description="Dify平台API密钥(在Dify应用设置中获取,格式为app-xxx)"
)
DIFY_TIMEOUT: int = Field(
default=30,
description="Dify API请求超时时间(秒)"
)
DIFY_DEFAULT_SCENE: str = Field(
default="通用聊天场景",
description="Dify默认场景描述(传给inputs.scene_description参数)"
)
DIFY_STREAM_CHUNK_SIZE: int = Field(
default=1024,
description="Dify流式响应读取块大小(字节)"
)
# Chat 服务配置(新增)
CHAT_CURRENT_VERSION: str = "DefaultChat" # 对应工厂类的注册名
CHAT_BASE_URL: str = "http://10.10.10.202/v1"
CHAT_API_KEY: str = "app-m7HZNV1aGiheh3wr6wNVHFxX" # 替换为真实 API-Key
CHAT_TIMEOUT: int = 300 # 流式请求超时时间(秒)
CHAT_RETRY_TIMES: int = 3 # 重试次数
CHAT_POOL_SIZE: int = 5 # 连接池大小
# TTS当前使用版本(对应TTS_REGISTRY中的key
TTS_CURRENT_VERSION: str = Field(default="pyttsx3", description="TTS服务当前版本")
@@ -64,6 +103,7 @@ class Settings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8")
settings = Settings()
# 确保模型目录存在(本地版本需要)
# settings.LOCAL_ASR_MODEL_PATH.parent.mkdir(parents=True, exist_ok=True)
@@ -0,0 +1,48 @@
# audio_ai_chat/asr/asr_manager.py
from typing import Optional,Tuple
from .base import ASRBase
from .factory import ASRFactory
from audio_ai_chat.config.settings import settings
class ASRManager: # 类名与文件名呼应
"""ASR 管理器:负责实例生命周期、连接池复用、资源管理"""
_instance: Optional[ASRBase] = None # 单例存储
@classmethod
async def initialize(cls) -> Tuple[bool, str]:
try:
# 1. 创建 ASR 实例
cls._instance = ASRFactory.get_asr_client()
# 2. 初始化连接池
pool_init_success = await cls._instance.initialize()
if not pool_init_success:
return False, "ASR 连接池初始化失败"
# 3. 异步调用获取有效连接数(添加 await)
valid_conn_count = await cls._instance.get_valid_connection_count()
if valid_conn_count == 0:
return False, f"有效连接数为 0(配置池大小:{settings.ASR_POOL_SIZE}"
return True, f"初始化成功:有效连接数 {valid_conn_count}"
except Exception as e:
return False, f"初始化失败:{str(e)}"
@classmethod
def get_instance(cls) -> Optional[ASRBase]:
"""获取全局 ASR 实例(业务代码调用)"""
return cls._instance
@classmethod
async def close(cls):
"""关闭 ASR 实例和连接池(FastAPI 关闭时调用)"""
if cls._instance:
await cls._instance.close()
cls._instance = None
print("ASR 管理器:实例和连接池已关闭")
# 未来可扩展的管理功能
@classmethod
def is_healthy(cls) -> bool:
"""检查 ASR 实例健康状态(管理功能扩展)"""
return cls._instance is not None
+38 -13
View File
@@ -1,7 +1,12 @@
# audio_ai_chat/asr/base.py
from abc import ABC, abstractmethod
from typing import Optional, Coroutine
from typing import Optional, Dict, Callable, Awaitable
from audio_ai_chat.config.settings import settings
# 定义回调函数类型(异步函数,接收 ASR 结果字典)
ASRResultCallback = Callable[[Dict], Awaitable[None]]
class ASRBase(ABC):
"""ASR服务统一抽象接口"""
def __init__(self):
@@ -10,16 +15,36 @@ class ASRBase(ABC):
self.retry_times = settings.ASR_RETRY_TIMES
@abstractmethod
async def recognize(
self,
voice_data: bytes,
user_id: Optional[str] = None,
**kwargs # 兼容不同版本的额外参数
) -> str:
"""
语音识别核心方法(所有ASR版本必须实现)
:param voice_data: 语音二进制数据
:param user_id: 用户ID(可选)
:return: 识别后的文本
"""
async def initialize(self) -> bool:
"""初始化ASR服务(如连接池初始化)"""
pass
@abstractmethod
async def get_connection(self) -> Optional[object]:
"""获取ASR连接对象"""
pass
@abstractmethod
async def push_audio(self, conn: object, audio_data: bytes) -> bool:
"""推送音频数据到ASR服务"""
pass
@abstractmethod
async def start_communication(self, conn: object, callback: ASRResultCallback) -> None:
"""启动ASR通信(发送音频+接收结果)"""
pass
@abstractmethod
async def release_connection(self, conn: object) -> None:
"""释放ASR连接"""
pass
@abstractmethod
async def close(self) -> None:
"""关闭ASR服务(释放所有连接)"""
pass
@abstractmethod
async def get_valid_connection_count(self) -> int:
"""获取有效连接数(异步方法,子类必须实现)"""
pass
+21 -21
View File
@@ -1,26 +1,26 @@
# audio_ai_chat/asr/factory.py
from typing import Type
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from .base import ASRBase
# from .version1 import LocalOfflineASR
# from .version2 import BaiduASR
#
# # 注册所有ASR版本:key=配置中的版本名,value=对应的类
# ASR_REGISTRY: dict[str, Type[ASRBase]] = {
# "local_v1": LocalOfflineASR,
# "baidu_v2": BaiduASR,
# # 新增版本时,只需在这里注册:"新版本名": 新类名
# }
from .fun_asr import FunASR # 对应原asr_client.py的实现类
# class ASRFactory:
# """ASR服务工厂类:根据配置创建对应版本的实例"""
# @staticmethod
# def get_asr_client() -> ASRBase:
# # 从配置中获取当前指定的ASR版本
# current_version = settings.ASR_CURRENT_VERSION
# if current_version not in ASR_REGISTRY:
# raise ServiceCallError(
# f"不支持的ASR版本:{current_version},可选版本:{list(ASR_REGISTRY.keys())}"
# )
# # 创建并返回对应版本的实例
# return ASR_REGISTRY[current_version]()
# 注册所有ASR版本:key=配置中的版本名,value=对应的类
ASR_REGISTRY: dict[str, Type[ASRBase]] = {
"FunASR": FunASR,
# 新增版本时,只需在这里注册:"新版本名": 新类名
}
class ASRFactory:
"""ASR服务工厂类:根据配置创建对应版本的实例"""
@staticmethod
def get_asr_client() -> ASRBase:
# 从配置中获取当前指定的ASR版本
current_version = settings.ASR_CURRENT_VERSION
if current_version not in ASR_REGISTRY:
raise ServiceCallError(
f"不支持的ASR版本:{current_version},可选版本:{list(ASR_REGISTRY.keys())}"
)
# 创建并返回对应版本的实例
return ASR_REGISTRY[current_version]()
@@ -0,0 +1 @@
from .fun_asr import FunASR
@@ -0,0 +1,264 @@
import asyncio
import json
import websockets
from typing import Optional, List, Dict, Callable, Awaitable, Any
from dataclasses import dataclass, field
from ..base import ASRBase, ASRResultCallback
from audio_ai_chat.config.settings import settings
# 从配置读取参数(替换原硬编码配置)
AUDIO_PARAMS = {
"sample_rate": settings.ASR_SAMPLE_RATE,
"channels": settings.ASR_CHANNELS,
"sample_width": settings.ASR_SAMPLE_WIDTH,
"frame_size": settings.ASR_FRAME_SIZE
}
ASR_CONFIG = {
"host": settings.ASR_HOST,
"port": settings.ASR_PORT,
"mode": settings.ASR_MODE,
"chunk_size": settings.ASR_CHUNK_SIZE,
"chunk_interval": settings.ASR_CHUNK_INTERVAL,
"use_itn": settings.ASR_USE_ITN,
"hotwords": settings.ASR_HOTWORDS,
"reconnect_max_times": settings.ASR_RECONNECT_MAX_TIMES,
"pool_size": settings.ASR_POOL_SIZE,
"audio_queue_size": settings.ASR_AUDIO_QUEUE_SIZE
}
@dataclass
class ASRConnection:
"""ASR 连接对象(内置音频队列)"""
ws: Optional[websockets.WebSocketClientProtocol] = None
is_busy: bool = False
is_alive: bool = False
reconnect_count: int = 0
audio_queue: asyncio.Queue = field(default_factory=lambda: asyncio.Queue(maxsize=ASR_CONFIG["audio_queue_size"]))
stop_event: asyncio.Event = field(default_factory=asyncio.Event)
class FunASR(ASRBase):
"""FunASR实现类"""
def __init__(self):
super().__init__()
self._connection_pool: List[ASRConnection] = []
self._pool_lock = asyncio.Lock() # 异步锁
async def initialize(self) -> bool:
"""初始化ASR连接池"""
print(f"开始初始化 FunASR 连接池,大小:{ASR_CONFIG['pool_size']}")
tasks = [self._create_single_connection() for _ in range(ASR_CONFIG["pool_size"])]
connections = await asyncio.gather(*tasks)
self._connection_pool = [conn for conn in connections if conn.is_alive]
print(f"FunASR 连接池初始化完成,有效连接数:{len(self._connection_pool)}")
return len(self._connection_pool) > 0
async def get_connection(self) -> Optional[ASRConnection]:
"""从连接池获取空闲连接(实现抽象方法)"""
async with self._pool_lock:
# 查找空闲连接
idle_conns = [
conn for conn in self._connection_pool
if not conn.is_busy and conn.is_alive
]
if idle_conns:
conn = idle_conns[0]
conn.is_busy = True
conn.stop_event.clear()
return conn
# 连接池未满时创建新连接
if len(self._connection_pool) < ASR_CONFIG["pool_size"]:
new_conn = await self._create_single_connection()
if new_conn.is_alive:
new_conn.is_busy = True
self._connection_pool.append(new_conn)
return new_conn
print("FunASR 连接池无空闲连接")
return None
async def push_audio(self, conn: ASRConnection, audio_data: bytes) -> bool:
"""推送音频数据到ASR连接队列(实现抽象方法)"""
if not isinstance(conn, ASRConnection):
print("无效的ASR连接对象")
return False
if not conn or not conn.is_alive or conn.stop_event.is_set():
return False
try:
conn.audio_queue.put_nowait(audio_data)
return True
except asyncio.QueueFull:
print("FunASR 音频队列已满,丢弃当前音频帧")
return False
async def start_communication(self, conn: ASRConnection, callback: ASRResultCallback) -> None:
"""启动ASR通信(发送音频+接收结果,实现抽象方法)"""
if not isinstance(conn, ASRConnection):
await callback({"error": "无效的ASR连接对象", "text": ""})
return
await self._handle_communication(conn, callback)
async def release_connection(self, conn: ASRConnection) -> None:
"""释放ASR连接(实现抽象方法)"""
if not isinstance(conn, ASRConnection):
print("无效的ASR连接对象,无法释放")
return
async with self._pool_lock:
conn.is_busy = False
conn.stop_event.set()
# 清空队列
while not conn.audio_queue.empty():
try:
conn.audio_queue.get_nowait()
except asyncio.QueueEmpty:
break
# 重连逻辑
if not conn.is_alive and conn.reconnect_count < ASR_CONFIG["reconnect_max_times"]:
print(f"尝试重连 FunASR 连接(次数:{conn.reconnect_count + 1}")
new_conn = await self._create_single_connection()
if new_conn.is_alive:
if conn in self._connection_pool:
idx = self._connection_pool.index(conn)
self._connection_pool[idx] = new_conn
else:
conn.reconnect_count += 1
elif conn.reconnect_count >= ASR_CONFIG["reconnect_max_times"]:
if conn in self._connection_pool:
self._connection_pool.remove(conn)
print("FunASR 连接重连次数耗尽,已移除")
async def close(self) -> None:
"""关闭所有ASR连接(实现抽象方法)"""
async with self._pool_lock:
# for conn in self._connection_pool:
# conn.stop_event.set()
# if conn.ws and not conn.ws.closed:
# try:
# await conn.ws.close()
# print("FunASR 连接已关闭")
# except Exception as e:
# print(f"关闭 FunASR 连接失败:{e}")
self._connection_pool.clear()
print("FunASR 连接池已清空")
async def _create_single_connection(self) -> ASRConnection:
"""创建单个ASR连接(内部私有方法)"""
asr_conn = ASRConnection()
asr_uri = f"ws://{ASR_CONFIG['host']}:{ASR_CONFIG['port']}"
try:
ws = await websockets.connect(
asr_uri,
subprotocols=["binary"],
ping_interval=None,
open_timeout=self.timeout # 使用基类的超时配置
)
asr_conn.ws = ws
asr_conn.is_alive = True
# 发送初始化配置
init_msg = json.dumps({
"mode": ASR_CONFIG["mode"],
"chunk_size": ASR_CONFIG["chunk_size"],
"chunk_interval": ASR_CONFIG["chunk_interval"],
"wav_name": "pool_connection",
"is_speaking": True,
"hotwords": ASR_CONFIG["hotwords"],
"itn": bool(ASR_CONFIG["use_itn"]),
"audio_fs": AUDIO_PARAMS["sample_rate"]
})
await ws.send(init_msg)
print("FunASR 连接初始化成功")
return asr_conn
except Exception as e:
print(f"创建 FunASR 连接失败:{e}")
asr_conn.is_alive = False
return asr_conn
async def _handle_communication(
self,
asr_conn: ASRConnection,
result_callback: ASRResultCallback
):
"""处理ASR通信细节(内部私有方法)"""
if not asr_conn or not asr_conn.ws:
await result_callback({"error": "无可用 FunASR 连接", "text": ""})
return
# 发送音频任务
async def send_audio():
while not asr_conn.stop_event.is_set() and asr_conn.is_alive:
try:
pcm_data = await asyncio.wait_for(asr_conn.audio_queue.get(), timeout=1.0)
if pcm_data and asr_conn.is_alive:
await asr_conn.ws.send(pcm_data)
await asyncio.sleep(0.005)
except asyncio.TimeoutError:
continue
except Exception as e:
print(f"发送音频到 FunASR 失败:{e}")
asr_conn.is_alive = False
await result_callback({"error": f"音频发送失败:{str(e)}", "text": ""})
asr_conn.stop_event.set()
break
# 接收结果任务
async def recv_result():
while not asr_conn.stop_event.is_set() and asr_conn.is_alive:
try:
asr_result = await asr_conn.ws.recv()
result_json = json.loads(asr_result)
# 打印asr实时结果
# print(result_json.get("text", ""))
if result_json.get("timestamp", "") == '':
continue
result = {
"text": result_json.get("text", ""),
"mode": result_json.get("mode", ""),
"timestamp": result_json.get("timestamp", ""),
"is_final": result_json.get("is_final", True),
"error": ""
}
await result_callback(result)
except websockets.exceptions.ConnectionClosed:
print("FunASR 连接已关闭")
asr_conn.is_alive = False
await result_callback({"error": "FunASR 连接断开", "text": ""})
asr_conn.stop_event.set()
break
except Exception as e:
print(f"接收 FunASR 结果失败:{e}")
asr_conn.is_alive = False
await result_callback({"error": f"接收结果失败:{str(e)}", "text": ""})
asr_conn.stop_event.set()
break
try:
send_task = asyncio.create_task(send_audio())
recv_task = asyncio.create_task(recv_result())
await asyncio.gather(send_task, recv_task)
finally:
# 确保任务被取消
send_task.cancel()
recv_task.cancel()
try:
await send_task
await recv_task
except asyncio.CancelledError:
pass
# 自动释放连接
await self.release_connection(asr_conn)
async def get_valid_connection_count(self) -> int:
"""获取有效连接数(异步方法,保证线程安全)"""
async with self._pool_lock: # 异步锁,自动 acquire/release
# 过滤出 "存活" 且 "在连接池内" 的连接
valid_conns = [conn for conn in self._connection_pool if conn.is_alive]
return len(valid_conns)
+359 -50
View File
@@ -1,9 +1,21 @@
from typing import Optional, Dict, Any, List
from typing import Optional, Dict, Any, List, TypedDict
import asyncio
from datetime import datetime
from audio_ai_chat.config.logger import logger
from audio_ai_chat.core.llm.dify.dify import LLMConversation
# from audio_ai_chat.core.llm.factory import LLMFactory
from audio_ai_chat.core.llm.base import LLMBase
# 定义对话历史条目类型(TypedDict 用于类型提示,更清晰)
class ChatHistoryItem(TypedDict):
"""对话历史条目结构(强类型定义)"""
role: str # 发言人角色:"user"(用户)、"assistant"(助手)、"system"(系统)
content: str # 对话内容(ASR转写结果/大模型回复/系统提示)
timestamp: str # 对话时间(ISO 8601格式,如 "2024-05-20T14:30:00.123Z"
source: str # 内容来源:"asr"(语音转写)、"text"(纯文本输入)、"llm"(大模型生成)、"system"(系统配置)
def get_current_iso_timestamp() -> str:
"""获取当前时间的ISO 8601格式字符串(UTC时间)"""
return datetime.utcnow().isoformat(timespec="milliseconds") + "Z"
class ConnectionContext:
"""
@@ -15,30 +27,43 @@ class ConnectionContext:
"""
初始化连接上下文
:param client_id: WebSocket连接唯一标识(如id(websocket)
:param user_id: 用户唯一标识(从前端请求中获取)
"""
self.MAX_CHAT_HISTORY = 100 # 单个连接最大对话历史条数
self.MAX_QUEUE_SIZE = 50 # 单个连接消息队列最大长度
self.client_id = client_id # 连接唯一ID
self.created_at = asyncio.get_event_loop().time() # 连接创建时间
self.created_at = asyncio.get_event_loop().time() # 连接创建时间(时间戳)
self.created_at_str = datetime.utcnow().isoformat() + "Z" # 连接创建时间(ISO格式)
# 1. 大模型独立Session(每个连接创建一个新的LLM客户端实例)
# self.llm_session: LLMBase = LLMFactory.get_llm_client() # 独立Session
self.chat_history: List[Dict[str, str]] = [] # 该连接的对话历史([(user: "...", assistant: "..."), ...]
self.llm_session: Optional[LLMConversation] = None # 实际是DifyLLMClient实例
# 优化后的对话历史:List[ChatHistoryItem]
self.chat_history: List[ChatHistoryItem] = [] # 该连接的完整对话历史
# 2. 异步消息队列(用于缓存TTS结果,有序推送给前端)
self.tts_client = None
self.message_queue: asyncio.Queue[bytes] = asyncio.Queue()
# 3. 连接状态(可选:如是否正在处理请求、是否断开等)
# 3. 连接状态
self.is_active: bool = True
self.is_processing: bool = False
self.name = None
self.user_id = None
self.token = None
self.name: Optional[str] = None # 用户名
self.user_id: Optional[str] = None # 用户唯一标识
self.token: Optional[str] = None # 用户令牌
self.disconnect_time: Optional[datetime] = None # 断开时间(None 表示活跃)
# 4. ASR临时缓存(处理流式结果,避免重复存储)
self._current_asr_text: str = "" # 当前正在拼接的ASR文本
self._current_asr_metadata: Optional[Dict[str, Any]] = None # 当前ASR元数据
self.is_processing: bool = False
def set_user_info(self, token: str, user_id: str, name: str = "匿名用户"):
"""
二次设置用户信息(身份校验通过后调用)
:param token:
:param token: 用户令牌
:param user_id: 用户唯一标识(必填)
:param name: 用户名(可选,默认匿名)
"""
@@ -49,6 +74,28 @@ class ConnectionContext:
self.token = token
logger.debug(f"客户端 {self.client_id} 设置用户信息:user_id={user_id}, name={name}")
def init_llm_session(self):
"""初始化Dify客户端(每个连接一个实例,存入context)"""
# if not self.user_id:
# raise InitError(f"客户端 {self.client_id} 未设置用户信息,无法初始化Dify客户端")
if self.llm_session:
logger.warning(f"客户端 {self.client_id} Dify客户端已存在,无需重复初始化")
return
# 工厂类创建Dify客户端实例,存入当前context
self.llm_session = LLMFactory.get_llm_client(version="dify")
logger.debug(f"客户端 {self.client_id}user_id={self.user_id}Dify客户端初始化完成")
def update_dify_context(self, user_text: str, assistant_text: str, conversation_id: Optional[str]):
"""更新Dify会话上下文(隔离存储)"""
if conversation_id:
self.dify_conversation_id = conversation_id
# 限制历史长度(最多50轮)
self.chat_history.append({"user": user_text, "assistant": assistant_text})
if len(self.chat_history) > 50:
self.chat_history.pop(0)
logger.debug(
f"客户端 {self.client_id} Dify上下文更新:conversation_id={self.dify_conversation_id},历史长度={len(self.dify_history)}")
async def add_message_to_queue(self, message: bytes):
"""将TTS结果添加到消息队列(异步安全)"""
if not self.is_active:
@@ -56,6 +103,15 @@ class ConnectionContext:
await self.message_queue.put(message)
logger.debug(f"消息队列添加数据:client_id={self.client_id},队列长度={self.message_queue.qsize()}")
def complete_initialization(self):
"""标记完成初始化(必须确保Dify客户端已创建)"""
# if not self.user_id:
# raise InitError(f"客户端 {self.client_id} 未设置用户信息")
# if not self.llm_session:
# raise InitError(f"客户端 {self.client_id} 未初始化Dify客户端")
self.is_initialized = True
logger.info(f"客户端 {self.client_id}user_id={self.user_id})完整初始化完成")
async def get_message_from_queue(self) -> Optional[bytes]:
"""从消息队列获取消息(异步阻塞,直到有消息或连接断开)"""
try:
@@ -65,81 +121,334 @@ class ConnectionContext:
logger.debug(f"消息队列超时:client_id={self.client_id},无新消息")
return None
def update_chat_history(self, user_text: str, assistant_text: str):
"""更新该连接的对话历史"""
self.chat_history.append({
"user": user_text,
"assistant": assistant_text
})
# 可选:限制历史长度(避免内存溢出)
if len(self.chat_history) > 50:
self.chat_history.pop(0) # 删除最早的历史
def add_chat_history(self, item: ChatHistoryItem):
"""
添加对话历史条目(统一接口,支持用户/助手/系统消息)
:param item: 符合 ChatHistoryItem 结构的对话条目
"""
# 补全必填字段(防止遗漏)
if "timestamp" not in item:
item["timestamp"] = datetime.utcnow().isoformat() + "Z"
if "source" not in item:
item["source"] = "unknown"
self.chat_history.append(item)
# 可选:限制历史长度(避免内存溢出,保留最近100条)
if len(self.chat_history) > 100:
removed_item = self.chat_history.pop(0)
logger.debug(
f"对话历史超出限制,删除最早条目:{removed_item['timestamp']} - {removed_item['role']}: {removed_item['content'][:20]}...")
logger.debug(
f"添加对话历史:client_id={self.client_id}"
f"role={item['role']}content={item['content'][:30]}..."
)
def add_asr_result(self, asr_result: Dict[str, Any]):
"""
处理ASR结果,拼接流式文本,最终结果存入对话历史
:param asr_result: ASR返回的结果字典(含text、is_final、timestamp等)
"""
if asr_result.get("error"):
logger.error(f"ASR错误:client_id={self.client_id}error={asr_result['error']}")
return
# 提取ASR核心信息
asr_text = asr_result.get("text", "").strip()
is_final = asr_result.get("is_final", False)
asr_timestamp = asr_result.get("timestamp", "")
asr_mode = asr_result.get("mode", "")
# 缓存ASR元数据(流式过程中更新)
self._current_asr_metadata = {
"timestamp": asr_timestamp,
"mode": asr_mode,
"is_final": is_final,
"source": "asr"
}
# 拼接流式文本(处理部分结果)
if asr_text:
# 避免重复拼接(如果ASR返回重复文本)
if not self._current_asr_text.endswith(asr_text) and self._current_asr_text != asr_text:
self._current_asr_text += asr_text if not self._current_asr_text else f" {asr_text}"
# 当ASR返回最终结果时,存入对话历史
if is_final:
if self._current_asr_text:
# 构造对话历史条目
chat_item: ChatHistoryItem = {
"role": "user", # ASR结果属于用户输入
"content": self._current_asr_text,
"timestamp": datetime.utcnow().isoformat() + "Z",
"source": "asr",
"asr_metadata": self._current_asr_metadata
}
# 添加到对话历史
self.add_chat_history(chat_item)
# 清空临时缓存
self._current_asr_text = ""
self._current_asr_metadata = None
else:
logger.warning(f"ASR最终结果为空:client_id={self.client_id}")
def add_llm_result(self, llm_text: str):
"""
添加大模型回复到对话历史
:param llm_text: 大模型生成的回复文本
"""
if not llm_text.strip():
logger.warning(f"大模型回复为空:client_id={self.client_id}")
return
chat_item: ChatHistoryItem = {
"role": "assistant", # 大模型回复属于助手角色
"content": llm_text.strip(),
"timestamp": datetime.utcnow().isoformat() + "Z",
"source": "llm",
"asr_metadata": None # 大模型回复无ASR元数据
}
self.add_chat_history(chat_item)
def add_system_message(self, system_text: str):
"""
添加系统消息到对话历史(如错误提示、系统通知)
:param system_text: 系统消息文本
"""
chat_item: ChatHistoryItem = {
"role": "system", # 系统角色
"content": system_text.strip(),
"timestamp": datetime.utcnow().isoformat() + "Z",
"source": "system",
"asr_metadata": None
}
self.add_chat_history(chat_item)
def get_chat_history(self, limit: Optional[int] = None) -> List[ChatHistoryItem]:
"""
获取对话历史(支持限制返回条数)
:param limit: 限制返回的最新条数,None表示返回全部
:return: 过滤后的对话历史
"""
if limit and isinstance(limit, int) and limit > 0:
return self.chat_history[-limit:] # 返回最近N条
return self.chat_history.copy() # 返回全部(拷贝,避免外部修改)
def close(self):
"""关闭连接上下文,释放资源"""
self.is_active = False
self.is_processing = False
# 清空消息队列(可选)
# 清空消息队列
while not self.message_queue.empty():
try:
self.message_queue.get_nowait()
except asyncio.QueueEmpty:
break
logger.info(f"连接上下文已关闭:client_id={self.client_id}user_id={self.user_id}")
# 记录连接关闭日志(包含对话历史统计)
logger.info(
f"连接上下文已关闭:client_id={self.client_id}"
f"user_id={self.user_id}"
f"对话历史条数={len(self.chat_history)}"
)
def mark_disconnected(self):
self.is_active = False
self.disconnect_time = datetime.utcnow()
logger.info(f"客户端 {self.client_id} 标记为断开,待延迟清理(user_id={self.user_id}")
def __del__(self):
"""析构函数:确保资源释放"""
self.close()
# -------------------------- 核心:调用Dify流式接口 --------------------------
async def call_dify_stream(
self,
user_text: str,
tts_client,
asr_metadata: Optional[Dict[str, Any]] = None # 接收ASR元数据
) -> str:
"""
调用Dify流式接口,使用ChatHistoryItem存储完整历史
:param user_text: ASR识别后的用户文本
:param tts_client: TTS客户端实例
:param asr_metadata: ASR元数据(如置信度、语音时长等)
:return: Dify完整回复文本
"""
if not self.is_initialized:
raise InitError(f"客户端 {self.client_id} 未完成初始化,无法调用Dify")
if not user_text:
raise ValueError("用户输入文本不能为空")
full_response = ""
llm_metadata: Dict[str, Any] = {} # 存储Dify元数据
# 1. 添加用户输入到对话历史(user角色,source=asr
user_history_item: ChatHistoryItem = {
"role": "user",
"content": user_text,
"timestamp": get_current_iso_timestamp(),
"source": "asr",
"asr_metadata": asr_metadata, # 传入ASR元数据
"llm_metadata": None
}
self.add_chat_history_item(user_history_item)
# -------------------------- 流式回调函数(闭包访问context --------------------------
async def stream_callback(chunk: str, conversation_id: str, is_finished: bool):
nonlocal full_response, llm_metadata
if chunk and not is_finished:
# 累加完整回复
full_response += chunk
logger.debug(
f"客户端 {self.client_id} Dify流式片段:content_len={len(chunk)}, "
f"累计_len={len(full_response)}"
)
# 实时调用TTS合成音频
try:
tts_audio = await tts_client.synthesize(
text=chunk,
user_id=self.user_id
)
await self.add_tts_to_queue(tts_audio)
except Exception as e:
logger.error(f"客户端 {self.client_id} TTS合成失败:{str(e)}")
return
# 流式结束:添加助手回复到对话历史
if is_finished and full_response:
# 更新Dify会话ID和元数据
self.dify_conversation_id = conversation_id
llm_metadata = {
"conversation_id": conversation_id,
"response_mode": "streaming",
"full_response_len": len(full_response),
"timestamp": get_current_iso_timestamp()
}
# 添加助手回复到对话历史(assistant角色,source=llm
assistant_history_item: ChatHistoryItem = {
"role": "assistant",
"content": full_response,
"timestamp": get_current_iso_timestamp(),
"source": "llm",
"asr_metadata": None,
"llm_metadata": llm_metadata # 存储Dify元数据
}
self.add_chat_history_item(assistant_history_item)
# -------------------------- 调用Dify流式接口 --------------------------
try:
await self.llm_session.chat(
text=user_text,
user_id=self.user_id,
history=self.get_dify_compatible_history(), # 传入Dify兼容格式的历史
stream_callback=stream_callback,
response_mode="streaming"
)
except Exception as e:
logger.error(f"客户端 {self.client_id} Dify流式调用失败:{str(e)}")
raise
return full_response
# -------------------------- 关键修改:ConnectionManager 全局单例 --------------------------
class ConnectionManager:
"""
WebSocket连接全局管理器:维护所有活跃连接的上下文
提供创建、查询、删除连接上下文的接口(线程/异步安全)
"""
"""全局连接上下文管理器(支持延迟清理和重连复用)"""
_instance: Optional["ConnectionManager"] = None
_lock = asyncio.Lock() # 单例锁
def __new__(cls):
raise NotImplementedError("请使用 ConnectionManager.get_instance() 获取实例")
def __init__(self):
# 存储所有活跃连接key=client_idint),value=ConnectionContext实例
self.connections: Dict[int, ConnectionContext] = {}
# 异步锁:确保多连接并发操作时的数据安全
self._lock = asyncio.Lock()
# 存储所有上下文key=client_id当前活跃连接的唯一标识)
self.active_contexts: Dict[str, ConnectionContext] = {}
# 存储待清理的上下文:key=user_id(用户唯一标识,用于重连匹配)
self.pending_clean_contexts: Dict[str, ConnectionContext] = {}
self._internal_lock = asyncio.Lock() # 操作锁
async def create_connection(self, client_id: int, user_id: Optional[str] = None) -> ConnectionContext:
"""创建新的连接上下文(线程安全)"""
async with self._lock:
# 避免重复创建(同一client_id不会重复连接)
@classmethod
async def get_instance(cls) -> "ConnectionManager":
"""获取全局唯一实例(异步安全)"""
if cls._instance is None:
async with cls._lock:
if cls._instance is None: # 双重检查锁定
cls._instance = super().__new__(cls)
cls._instance.__init__()
logger.info("ConnectionManager 全局单例初始化成功")
return cls._instance
async def create_connection(self, client_id: str) -> ConnectionContext:
"""创建连接上下文(异步安全)"""
async with self._internal_lock:
if client_id in self.connections:
logger.warning(f"连接已存在:client_id={client_id},将覆盖旧连接")
self.connections[client_id].close()
# 创建新的连接上下文(包含独立LLM Session和消息队列)
context = ConnectionContext(client_id=client_id, user_id=user_id)
context = ConnectionContext(client_id=client_id)
self.connections[client_id] = context
logger.info(
f"创建新连接上下文:client_id={client_id}user_id={user_id},当前活跃连接数={len(self.connections)}")
logger.info(f"创建连接上下文:client_id={client_id},活跃连接数={len(self.connections)}")
return context
async def get_connection(self, client_id: int) -> Optional[ConnectionContext]:
"""获取指定client_id的连接上下文(线程安全)"""
async with self._lock:
async def get_connection(self, client_id: str) -> Optional[ConnectionContext]:
"""获取连接上下文(异步安全)"""
async with self._internal_lock:
context = self.connections.get(client_id)
if context and not context.is_active:
# 清理已断开的连接
del self.connections[client_id]
return None
return context
async def remove_connection(self, client_id: int):
"""除连接上下文(线程安全)"""
async with self._lock:
async def remove_connection(self, client_id: str):
"""除连接上下文(异步安全)"""
async with self._internal_lock:
context = self.connections.pop(client_id, None)
if context:
context.close()
logger.info(f"移除连接上下文:client_id={client_id}当前活跃连接数={len(self.connections)}")
logger.info(f"移除连接上下文:client_id={client_id},活跃连接数={len(self.connections)}")
async def get_active_connections_count(self) -> int:
"""获取当前活跃连接数(线程安全)"""
async with self._lock:
# 过滤已断开的连接
"""获取活跃连接数(异步安全)"""
async with self._internal_lock:
self.connections = {k: v for k, v in self.connections.items() if v.is_active}
return len(self.connections)
# 检查是否在可重连时间窗口内(30分钟)
# 标记为断开连接(不立即清理)
def is_reconnectable(self, timeout: int = 30) -> bool:
if self.is_active or not self.disconnect_time:
return False # 活跃连接或未记录断开时间,不可重连
# 计算断开时间是否在 timeout 分钟内
return datetime.utcnow() - self.disconnect_time <= timedelta(minutes=timeout)
async def create_or_reconnect_context(self, new_client_id: str, user_id: Optional[str] = None) -> ConnectionContext:
"""
创建新上下文或重连复用旧上下文
:param new_client_id: 新 WebSocket 连接的 client_id
:param user_id: 用户唯一标识(用于匹配旧上下文)
:return: 新上下文或复用的旧上下文
"""
async with self._internal_lock:
# 1. 如果用户已登录(有 user_id),先尝试重连复用
if user_id and user_id in self.pending_clean_contexts:
old_context = self.pending_clean_contexts.pop(user_id)
if old_context.is_reconnectable():
# 重连激活旧上下文,更新 client_id
old_context.reconnect(new_client_id=new_client_id)
# 加入活跃上下文列表
self.active_contexts[new_client_id] = old_context
return old_context
else:
# 旧上下文已超时,清理并创建新的
old_context.close()
# 2. 无旧上下文可复用,创建新上下文
new_context = ConnectionContext(client_id=new_client_id)
if user_id:
new_context.user_id = user_id # 绑定用户标识(如果已提供)
self.active_contexts[new_client_id] = new_context
logger.info(f"创建新上下文:client_id={new_client_id}user_id={user_id}")
return new_context
+64 -19
View File
@@ -1,26 +1,71 @@
from abc import ABC, abstractmethod
from typing import Optional, Coroutine
from typing import Optional, Dict, Callable, Awaitable
from audio_ai_chat.config.settings import settings
class LLMBase(ABC):
"""LLM服务统一抽象接口"""
# 定义回调函数类型(异步函数,与 ASR 回调风格一致)
# ChatStreamCallback = Callable[[str, bool, ConnectionContext], Awaitable[None]]
"""
Chat 流式回调函数类型:
- 第一个参数:流式文本块
- 第二个参数:是否结束标记
- 第三个参数:上下文对象
"""
class ChatBase(ABC):
"""Chat 服务统一抽象接口(与 ASRBase 接口风格完全对齐)"""
def __init__(self):
self.timeout = settings.LLM_TIMEOUT
self.retry_times = settings.LLM_RETRY_TIMES
self.model = settings.LLM_MODEL # 模型版本(不同LLM可能支持不同模型)
# 公共配置(所有 Chat 实现共享)
self.timeout = settings.CHAT_TIMEOUT # 需在配置中添加 CHAT_TIMEOUT
self.retry_times = settings.CHAT_RETRY_TIMES # 需在配置中添加 CHAT_RETRY_TIMES
self.base_url = settings.CHAT_BASE_URL # 配置中添加:http://10.10.10.202/v1
self.api_key = settings.CHAT_API_KEY # 配置中添加 API-Key
@abstractmethod
async def chat(
async def initialize(self) -> bool:
"""初始化 Chat 服务(如连接池、全局配置)"""
pass
@abstractmethod
async def get_connection(self) -> Optional[object]:
"""获取 Chat 连接对象(与 ASR 的 get_connection 对应)"""
pass
@abstractmethod
async def send_message(
self,
text: str,
user_id: Optional[str] = None,
history: Optional[list] = None, # 对话历史(部分LLM支持)
**kwargs
) -> str:
"""
大模型对话核心方法
:param text: 用户输入文本(ASR识别结果)
:param user_id: 用户ID(可选)
:param history: 对话历史(可选,格式:[(用户输入, 模型回答), ...])
:return: 模型生成的回答文本
"""
conn: object,
query: str,
context: ConnectionContext,
inputs: Optional[Dict] = None
) -> bool:
"""发送聊天消息(类似 ASR 的 push_audio"""
pass
@abstractmethod
async def start_communication(
self,
conn: object,
callback: ChatStreamCallback,
query: str,
context: ConnectionContext,
inputs: Optional[Dict] = None
) -> None:
"""启动 Chat 通信(流式接收结果,与 ASR 的 start_communication 对应)"""
pass
@abstractmethod
async def release_connection(self, conn: object) -> None:
"""释放 Chat 连接(与 ASR 的 release_connection 对应)"""
pass
@abstractmethod
async def close(self) -> None:
"""关闭 Chat 服务(释放所有连接,与 ASR 的 close 对应)"""
pass
@abstractmethod
async def get_valid_connection_count(self) -> int:
"""获取有效连接数(与 ASR 的接口完全一致)"""
pass
@@ -0,0 +1 @@
# from .dify import DifyLLMClient
@@ -0,0 +1,175 @@
import asyncio
import json
from typing import Optional, Dict, Callable, Awaitable
from dataclasses import dataclass, field
import aiohttp # 新增:异步HTTP库
# 大模型配置(集中管理)
LLM_CONFIG = {
"base_url": "http://10.10.10.202:8088/v1",
"api_key": "app-m7HZNV1aGiheh3wr6wNVHFxX",
"timeout": 30, # 请求超时时间(秒)
"default_scene": "通用聊天场景", # 默认场景描述
"stream_chunk_size": 1024 # 流式接收块大小
}
# 定义流式回调函数类型(异步)
LLMStreamCallback = Callable[[str, Optional[str], bool], Awaitable[None]]
"""
回调函数参数说明:
- chunk: 单次流式返回的文本片段
- conversation_id: 会话ID(首次返回,后续复用)
- is_finished: 是否结束(True=流式结束/同步返回完成)
"""
@dataclass
class LLMConversation:
"""会话对象(管理会话ID和上下文)"""
conversation_id: Optional[str] = None
user_id: str = ""
scene_description: str = LLM_CONFIG["default_scene"]
# 可选:存储会话历史(如需上下文管理)
history: list = field(default_factory=list)
class LLMClient:
"""大模型客户端封装(异步修复版)"""
def __init__(self):
self.base_url = LLM_CONFIG["base_url"]
self.headers = {
"Authorization": f"Bearer {LLM_CONFIG['api_key']}",
"Content-Type": "application/json"
}
self.timeout = aiohttp.ClientTimeout(total=LLM_CONFIG["timeout"]) # 异步超时
self._session: Optional[aiohttp.ClientSession] = None # 异步会话(复用连接)
async def _get_session(self) -> aiohttp.ClientSession:
"""获取/复用异步HTTP会话"""
if self._session is None or self._session.closed:
self._session = aiohttp.ClientSession(timeout=self.timeout)
return self._session
async def send_message(
self,
query: str,
conversation: LLMConversation,
stream_callback: Optional[LLMStreamCallback] = None,
response_mode: str = "streaming"
) -> tuple[Optional[str], str]:
"""
发送消息到大模型(纯异步版,无线程池阻塞)
"""
payload = {
"query": query,
"inputs": {"scene_description": conversation.scene_description},
"response_mode": response_mode,
"user": conversation.user_id
}
if conversation.conversation_id:
payload["conversation_id"] = conversation.conversation_id
url = f"{self.base_url}/chat-messages"
full_response = ""
res_conversation_id = conversation.conversation_id
try:
session = await self._get_session()
if response_mode == "streaming":
# 异步流式请求(无线程池,纯异步IO)
async with session.post(url, headers=self.headers, json=payload) as response:
response.raise_for_status()
# 实时迭代流式响应
async for line in response.content.iter_chunked(LLM_CONFIG["stream_chunk_size"]):
if not line:
continue
line_data = line.decode("utf-8")
if line_data.startswith("data: "):
json_str = line_data[6:].strip()
if json_str == "[DONE]":
if stream_callback:
await stream_callback("", res_conversation_id, True)
break
try:
data = json.loads(json_str)
# 更新会话ID
if not res_conversation_id and "conversation_id" in data:
res_conversation_id = data["conversation_id"]
# 提取内容
chunk = data.get("content", data.get("answer", data.get("message", "")))
# print('大模型返回的', chunk)
if chunk:
full_response += chunk
if stream_callback:
await stream_callback(chunk, res_conversation_id, False)
await asyncio.sleep(0) # 让出调度权
except json.JSONDecodeError as e:
# print(f"大模型解析流式数据失败: {e}")
continue
else:
# 异步非流式请求
async with session.post(url, headers=self.headers, json=payload) as response:
response.raise_for_status()
data = await response.json()
res_conversation_id = data.get("conversation_id", conversation.conversation_id)
full_response = data.get("content", data.get("answer", data.get("message", "")))
if stream_callback:
await stream_callback(full_response, res_conversation_id, True)
conversation.conversation_id = res_conversation_id
return res_conversation_id, full_response
except aiohttp.ClientError as e:
error_msg = f"大模型请求失败: {str(e)}"
print(error_msg)
if stream_callback:
await stream_callback(f"[错误] {error_msg}", res_conversation_id, True)
return res_conversation_id, ""
except Exception as e:
error_msg = f"大模型处理异常: {str(e)}"
print(error_msg)
if stream_callback:
await stream_callback(f"[错误] {error_msg}", res_conversation_id, True)
return res_conversation_id, ""
async def close(self):
"""关闭异步会话(程序退出时调用)"""
if self._session and not self._session.closed:
await self._session.close()
# 全局单例客户端(异步版)
llm_client = LLMClient()
# 快捷调用函数(保持原有接口不变)
async def call_llm(
query: str,
user_id: str,
scene_description: str = LLM_CONFIG["default_scene"],
conversation_id: Optional[str] = None,
stream_callback: Optional[LLMStreamCallback] = None,
response_mode: str = "streaming"
) -> tuple[Optional[str], str]:
"""
快捷调用大模型(无需手动创建会话对象)
:param query: 用户提问
:param user_id: 用户ID
:param scene_description: 场景描述
:param conversation_id: 会话ID(续聊用)
:param stream_callback: 流式回调
:param response_mode: 响应模式
:return: (conversation_id, 完整回复)
"""
conversation = LLMConversation(
conversation_id=conversation_id,
user_id=user_id,
scene_description=scene_description
)
return await llm_client.send_message(
query=query,
conversation=conversation,
stream_callback=stream_callback,
response_mode=response_mode
)
# 可选:程序退出时关闭会话(如FastAPI的shutdown事件)
async def shutdown_llm_client():
await llm_client.close()
+22 -19
View File
@@ -1,22 +1,25 @@
from typing import Type
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from .base import LLMBase
#
# from .openai_llm import OpenAILLM
# from .local_llm import LocalLLM
#
# LLM_REGISTRY: dict[str, Type[LLMBase]] = {
# "openai": OpenAILLM,
# "local": LocalLLM,
# }
#
# class LLMFactory:
# @staticmethod
# def get_llm_client() -> LLMBase:
# current_version = settings.LLM_CURRENT_VERSION
# if current_version not in LLM_REGISTRY:
# raise ServiceCallError(
# f"不支持的LLM版本:{current_version},可选版本:{list(LLM_REGISTRY.keys())}"
# )
# return LLM_REGISTRY[current_version]()
from .base import ChatBase
from .dify import DifyLLMClient # 具体实现类(对应 ASR 的 FunASR)
# 注册所有 Chat 实现:key=配置中的版本名,value=对应的类
CHAT_REGISTRY: dict[str, Type[ChatBase]] = {
"DefaultChat": DifyLLMClient,
# 新增 Chat 实现时,只需在这里注册
}
class ChatFactory:
"""Chat 服务工厂类(与 ASRFactory 逻辑完全一致)"""
@staticmethod
def get_chat_client() -> ChatBase:
# 从配置中获取当前指定的 Chat 版本
current_version = settings.CHAT_CURRENT_VERSION # 配置中添加该字段
if current_version not in CHAT_REGISTRY:
raise ServiceCallError(
f"不支持的 Chat 版本:{current_version},可选版本:{list(CHAT_REGISTRY.keys())}"
)
# 创建并返回对应版本的实例
return CHAT_REGISTRY[current_version]()
@@ -0,0 +1,75 @@
from typing import Optional, Tuple, Dict
from .base import ChatBase
from .factory import ChatFactory
from audio_ai_chat.config.settings import settings
# # 上下文管理类(集成到管理器中,与业务逻辑解耦)
# @dataclass
# class ConnectionContext:
# """Chat 连接上下文管理类(与你的原有 Context 兼容)"""
# user_id: str
# conversation_id: Optional[str] = None
# current_context: Dict = None
# # 可添加更多业务字段(如请求ID、会话状态等)
#
# def __post_init__(self):
# if self.current_context is None:
# self.current_context = {}
#
# def update_conversation_id(self, conversation_id: str):
# """更新会话ID(历史对话用)"""
# self.conversation_id = conversation_id
# self.current_context['conversation_id'] = conversation_id
#
# def add_context_data(self, key: str, value):
# """添加上下文数据"""
# self.current_context[key] = value
#
# def get_context_data(self, key: str, default=None):
# """获取上下文数据"""
# return self.current_context.get(key, default)
class LLMManager:
"""Chat 管理器(与 ASRManager 结构、接口完全一致)"""
_instance: Optional[ChatBase] = None # 单例存储(对应 ASRManager._instance
@classmethod
async def initialize(cls) -> Tuple[bool, str]:
"""初始化 Chat 服务(与 ASRManager.initialize 接口一致)"""
try:
# 1. 通过工厂创建 Chat 实例
cls._instance = ChatFactory.get_chat_client()
# 2. 初始化 Chat 服务(如连接池)
init_success = await cls._instance.initialize()
if not init_success:
return False, "Chat 服务初始化失败"
# 3. 检查有效连接数
valid_conn_count = await cls._instance.get_valid_connection_count()
if valid_conn_count == 0:
return False, f"Chat 有效连接数为 0(配置池大小:{settings.CHAT_POOL_SIZE}"
return True, f"Chat 初始化成功:有效连接数 {valid_conn_count}"
except Exception as e:
return False, f"Chat 初始化失败:{str(e)}"
@classmethod
def get_instance(cls) -> Optional[ChatBase]:
"""获取全局 Chat 实例(业务代码调用,与 ASR 用法一致)"""
return cls._instance
@classmethod
async def close(cls):
"""关闭 Chat 服务(FastAPI 关闭时调用,与 ASR 一致)"""
if cls._instance:
await cls._instance.close()
cls._instance = None
print("Chat 管理器:实例和连接池已关闭")
@classmethod
def is_healthy(cls) -> bool:
"""检查 Chat 服务健康状态(与 ASR 一致)"""
return cls._instance is not None
@@ -0,0 +1,692 @@
import asyncio
import json
import websockets
import numpy as np
import sounddevice as sd
from typing import Optional, Callable, Dict, Any, List
from dataclasses import dataclass, field
import uuid
import copy
from protocols import (
EventType,
MsgType,
finish_connection,
finish_session,
receive_message,
start_connection,
start_session,
task_request,
wait_for_event,
)
# ------------------------------
# 配置常量(可根据需求调整)
# ------------------------------
DEFAULT_APPID = "7069844318"
DEFAULT_ACCESS_TOKEN = "osFMEJr20SSTWRql43cJlZkAOg7iwvxu"
DEFAULT_ENDPOINT = "wss://openspeech.bytedance.com/api/v3/tts/bidirection"
DEFAULT_VOICE_TYPE = "zh_female_gaolengyujie_emo_v2_mars_bigtts"
DEFAULT_ENCODING = "pcm"
DEFAULT_SAMPLE_RATE = 16000
@dataclass
class TTSRequest:
"""TTS请求对象(带唯一标识)"""
tts_text: str
voice_type: str = DEFAULT_VOICE_TYPE
encoding: str = DEFAULT_ENCODING
speed: float = 1.0 # 语速(字节跳动TTS支持,需服务端兼容)
stream: bool = True # 是否流式合成
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
session_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 会话ID(每个请求一个会话)
class ByteDanceTTSSocketClient:
"""字节跳动 TTS WebSocket 客户端(异步/流式/带任务队列)"""
def __init__(
self,
appid: str = DEFAULT_APPID,
access_token: str = DEFAULT_ACCESS_TOKEN,
endpoint: str = DEFAULT_ENDPOINT,
max_queue_size: int = 100
):
"""
初始化客户端
:param appid: 字节跳动APP ID
:param access_token: 访问令牌
:param endpoint: WebSocket 服务端地址
:param max_queue_size: 最大队列长度(防止内存溢出)
"""
# 基础配置
self.appid = appid
self.access_token = access_token
self.endpoint = endpoint
self.max_queue_size = max_queue_size
# WebSocket 连接状态
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
self.is_connected = False
self.is_processing = False # 是否正在处理请求
self.logid: Optional[str] = None # 服务端返回的日志ID
# 异步任务队列(FIFO
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
# 回调函数定义(所有回调都带request_id,方便关联请求)
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
self.on_audio_chunk: Callable[[str, bytes], None] = lambda req_id, chunk: None # 音频块回调(原始字节)
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
# 音频播放相关(支持MP3格式直接播放)
self.play_stream: Optional[sd.OutputStream] = None
self.current_req_id: Optional[str] = None
self.enable_playback: bool = True # 是否启用实时播放
# 外部回调函数(返回完整结果)
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
# 存储每个请求的完整音频数据(原始字节)
self.audio_buffers: Dict[str, List[bytes]] = {}
def _get_resource_id(self, voice_type: str) -> str:
"""根据音色类型获取资源ID(字节跳动TTS协议要求)"""
if voice_type.startswith("S_"):
return "volc.megatts.default"
return "volc.service_type.10029"
async def _create_websocket_connection(self):
"""创建WebSocket连接(内部使用)"""
headers = {
"X-Api-App-Key": self.appid,
"X-Api-Access-Key": self.access_token,
"X-Api-Resource-Id": self._get_resource_id(DEFAULT_VOICE_TYPE), # 用默认音色获取资源ID
"X-Api-Connect-Id": str(uuid.uuid4()),
}
print(f"连接到 TTS 服务端: {self.endpoint}")
self.websocket = await websockets.connect(
self.endpoint,
additional_headers=headers,
max_size=10 * 1024 * 1024 # 10MB缓冲区
)
self.is_connected = True
self.logid = self.websocket.response.headers.get("x-tt-logid")
print(f"连接成功,LogID: {self.logid}")
# 发送连接启动指令
await start_connection(self.websocket)
await wait_for_event(
self.websocket, MsgType.FullServerResponse, EventType.ConnectionStarted
)
print("TTS连接已初始化完成")
async def connect(self):
"""建立WebSocket连接(外部调用,初始化一次)"""
if not self.is_connected:
try:
await self._create_websocket_connection()
# 启动队列消费协程(后台运行)
asyncio.create_task(self._consume_queue())
except Exception as e:
error_msg = f"连接失败: {str(e)}"
print(error_msg)
raise ConnectionError(error_msg)
async def disconnect(self):
"""关闭WebSocket连接"""
if self.is_connected and self.websocket:
try:
# 发送连接结束指令
await finish_connection(self.websocket)
await wait_for_event(
self.websocket, MsgType.FullServerResponse, EventType.ConnectionFinished
)
except Exception as e:
print(f"关闭连接时异常: {str(e)}")
finally:
await self.websocket.close()
self.is_connected = False
self.websocket = None
print("已断开与TTS服务端的连接")
# 清理音频播放流
if self.play_stream:
self.play_stream.stop()
self.play_stream.close()
self.play_stream = None
def set_external_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
"""设置外部回调函数,用于返回完整结果"""
self.external_callback = callback
def set_playback_enabled(self, enabled: bool):
"""设置是否启用音频实时播放"""
self.enable_playback = enabled
print(f"音频实时播放已{'启用' if enabled else '禁用'}")
async def synthesize(self, tts_text: str, **kwargs) -> str:
"""
异步非阻塞添加TTS请求到队列
:param tts_text: 要合成的文本
:param kwargs: 其他TTS参数(voice_type, encoding, speed等)
:return: 唯一请求ID
"""
# 创建请求对象(支持覆盖默认参数)
request = TTSRequest(tts_text=tts_text, **kwargs)
req_id = request.request_id
# 初始化音频缓冲区
self.audio_buffers[req_id] = []
# 异步入队(非阻塞)
try:
await self.request_queue.put(request)
self.on_task_enqueue(req_id)
print(f"请求 [{req_id[:8]}] 已加入队列,当前队列长度: {self.request_queue.qsize()}")
return req_id
except asyncio.QueueFull:
self.on_queue_full(req_id)
error_msg = f"队列已满(最大长度{self.max_queue_size}),请求 [{req_id[:8]}] 入队失败"
print(error_msg)
raise Exception(error_msg)
async def _consume_queue(self):
"""消费队列(后台协程,自动处理排队请求)"""
print("队列消费协程已启动")
while True:
try:
# 等待队列中有请求(阻塞,直到有任务)
request = await self.request_queue.get()
req_id = request.request_id
# 标记为处理中
self.is_processing = True
# print(f"\n开始处理请求 [{req_id[:8]}],剩余队列长度: {self.request_queue.qsize()}")
# 处理单个请求
await self._process_single_request(request)
# 标记任务完成(让Queue知道可以继续)
self.request_queue.task_done()
self.is_processing = False
except Exception as e:
error_msg = f"队列消费异常: {str(e)}"
print(error_msg)
self.is_processing = False
# 短暂等待,避免死循环占用CPU
await asyncio.sleep(0.1)
def _build_base_request(self, request: TTSRequest) -> Dict[str, Any]:
"""构建字节跳动TTS基础请求参数"""
aaa = {
"user": {"uid": str(uuid.uuid4())},
"namespace": "BidirectionalTTS",
"req_params": {
"speaker": request.voice_type,
"audio_params": {
"format": request.encoding,
"sample_rate": DEFAULT_SAMPLE_RATE,
"enable_timestamp": True,
},
"additions": json.dumps({"disable_markdown_filter": False}),
"speed": request.speed, # 语速参数(需服务端支持)
},
}
print('aaa', aaa)
return aaa
async def _send_text_stream(self, request: TTSRequest, session_id: str):
"""流式发送文本(逐字符发送,字节跳动TTS流式协议要求)"""
base_request = self._build_base_request(request)
text = request.tts_text.strip()
if not text:
print(f"请求 [{request.request_id[:8]}] 文本为空,跳过发送")
return
# 逐字符发送(控制发送速率,避免拥塞)
for char in text:
if not self.is_connected or not self.websocket:
raise ConnectionError("连接已断开,无法继续发送文本")
# 构建单个字符的任务请求
task_req = copy.deepcopy(base_request)
task_req["event"] = EventType.TaskRequest
task_req["req_params"]["text"] = char
# 发送任务请求
await task_request(
self.websocket,
json.dumps(task_req).encode("utf-8"),
session_id
)
# 控制发送速率(5ms/字符,可调整)
await asyncio.sleep(0.005)
# 发送会话结束指令
await finish_session(self.websocket, session_id)
print(f"请求 [{request.request_id[:8]}] 文本发送完成")
async def _handle_audio_response(self, req_id: str, session_id: str, request: TTSRequest) -> Dict[str, Any]:
"""处理服务端的流式音频响应"""
if not self.websocket:
raise ConnectionError("WebSocket连接未建立")
sample_rate = DEFAULT_SAMPLE_RATE
audio_received = False
try:
while True:
# 接收服务端消息(异步阻塞)
msg = await receive_message(self.websocket)
if msg.type == MsgType.FullServerResponse:
# 完整响应(开始/结束/错误)
if msg.event == EventType.SessionStarted:
# 会话开始回调
start_data = {
"session_id": session_id,
"sample_rate": sample_rate,
"encoding": request.encoding,
"voice_type": request.voice_type,
"logid": self.logid
}
self.on_start(req_id, start_data)
print(f"请求 [{req_id[:8]}] 合成开始")
elif msg.event == EventType.SessionFinished:
# 会话结束,退出循环
end_data = {"session_id": session_id, "message": "合成完成"}
self.on_end(req_id, end_data)
print(f"请求 [{req_id[:8]}] 合成结束")
break
elif msg.type == MsgType.AudioOnlyServer:
# 流式音频数据(原始字节)
audio_chunk = msg.payload
if audio_chunk:
audio_received = True
# 保存到缓冲区
self.audio_buffers[req_id].append(audio_chunk)
# 音频块回调
self.on_audio_chunk(req_id, audio_chunk)
# 实时播放(如果启用)
await self._play_audio_chunk(req_id, audio_chunk, sample_rate)
else:
# 未知消息类型
raise RuntimeError(f"收到未知消息类型: {msg.type}, 内容: {msg}")
# 组装完整结果
full_audio = b"".join(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else b""
return {
"status": "completed",
"request_id": req_id,
"session_id": session_id,
"sample_rate": sample_rate,
"encoding": request.encoding,
"audio_data": full_audio, # 完整音频字节数据
"audio_length": len(full_audio),
"message": "合成成功" if audio_received else "合成完成但未收到音频数据"
}
except Exception as e:
error_msg = f"处理音频响应异常: {str(e)}"
self.on_error(req_id, error_msg)
return {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
}
async def _play_audio_chunk(self, req_id: str, chunk: bytes, sample_rate: int):
"""实时播放音频块(支持MP3格式)"""
if not self.enable_playback:
return
# 确保当前请求是正在播放的请求
if self.current_req_id is None:
self.current_req_id = req_id
if req_id != self.current_req_id:
# 切换请求时,重置播放流
if self.play_stream:
self.play_stream.stop()
self.play_stream.close()
self.current_req_id = req_id
try:
# 初始化播放流(如果未初始化)
if not self.play_stream:
self.play_stream = sd.OutputStream(
samplerate=sample_rate,
channels=1, # 单声道
dtype=np.float32
)
self.play_stream.start()
# MP3字节 → 音频数组(直接播放)
# 注意:sounddevice默认支持PCM格式,如果是MP3需要解码,这里简化处理(实际使用建议用pydub解码)
# 如需支持MP3播放,请安装 pydub: pip install pydub ffmpeg
try:
# 简化处理:假设服务端返回PCM(如果是MP3,需替换为解码逻辑)
audio_array = np.frombuffer(chunk, dtype=np.float32)
if audio_array.size > 0:
self.play_stream.write(audio_array)
except Exception as e:
print(f"音频播放异常: {str(e)},请确保音频格式正确")
except Exception as e:
print(f"播放流初始化失败: {str(e)}")
async def _process_single_request(self, request: TTSRequest):
"""处理单个TTS请求(完整流程:连接→启动会话→流式发送文本→接收音频→回调结果)"""
req_id = request.request_id
session_id = request.session_id
# 参数校验
if not request.tts_text.strip():
error_msg = "合成文本不能为空"
self.on_error(req_id, error_msg)
self._send_external_callback(req_id, {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
})
return
# 确保连接已建立(断开时自动重连)
if not self.is_connected:
print(f"请求 [{req_id[:8]}] 处理时连接已断开,尝试重连...")
try:
await self._create_websocket_connection()
except Exception as e:
error_msg = f"重连失败: {str(e)}"
self.on_error(req_id, error_msg)
self._send_external_callback(req_id, {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
})
return
try:
# 1. 启动会话
base_request = self._build_base_request(request)
start_session_req = copy.deepcopy(base_request)
start_session_req["event"] = EventType.StartSession
await start_session(
self.websocket,
json.dumps(start_session_req).encode("utf-8"),
session_id
)
# 等待会话启动成功
await wait_for_event(
self.websocket, MsgType.FullServerResponse, EventType.SessionStarted
)
# 2. 异步流式发送文本(后台任务,不阻塞接收音频)
send_task = asyncio.create_task(self._send_text_stream(request, session_id))
# 3. 接收并处理音频响应
result_data = await self._handle_audio_response(req_id, session_id, request)
# 4. 等待文本发送任务完成
await send_task
# 5. 发送外部回调
self._send_external_callback(req_id, result_data)
except Exception as e:
error_msg = f"处理请求 [{req_id[:8]}] 异常: {str(e)}"
self.on_error(req_id, error_msg)
self._send_external_callback(req_id, {
"status": "error",
"request_id": req_id,
"session_id": session_id,
"message": error_msg
})
finally:
# 清理缓冲区
if req_id in self.audio_buffers:
del self.audio_buffers[req_id]
def _send_external_callback(self, req_id: str, result_data: Dict[str, Any]):
"""发送外部回调(支持同步/异步回调函数)"""
if not self.external_callback:
return
try:
# 异步回调:直接await
if asyncio.iscoroutinefunction(self.external_callback):
asyncio.create_task(self.external_callback(req_id, result_data))
# 同步回调:在线程池中执行(避免阻塞事件循环)
else:
asyncio.get_event_loop().run_in_executor(
None, self.external_callback, req_id, result_data
)
except Exception as e:
print(f"外部回调执行异常: {str(e)}")
async def wait_all_completed(self):
"""等待队列中所有任务处理完成(阻塞)"""
await self.request_queue.join()
print("\n所有队列任务已处理完成")
# ------------------------------
# 使用示例(与你提供的风格完全一致)
# ------------------------------
class TTSManager:
"""TTS管理器 - 供外部代码调用(封装客户端,简化使用)"""
def __init__(
self,
appid: str = DEFAULT_APPID,
access_token: str = DEFAULT_ACCESS_TOKEN,
endpoint: str = DEFAULT_ENDPOINT
):
self.client = ByteDanceTTSSocketClient(
appid=appid,
access_token=access_token,
endpoint=endpoint
)
self._setup_internal_callbacks()
def _setup_internal_callbacks(self):
"""设置内部回调(日志/状态提示)"""
def on_task_enqueue(req_id: str):
"""任务入队回调"""
print(f"📥 任务 [{req_id[:8]}] 已入队")
def on_tts_start(req_id: str, data: Dict[str, Any]):
"""合成开始回调"""
print(f"🎤 合成开始 [{req_id[:8]}] - 采样率: {data['sample_rate']}, 编码: {data['encoding']}")
def on_audio_chunk(req_id: str, chunk: bytes):
"""音频块回调(内部仅打印日志,外部通过external_callback获取)"""
print(f"🔊 收到音频块 [{req_id[:8]}] - 大小: {len(chunk)}字节", end="\r")
def on_tts_end(req_id: str, data: Dict[str, Any]):
"""合成结束回调"""
print(f"\n🏁 合成结束 [{req_id[:8]}] - 会话ID: {data['session_id']}")
def on_tts_error(req_id: str, msg: str):
"""错误回调"""
print(f"\n❌ 合成失败 [{req_id[:8]}] - 错误: {msg}")
def on_queue_full(req_id: str):
"""队列满回调"""
print(f"⚠️ 队列已满,请求 [{req_id[:8]}] 入队失败")
# 绑定内部回调
self.client.on_task_enqueue = on_task_enqueue
self.client.on_start = on_tts_start
self.client.on_audio_chunk = on_audio_chunk
self.client.on_end = on_tts_end
self.client.on_error = on_tts_error
self.client.on_queue_full = on_queue_full
async def initialize(self):
"""初始化连接"""
await self.client.connect()
async def shutdown(self):
"""关闭连接"""
await self.client.disconnect()
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
"""设置外部结果回调(获取完整音频数据)"""
self.client.set_external_callback(callback)
def set_playback_enabled(self, enabled: bool):
"""设置是否启用实时播放"""
self.client.set_playback_enabled(enabled)
async def synthesize(self, text: str, **kwargs) -> str:
"""
异步非阻塞合成文本
:param text: 要合成的文本
:param kwargs: 其他参数(voice_type, encoding, speed等)
:return: 请求ID
"""
return await self.client.synthesize(text, **kwargs)
async def wait_all_completed(self):
"""等待所有任务完成"""
await self.client.wait_all_completed()
# ------------------------------
# 外部调用示例
# ------------------------------
async def external_usage_example():
"""外部代码使用示例"""
# 1. 创建TTS管理器(可替换为自己的appid和access_token
tts_manager = TTSManager(
appid=DEFAULT_APPID,
access_token=DEFAULT_ACCESS_TOKEN,
endpoint=DEFAULT_ENDPOINT
)
# 2. 设置外部结果回调(获取完整音频数据)
def handle_tts_result(req_id: str, result: Dict[str, Any]):
"""处理TTS完整结果(同步回调)"""
status = result.get("status")
if status == "completed":
audio_data = result.get("audio_data")
encoding = result.get("encoding")
audio_length = result.get("audio_length")
print(f"\n✅ 收到完整结果 [{req_id[:8]}] - 长度: {audio_length}字节, 编码: {encoding}")
# 保存音频文件
filename = f"tts_output_{req_id[:8]}.{encoding}"
with open(filename, "wb") as f:
f.write(audio_data)
print(f"💾 音频文件已保存: {filename}")
elif status == "error":
error_msg = result.get("message")
print(f"\n❌ 请求 [{req_id[:8]}] 处理失败: {error_msg}")
# 绑定外部回调
tts_manager.set_result_callback(handle_tts_result)
# 3. 设置是否启用实时播放(默认True)
tts_manager.set_playback_enabled(True)
# 4. 初始化连接
await tts_manager.initialize()
# 5. 异步提交多个TTS请求(非阻塞)
texts = [
"你好,这是字节跳动TTS的流式合成测试。",
"我支持异步非阻塞调用,多个请求可以排队处理。",
"每个请求都会返回唯一的ID,方便你跟踪结果。",
"音频数据会通过回调函数返回,支持实时播放和保存文件。",
"最后一个测试句子,演示队列的自动消费功能。"
]
req_ids = []
for i, text in enumerate(texts):
# 提交请求(非阻塞,立即返回)
req_id = await tts_manager.synthesize(
text,
voice_type=DEFAULT_VOICE_TYPE,
encoding=DEFAULT_ENCODING,
speed=1.0
)
req_ids.append(req_id)
print(f"📤 已提交请求 {i+1}: ID={req_id[:8]}")
# 模拟其他业务逻辑(无需等待TTS完成)
await asyncio.sleep(0.3)
# 6. 等待所有TTS任务完成(可选,根据业务需求决定是否等待)
await tts_manager.wait_all_completed()
# 7. 关闭连接(程序退出前调用)
await tts_manager.shutdown()
# ------------------------------
# 异步结果回调示例(高级用法)
# ------------------------------
async def async_result_callback(req_id: str, result: Dict[str, Any]):
"""异步结果回调(支持异步操作,如上传音频到服务器)"""
if result["status"] == "completed":
print(f"\n⚡ 异步处理结果 [{req_id[:8]}] - 开始上传音频...")
# 模拟异步上传操作
await asyncio.sleep(0.5)
print(f"⚡ 异步处理结果 [{req_id[:8]}] - 音频上传完成")
async def advanced_usage_example():
"""高级使用示例:异步回调 + 禁用播放 + 批量请求"""
tts_manager = TTSManager()
# 设置异步结果回调
tts_manager.set_result_callback(async_result_callback)
# 禁用实时播放(只获取音频数据)
tts_manager.set_playback_enabled(False)
await tts_manager.initialize()
# 批量提交请求(并行提交)
tasks = []
for i in range(3):
text = f"这是第{i+1}个高级测试文本,使用异步回调处理结果。"
task = tts_manager.synthesize(text, speed=0.9)
tasks.append(task)
# 并行提交所有请求
req_ids = await asyncio.gather(*tasks)
print(f"\n已并行提交 {len(req_ids)} 个请求")
# 等待所有任务完成
await tts_manager.wait_all_completed()
await tts_manager.shutdown()
if __name__ == "__main__":
try:
# 运行基础使用示例
asyncio.run(external_usage_example())
# 运行高级使用示例(取消注释)
# asyncio.run(advanced_usage_example())
except KeyboardInterrupt:
print("\n程序被用户中断")
except Exception as e:
print(f"程序异常: {str(e)}")
@@ -1,113 +1,91 @@
from fastapi import WebSocket
from typing import Dict, List, Optional
from pyexpat.errors import messages
from audio_ai_chat.config.logger import logger
# from audio_ai_chat.core.asr.factory import ASRFactory # 导入ASR工厂
# from audio_ai_chat.core.llm.factory import LLMFactory # 导入LLM工厂
# from audio_ai_chat.core.tts.factory import TTSFactory # 导入TTS工厂
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from fastapi import WebSocket, WebSocketDisconnect
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec
from typing import Optional, Dict, Callable, Awaitable, List, Any, Coroutine
from typing import Dict, List, Optional, Callable,Any
from dataclasses import dataclass, field
import sys
import json
import asyncio
import websockets
import uuid
import logging
from audio_ai_chat.core.connection import ConnectionContext
from audio_ai_chat.config.logger import logger
from audio_ai_chat.core.asr.asr_manager import ASRManager
from audio_ai_chat.core.connection import ConnectionManager, ConnectionContext
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec, MessageType
from audio_ai_chat.utils.exceptions import ServiceCallError
from audio_ai_chat.core.llm.dify.dify import LLMConversation, llm_client
from audio_ai_chat.core.tts.tts_client import TTSManager
from functools import partial
# 全局WebSocket连接管理器(单例模式,确保全局统一)
class WebSocketConnectionManager:
"""WebSocket连接管理器"""
_instance: Optional["WebSocketConnectionManager"] = None
def __init__(self):
# 活跃连接列表
self.active_connections: List[WebSocket] = []
# 关键映射:client_id -> ConnectionContext(快速获取用户专属上下文)
self.client_context_map: Dict[str, ConnectionContext] = {}
# self.asr_client = ASRFactory.get_asr_client()
# self.tts_client = TTSFactory.get_tts_client()
# 用户LLM会话存储
# self.user_llm_conversations: Dict[str, LLMConversation] = {}
# 全局唤醒事件
self.consume_wakeup = asyncio.Event()
self.connection_manager = None
async def connect(self, client_id, websocket: WebSocket) -> ConnectionContext:
"""
建立连接+身份校验(前端主动发送身份信息)
超时逻辑:5秒内未收到前端身份信息,自动关闭连接
返回:校验通过的 ConnectionContext(保证非空)
"""
# 1. 接受连接并加入活跃列表
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance.active_connections: List[WebSocket] = []
cls._instance.client_context_map: Dict[str, ConnectionContext] = {}
cls._instance.connection_manager: Optional[ConnectionManager] = None
cls._instance.consume_wakeup = asyncio.Event()
return cls._instance
async def initialize(self):
"""初始化:获取ConnectionManager全局单例"""
self.connection_manager = await ConnectionManager.get_instance()
logger.info("WebSocketConnectionManager 初始化成功")
async def connect(self, client_id: str, websocket: WebSocket) -> ConnectionContext:
"""建立连接+身份校验"""
await websocket.accept()
context = ConnectionContext(client_id=client_id) # 提前创建上下文(保证最终返回非空)
self.active_connections.append(websocket)
logger.info(
f"连接 {client_id} 已接受,等待前端发送身份信息(5秒超时)...,当前连接数: {len(self.active_connections)}")
f"连接 {client_id} 已接受,等待身份信息(5秒超时),当前连接数: {len(self.active_connections)}"
)
# 2. 超时控制:5秒内未收到身份信息 -> 关闭连接
# 超时接收身份包
try:
ping_packet = await asyncio.wait_for(
websocket.receive_bytes(),
timeout=5.0
)
ping_packet = await asyncio.wait_for(websocket.receive_bytes(), timeout=5.0)
except asyncio.TimeoutError:
error_msg = f"连接 {client_id} 身份校验超时5秒未收到消息)"
error_msg = f"连接 {client_id} 身份校验超时"
logger.warning(error_msg)
# 发送超时错误响应(二进制格式)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 1008, "message": "身份校验超时,请重试"}
MessageType.ERROR, {"code": 1008, "message": "身份校验超时,请重试"}
)
await websocket.send_bytes(error_packet)
raise TimeoutError(error_msg) # 抛出异常,进入后续清理逻辑
raise TimeoutError(error_msg)
# 3. 解包并验证包类型
print('ping_packet', ping_packet)
# 解包并验证包类型
msg_type, _, identity_data = ProtocolCodec.unpack(ping_packet)
if msg_type != MessageType.IDENTITY:
error_msg = f"连接 {client_id} 首个包类型错误(期望{MessageType.IDENTITY.value},实际{msg_type.value}"
error_msg = f"连接 {client_id} 首个包类型错误"
logger.error(error_msg)
# 发送类型错误响应
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
MessageType.ERROR, {"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# 4. 身份信息并校验
# todo
# 提取核心字段(必选字段校验)
# 校验身份信息
user_id = identity_data.get("user_id")
token = identity_data.get("token")
name = identity_data.get("name") or f"用户{user_id}" # 提供默认名称
name = identity_data.get("name") or f"用户{user_id}"
if not all([user_id, token]):
error_msg = f"连接 {client_id} 身份信息不完整(缺少user_id或token"
error_msg = f"连接 {client_id} 身份信息不完整"
logger.error(error_msg)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4003, "message": "身份信息不完整:必须包含user_id和token"}
MessageType.ERROR, {"code": 4003, "message": "身份信息不完整:必须包含user_id和token"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# TODO: 实际身份校验逻辑(根据你的业务扩展)
# 5. 校验通过:更新上下文并响应前端
# 创建/获取连接上下文
context = await self.connection_manager.create_or_reconnect_context(
new_client_id=client_id, user_id=user_id
)
context.set_user_info(token, user_id, name)
self.client_context_map[client_id] = context # 加入上下文映射
self.client_context_map[client_id] = context
# 发送成功响应
# 响应身份校验成功
success_packet = ProtocolCodec.pack(
MessageType.IDENTITY,
{
@@ -117,309 +95,225 @@ class WebSocketConnectionManager:
}
)
await websocket.send_bytes(success_packet)
logger.info(f"用户 {user_id}{name})身份校验通过,连接就绪client_id: {client_id}")
logger.info(f"用户 {user_id}{name})身份校验通过(client_id: {client_id}")
return context
# 初始化LLM会话
# self._init_llm_conversation(user_id)
# return conn_id, user_id, conn_id # conn_id 同时作为 tts_session_id
# def _init_llm_conversation(self, user_id: str):
# """初始化用户LLM会话"""
# if user_id not in self.user_llm_conversations:
# self.user_llm_conversations[user_id] = LLMConversation(
# user_id=user_id,
# scene_description="语音识别对话场景"
# )
def disconnect(self, websocket: WebSocket, conn_id: str):
"""断开连接并清理资源"""
if websocket in self.active_connections:
self.active_connections.remove(websocket)
logger.info(f"连接 {conn_id} 已断开,当前连接数: {len(self.active_connections)}")
# async def setup_tts_manager(self, result_queue: asyncio.Queue) -> TTSManager:
# """初始化TTS管理器"""
#
# def handle_tts_result(req_id: str, result: Dict[str, Any]):
# """TTS结果回调处理"""
# try:
# status = result.get("status")
# if status == "completed":
# audio_data = result.get("audio_data")
# if audio_data is not None and len(audio_data) > 0:
# # 转换为PCM格式
# pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
# pcm_bytes = pcm_data.tobytes()
# result_queue.put_nowait(pcm_bytes)
# except Exception as e:
# logger.error(f"TTS结果处理失败: {str(e)}")
#
# self.tts_client = TTSFactory.get_tts_client()
# tts_manager = TTSManager()
# tts_manager.set_result_callback(handle_tts_result)
# tts_manager.set_playback_enabled(False)
# await tts_manager.initialize()
# return tts_manager
async def asr_result_callback(self, result: dict, websocket: WebSocket,
user_id: str, result_queue: asyncio.Queue):
"""ASR结果回调处理"""
try:
logger.info(f"ASR识别结果: {result}")
final_asr_text = result.get("text", "").strip()
# 转发ASR结果到前端队列
if final_asr_text:
print(f"插入ASR结果时队列大小: {result_queue.qsize()}")
self.consume_wakeup.set() # 唤醒消费协程
# 异步调用大模型
llm_conversation = self.user_llm_conversations.get(user_id)
if llm_conversation:
asyncio.create_task(
self.call_llm_and_send(
query=final_asr_text,
conversation=llm_conversation,
websocket=websocket
)
def _create_asr_callback(self, context: ConnectionContext) -> Callable[[dict], None]:
"""
闭包:为当前连接创建专属的ASR回调函数
回调内部持有ConnectionContext引用,直接操作其消息队列
"""
async def asr_result_callback(result: dict):
"""专属回调:将ASR结果打包后插入当前连接的消息队列"""
try:
print('result', 'result', result)
# 处理错误结果
if result.get("error"):
logger.error(f"ASR错误(client_id: {context.client_id}):{result['error']}")
# 打包错误消息
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 5001, "message": f"ASR服务错误:{result['error']}"}
)
context.message_queue.put_nowait(error_packet)
else:
# 3. 调用Dify流式接口(传入ASR元数据,用于存储到对话历史)
print('context', context.user_id)
final_asr_text = result.get("text", "")
# 2. 调用大模型(异步)
if final_asr_text:
asyncio.create_task(
self.call_llm_and_send(
context=context,
query=final_asr_text,
conversation=context.llm_session
)
)
logger.debug(
f"ASR结果入队(client_id: {context.client_id}):"
f"文本={result['text']},最终结果={result['is_final']}"
)
except Exception as e:
logger.error(f"ASR回调处理失败(client_id: {context.client_id}):{str(e)}")
return asr_result_callback
# ====================== 调用大模型 ======================
# ====================== 大模型流式回调 ======================
@staticmethod
async def llm_stream_callback(context, chunk: str, conversation_id: str, is_finished: bool):
"""大模型流式回调(纯异步,无阻塞)"""
if not chunk:
return
print('大模型流式回调', chunk)
req_id = await context.tts_client.synthesize(chunk)
print('req_id', req_id)
async def call_llm_and_send(self,context ,query: str, conversation: LLMConversation):
"""调用大模型,流式结果转发前端 + TTS"""
if not query:
return
logger.info(f"调用大模型 - 用户(): {query}")
try:
stream_callback = partial(WebSocketConnectionManager.llm_stream_callback, context)
conv_id, full_reply = await llm_client.send_message(
query=query,
conversation=conversation,
stream_callback=stream_callback, # 传递绑定后的回调
response_mode="streaming"
)
logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
except Exception as e:
logger.error(f"ASR回调执行失败: {str(e)}")
logger.error(f"大模型调用失败: {str(e)}")
# async def llm_stream_callback(self, chunk: str, tts_manager: TTSManager):
# """大模型流式回调处理"""
# if not chunk:
# return
# try:
# # 提交TTS合成请求
# await tts_manager.synthesize(chunk)
# await asyncio.sleep(0) # 让出调度权
# except Exception as e:
# logger.error(f"LLM流式回调处理失败: {str(e)}")
# async def call_llm_and_send(self, query: str, conversation: LLMConversation, websocket: WebSocket):
# """调用大模型并处理结果"""
# logger.info(f"调用大模型 - 用户({conversation.user_id}): {query}")
# try:
# conv_id, full_reply = await llm_client.send_message(
# query=query,
# conversation=conversation,
# stream_callback=self.llm_stream_callback,
# response_mode="streaming"
# )
# logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
# except Exception as e:
# logger.error(f"大模型调用失败: {str(e)}")
# if not websocket.client_state.disconnected:
# await websocket.send_json({
# "type": "llm_error",
# "data": {"error": str(e)}
# })
async def recv_frontend_data(self, websocket: WebSocket, asr_conn):
"""接收前端音频数据并推送到ASR"""
while not asr_conn.stop_event.is_set():
try:
raw_bytes = await websocket.receive_bytes()
# success = await push_audio_data(asr_conn, raw_bytes)
# if not success:
# logger.warning("音频数据插入ASR失败(队列满/连接失效)")
except WebSocketDisconnect:
logger.info("前端主动断开连接")
asr_conn.stop_event.set()
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
# break
async def send_results(self, websocket: WebSocket, result_queue: asyncio.Queue, asr_conn):
"""从结果队列发送数据到前端"""
while True:
try:
# 等待队列数据或超时
result = await asyncio.wait_for(result_queue.get(), timeout=0.05)
# if not websocket.client_state.disconnected:
await websocket.send_bytes(result)
except asyncio.TimeoutError:
if asr_conn.stop_event.is_set():
break
continue
except Exception as e:
logger.error(f"发送结果到前端失败: {str(e)}")
asr_conn.stop_event.set()
break
async def handle_connection(self, websocket: WebSocket):
"""处理单个WebSocket连接的完整生命周期"""
client_id = str(id(websocket))
context = None
client_id = str(id(websocket)) # 生成唯一连接ID
logger.info(f"新WebSocket连接:client_id={client_id}")
context: Optional[ConnectionContext] = None
asr_conn = None
llm_conn = None # 新增:LLM连接变量
try:
# 1. 建立连接并获取上下文
context = await self.connect(client_id, websocket)
if not context:
logger.error(f"连接 {client_id} 上下文创建失败")
return
# 接收前端数据
# 2. 获取ASR连接和专属回调
asr_client = ASRManager.get_instance()
asr_conn = await asr_client.get_connection()
if not asr_conn:
raise ServiceCallError("获取ASR连接失败")
# 创建当前连接的专属ASR回调(闭包绑定context)
asr_callback = self._create_asr_callback(context)
# 3. 启动ASR通信任务(传入专属回调)
communication_task = asyncio.create_task(
asr_client.start_communication(asr_conn, asr_callback)
)
context.llm_session = LLMConversation(
user_id=context.user_id,
scene_description="语音识别对话场景"
)
context.tts_client = TTSManager()
def handle_tts_result(context, req_id: str, result: Dict[str, Any]):
"""处理TTS结果回调"""
status = result.get("status")
print('处理TTS结果回调')
if status == "completed":
audio_data = result.get("audio_data")
pack_data = ProtocolCodec.pack(MessageType.AUDIO_DATA, audio_data)
context.message_queue.put_nowait(pack_data)
print('插入', len(audio_data))
stream_callback = partial(handle_tts_result, context)
context.tts_client.set_result_callback(stream_callback)
# 3. 设置是否播放(可选,默认True)
context.tts_client.set_playback_enabled(False) # 设置为False则不播放
# 4. 初始化连接
await context.tts_client.initialize()
# 4. 定义前端数据接收任务
async def recv_frontend_data():
"""接收前端音频/控制指令"""
# while not asr_conn.stop_event.is_set():
while True:
try:
if not context.message_queue.empty():
await asyncio.sleep(0) # 立即让权
continue
raw_bytes = await websocket.receive_bytes()
unpack_bytes = ProtocolCodec.unpack(raw_bytes)
success = await push_audio_data(asr_conn, unpack_bytes)
# if not success:
# print("音频数据插入失败(队列满/连接失效)")
msg_type, sequence, data = ProtocolCodec.unpack(raw_bytes)
if msg_type == MessageType.AUDIO_DATA:
# 推送音频数据到ASR
success = await asr_client.push_audio(asr_conn, data)
if not success:
logger.warning(f"连接 {client_id} 音频推送失败(队列满/连接失效)")
elif msg_type == MessageType.CONTROL:
# 处理控制指令(如暂停/继续ASR
logger.info(f"连接 {client_id} 收到控制指令:{data}")
if data.get("action") == "stop_asr":
asr_conn.stop_event.set()
else:
logger.warning(f"连接 {client_id} 收到未知消息类型:{msg_type.value}")
except WebSocketDisconnect:
logger.info(f"前端 {conn_id} 主动断开连接")
asr_conn.stop_event.set()
logger.info(f"前端 {client_id} 主动断开连接")
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
asr_conn.stop_event.set()
await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
logger.error(f"连接 {client_id} 接收前端数据失败{str(e)}")
break
# 发送 ASR 结果
# 5. 定义ASR结果发送任务(从上下文队列取数据)
async def send_asr_result():
"""从结果队列发送 ASR 结果到前端(二进制格式)"""
while True:
try:
result = await asyncio.wait_for(context.message_queue.get(), timeout=0.05)
await websocket.send_bytes(result)
# 从当前连接的消息队列获取ASR结果(超时0.05秒避免阻塞)
result_packet = await asyncio.wait_for(
context.message_queue.get(), timeout=0.05
)
print('发送', result_packet)
await websocket.send_bytes(result_packet)
except asyncio.TimeoutError:
continue
continue # 无数据时继续等待
except Exception as e:
logger.error(f"发送 ASR 结果失败: {str(e)}")
logger.error(f"连接 {client_id} 发送ASR结果失败{str(e)}")
break
# 6. 启动任务并等待完成
task_send = asyncio.create_task(send_asr_result())
task_recv = asyncio.create_task(recv_frontend_data())
try:
# 等待两个任务,只要有一个完成就返回(比如前端断开/发送出错)
done, pending = await asyncio.wait(
[task_recv, task_send],
return_when=asyncio.FIRST_COMPLETED,
timeout=None # 无限等待,直到有任务完成
)
finally:
# 确保协程正确退出
# 等待剩余任务完成
for task in pending:
task.cancel()
await asyncio.gather(task_recv, task_send, return_exceptions=True)
pass
done, pending = await asyncio.wait(
[task_recv, task_send, communication_task],
return_when=asyncio.FIRST_COMPLETED
)
# 取消未完成的任务
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
except Exception as e:
logger.error(f"WebSocket连接处理异常: {str(e)}")
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射(如果已添加)
if context is not None and context.client_id in self.client_context_map:
del self.client_context_map[client_id]
logger.error(f"连接 {client_id} 处理异常{str(e)}")
# 异常时发送错误消息给前端
if websocket.state == "CONNECTED":
error_packet = ProtocolCodec.pack(
MessageType.ERROR, {"code": 5000, "message": f"服务异常:{str(e)}"}
)
await websocket.send_bytes(error_packet)
finally:
# 6. 统一资源清理(无论成功/失败,都执行
# 关闭WebSocket连接
try:
if hasattr(websocket, "state") and websocket.state == "CONNECTED":
await websocket.close(code=1008, reason="连接终止")
except Exception as close_e:
logger.warning(f"关闭连接失败 (client_id: {client_id}): {str(close_e)}")
# 7. 资源清理(关键
# 停止ASR通信任务
# if communication_task and not communication_task.done():
# communication_task.cancel()
# try:
# await communication_task
# except Exception as e:
# logger.warning(f"连接 {client_id} ASR任务取消异常:{str(e)}")
# 移除活跃连接
# 关闭ASR连接
# if asr_conn:
# await asr_client.close_connection(asr_conn)
# 关闭WebSocket连接
if websocket.state == "CONNECTED":
await websocket.close(code=1008, reason="连接终止")
# 移除连接和上下文
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射
if client_id in self.client_context_map:
del self.client_context_map[client_id]
if context:
pass
logger.info(f"连接资源清理完成 (client_id: {client_id}),当前连接数: {len(self.active_connections)}")
# 1. 建立连接
# 2. 初始化TTS
# tts_manager = await self.setup_tts_manager(result_queue)
# 3. 获取ASR连接
asr_conn = await get_idle_asr_connection()
if not asr_conn:
await websocket.send_json({"error": "ASR服务暂时不可用", "text": ""})
return
# 4. 启动ASR通信协程
# asr_callback = lambda res: self.asr_result_callback(res, websocket, user_id, result_queue)
# asr_task = asyncio.create_task(handle_asr_communication(asr_conn, asr_callback))
#
# # 5. 启动数据接收和发送协程
# task_recv = asyncio.create_task(self.recv_frontend_data(websocket, asr_conn))
# task_send = asyncio.create_task(self.send_results(websocket, result_queue, asr_conn))
#
# # 6. 等待任一任务完成
# done, pending = await asyncio.wait(
# [task_recv, task_send],
# return_when=asyncio.FIRST_COMPLETED
# )
# except Exception as e:
# logger.error(f"WebSocket连接处理异常: {str(e)}")
# if asr_conn:
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": str(e)})
# logger.error(f"连接 {client_id} 建立失败: {type(e).__name__}: {e}")
# try:
# # 确保连接已关闭(处理未正常关闭的情况)
# if websocket.client_state == "CONNECTED": # 根据实际WebSocket类型调整状态判断
# await websocket.close(code=1008, reason=str(e))
# except:
# pass
# 移除活跃连接(避免内存泄漏)
# if websocket in self.active_connections:
# self.active_connections.remove(websocket)
# # 移除上下文映射(如果已添加)
# if context is not None and context.client_id in self.client_context_map:
# del self.client_context_map[client_id]
# finally:
# pass
# 7. 资源清理
# logger.info(f"开始清理连接 {conn_id} 的资源")
# # 停止ASR
# if asr_conn:
# asr_conn.stop_event.set()
#
# # 取消任务
# if asr_task and not asr_task.done():
# asr_task.cancel()
# try:
# await asr_task
# except asyncio.CancelledError:
# pass
#
# # 清理TTS
# if tts_manager:
# await tts_manager.cleanup() # 假设TTSManager有cleanup方法,无则忽略
#
# # 断开连接
# if websocket:
# self.disconnect(websocket, conn_id)
# try:
# await websocket.close()
# except Exception:
# pass
#
# logger.info(f"连接 {conn_id} 资源清理完成")
logger.info(
f"连接 {client_id} 资源清理完成,当前连接数: {len(self.active_connections)}"
)
@@ -0,0 +1,449 @@
from fastapi import WebSocket
from typing import Dict, List, Optional
from pyexpat.errors import messages
from audio_ai_chat.config.logger import logger
from audio_ai_chat.core.asr.factory import ASRFactory # 导入ASR工厂
# from audio_ai_chat.core.llm.factory import LLMFactory # 导入LLM工厂
# from audio_ai_chat.core.tts.factory import TTSFactory # 导入TTS工厂
from audio_ai_chat.config.settings import settings
from audio_ai_chat.utils.exceptions import ServiceCallError
from fastapi import WebSocket, WebSocketDisconnect
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec
from typing import Optional, Dict, Callable, Awaitable, List, Any, Coroutine
from dataclasses import dataclass, field
import sys
import json
import asyncio
import websockets
import uuid
import logging
from audio_ai_chat.core.connection import ConnectionManager,ConnectionContext
from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec, MessageType
from audio_ai_chat.core.asr.asr_manager import ASRManager
async def asr_result_callback(result: dict):
if result.get("error"):
print(f"ASR错误:{result['error']}")
else:
print(f"ASR结果:{result['text']}(最终结果:{result['is_final']}")
class WebSocketConnectionManager:
"""WebSocket连接管理器"""
def __init__(self):
# 活跃连接列表
self.active_connections: List[WebSocket] = []
# 关键映射:client_id -> ConnectionContext(快速获取用户专属上下文)
self.client_context_map: Dict[str, ConnectionContext] = {}
self.connection_manager: Optional[ConnectionManager] = None
# self.asr_client = ASRFactory.get_asr_client()
# self.tts_client = TTSFactory.get_tts_client()
# 用户LLM会话存储
# self.user_llm_conversations: Dict[str, LLMConversation] = {}
# 全局唤醒事件
self.consume_wakeup = asyncio.Event()
async def initialize(self):
"""初始化:获取ConnectionManager全局单例(在FastAPI启动时调用)"""
self.connection_manager = await ConnectionManager.get_instance()
logger.info("WebSocketConnectionManager 初始化成功(绑定全局ConnectionManager")
async def connect(self, client_id, websocket: WebSocket) -> ConnectionContext:
"""
建立连接+身份校验(前端主动发送身份信息)
超时逻辑:5秒内未收到前端身份信息,自动关闭连接
返回:校验通过的 ConnectionContext(保证非空)
"""
# 1. 接受连接并加入活跃列表
await websocket.accept()
self.active_connections.append(websocket)
logger.info(
f"连接 {client_id} 已接受,等待前端发送身份信息(5秒超时)...,当前连接数: {len(self.active_connections)}")
# 2. 超时控制:5秒内未收到身份信息 -> 关闭连接
try:
ping_packet = await asyncio.wait_for(
websocket.receive_bytes(),
timeout=5.0
)
except asyncio.TimeoutError:
error_msg = f"连接 {client_id} 身份校验超时(5秒未收到消息)"
logger.warning(error_msg)
# 发送超时错误响应(二进制格式)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 1008, "message": "身份校验超时,请重试"}
)
await websocket.send_bytes(error_packet)
raise TimeoutError(error_msg) # 抛出异常,进入后续清理逻辑
# 3. 解包并验证包类型
print('ping_packet', ping_packet)
msg_type, _, identity_data = ProtocolCodec.unpack(ping_packet)
if msg_type != MessageType.IDENTITY:
error_msg = f"连接 {client_id} 首个包类型错误(期望{MessageType.IDENTITY.value},实际{msg_type.value}"
logger.error(error_msg)
# 发送类型错误响应
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4001, "message": "非法请求:首个包必须是身份校验包"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# 4. 身份信息并校验
# todo
# 提取核心字段(必选字段校验)
user_id = identity_data.get("user_id")
token = identity_data.get("token")
name = identity_data.get("name") or f"用户{user_id}" # 提供默认名称
if not all([user_id, token]):
error_msg = f"连接 {client_id} 身份信息不完整(缺少user_id或token"
logger.error(error_msg)
error_packet = ProtocolCodec.pack(
MessageType.ERROR,
{"code": 4003, "message": "身份信息不完整:必须包含user_id和token"}
)
await websocket.send_bytes(error_packet)
raise ValueError(error_msg)
# TODO: 实际身份校验逻辑(根据你的业务扩展)
# 5. 校验通过:更新上下文并响应前端
context = await self.connection_manager.create_or_reconnect_context(
new_client_id=client_id,
user_id=user_id
)
context = ConnectionContext(client_id=client_id) # 提前创建上下文(保证最终返回非空)
context.set_user_info(token, user_id, name)
self.client_context_map[client_id] = context # 加入上下文映射
# 发送成功响应
success_packet = ProtocolCodec.pack(
MessageType.IDENTITY,
{
"code": 200,
"message": "身份校验成功,连接已就绪",
"data": {"client_id": client_id, "user_id": user_id, "name": name}
}
)
await websocket.send_bytes(success_packet)
logger.info(f"用户 {user_id}{name})身份校验通过,连接就绪(client_id: {client_id}")
return context
# 初始化LLM会话
# self._init_llm_conversation(user_id)
# return conn_id, user_id, conn_id # conn_id 同时作为 tts_session_id
# def _init_llm_conversation(self, user_id: str):
# """初始化用户LLM会话"""
# if user_id not in self.user_llm_conversations:
# self.user_llm_conversations[user_id] = LLMConversation(
# user_id=user_id,
# scene_description="语音识别对话场景"
# )
def disconnect(self, websocket: WebSocket, conn_id: str):
"""断开连接并清理资源"""
if websocket in self.active_connections:
self.active_connections.remove(websocket)
logger.info(f"连接 {conn_id} 已断开,当前连接数: {len(self.active_connections)}")
# async def setup_tts_manager(self, result_queue: asyncio.Queue) -> TTSManager:
# """初始化TTS管理器"""
#
# def handle_tts_result(req_id: str, result: Dict[str, Any]):
# """TTS结果回调处理"""
# try:
# status = result.get("status")
# if status == "completed":
# audio_data = result.get("audio_data")
# if audio_data is not None and len(audio_data) > 0:
# # 转换为PCM格式
# pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
# pcm_bytes = pcm_data.tobytes()
# result_queue.put_nowait(pcm_bytes)
# except Exception as e:
# logger.error(f"TTS结果处理失败: {str(e)}")
#
# self.tts_client = TTSFactory.get_tts_client()
# tts_manager = TTSManager()
# tts_manager.set_result_callback(handle_tts_result)
# tts_manager.set_playback_enabled(False)
# await tts_manager.initialize()
# return tts_manager
# async def asr_result_callback(self, result: dict, websocket: WebSocket,
# user_id: str, result_queue: asyncio.Queue):
# """ASR结果回调处理"""
# try:
# logger.info(f"ASR识别结果: {result}")
# final_asr_text = result.get("text", "").strip()
#
# # 转发ASR结果到前端队列
# if final_asr_text:
# print(f"插入ASR结果时队列大小: {result_queue.qsize()}")
# self.consume_wakeup.set() # 唤醒消费协程
#
# # 异步调用大模型
# llm_conversation = self.user_llm_conversations.get(user_id)
# if llm_conversation:
# asyncio.create_task(
# self.call_llm_and_send(
# query=final_asr_text,
# conversation=llm_conversation,
# websocket=websocket
# )
# )
# except Exception as e:
# logger.error(f"ASR回调执行失败: {str(e)}")
# async def llm_stream_callback(self, chunk: str, tts_manager: TTSManager):
# """大模型流式回调处理"""
# if not chunk:
# return
# try:
# # 提交TTS合成请求
# await tts_manager.synthesize(chunk)
# await asyncio.sleep(0) # 让出调度权
# except Exception as e:
# logger.error(f"LLM流式回调处理失败: {str(e)}")
# async def call_llm_and_send(self, query: str, conversation: LLMConversation, websocket: WebSocket):
# """调用大模型并处理结果"""
# logger.info(f"调用大模型 - 用户({conversation.user_id}): {query}")
# try:
# conv_id, full_reply = await llm_client.send_message(
# query=query,
# conversation=conversation,
# stream_callback=self.llm_stream_callback,
# response_mode="streaming"
# )
# logger.info(f"大模型回复完成 - 会话ID: {conv_id}, 完整回复: {full_reply}")
# except Exception as e:
# logger.error(f"大模型调用失败: {str(e)}")
# if not websocket.client_state.disconnected:
# await websocket.send_json({
# "type": "llm_error",
# "data": {"error": str(e)}
# })
async def recv_frontend_data(self, websocket: WebSocket, asr_conn):
"""接收前端音频数据并推送到ASR"""
while not asr_conn.stop_event.is_set():
try:
raw_bytes = await websocket.receive_bytes()
# success = await push_audio_data(asr_conn, raw_bytes)
# if not success:
# logger.warning("音频数据插入ASR失败(队列满/连接失效)")
except WebSocketDisconnect:
logger.info("前端主动断开连接")
asr_conn.stop_event.set()
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
# break
async def send_results(self, websocket: WebSocket, result_queue: asyncio.Queue, asr_conn):
"""从结果队列发送数据到前端"""
while True:
try:
# 等待队列数据或超时
result = await asyncio.wait_for(result_queue.get(), timeout=0.05)
# if not websocket.client_state.disconnected:
await websocket.send_bytes(result)
except asyncio.TimeoutError:
if asr_conn.stop_event.is_set():
break
continue
except Exception as e:
logger.error(f"发送结果到前端失败: {str(e)}")
asr_conn.stop_event.set()
break
async def handle_connection(self, websocket: WebSocket):
"""处理单个WebSocket连接的完整生命周期"""
new_client_id = str(id(websocket))
logger.info(f"新WebSocket连接:client_id={new_client_id}")
context = None
try:
context = await self.connect(new_client_id, websocket)
asr_client = ASRManager.get_instance()
asr_conn = await asr_client.get_connection()
if not asr_conn:
print("获取ASR连接失败")
raise
communication_task = asyncio.create_task(
asr_client.start_communication(asr_conn, asr_result_callback)
)
# 接收前端数据
async def recv_frontend_data():
"""接收前端音频/控制指令"""
# while not asr_conn.stop_event.is_set():
while True:
try:
if not context.message_queue.empty():
await asyncio.sleep(0) # 立即让权
continue
raw_bytes = await websocket.receive_bytes()
msg_type, sequence, unpack_bytes = ProtocolCodec.unpack(raw_bytes)
if msg_type == MessageType.AUDIO_DATA:
success = await asr_client.push_audio(asr_conn, unpack_bytes)
if not success:
print(f"音频数据插入失败(队列满/连接失效)")
else:
print(f"其他类型数据", msg_type)
except WebSocketDisconnect:
# logger.info(f"前端 {conn_id} 主动断开连接")
# asr_conn.stop_event.set()
break
except Exception as e:
logger.error(f"接收前端数据失败: {str(e)}")
# asr_conn.stop_event.set()
# await websocket.send_json({"error": f"接收数据失败: {str(e)}"})
break
# 发送 ASR 结果
async def send_asr_result():
"""从结果队列发送 ASR 结果到前端(二进制格式)"""
while True:
try:
result = await asyncio.wait_for(context.message_queue.get(), timeout=0.05)
await websocket.send_bytes(result)
except asyncio.TimeoutError:
continue
except Exception as e:
logger.error(f"发送 ASR 结果失败: {str(e)}")
break
task_send = asyncio.create_task(send_asr_result())
task_recv = asyncio.create_task(recv_frontend_data())
try:
# 等待两个任务,只要有一个完成就返回(比如前端断开/发送出错)
done, pending = await asyncio.wait(
[task_recv, task_send],
return_when=asyncio.FIRST_COMPLETED,
timeout=None # 无限等待,直到有任务完成
)
finally:
# 确保协程正确退出
# 等待剩余任务完成
for task in pending:
task.cancel()
await asyncio.gather(task_recv, task_send, return_exceptions=True)
pass
except Exception as e:
logger.error(f"WebSocket连接处理异常: {str(e)}")
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射(如果已添加)
if context is not None and context.client_id in self.client_context_map:
del self.client_context_map[new_client_id]
finally:
# 6. 统一资源清理(无论成功/失败,都执行)
# 关闭WebSocket连接
try:
if hasattr(websocket, "state") and websocket.state == "CONNECTED":
await websocket.close(code=1008, reason="连接终止")
except Exception as close_e:
logger.warning(f"关闭连接失败 (client_id: {new_client_id}): {str(close_e)}")
# 移除活跃连接
if websocket in self.active_connections:
self.active_connections.remove(websocket)
# 移除上下文映射
if new_client_id in self.client_context_map:
del self.client_context_map[new_client_id]
if context:
pass
logger.info(f"连接资源清理完成 (client_id: {new_client_id}),当前连接数: {len(self.active_connections)}")
# 1. 建立连接
# 2. 初始化TTS
# tts_manager = await self.setup_tts_manager(result_queue)
# 3. 获取ASR连接
# asr_conn = await get_idle_asr_connection()
# if not asr_conn:
# await websocket.send_json({"error": "ASR服务暂时不可用", "text": ""})
# return
# 4. 启动ASR通信协程
# asr_callback = lambda res: self.asr_result_callback(res, websocket, user_id, result_queue)
# asr_task = asyncio.create_task(handle_asr_communication(asr_conn, asr_callback))
#
# # 5. 启动数据接收和发送协程
# task_recv = asyncio.create_task(self.recv_frontend_data(websocket, asr_conn))
# task_send = asyncio.create_task(self.send_results(websocket, result_queue, asr_conn))
#
# # 6. 等待任一任务完成
# done, pending = await asyncio.wait(
# [task_recv, task_send],
# return_when=asyncio.FIRST_COMPLETED
# )
# except Exception as e:
# logger.error(f"WebSocket连接处理异常: {str(e)}")
# if asr_conn:
# asr_conn.stop_event.set()
# if not websocket.client_state.disconnected:
# await websocket.send_json({"error": str(e)})
# logger.error(f"连接 {client_id} 建立失败: {type(e).__name__}: {e}")
# try:
# # 确保连接已关闭(处理未正常关闭的情况)
# if websocket.client_state == "CONNECTED": # 根据实际WebSocket类型调整状态判断
# await websocket.close(code=1008, reason=str(e))
# except:
# pass
# 移除活跃连接(避免内存泄漏)
# if websocket in self.active_connections:
# self.active_connections.remove(websocket)
# # 移除上下文映射(如果已添加)
# if context is not None and context.client_id in self.client_context_map:
# del self.client_context_map[client_id]
# finally:
# pass
# 7. 资源清理
# logger.info(f"开始清理连接 {conn_id} 的资源")
# # 停止ASR
# if asr_conn:
# asr_conn.stop_event.set()
#
# # 取消任务
# if asr_task and not asr_task.done():
# asr_task.cancel()
# try:
# await asr_task
# except asyncio.CancelledError:
# pass
#
# # 清理TTS
# if tts_manager:
# await tts_manager.cleanup() # 假设TTSManager有cleanup方法,无则忽略
#
# # 断开连接
# if websocket:
# self.disconnect(websocket, conn_id)
# try:
# await websocket.close()
# except Exception:
# pass
#
# logger.info(f"连接 {conn_id} 资源清理完成")
+19 -2
View File
@@ -8,18 +8,35 @@ from audio_ai_chat.codec.ProtocolCodec import ProtocolCodec # 已有加密类
from audio_ai_chat.utils.exceptions import CodecError, ServiceCallError
from contextlib import asynccontextmanager
from frontend_ws import frontend_websocket_handler
from audio_ai_chat.core.asr.asr_manager import ASRManager
# FastAPI 启动时初始化 ASR 连接池
@asynccontextmanager
async def lifespan(app: FastAPI):
# 启动时执行(原 startup 逻辑)
print(' FastAPI 启动时初始化 ASR 连接池')
# await init_asr_pool()
yield # 应用运行中
# 关闭时执行(可选,比如清理连接池)
print("应用关闭,开始清理 ASR 连接池...")
# await close_asr_pool()
# 生命周期函数
@asynccontextmanager
async def lifespan(app: FastAPI):
print("=== 开始初始化 ASR 服务 ===")
init_success, init_msg = await ASRManager.initialize()
print(f"ASR 初始化结果:{init_msg}")
# if not init_success:
# 连接池为空/初始化失败,终止服务启动
# raise ServiceInitError(f"服务启动失败:{init_msg}")
# 2. 初始化WebSocketConnectionManager(绑定全局ConnectionManager
await ws_manager.initialize()
print("=== ASR 服务初始化完成 ===")
yield # 应用运行中
# 关闭时清理
print("=== 开始关闭 ASR 服务 ===")
await ASRManager.close()
print("=== ASR 服务关闭完成 ===")
app = FastAPI(
title="语音AI对话系统",
File diff suppressed because it is too large Load Diff
+238
View File
@@ -0,0 +1,238 @@
# audio_ai_chat/websocket/manager.py
from fastapi import WebSocket
from typing import Dict, Optional
from datetime import datetime
import base64
from audio_ai_chat.asr.base import ASRBase, ASRResultCallback
from audio_ai_chat.asr.asr_manager import ASRManager
from audio_ai_chat.websocket.connection_context import ConnectionManager, ConnectionContext # 导入全局单例类
from audio_ai_chat.config.logger import logger
class WebSocketConnectionManager:
"""全局唯一的WebSocket连接处理器(管理WebSocket连接生命周期)"""
def __init__(self):
self.asr_conn_map: Dict[str, Optional[object]] = {} # key=client_idvalue=ASR连接
# 不实例化新的ConnectionManager,而是使用全局单例
self.connection_manager: Optional[ConnectionManager] = None
async def initialize(self):
"""初始化:获取ConnectionManager全局单例(在FastAPI启动时调用)"""
self.connection_manager = await ConnectionManager.get_instance()
logger.info("WebSocketConnectionManager 初始化成功(绑定全局ConnectionManager")
async def handle_connection(self, websocket: WebSocket):
"""处理单个WebSocket连接的完整生命周期"""
# 校验ConnectionManager是否初始化
if not self.connection_manager:
await websocket.accept()
await websocket.send_text("服务未初始化完成,请稍后重试")
await websocket.close()
logger.error("WebSocketConnectionManager 未初始化,拒绝连接")
return
# 1. 接受连接,生成client_id(用字符串类型,避免int溢出)
await websocket.accept()
client_id = str(id(websocket)) # client_id为字符串,与ConnectionManager的key类型一致
logger.info(f"新WebSocket连接:client_id={client_id}")
try:
# 2. 创建连接上下文(通过全局ConnectionManager
context = await self.connection_manager.create_connection(client_id=client_id)
if not context:
await websocket.send_text("连接上下文创建失败")
await websocket.close()
return
# 3. 获取ASR实例
asr_client = ASRManager.get_instance()
if not asr_client or not ASRManager.is_available():
await websocket.send_json({
"type": "error",
"message": "ASR服务未初始化,无法提供转写服务",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
await self.connection_manager.remove_connection(client_id=client_id)
await websocket.close()
return
# 4. 获取ASR连接
asr_conn = await asr_client.get_connection()
if not asr_conn:
await websocket.send_json({
"type": "error",
"message": "ASR无空闲连接,连接失败",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
await self.connection_manager.remove_connection(client_id=client_id)
await websocket.close()
return
self.asr_conn_map[client_id] = asr_conn
# 5. 定义ASR结果回调(绑定当前上下文)
async def asr_callback(result: Dict[str, Any]):
if not context.is_active:
logger.warning(f"连接已关闭,忽略ASR结果:client_id={client_id}")
return
# 处理ASR结果并存入上下文
context.add_asr_result(result)
# 推送给前端
if result.get("error"):
await websocket.send_json({
"type": "asr_error",
"message": result["error"],
"timestamp": datetime.utcnow().isoformat() + "Z"
})
else:
await websocket.send_json({
"type": "asr_progress" if not result["is_final"] else "asr_final",
"text": result["text"],
"is_final": result["is_final"],
"timestamp": result.get("timestamp", datetime.utcnow().isoformat() + "Z")
})
# 6. 启动ASR通信
asr_task = asyncio.create_task(
asr_client.start_communication(conn=asr_conn, callback=asr_callback)
)
# 7. 循环接收前端数据
while context.is_active:
try:
# 假设前端发送JSON格式数据(区分音频/文本/用户信息)
data = await websocket.receive_json()
data_type = data.get("type")
# 处理用户信息(登录后发送)
if data_type == "user_info":
try:
token = data.get("token")
user_id = data.get("user_id")
name = data.get("name", "匿名用户")
context.set_user_info(token=token, user_id=user_id, name=name)
await websocket.send_json({
"type": "info",
"message": "用户信息设置成功",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
err_msg = f"用户信息设置失败:{str(e)}"
context.add_system_message(err_msg)
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
# 处理Base64编码的音频数据
elif data_type == "audio_data":
audio_base64 = data.get("audio_data")
if not audio_base64:
continue
try:
audio_data = base64.b64decode(audio_base64)
success = await asr_client.push_audio(asr_conn, audio_data)
if not success:
await websocket.send_json({
"type": "warning",
"message": "ASR音频队列已满,部分数据丢失",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
err_msg = f"音频解码失败:{str(e)}"
logger.error(f"client_id={client_id}{err_msg}")
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
# 处理纯文本输入
elif data_type == "text_input":
text = data.get("text", "").strip()
if text:
context.add_chat_history({
"role": "user",
"content": text,
"source": "text",
"asr_metadata": None
})
await websocket.send_json({
"type": "info",
"message": f"已接收文本:{text}",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
# 处理大模型请求
elif data_type == "request_llm":
if context.is_processing:
await websocket.send_json({
"type": "warning",
"message": "正在处理上一个请求,请稍后再试",
"timestamp": datetime.utcnow().isoformat() + "Z"
})
continue
# 获取对话历史
chat_history = context.get_chat_history(limit=20)
logger.debug(f"请求大模型:client_id={client_id},历史条数={len(chat_history)}")
# 模拟大模型调用(实际替换为真实LLM调用)
context.is_processing = True
try:
# llm_response = await context.llm_session.generate(chat_history=chat_history)
llm_response = f"模拟大模型回复:已收到你的{len(chat_history)}条对话历史"
context.add_llm_result(llm_response)
await websocket.send_json({
"type": "llm_response",
"text": llm_response,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
err_msg = f"大模型调用失败:{str(e)}"
context.add_system_message(err_msg)
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
finally:
context.is_processing = False
# 未知数据类型
else:
err_msg = f"未知数据类型:{data_type}"
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
except Exception as e:
# 捕获前端发送数据异常(如断开连接)
logger.error(f"接收前端数据异常:client_id={client_id}error={str(e)}")
break
except Exception as e:
# 其他异常
err_msg = f"连接处理异常:{str(e)}"
logger.error(f"client_id={client_id}{err_msg}")
await websocket.send_json({
"type": "error",
"message": err_msg,
"timestamp": datetime.utcnow().isoformat() + "Z"
})
finally:
# 8. 资源清理
# 取消ASR任务
asr_task.cancel()
try:
await asr_task
except asyncio.CancelledError:
pass
# 释放ASR连接
if client_id in self.asr_conn_map:
asr_conn = self.asr_conn_map.pop(client_id)
await asr_client.release_connection(asr_conn)
# 移除连接上下文
await self.connection_manager.remove_connection(client_id=client_id)
# 关闭WebSocket
await websocket.close()
logger.info(f"WebSocket连接关闭:client_id={client_id}")
+1 -1
View File
@@ -2,7 +2,7 @@
<module type="PYTHON_MODULE" version="4">
<component name="NewModuleRootManager">
<content url="file://$MODULE_DIR$" />
<orderEntry type="jdk" jdkName="AIStreamTest" jdkType="Python SDK" />
<orderEntry type="jdk" jdkName="audio-ai-chat" jdkType="Python SDK" />
<orderEntry type="sourceFolder" forTests="false" />
</component>
</module>
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+39
View File
@@ -0,0 +1,39 @@
# 开发环境
ENV = 'development'
# 'development'
# VITE_APP_BASE_API_Url = 'https://aits.jlbank.com.cn:7001'
VITE_APP_BASE_API_Url = 'https://aitstest.jlbank.com.cn:7002'
# VITE_APP_BASE_API_Url = 'https://aitstest.jlbank.com.cn:7002'
# VITE_APP_BASE_API_Url = 'http://192.168.247.200'
# VITE_APP_BASE_API_Url = 'http://aitscdn.jlbank.com.cn:7001'
# VITE_APP_BASE_API_Url = 'http://192.168.108.129'
# dev
# VITE_APP_BASE_API_Url = 'http://25.18.122.65:7001'
# sit
# VITE_APP_BASE_API_Url = 'http://25.18.122.91:9786'
# UAT
# VITE_APP_BASE_API_Url = 'http://25.18.122.78:9786'
# DEV
# app入口
# VITE_APP_BASE_API_Url = 'http://25.18.122.65:7001'
# VITE_APP_BASE_API_Url = 'http://25.16.122.91:9786'
# VITE_APP_BASE_API_Url = 'http://25.18.122.66:9786'
# h5专用地址x2
# VITE_APP_BASE_H5_API_Url_TRAAPP = 'http://25.64.32.150:9601'
# dev
#VITE_APP_BASE_H5_API_Url = 'http://25.18.122.65:7001'
# sit
# VITE_APP_BASE_H5_API_Url = 'http://25.18.122.91:9786'
+8
View File
@@ -0,0 +1,8 @@
# 任东专用开发环境
ENV = 'development.rd'
# 'development.rd'
#VITE_APP_BASE_API_Url = 'http://25.18.122.65:7001'
VITE_APP_BASE_API_Url = 'http://25.18.122.91:9786'
VITE_APP_BASE_H5_API_Url_TRAEXAM = 'http://25.64.32.154:9604'
+9
View File
@@ -0,0 +1,9 @@
# 田岩开发环境
ENV = 'development.ty'
# 'development.ty'
VITE_APP_BASE_API_Url = 'https://aitstest.jlbank.com.cn:7001'
# VITE_APP_BASE_H5_API_Url_TRAEXAM = 'http://192.168.108.129'
+14
View File
@@ -0,0 +1,14 @@
# 杨航开发环境
ENV = 'development.yh'
# 'development.yh'
# VITE_APP_BASE_H5_API_Url = 'http://25.18.122.65:7001'
VITE_APP_BASE_H5_API_Url = 'http://25.18.122.91:9786'
# VITE_APP_BASE_H5_API_Url = 'http://25.18.122.78:9786'
# VITE_APP_BASE_API_Url = 'https://aitstest.jlbank.com.cn:7002'
# VITE_APP_BASE_H5_API_Url_TRAEXAM = 'http://25.64.32.154:9604'
# VITE_APP_BASE_H5_API_Url_TRAPRACTICE = 'http://25.64.32.156:9603'
# VITE_APP_BASE_H5_API_Url_TRAAPP = 'http://25.64.32.154:9601'
# VITE_APP_BASE_H5_API_Url_TRAAPP = 'http://25.64.32.158:9601'
+5
View File
@@ -0,0 +1,5 @@
# 开发环境
ENV = 'development'
# 'development'
VITE_APP_BASE_API_Url = ''
+5
View File
@@ -0,0 +1,5 @@
# 生产环境
ENV = 'production'
# base api
VITE_APP_BASE_API_Url = 'https://aits.jlbank.com.cn:7001'
+5
View File
@@ -0,0 +1,5 @@
# SIT环境
ENV = 'sit'
# base api
VITE_APP_BASE_API_Url = 'https://aitstest.jlbank.com.cn:7002'
+5
View File
@@ -0,0 +1,5 @@
# UAT环境
ENV = 'uat'
# base api
VITE_APP_BASE_API_Url = 'https://aitstest.jlbank.com.cn:7003'
+26
View File
@@ -0,0 +1,26 @@
# Logs
logs
*.log
npm-debug.log*
yarn-debug.log*
yarn-error.log*
pnpm-debug.log*
lerna-debug.log*
package-lock
package-lock.json
node_modules
.DS_Store
dist
unpackage
*.local
build
release
# Editor directories and files
.idea
.svn
*.suo
*.ntvs*
*.njsproj
*.sln
*.sw?
+9
View File
@@ -0,0 +1,9 @@
npm install --registry=http://25.12.10.69:8081/repository/aliyun-npm/
npm 淘宝源下载 npm config set registry https://registry.npmmirror.com
npm config set registry http://25.12.10.69:8081/repository/aliyun-npm/
npm install --registry=https://registry.npmmirror.com
+20
View File
@@ -0,0 +1,20 @@
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<script>
var coverSupport = 'CSS' in window && typeof CSS.supports === 'function' && (CSS.supports('top: env(a)') ||
CSS.supports('top: constant(a)'))
document.write(
'<meta name="viewport" content="width=device-width, user-scalable=no, initial-scale=1.0, maximum-scale=1.0, minimum-scale=1.0' +
(coverSupport ? ', viewport-fit=cover' : '') + '" />')
</script>
<title></title>
<!--preload-links-->
<!--app-context-->
</head>
<body>
<div id="app"><!--app-html--></div>
<script type="module" src="/src/main.js"></script>
</body>
</html>
+57
View File
@@ -0,0 +1,57 @@
{
"name": "tra-app",
"version": "0.0.1",
"description": "",
"scripts": {
"dev:h5": "uni",
"dev:rd": "uni -p h5 --mode development.rd",
"dev:yh": "uni -p h5 --mode development.yh",
"dev:ty": "uni -p h5 --mode development.ty",
"dev:h5:sit": "uni -p h5 --mode sit",
"dev:h5:uat": "uni -p h5 --mode uat",
"dev:h5:production": "uni -p h5 --mode production",
"build:app-plus:dev": "uni build -p app-plus --mode development",
"build:app-plus:develop": "uni build -p app-plus --mode development",
"build:app-plus:sit": "uni build -p app-plus --mode sit",
"build:app-plus:uat": "uni build -p app-plus --mode uat",
"build:app-plus:prod": "uni build -p app-plus",
"build:h5:dev": "uni build -p h5 --mode development_h5",
"build:h5": "uni build"
},
"dependencies": {
"@dcloudio/uni-app": "3.0.0-4060620250520001",
"@dcloudio/uni-app-harmony": "3.0.0-4060620250520001",
"@dcloudio/uni-app-plus": "3.0.0-4060620250520001",
"@dcloudio/uni-components": "3.0.0-4060620250520001",
"@dcloudio/uni-h5": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-alipay": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-baidu": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-harmony": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-jd": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-kuaishou": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-lark": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-qq": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-toutiao": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-weixin": "3.0.0-4060620250520001",
"@dcloudio/uni-mp-xhs": "3.0.0-4060620250520001",
"@dcloudio/uni-quickapp-webview": "3.0.0-4060620250520001",
"crypto-js": "^3.1.9-1",
"dompurify": "^3.3.0",
"markdown-it": "^14.1.0",
"pinia": "^2.3.1",
"terser": "^5.42.0",
"uuid": "^11.1.0",
"vue": "^3.5.11",
"vue-i18n": "^9.1.9"
},
"devDependencies": {
"@dcloudio/types": "^3.4.8",
"@dcloudio/uni-automator": "3.0.0-4060620250520001",
"@dcloudio/uni-cli-shared": "3.0.0-4060620250520001",
"@dcloudio/uni-stacktracey": "3.0.0-4060620250520001",
"@dcloudio/vite-plugin-uni": "3.0.0-4060620250520001",
"@vue/runtime-core": "^3.4.21",
"sass": "1.77.0",
"vite": "5.2.8"
}
}
+10
View File
@@ -0,0 +1,10 @@
/// <reference types='@dcloudio/types' />
import 'vue'
declare module '@vue/runtime-core' {
type Hooks = App.AppInstance & Page.PageInstance;
interface ComponentCustomOptions extends Hooks {
}
}
+153
View File
@@ -0,0 +1,153 @@
<script>
export default {
onLaunch: function () {
uni.onTabBarMidButtonTap(() => {
uni.navigateTo({
url: '/pages/index/dialog',
animationType: 'slide-in-bottom'
});
});
},
onShow: function () {
},
onHide: function () {
}
};
</script>
<style lang="scss">
/* 注意要写在第一行,同时给style标签加入lang="scss"属性 */
// @import "@/uni_modules/uview-ui/index.scss";
@font-face {
font-family: 'PingFangSC';
src: url('@/static/font/PingFang-SC.ttf') format('truetype');
font-weight: normal;
font-style: normal;
}
view {
box-sizing: border-box;
font-family: 'PingFangSC';
}
// uview线
@font-face {
font-family: 'uicon-iconfont';
src: url('@/static/font/font_2225171_8kdcwk4po24.ttf') format('truetype');
font-weight: normal;
font-style: normal;
}
.max_page {
min-height: 100vh;
}
.mr-10 {
margin-right: 10rpx;
}
.mb-8 {
margin-bottom: 8rpx;
}
.pt-8 {
padding-top: 8rpx;
}
.mb-10 {
margin-bottom: 10rpx;
}
.mb-18 {
margin-bottom: 18rpx;
}
.mb-64 {
margin-bottom: 64rpx;
}
.mb-30 {
margin-bottom: 30rpx;
}
.mb-36 {
margin-bottom: 36rpx;
}
.mb-4 {
margin-bottom: 4rpx;
}
.fz-30 {
font-size: 30rpx;
}
.fz-28 {
font-size: 28rpx;
}
.fz-24 {
font-size: 24rpx;
}
.w_100 {
width: 100%;
}
.h_100 {
height: 100%;
}
.flex {
display: flex;
}
.flex-jcsb {
display: flex;
justify-content: space-between;
}
.flex-c-c {
display: flex;
justify-content: center;
align-items: center;
}
.tar {
text-align: right;
}
.fw-600 {
font-weight: 600;
}
.fw-400 {
font-weight: 400;
}
.img_100 {
width: 100%;
height: 100%;
}
.color-666{
color: #666666;
}
.fl1 {
flex: 1;
}
.font_pf {
font-family: 'PingFangSC';
}
.ellipsis-text {
white-space: nowrap;
/* 强制不换行 */
overflow: hidden;
/* 超出部分隐藏 */
text-overflow: ellipsis;
/* 超出部分显示省略号 */
}
/* 多行文本溢出 */
.multi-line-ellipsis {
display: -webkit-box;
-webkit-box-orient: vertical;
-webkit-line-clamp: 3;
/* 显示的行数 */
overflow: hidden;
}
.no-overflow {
overflow: hidden;
}
.bold {
font-weight: bold;
}
.maoh{
margin-right: 8rpx;
margin-left: 2rpx;
}
</style>
+38
View File
@@ -0,0 +1,38 @@
{
"version": "1",
"prompt": "template",
"title": "个人信息保护提示",
"message": "欢迎使用吉AI学!<br/>  请你务必审慎阅读、充分理解<a href=\"https://aits.jlbank.com.cn:7001/traapp/dfAgtInfo/queryValidDfAgtHtml\">《用户隐私协议》</a>各条款,帮助您了解我们为您提供的服务、我们如何处理个人信息以及您享有的权利。我们会严格按照相关法律法规要求,采取各种安全措施来保护您的个人信息。<br/>  如果你同意,请点击下面按钮开始接受我们的服务。<br/>1.为了保障软件的安全运行和账户安全,我们会申请手机您的设备信息、IP地址、WLAN MAC地址。<br/>2.上传或拍摄图片,需要使用您的媒体影音、图片、视频、音频、相机、等权限。<br/>3.为了实现AI问答、课程学习等APP内功能,我们需要使用您的麦克风权限。",
"buttonAccept": "同意并接受",
"buttonRefuse": "暂不同意",
"hrefLoader": "system",
"backToExit":"true",
"second": {
"title": "确认提示",
"message": "  进入应用前,你需先同意<a href=\"https://aits.jlbank.com.cn:7001/traapp/dfAgtInfo/queryValidDfAgtHtml\">《用户隐私协议》</a>,否则将退出应用。",
"buttonAccept": "同意并继续",
"buttonRefuse": "退出应用"
},
"disagreeMode":{
"support": false,
"loadNativePlugins": false,
"visitorEntry": false,
"showAlways": false
},
"styles": {
"backgroundColor": "#fff",
"borderRadius":"5px",
"title": {
"color": "#000"
},
"buttonAccept": {
"color": "#F91B59"
},
"buttonRefuse": {
"color": "#333"
},
"buttonVisitor": {
"color": "#00ffff"
}
}
}
+48
View File
@@ -0,0 +1,48 @@
import request from '@/api/request'
// 1.本月/本年/累计 学习信息查询接口
export const queryStudyStatistics = (data) => {
return request({
url: '/traapp/traStudyAnalysis/queryStudyStatistics',
method: 'post',
data
});
};
// 课程推荐查询接口
export const queryRecmdCrsInfo = (data) => {
return request({
url: '/traapp/traStudyAnalysis/queryRecmdCrsInfo',
method: 'post',
data
});
};
// 4.维度分析查询接口
export const queryDimesionInfo = (data) => {
return request({
url: '/traapp/traStudyAnalysis/queryDimesionInfo',
method: 'post',
timeout: 60000,
data
});
};
// 5.学习习惯查询接口
export const queryHabitInfo = (data) => {
return request({
url: '/traapp/traStudyAnalysis/queryHabitInfo',
method: 'post',
timeout: 60000,
data
});
};
// 6.学习建议查询接口
export const querySuggestInfo = (data) => {
return request({
url: '/traapp/traStudyAnalysis/querySuggestInfo',
method: 'post',
timeout: 60000,
data
});
};
+90
View File
@@ -0,0 +1,90 @@
import { getToken } from '@/common/common.js'
import { get_base_url } from '@/api/request'
const base_url = get_base_url();
const token = getToken();
// 基础配置
const baseConfig = {
enableChunked: true, // 必须开启分块传输
responseType: 'arraybuffer', // 微信小程序需用arraybuffer
timeout: 30000 // 超时时间延长
}
// 流式请求核心实现
function streamRequest(options) {
const {
url,
method = 'GET',
data,
headers = {},
onChunk,
onComplete,
onError
} = options
// headers['Accept'] = 'text/event-stream'
headers['Content-Type'] = 'text/event-stream'
headers['summary'] = token
// 创建请求任务
const requestTask = uni.request({
url,
method,
data,
header: headers,
responseType: 'arraybuffer',
...baseConfig,
success: (res) => {
console.log(res.data)
onComplete?.(res)
},
fail: (err) => {
onError?.(err)
}
})
console.log(requestTask)
// 流式数据处理器
let bufferCache = ''
requestTask.onChunkReceived((res) => {
try {
// ArrayBuffer转字符串(兼容多平台)
const uint8Array = new Uint8Array(res.data)
const chunkText = bufferCache +
String.fromCharCode.apply(null, uint8Array)
// 处理SSE格式数据(data:开头)
const events = chunkText.split('\n\n')
bufferCache = events.pop() || '' // 缓存不完整数据
events.forEach(event => {
if (event.startsWith('data:')) {
onChunk?.(event.substring(5).trim())
}
})
} catch (e) {
onError?.(e)
}
})
return {
abort: () => requestTask.abort()
}
}
// 使用示例
export const demoUsage = (data) => {
const controller = streamRequest({
url: `${base_url}/trastudy/intgask/askAbout`,
data: data,
onChunk: (chunk) => {
console.log('收到数据块:', chunk)
// 实时更新UI逻辑
},
onComplete: () => console.log('传输完成'),
onError: (err) => console.error('错误:', err)
})
// 需要中断时调用
// controller.abort()
}
+10
View File
@@ -0,0 +1,10 @@
import request from '@/api/request'
// 获取list
export const leaderBoardApi = (data) => {
return request({
url: '/traapp/myLeaderBoard/leaderBoard ',
method: 'post',
data
});
};
+272
View File
@@ -0,0 +1,272 @@
import {
get_base_url,
REQUES_ERROR_CODES,
goto_login_fun
} from '@/api/request'
import {
getToken
} from '@/common/common.js'
import common from '@/common/common';
export const uploadFile = (file, data) => {
const url = '/traoss/tras3/upload';
const base_url = get_base_url(url);
const token = getToken();
const header = {};
console.log('file', file);
if (token)
header['summary'] = token
return new Promise((resolve, reject) => {
uni.uploadFile({
url: `${base_url}${url}`,
header: header,
filePath: file,
name: 'file',
formData: data,
success(result) {
try {
const res = JSON.parse(result.data);
if (res.rtnCode === '0000') {
resolve(res.body);
} else {
reject(new Error(res.msg));
}
} catch (error) {
reject(new Error('Failed to parse response'));
}
},
fail(error) {
console.log('error', error);
reject(new Error('Upload failed'));
}
});
});
};
export const getPreviewFileUrl = (id) => {
return get_base_url() + `/traoss/tras3/show/${id}`
}
// 对象存储下载 /traoss/tras3/show/{文件ID}
export function downloadFile(id) {
const token = getToken();
const header = {
"Content-Type": "application/x-www-form-urlencoded; charset=UTF-8",
};
if (token)
header['summary'] = token
let baseUrl = get_base_url();
return new Promise((resolve, reject) => {
uni.downloadFile({
url: `${baseUrl}/traoss/tras3/show/${id}`,
header,
success(result) {
resolve(result)
},
fail(error) {
reject(new Error(error));
}
});
});
}
// 上传音频 通用
export const commonUploadVoiceFile = (file, url, data = {}, fileKey = 'voice') => {
const base_url = get_base_url(url);
const token = getToken();
const header = {
"Content-Type": "multipart/form-data; charset=UTF-8",
};
if (token)
header['summary'] = token
return new Promise((resolve, reject) => {
uni.uploadFile({
url: `${base_url}${url}`,
header: header,
filePath: file,
name: fileKey,
formData: data,
timeout: 60000,
success(result) {
try {
console.log(result)
const res = JSON.parse(result.data);
if (res.rtnCode === '0000') {
return resolve(res);
} else if (REQUES_ERROR_CODES['NO_LOGIN'].includes(res.rtnCode)) {
goto_login_fun()
return reject('NO_LOGIN');
} else if (res.rtnCode === '0004') {
return reject('N');
}
return reject('Other');
} catch (error) {
reject(new Error('Failed to parse response'));
}
},
fail(error) {
return reject('Other');
}
});
});
};
// 上传音频 翻译文本
export const uploadVoiceFile = (file) => {
const url = '/trastudy/traStdyInfo/audioTranscriptions';
const base_url = get_base_url(url);
const token = getToken();
const header = {
"Content-Type": "multipart/form-data; charset=UTF-8",
};
if (token)
header['summary'] = token
return new Promise((resolve, reject) => {
uni.uploadFile({
url: `${base_url}${url}`,
header: header,
filePath: file,
name: 'voice',
success(result) {
try {
const res = JSON.parse(result.data);
if (res.rtnCode === '0000') {
resolve(result);
} else {
reject(new Error(res.message));
}
} catch (error) {
reject(new Error('Failed to parse response'));
}
},
fail(error) {
reject(new Error('Upload failed'));
}
});
});
};
// 获取音频 通用
export function downloadVoiceFile(url) {
const token = getToken();
const header = {
"Content-Type": "application/x-www-form-urlencoded; charset=UTF-8",
};
if (token)
header['summary'] = token
let baseUrl = get_base_url(url);
return new Promise((resolve, reject) => {
uni.downloadFile({
url: `${baseUrl}${url}`,
header,
timeout: 180000,
success(result) {
console.log(result)
resolve(result)
},
fail(error) {
console.log(error)
reject(new Error(error));
}
});
});
}
// 获取音频 通用参数
export function downloadVoiceParams(url) {
const token = getToken();
const header = {
"Content-Type": "application/x-www-form-urlencoded; charset=UTF-8",
};
if (token)
header['summary'] = token
let baseUrl = get_base_url(url);
return {
url: `${baseUrl}${url}`,
header,
timeout: 180000,
}
}
export function getAudioByText(data) {
const token = getToken();
const header = {
"Content-Type": "application/x-www-form-urlencoded; charset=UTF-8",
};
if (token)
header['summary'] = token
let baseUrl = get_base_url();
return new Promise((resolve, reject) => {
uni.request({
url: `${baseUrl}/trastudy/traStdyInfo/readContent`,
header,
method: "POST",
data,
responseType: "arraybuffer",
success(result) {
if (result.statusCode === 200) {
resolve(result.data)
} else {
reject(new Error(result))
}
},
fail(error) {
reject(new Error(error));
}
});
});
}
// 对象存储预览 /traoss/tras3/show/{文件ID}
export function previewFile(id, responseType = 'arraybuffer') {
const token = getToken();
const header = {
"Content-Type": "application/x-www-form-urlencoded; charset=UTF-8",
};
if (token)
header['summary'] = token
let baseUrl = get_base_url();
return new Promise((resolve, reject) => {
uni.request({
url: `${baseUrl}/traoss/tras3/show/${id}`,
header,
responseType,
success(result) {
resolve(result)
},
fail(error) {
reject(new Error(error));
}
});
});
}
// 缩略图预览 /traoss/tras3/getCacheThumbImage/{fileId}/{width}/{height}
export function previewCacheThumbImage(id, width = 30, height = 30) {
const token = getToken();
const header = {
"Content-Type": "application/x-www-form-urlencoded; charset=UTF-8",
};
if (token)
header['summary'] = token
let baseUrl = get_base_url();
return new Promise((resolve, reject) => {
uni.request({
url: `${baseUrl}/traoss/tras3/getCacheThumbImage/${id}/${width}/${height}`,
header,
success(result) {
if (result.data.rtnCode === '0000') {
let base64 = `data:image/png;base64,` + result.data.body
resolve(base64)
} else {
reject({
error: result.data.message
})
}
},
fail(error) {
reject(new Error(error));
}
});
});
}
+85
View File
@@ -0,0 +1,85 @@
import request from '@/api/request'
const module_url = '/traexam'
// 查询竞赛列表
export const queryTraCompetitionInfoPaging = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryTraCompetitionInfoPaging',
method: 'post',
data
});
};
// 查询试卷
export const queryTestPapers = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryTestPapers',
method: 'post',
data
});
};
// 查询PK人员+试题
export const queryTraCompetitionInfoByPkUser = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryTraCompetitionInfoByPkUser',
method: 'post',
data
});
};
// 循环获取简答题结果
export const queryCompetitionPracticeAnswerResult = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryCompetitionPracticeAnswerResult',
method: 'post',
data
});
};
// 答题结束
export const competitionAnswerEnd = (data) => {
return request({
url: module_url + '/traCompetitionInfo/competitionAnswerEnd',
method: 'post',
data,
suppressErrors: true // 屏蔽报错信息
});
};
// 竞赛答题
export const competitionAnswer = (data) => {
return request({
url: module_url + '/traCompetitionInfo/competitionAnswer',
method: 'post',
data
});
};
// 查询答题结果
export const queryCompetitionAnswerResult = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryCompetitionAnswerResult',
method: 'post',
data
});
};
// 竞赛排行榜
export const queryTraCompetitionRankingList = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryTraCompetitionRankingList',
method: 'post',
data
});
};
// 全行的机构数据
export const queryOrgTree = (data) => {
return request({
url: module_url + '/traCompetitionInfo/queryOrgTree',
method: 'post',
data
});
};
+72
View File
@@ -0,0 +1,72 @@
import request from '@/api/request'
const base_url = '/traapp'
const course_base_url = '/trastudy'
// 获取课程分类树接口
export const getCourseTypeListApi = (data) => {
return request({
url: base_url + '/app/traCrsCatalogInfo/courseQueryTraCrsCatalogInfoList',
method: 'post',
data
});
};
// 获取生产线分类树接口
export const getPrdLineTypeListApi = (data) => {
return request({
url: base_url + '/app/traPrdLineInfo/courseQueryTraPrdLineInfoList',
method: 'post',
data
});
};
// 获取标签分类树 接口
export const getTraTagListApi = (data) => {
return request({
url: base_url + '/app/traTagCatalogInfo/queryTraTagCatalogInfoList',
method: 'post',
data
});
};
// 全部模块
export const queryCrsInfoByCatalogApi = (data) => {
return request({
url: course_base_url + '/crsInfo/queryCrsInfoByCatalog',
method: 'post',
data
});
};
// 生产线模块-
export const queryCrsInfoByPrdLineApi = (data) => {
return request({
url: course_base_url + '/crsInfo/queryCrsInfoByPrdLine',
method: 'post',
data
});
};
// 最新
export const queryCrsInfoByLastnewApi = (data) => {
return request({
url: course_base_url + '/crsInfo/queryCrsInfoByLastnew',
method: 'post',
data
});
};
// 最热
export const queryCrsInfoByPopularApi = (data) => {
return request({
url: course_base_url + '/crsInfo/queryCrsInfoByPopular',
method: 'post',
data
});
};
// 推荐
export const queryCrsInfoByRecmdApi = (data) => {
return request({
url: course_base_url + '/crsInfo/queryCrsInfoByRecmd',
method: 'post',
data
});
};
+72
View File
@@ -0,0 +1,72 @@
import request from '@/api/request'
const module_url = '/traapp'
// AITS-A-4006--APP课程详情查询-含知识点学习状态数据
export const getCourseDetailApi = (data) => {
return request({
url: module_url + '/crsInfo/queryCrsInfoDetail',
method: 'post',
data
});
};
// AITS-A-4007--APP课程评价列表数据查询(20250723预计废弃)
export const qryTraCrsBbsDataByCrsNum = (data) => {
return request({
url: module_url + '/traCrsBbsInfo/qryTraCrsBbsDataByCrsNum',
method: 'post',
data
});
};
//王洋新写
export const queryTraCrsBbsInfoByCrsId = (data) => {
return request({
url: module_url + '/traCrsBbsInfo/queryTraCrsBbsInfoByCrsId',
method: 'post',
data
});
};
// AITS-A-4008--APP课程评价下回复评论时刷新二级评论列表数据接口
export const qryChildTraCrsBbsInfo = (data) => {
return request({
url: module_url + '/traCrsBbsInfo/qryChildTraCrsBbsInfo',
method: 'post',
data
});
};
// AITS-A-4009--APP课程评价回复评论接口
export const replyTraCrsBbsInfo = (data) => {
return request({
url: module_url + '/traCrsBbsInfo/replyTraCrsBbsInfo',
method: 'post',
data
});
};
// AITS-A-4010--APP课程评价接口
export const publishCrsEvaluation = (data) => {
return request({
url: module_url + '/traCrsBbsInfo/publishCrsEvaluation',
method: 'post',
data
});
};
// 考试排行榜/traexam/traCrsPapers/queryTraCrsRankingList
export const queryTraCrsRankingList = (data) => {
return request({
url: '/traexam/traCrsPapers/queryTraCrsRankingList',
method: 'post',
data
});
};
+124
View File
@@ -0,0 +1,124 @@
import request from '@/api/request'
const base_url = '/traexam'
/**学习*/
// 学习列表页
export const queryStdyRecordPaging = (data) => {
return request({
url: base_url + '/traRecord/queryStdyRecordPaging',
method: 'post',
toastErrors: true,
data
});
};
// 学习单项列表 crsId
export const queryStdyBatchByCrsId = (data) => {
return request({
url: base_url + '/traRecord/queryStdyBatchByCrsId',
method: 'post',
toastErrors: true,
data
});
};
// 学习详情 通过 stdyId 查询
export const queryTraStudyExecuteRecord = (data) => {
return request({
url: base_url + '/traCrsPapers/queryTraStudyExecuteRecord',
method: 'post',
toastErrors: true,
data
});
};
/**练习*/
// 练习列表页
export const queryPracticeRecordPaging = (data) => {
return request({
url: base_url + '/traRecord/queryPracticeRecordPaging',
method: 'post',
toastErrors: true,
data
});
};
// 练习每一项列表
export const queryPracticeHis = (data) => {
return request({
url: base_url + '/traRecord/queryPracticeHis',
method: 'post',
toastErrors: true,
data
});
};
// 练习详情 通过 exrId 查询
export const queryPracticeRecordByExrId = (data) => {
return request({
url: base_url + '/traCrsPractice/queryPracticeRecordByExrId',
method: 'post',
toastErrors: true,
data
});
};
/**考试*/
export const queryCrsExamRecordPaging = (data) => {
return request({
url: base_url + '/traRecord/queryCrsExamRecordPaging',
method: 'post',
toastErrors: true,
data
});
};
// 考试每一项列表
export const queryCrsExamHisByExamId = (data) => {
return request({
url: base_url + '/traRecord/queryCrsExamHisByExamId',
method: 'post',
toastErrors: true,
data
});
};
/**考试中心考试*/
// 考试列表页
export const queryExamRecordPaging = (data) => {
return request({
url: base_url + '/traRecord/queryExamRecordPaging',
method: 'post',
toastErrors: true,
data
});
};
// 考试每一项列表
export const queryExamHisByExamId = (data) => {
return request({
url: base_url + '/traRecord/queryExamHisByExamId',
method: 'post',
toastErrors: true,
data
});
};
/**竞赛*/
// 查询竞赛记录
export const queryTraCompetitionRecordPaging = (data) => {
return request({
url: base_url + '/traCompetitionInfo/queryTraCompetitionRecordPaging',
method: 'post',
toastErrors: true,
data
});
};
// 查询竞赛答题列表
export const queryTraCompetitionRecordInfoPaging = (data) => {
return request({
url: base_url + '/traCompetitionInfo/queryTraCompetitionRecordInfoPaging',
method: 'post',
toastErrors: true,
data
});
};
+64
View File
@@ -0,0 +1,64 @@
import request from '@/api/request'
// 获取问题list
export const queryHotQuestionApi = (data) => {
return request({
// url: '/traask/traAskchat/queryHotQuestion',
url: '/traask/traAskchat/queryHotQuestions',
method: 'post',
data
});
};
// 获取问答历史
export const queryAskBatchSessionPaging = (data) => {
return request({
url: '/traask/traAskchatSession/queryAskBatchSessionPaging',
method: 'post',
data
});
};
// 问答详情 /traAskchatSession/queryAskHisByBatchIdPaging
export const queryAskHisByBatchIdPaging = (data) => {
return request({
url: '/traask/traAskchatSession/queryAskHisByBatchId',
method: 'post',
data
});
};
// 问答反馈
export const insertTraAskFeedback = (data) => {
return request({
url: '/traask/traAskFeedback/insertTraAskFeedback',
method: 'post',
data
});
};
// 取消反馈 取消踩 /traAskFeedback/updateDownvoteStat
export const updateDownvoteStat = (data) => {
return request({
url: '/traask/traAskFeedback/updateDownvoteStat',
method: 'post',
data
});
};
// 赞 /取消赞/traAskFeedback/updateDownvoteStat
export const updateLikeStat = (data) => {
return request({
url: '/traask/traAskFeedback/updateLikeStat',
method: 'post',
data
});
};
// example /traask/traTestchat/chatVoice
export const chatVoice = (data) => {
return request({
url: '/traask/traTestchat/chatVoice',
method: 'post',
timeout: 120000,
data
});
};
+90
View File
@@ -0,0 +1,90 @@
import CryptoJS from 'crypto-js'// 引入AES加密库 npm install crypto-js@3.1.9-1 -save
// import { JSEncrypt } from 'jsencrypt'; // npm install jsencrypt -save
const RsaPublicKey = "MIGfMA0GCSqGSIb3DQEBAQUAA4GNADCBiQKBgQCZf8e1tzYTSr7KciAHH2KrE3O13ftuv3rHRRz7MztYXirMdLquocTNSJ5BAoj1H3V8fvFtfmN6BkuB6XnQrbY5heOzZZQWvO4qAa0EktwkLQUwCkYJfaQNc1Iw4wZFt4VO7U8+LLo+IO7jtxmaNIBqJm8laDc1C6zw2ETZD4bQ6wIDAQAB"
const iv = CryptoJS.enc.Utf8.parse('1234567812345678');
//定义rsa加密类
// const crypt = new JSEncrypt();
// crypt.setPublicKey(RsaPublicKey); // 设置公钥
// 生成一个随机的128位(16字节)AES密钥
const generateAESKey = () => {
return CryptoJS.lib.WordArray.random(16);
}
// 转化生成AES密钥成base64字符串
const getGenerateAESKeyBase64 = (aesKey) => aesKey.toString(CryptoJS.enc.Base64)
// 执行rsa加密
const encryptRSA = (str) => crypt.encrypt(str)
// 执行rsa解密
const decryptRSA = (str) => crypt.decrypt(str)
// 执行aes加密
export const encryptAES = (data, aesKeyBase64, iv) => {
const key = CryptoJS.enc.Utf8.parse(aesKeyBase64);
return CryptoJS.AES.encrypt(JSON.stringify(data), key, {
iv: iv,
mode: CryptoJS.mode.CBC,
padding: CryptoJS.pad.Pkcs7
}).toString();
}
// aes解密
const decryptAES = (ciphertext, aesKeyBase64, iv) => {
const key = CryptoJS.enc.Utf8.parse(aesKeyBase64);
const bytes = CryptoJS.AES.decrypt(ciphertext, key, {
iv: iv,
mode: CryptoJS.mode.CBC,
padding: CryptoJS.pad.Pkcs7
});
return JSON.parse(bytes.toString(CryptoJS.enc.Utf8));
}
// 加密过程
export const encryptData = (data={}) => {
const aesKey = generateAESKey(); // 获取随机密钥key
const aesKeyBase64 = getGenerateAESKeyBase64(aesKey) // 变成base64
const key = encryptRSA(aesKeyBase64)
const ciphertext = encryptAES(data, aesKeyBase64, iv);
const decryptedData = decryptAES(ciphertext, aesKeyBase64, iv);
return {text:ciphertext, key:key, aesKeyBase64:aesKeyBase64}
}
// 解密过程
export const decryptData= (aesStr, aesKey) => {
return decryptAES(aesStr, aesKey, iv)
}
/**
* AES加密
* @param plainText 明文
* @param keyInBase64Str base64编码后的key
* @returns {string} base64编码后的密文
*/
export function encryptByAES(plainText, keyInBase64Str) {
let key = CryptoJS.enc.Base64.parse(keyInBase64Str);
let encrypted = CryptoJS.AES.encrypt(plainText, key, {
mode: CryptoJS.mode.ECB,
padding: CryptoJS.pad.Pkcs7,
});
// 这里的encrypted不是字符串,而是一个CipherParams对象
return encrypted.ciphertext.toString(CryptoJS.enc.Base64);
}
/**
* AES解密
* @param cipherText 密文
* @param keyInBase64Str base64编码后的key
* @return 明文
*/
export function decryptByAES(cipherText, keyInBase64Str) {
let key = CryptoJS.enc.Base64.parse(keyInBase64Str);
// 返回的是一个Word Array Object,其实就是Java里的字节数组
let decrypted = CryptoJS.AES.decrypt(cipherText, key, {
mode: CryptoJS.mode.ECB,
padding: CryptoJS.pad.Pkcs7,
});
return decrypted.toString(CryptoJS.enc.Utf8);
}
+110
View File
@@ -0,0 +1,110 @@
import request from '@/api/request'
import { getToken } from '@/common/common.js'
import { get_base_url } from '@/api/request'
const base_url = '/traexam'
// 开始考试(试卷信息)
export const startTraExamPapersInfo = (data) => {
return request({
url: base_url + '/traExamPapers/startTraExamPapersInfo',
method: 'post',
toastErrors: true,
data
});
};
// 提交试卷
export const commitPaperExam = (data) => {
return request({
url: base_url + '/traExamPapers/commitPaperExam',
method: 'post',
toastErrors: true,
data
});
};
// 考试中心考试查分
export const queryPaperExamScore = (data) => {
return request({
url: base_url + '/traExamPapers/queryPaperExamScore',
method: 'post',
toastErrors: true,
data
});
};
// 根据试卷id查询详情
export const queryTraExamPapersByPapersId = (data) => {
return request({
url: base_url + '/traExamPapers/queryTraExamPapersByPapersId',
method: 'post',
toastErrors: true,
data
});
};
// ↑↑↑↑↑↑↑↑↑考试中心考试↑↑↑↑↑↑↑
// ↓↓↓↓↓↓↓↓↓课程考试↓↓↓↓↓↓↓↓
export const startExamByCrs = (data) => {
return request({
url: base_url + '/traCrsPapers/startExamByCrs',
method: 'post',
toastErrors: true,
data
});
};
// 提交课程考试
export const commitCrsExam = (data) => {
return request({
url: base_url + '/traCrsPapers/commitCrsExam',
method: 'post',
toastErrors: true,
data
});
};
// 试卷考试查分
export const queryCrsExamScore = (data) => {
return request({
url: base_url + '/traCrsPapers/queryCrsExamScore',
method: 'post',
toastErrors: true,
data
});
};
// 试卷考试 排行榜
export const queryPaperExamRankingList = (data) => {
return request({
url: base_url + '/traExamPapers/queryPaperExamRankingList',
method: 'post',
toastErrors: true,
data
});
};
// 反馈问题
export const saveTraQnsFeedback = (data) => {
return request({
url: base_url + '/traCrsPapers/saveTraQnsFeedback',
method: 'post',
toastErrors: true,
data
});
};
// 考试过程中提交结果,后端用来缓存(考试中心考试和课程考试共用这个缓存接口)
export const commitPaperCache = (data) => {
return request({
url: base_url + '/traExamPapers/commitPaperCache',
method: 'post',
toastErrors: true,
data
});
};
+49
View File
@@ -0,0 +1,49 @@
import request from '@/api/request'
const base_url = '/trastudy'
// 课程收藏分页查询
export const queryTraPersonalCollectionCrsPaging = (data) => {
return request({
url: base_url + '/traPersonalCollection/queryTraPersonalCollectionCrsPaging',
method: 'post',
toastErrors: true,
data
});
};
// 任务收藏分页查询
export const queryTraPersonalCollectionTaskPaging = (data) => {
return request({
url: base_url + '/traPersonalCollection/queryTraPersonalCollectionTaskPaging',
method: 'post',
toastErrors: true,
data
});
};
// 根据收藏ID删除收藏
export const deleteTraPersonalCollectionById = (data) => {
return request({
url: base_url + '/traPersonalCollection/deleteTraPersonalCollectionById',
method: 'post',
toastErrors: true,
data
});
};
/*
新增任务收藏
入参
collectTyp 收藏类型 必须 String 01--课程02--任务
collectBusId 收藏业务ID 必须 String
**/
export const insertTraPersonalCollectionTask = (data) => {
return request({
url: base_url + '/traPersonalCollection/insertTraPersonalCollectionTask',
method: 'post',
toastErrors: true,
data
});
};
+91
View File
@@ -0,0 +1,91 @@
import request from '@/api/request'
// 获取首页统计数据方法(查询个人学习统计信息 )
export const getStatisticsData = (data) => {
/**
userId 学员ID String
stdtTmLen 本年学习时长 double
accmPracticeCnt 累计训练次数 int
accmExamCnt 累计考试次数 int
**/
return request({
url: '/traapp/traStdyInfo/queryTraStdyStatisticsPerson',
method: 'post',
data
});
};
// 最新课程查询接口
export const queryCrsInfoByLastnew = (data) => {
/**
crsNum 课程编号 String
crsName 课程名称 String
tagId 课程分类 String
prdLineId 生产线分类 String
issuTm 发布日期 String
imgId 课程封面ID String
accmStdyCnt 累计学习次数 int
evalGrade 课程评分 Bigdecimal
tagCataList 课程标签集 List<Object>
tagId 标签ID String
tagName 标签名称 String
**/
return request({
url: '/traapp/mycrs/queryCrsInfoByLastnew',
method: 'post',
data
});
};
// 热门课程查询接口
export const queryCrsInfoByPopular = (data) => {
/**
tagId 标签ID(可选填)为空时查全部 必输项:false 类型:String
出参:
tagId 标签ID(可选填)为空时查全部 必输项:false 类型:String
出参:
crsId 课程ID StringcrsNum 课程编号 String
crsName 课程名称 String
tagId 课程分类 String
prdLineId 生产线分类 String
issuTm 发布日期 String
imgId 课程封面ID String
accmStdyCnt 累计学习次数 int
evalGrade 课程评分 Bigdecimal
tagCataList 课程标签集 List<Object>
tagId 标签ID String
tagName 标签名称 String
**/
return request({
url: '/traapp/mycrs/queryCrsInfoByPopular',
method: 'post',
data
});
};
// 推荐课程查询接口
export const queryCrsInfoByRecmd = (data) => {
/**
crsNum 课程编号 String
crsName 课程名称 String
tagId 课程分类 String
prdLineId 生产线分类 String
issuTm 发布日期 String
imgId 课程封面ID String
accmStdyCnt 累计学习次数 int
evalGrade 课程评分 Bigdecimal
tagCataList 课程标签集 List<Object>
tagId 标签ID String
tagName 标签名称 String
**/
return request({
url: '/traapp/mycrs/queryCrsInfoByRecmd',
method: 'post',
data
});
};
+38
View File
@@ -0,0 +1,38 @@
import request from '@/api/request'
const base_url = '/traapp'
// 学习时长
export const getLeaderBoardStdyTmApi = (data) => {
return request({
url: base_url + '/myLeaderBoard/stdyTmLen',
method: 'post',
data
});
};
// 练习次数
export const getLeaderBoardPractCntApi = (data) => {
return request({
url: base_url + '/myLeaderBoard/practCnt',
method: 'post',
data
});
};
// 考试次数
export const getLeaderBoardExamCntApi = (data) => {
return request({
url: base_url + '/myLeaderBoard/examCnt',
method: 'post',
data
});
};
// 通关课程数
export const getLeaderBoardCrsCntApi = (data) => {
return request({
url: base_url + '/myLeaderBoard/crsCnt',
method: 'post',
data
});
};
+31
View File
@@ -0,0 +1,31 @@
import request from '@/api/request'
const base_url = '/trastudy'
// 查询所有、未完成、已完成的任务(分页)
export const queryAllTraTaskInfoByUserId = (data) => {
return request({
url: base_url + '/traTaskInfo/queryAllTraTaskInfoByUserId',
method: 'post',
data
});
};
// 查询当前用户的课程完成情况
export const queryAllTraTaskCrsInfo = (data) => {
return request({
url: base_url + '/traTaskInfo/queryAllTraTaskCrsInfo',
method: 'post',
data
});
};
// 根据任务id查询这一条
export const queryTraTaskInfoByTaskId = (data) => {
return request({
url: base_url + '/traTaskInfo/queryTraTaskInfoByTaskId',
method: 'post',
data
});
};
+239
View File
@@ -0,0 +1,239 @@
import request from '@/api/request'
import common from "@/common/common";
import {
queryCurrentValidMascotInfo
} from "@/api/mascot.js"
import {
queryCurrentStatus
} from '@/api/pointsAndRank.js';
import {
setToken,
setUserInfo,
getUserInfo,
getToken
} from "@/common/common";
import {
get_base_url
} from '@/api/request.js'
// 账号登录接口
export const useAccountLoginApp = (data) => request({
url: '/traapp/access/userLogin',
method: 'post',
data
}).then(async ({
rtnCode,
body
}) => {
if (rtnCode === '0000') {
setToken(body)
const res = await Promise.all([queryCurrentUser(), queryCurrentValidMascotInfo()
// queryCurrentStatus()
])
const userInfo = getUserInfo()
const mascotInfo = res[1]
setUserInfo({
...userInfo,
'mascotInfo': mascotInfo
})
// 积分状态
// const pointsStatus = res[2].body
const pointsStatus = {}
common.setValue('points_data', {
badgeNum: pointsStatus['badgeNum'] ?? 0, // 徽章数
currentPoints: pointsStatus['currentPoints'] ?? 0, // 当前积分
currentRankAddr: pointsStatus['ossAddr'] ?? '', // 当前段位图片
currentRankName: pointsStatus['currentRankName'] ?? '未知', // 当前段位名称
})
return {
mascotState: !mascotInfo.mascotId
}
}
return Promise.reject({
rtnCode,
body
})
}).catch(res => {
return Promise.reject(res)
})
/**
* 获取个人信息
* 入参
gender f是女 m男
imageAddr 图片
loginName ID
userName 姓名
*/
export const queryCurrentUser = (data) => request({
url: '/traapp/access/queryCurrentUser',
method: 'post',
data
}).then(res => {
// deptId: ""
// deptName: null
// loginName: "ADMIN" // todo id
// orgId: "00E200"
// orgName: "金融科技部"
// userId: "USER001"
// userName: "ADMIN" / todo name
// userStatus: "01"
const userInfo = getUserInfo() ?? {}
setUserInfo({
...userInfo,
...res.body
})
// console.log('预加载头像');
const base_url = get_base_url()
const token = getToken()
// 预加载头像
// #ifdef APP-PLUS
// 更新用户信息中的头像路径
const updateUserAvatar = (imagePath) => {
const _userInfo = getUserInfo() ?? {}
setUserInfo({
..._userInfo,
...res.body,
'imagePath': imagePath
})
}
// 确保目录存在
const ensureDirectoryExists = (dirPath) => {
return new Promise((resolve) => {
plus.io.requestFileSystem(plus.io.PRIVATE_DOC, (fs) => {
fs.root.getDirectory(dirPath, {
create: true
}, resolve, resolve)
})
})
}
// 生成唯一文件名
const generateUniqueFilename = (pathName) => {
if (!pathName || typeof pathName !== 'string') {
throw new Error('无效的路径');
}
// 直接使用原始文件名(带扩展名)
const filename = pathName.split('/').pop();
if (!filename) {
throw new Error('无法从路径中提取文件名');
}
// 确保文件扩展名为.png (可选,根据实际需求)
const ext = filename.split('.').pop().toLowerCase();
if (ext === 'png' || ext === 'jpg' || ext === 'jpeg' || ext === 'webp') {
return filename;
}
// 如果没有有效扩展名,添加.png
return `${filename}.png`;
};
// 下载并保存头像
const downloadAndSaveAvatar = async (fileUrl, localFilePath) => {
try {
await ensureDirectoryExists('avatar')
return new Promise((resolve, reject) => {
console.log('开始下载头像:', fileUrl)
const dtask = plus.downloader.createDownload(fileUrl, {
filename: localFilePath
}, (d, status) => {
if (status === 200) {
console.log("保存路径:", d.filename)
resolve(d.filename)
} else {
console.error("文件下载失败:", status)
reject(new Error(`下载失败,状态码: ${status}`))
}
})
dtask.start()
})
} catch (error) {
console.error('创建目录失败:', error)
throw error
}
}
// 检查文件是否存在
const checkFileExists = (filePath) => {
return new Promise((resolve) => {
uni.getFileInfo({
filePath,
success: () => resolve(true),
fail: () => resolve(false)
})
})
}
// 主逻辑
(async () => {
try {
const pathName = res.body.imageAddr
if (pathName) {
const fileUrl = base_url + pathName + '?tk=' + token
console.log('fileUrl', fileUrl);
// 生成保存路径
const filename = generateUniqueFilename(pathName)
const localFilePath = '_doc/avatar/' + filename
// 检查文件是否存在
const exists = await checkFileExists(localFilePath)
if (exists) {
// console.log('头像已存在,无需下载')
updateUserAvatar(localFilePath)
} else {
const savedPath = await downloadAndSaveAvatar(fileUrl, localFilePath)
updateUserAvatar(savedPath)
}
}
} catch (error) {
console.error('头像预加载过程发生错误:', error)
}
})()
// #endif
return res.body
})
//获取隐私协议
export const getYsxy = (data) => request({
url: '/traapp/dfAgtInfo/queryValidDfAgt',
method: 'post',
data
})
//获取服务条款
export const getFutk = (data) => request({
url: '/traapp/dfAgtInfo/queryValidDfAgtTerm',
method: 'post',
data
})
// 获取验证码(作废)
export const getLoginAppCheckCode = (data) => request({
url: '/traapp/access/checkCode',
method: 'get',
data
})
// 退出登录
export const loginOut = (data) => request({
url: '/traapp/access/loginOut',
method: 'post',
data
}).then(() => {
setToken(null)
setUserInfo(null)
common.setValue('me_statistics_data', null) // 我的页统计数据缓存
common.setValue('index_statistics_data', null) // 首页统计数据缓存
common.setValue('points_data', null) // 首页统计数据缓存
})
+112
View File
@@ -0,0 +1,112 @@
import request from '@/api/request'
import {
get_base_url
} from '@/api/request.js'
import { getToken } from '@/common/common.js'
const app_url = '/traapp'
/**
* 查询吉祥物列表
* 出参:
mascotId 吉祥物ID String
mascotName 吉祥物名称 String
mascotNameEn 吉祥物英文名称 String
mascotDesc 吉祥物简介 String
mascotNo 吉祥物编号 String
imgAddr 吉祥物图片 String
imgThumbAddr 吉祥物缩略图 String
showOrder 显示顺序 int
*/
export const queryValidTraMascotInfo = (data) => {
return request({
url: app_url + '/traMascotInfo/queryValidTraMascotInfo',
method: 'post',
toastErrors: true,
data
});
};
/**
* 设置吉祥物
* 入参:
mascotId 必输项:true 类型:String
*/
export const setTraMascotInfo = (data) => {
return request({
url: app_url + '/traMascotInfo/setTraMascotInfo',
method: 'post',
toastErrors: true,
data
});
};
// 查询当前已经选择的吉祥物
// 检查文件是否存在
const checkFileExists = (filePath) => {
return new Promise((resolve) => {
uni.getFileInfo({
filePath,
success: () => resolve(true),
fail: () => resolve(false)
})
})
}
/**
* 查询当前已经选择的吉祥物
* 出参:
mascotId 吉祥物ID String
mascotName 吉祥物名称 String
mascotNameEn 吉祥物英文名称 String
mascotDesc 吉祥物简介 String
mascotNo 吉祥物编号 String
imgAddr 吉祥物图片 String
imgThumbAddr 吉祥物缩略图 String
showOrder 显示顺序 int
*/
export const queryCurrentValidMascotInfo = (data) => {
return request({
url: app_url + '/traMascotInfo/queryCurrentValidMascotInfo',
method: 'post',
toastErrors: true,
data
}).then(({
body
}) => {
if (body.imgAddr) {
console.log('body.imgAddrmascotNo');
const base_url = get_base_url()
const token = getToken()
const pathName = body.imgAddr
const fileUrl = base_url + pathName + '?tk=' + token
const localFilePath = '_doc/mascot/' + body.mascotNo + '.png'
console.log('fileUrl', fileUrl);
// 将 Promise 回调函数声明为 async
return new Promise(async (resolve, reject) => {
// 现在可以在这里使用 await 了
const exists = await checkFileExists(localFilePath)
if (exists) {
// console.log('已存在,无需下载吉祥物')
resolve({...body, imgAddr:localFilePath})
} else {
// console.log('开始下载吉祥物:', fileUrl)
const dtask = plus.downloader.createDownload(fileUrl, {
filename: localFilePath
}, (d, status) => {
if (status === 200) {
console.log("保存路径:", d.filename)
resolve({...body, imgAddr:d.filename})
} else {
console.error("文件下载失败:", status)
reject(new Error(`下载失败,状态码: ${status}`))
}
})
dtask.start()
}
})
}
return body
})
};
+62
View File
@@ -0,0 +1,62 @@
import request from '@/api/request'
const base_url = '/trastudy'
// 查询当前用户的所有消息通知 (分页)
export const queryTraPersonalMessageNoticePaging = (data) => {
return request({
url: base_url + '/traPersonalMessageNotice/queryTraPersonalMessageNoticePaging',
method: 'post',
toastErrors: true,
data
});
};
// 读取单一消息
export const readNotice = (data) => {
return request({
url: base_url + '/traPersonalMessageNotice/readNotice',
method: 'post',
toastErrors: true,
data
});
};
// 一键读取所有消息(不展示)
export const oneTimeRead = (data) => {
return request({
url: base_url + '/traPersonalMessageNotice/oneTimeRead',
method: 'post',
toastErrors: true,
data
});
};
export const noticeMessageCount = (data) => {
return request({
url: base_url + '/traPersonalMessageNotice/noticeMessageCount',
method: 'post',
data
});
};
export const queryMessageNoticeClass = (data) => {
return request({
url: base_url + '/traPersonalMessageNotice/queryMessageNoticeClass',
method: 'post',
data
});
};
// 按照分类一键读取消息
export const oneTimeReadByNoticeTyp = (data) => {
return request({
url: base_url + '/traPersonalMessageNotice/oneTimeReadByNoticeTyp',
method: 'post',
data
});
};
+129
View File
@@ -0,0 +1,129 @@
import request from '@/api/request'
const base_url = '/traapp'
/**
* 查询签到
* userId 学员ID String
* checkInToday 今日是否签到 String
* days 连续签到天数 int
*/
export const queryCheckIn = (data) => {
return request({
url: base_url + '/traPointsInfo/queryCheckIn',
method: 'post',
data
});
};
/**
* 点击签到
*/
export const addPointsByCheckIn = (data) => {
return request({
url: base_url + '/traPointsInfo/addPointsByCheckIn',
method: 'post',
data
});
};
/**
* 每次完成任务获得弹窗奖励
* 入参:
taskId 必输项:true 类型:String
type 必输项:true 类型:String
出参:
userId 学员ID String
badgeName 徽章名称 String
badgeImg 徽章图片 String
num 次数 int
name 描述 String
pointsBadgeName 积分徽章名称 String
pointsBadgeImg 徽章图片 String
currentPoints 当前积分 int
changePoints 本次获得积分 int
*/
export const addPointsByTask = (data) => {
// base_url
return request({
url: base_url + '/traPointsInfo/addPointsByTask',
method: 'post',
data
});
// return request({
// url: '/test/traPointsInfo/addPointsByTask',
// method: 'post',
// data
// });
};
/**
* 查询积分明细
* 出参:
detailId 明细ID String
userId 学员ID String
mattr 积分来源事项 String
mattrDesc 积分来源描述 String
changeValue 积分变化值 String
currentPoints 当前积分 int
countPoints 累计积分 int
ctTime 创建时间 String
*/
export const queryPointsDetail = (data) => {
return request({
url: base_url + '/traPointsInfo/queryPointsDetail',
method: 'post',
data
});
};
/**
* 查询段位明细
* 出参:
rankId 段位ID String
rankName 段位名称 String
currentRankName 当前段位名称 String
ossAddr 段位图片 String
pointsStart 段位开始区间 int
pointsEnd 段位结束区间 int
countPoints 累计积分 int
*/
export const queryRankDetail = (data) => {
return request({
url: base_url + '/traPointsInfo/queryRankDetail',
method: 'post',
data
});
};
/**
* 查询当前状态
* 出参:
userId 学员ID String
orgId 机构ID String
countPoints 累计积分 int
currentPoints 当前积分 int
currentRankName 当前段位 String
badgeNum 徽章数 int
*/
export const queryCurrentStatus = (data) => {
return request({
url: base_url + '/traPointsInfo/queryCurrentStatus',
method: 'post',
data
});
};
/**
* 徽章墙
*/
export const queryRankWall = (data) => {
return request({
url: base_url + '/traPointsInfo/queryRankWall',
method: 'post',
data
});
};
+20
View File
@@ -0,0 +1,20 @@
import request from '@/api/request'
const base_url = '/traapp'
// String actId,活动id
export const queryPointsActivityDetail = (data) => {
return request({
url: base_url + '/traPointsActivity/queryPointsActivityDetail',
method: 'post',
data
});
};
export const addPointsExchange = (data) => {
return request({
url: base_url + '/traPointsActivity/addPointsExchange',
method: 'post',
data
});
};
+49
View File
@@ -0,0 +1,49 @@
import request from '@/api/request'
const base_url = '/traexam'
// 1.开始练习(课程练习)
export const startPracticeByCrs = (data) => {
return request({
url: base_url + '/traCrsPractice/startPracticeByCrs',
method: 'post',
data
});
};
// 2.课程练习答题
export const practiceAnswer = (data) => {
return request({
url: base_url + '/traCrsPractice/practiceAnswer',
method: 'post',
data
});
};
// 3.课程练习状态变更
export const practiceCommit = (data) => {
return request({
url: base_url + '/traCrsPractice/practiceCommit',
method: 'post',
suppressErrors:true,
data
});
};
// 4.课程练习记录
export const queryPracticeRecordByExrId = (data) => {
return request({
url: base_url + '/traCrsPractice/queryPracticeRecordByExrId',
method: 'post',
data
});
};
// 5.练习时候查询问答题结果
export const queryCrsPracticeAnswerResult = (data) => {
return request({
url: base_url + '/traCrsPractice/queryCrsPracticeAnswerResult',
method: 'post',
toastErrors: true,
data
});
};
+20
View File
@@ -0,0 +1,20 @@
import request from '@/api/request'
const base_url = '/trastudy'
export const queryTraStdyInfoPreviewPaging = (data) => {
return request({
url: base_url + '/traStdyInfoPreview/queryTraStdyInfoPreviewPaging',
method: 'post',
toastErrors: true,
data
});
};
export const queryTraListPreviewPaging = (data) => {
return request({
url: '/trapractice/traPartnerInfo/queryTraPartnerPreviewPaging',
method: 'post',
toastErrors: true,
data
});
};
+154
View File
@@ -0,0 +1,154 @@
import {
getToken
} from '@/common/common.js'
import common from '@/common/common.js'
const ENV = import.meta.env
// function getBrowserInfo() {
// const userAgent = navigator.userAgent;
// let browserName = '未知浏览器';
// let version = '未知版本';
// // 检测 Chrome
// if (userAgent.indexOf('Chrome') > -1 && userAgent.indexOf('Edg') === -1) {
// browserName = 'Chrome';
// version = userAgent.match(/Chrome\/(\d+\.\d+)/)[1];
// }
// return {
// browser: browserName,
// version: version,
// userAgent: userAgent
// };
// }
export const get_base_url = (url = '') => {
if (['development', 'development.rd', 'development.yh', 'development.ty'].includes(ENV.MODE)) { // 开发环境
// #ifdef H5
if (url) {
let _str = (url.split('/')[0] || url.split('/')[1])?.toUpperCase()
return ENV[`VITE_APP_BASE_H5_API_Url_${_str}`] || ENV.VITE_APP_BASE_H5_API_Url || ENV
.VITE_APP_BASE_API_Url
} else {
return ENV.VITE_APP_BASE_H5_API_Url || ENV.VITE_APP_BASE_API_Url
}
// #endif
// #ifdef APP-PLUS
return ENV.VITE_APP_BASE_API_Url
// #endif
} else { // 其他环境,包括生产环境,sit环境,uat环境
return ENV.VITE_APP_BASE_API_Url
// return ENV.VITE_APP_BASE_API_Url
}
}
// 请求错误码
export const REQUES_ERROR_CODES = {
NO_LOGIN: ['QQ0005', '0005'] //登录
};
// 去登录弹窗状态
let no_login_show_modal_state = false
// 去登录方法
export const goto_login_fun = (from = 'http') => {
if (no_login_show_modal_state) return;
no_login_show_modal_state = true
let page_route = ''
const pages = getCurrentPages();
if (pages.length >= 1) {
const page = pages[pages.length - 1];
page_route = page.route
}
console.log('当前页路由', page_route);
if (page_route === 'pages/login/login') return;
common.hideLoading()
common.show('未登录', '现在去登录', false).then(() => {
// common.navigateTo('/pages/login/login')
if (from === 'socket' && page_route !== '') {
common.redirectTo('/pages/login/login?path=' + page_route)
} else {
common.redirectTo('/pages/login/login')
// common.navigateTo('/pages/login/login')
}
}).finally(() => {
no_login_show_modal_state = false
})
}
export default function request({
url,
method,
data,
meta,
isStream,
suppressErrors = false,
toastErrors = true,
header = {},
timeout = 6000,
baseUrl = ''
}) {
const base_url = get_base_url(url);
const headers_base = {
'content-type': 'application/x-www-form-urlencoded; charset=UTF-8',
...(header || {})
}
let token = getToken()
if (token) {
headers_base['summary'] = token // 让每个请求携带令牌
}
let that = this
return new Promise((resolve, reject) => {
uni.request({
url: `${baseUrl || base_url}${url}`,
data: data,
method: method,
header: headers_base,
timeout: timeout,
success: (response) => {
if (isStream) { // 流式传输,直接返回
return resolve(response);
}
// suppressErrors
// 模拟拦截器功能
const res = response.data
if (res?.rtnCode === '0000') {
return resolve(res)
} else if (res?.rtnCode === '1001') { // 为前后端约定不公共拦截的码值(小树20250903定)
return reject(res)
} else if (res?.rtnCode === '0002' && res?.message) {
if (toastErrors) {
common.msg(res.message)
}
return reject(response)
}
//到这里下面全是异常的处理
if (suppressErrors) { // 静默的接口,报错也不处理,因为下面有公共处理
return reject(response)
}
if (!res) {
common.msg("系统异常.")
return reject()
}
// 0005 QQ0005 登录已过期
if (REQUES_ERROR_CODES['NO_LOGIN'].includes(res.rtnCode)) {
goto_login_fun()
}
if (res?.rtnCode === '9999') {
common.msg("系统异常.")
}
return reject(res)
},
fail(res) {
console.log('请求错误', res, url, base_url);
return reject(res)
}
});
})
};
+12
View File
@@ -0,0 +1,12 @@
import request from '@/api/request'
const base_url = '/traapp'
// 搜索
export const appSearchCrsApi = (data) => {
return request({
url: base_url + '/traSearchHis/appSearchCrs',
method: 'post',
toastErrors: true,
data
});
};
+242
View File
@@ -0,0 +1,242 @@
import { goto_login_fun } from '@/api/request'
export default class WebSocketUtil {
/**
* WebSocket工具类用于管理WebSocket连接
* @param {string} url - WebSocket服务器地址
* @param {Object} options - 配置选项
* @param {number} [options.maxReconnectCount=5] - 最大重连次数
* @param {number} [options.reconnectInterval=3000] - 重连间隔时间(ms)
* @param {number} [options.heartbeatInterval=30000] - 心跳间隔时间(ms)
* @param {string|Object} [options.heartbeatMsg='ping'] - 心跳消息
*/
constructor(url, options = {}) {
this.url = url;
this.options = options;
this.socketTask = null; // WebSocket任务实例
this.reconnectTimer = null; // 重连计时器
this.reconnectCount = 0; // 当前重连次数
this.maxReconnectCount = options.maxReconnectCount || 5; // 最大重连次数
this.reconnectInterval = options.reconnectInterval || 3000; // 重连间隔(ms)
this.heartbeatTimer = null; // 心跳计时器
this.heartbeatInterval = options.heartbeatInterval || 30000; // 心跳间隔(ms)
this.heartbeatMsg = options.heartbeatMsg || 'ping'; // 心跳消息
this.isClose = false; // 手动关闭
this.callbacks = {
open: [], // 连接打开回调
message: [], // 消息接收回调
close: [], // 连接关闭回调
error: [] // 错误处理回调
};
this.init(true); // 初始化WebSocket连接
}
/**
* 初始化WebSocket连接
*/
init(flag = false) {
if(flag) this.reconnectCount = 0
this.isClose = false; // 手动关闭
// 创建WebSocket连接
this.socketTask = uni.connectSocket({
url: this.url,
header: this.options.header || {},
protocols: this.options.protocols || [],
success: () => {
// console.log('WebSocket连接创建成功');
},
fail: (err) => {
// console.error('WebSocket连接创建失败', err);
this.triggerEvent('error', err);
this.tryReconnect(); // 连接失败时尝试重连
}
});
// 监听WebSocket连接打开事件
this.socketTask.onOpen(() => {
// console.log('WebSocket连接已打开');
this.reconnectCount = 0; // 重置重连计数
clearInterval(this.reconnectTimer); // 清除重连计时器
// this.startHeartbeat(); // 启动心跳机制
this.triggerEvent('open'); // 触发连接打开事件
});
// 监听WebSocket消息接收事件
this.socketTask.onMessage((res) => {
if(res.data){
const { rtnCode } = JSON.parse(res.data)
// console.log('收到WebSocket消息', JSON.parse(res.data) );
if(['0005', 'QQ0005'].includes(rtnCode)) {
clearInterval(this.reconnectTimer);
this.reconnectTimer = null
this.close()
goto_login_fun('socket');
return
}
}else{
clearInterval(this.reconnectTimer); // 清除重连计时器
this.reconnectTimer = null
goto_login_fun('socket');
return
}
this.triggerEvent('message', res); // 触发消息接收事件
});
// 监听WebSocket连接关闭事件
this.socketTask.onClose((res) => {
// console.log('WebSocket连接已关闭', res);
clearInterval(this.heartbeatTimer); // 清除心跳计时器
this.triggerEvent('close', res); // 触发连接关闭事件
this.tryReconnect(); // 尝试重连
});
// 监听WebSocket错误事件
this.socketTask.onError((err) => {
console.error('WebSocket发生错误', err);
clearInterval(this.heartbeatTimer); // 清除心跳计时器
this.triggerEvent('error', err); // 触发错误事件
this.tryReconnect(); // 尝试重连
});
}
/**
* 注册事件回调
* @param {string} event - 事件名称(open/message/close/error)
* @param {Function} callback - 回调函数
* @returns {WebSocketUtil} - 返回当前实例支持链式调用
*/
on(event, callback) {
if (this.callbacks[event]) {
this.callbacks[event].push(callback);
}
return this;
}
/**
* 移除事件回调
* @param {string} event - 事件名称(open/message/close/error)
* @param {Function} callback - 要移除的回调函数
* @returns {WebSocketUtil} - 返回当前实例支持链式调用
*/
off(event, callback) {
if (this.callbacks[event]) {
this.callbacks[event] = this.callbacks[event].filter(cb => cb !== callback);
}
return this;
}
/**
* 触发事件回调
* @param {string} event - 事件名称
* @param {any} [data] - 传递给回调函数的数据
*/
triggerEvent(event, data) {
if (this.callbacks[event]) {
this.callbacks[event].forEach(callback => callback(data));
}
}
/**
* 通过WebSocket发送消息
* @param {string|Object} message - 要发送的消息可以是字符串或对象
* @returns {WebSocketUtil} - 返回当前实例支持链式调用
*/
send(message, num = 1) {
if (this.socketTask && this.getStatus() === 1) {
// 发送消息,对象会自动转换为JSON字符串
this.socketTask.send({
data: typeof message === 'string' ? message : JSON.stringify(message),
success: () => {
// console.log('WebSocket消息发送成功');
},
fail: (err) => {
console.error('WebSocket消息发送失败', err);
this.triggerEvent('error', err);
}
});
} else {
console.error('WebSocket连接未打开,无法发送消息');
this.triggerEvent('error', new Error('WebSocket连接未打开'));
if(num === 10) return
setTimeout(()=>{
this.send(message, num + 1)
},500)
}
return this;
}
/**
* 关闭WebSocket连接
* @param {number} [code=1000] - 关闭码
* @param {string} [reason=''] - 关闭原因
* @returns {WebSocketUtil} - 返回当前实例支持链式调用
*/
close(code = 1000, reason = '') {
this.isClose = true
if (this.socketTask) {
clearInterval(this.heartbeatTimer); // 清除心跳计时器
clearInterval(this.reconnectTimer); // 清除重连计时器
// 关闭WebSocket连接
this.socketTask.close({
code,
reason,
success: () => {
console.log('WebSocket连接正在关闭');
},
fail: (err) => {
console.error('WebSocket关闭失败', err);
this.triggerEvent('error', err);
}
});
}
return this;
}
/**
* 获取WebSocket连接状态
* @returns {number} - 连接状态0-连接中1-已连接2-连接关闭中3-已关闭-1-未知
*/
getStatus() {
if (this.socketTask) {
try {
return this.socketTask.readyState;
} catch (e) {
console.error('获取WebSocket状态失败', e);
return -1;
}
}
return -1;
}
/**
* 尝试重新连接WebSocket
*/
tryReconnect() {
if(this.isClose) return
if (this.reconnectCount < this.maxReconnectCount) {
clearInterval(this.reconnectTimer);
this.reconnectTimer = setTimeout(() => {
this.reconnectCount++;
console.log(`尝试重新连接WebSocket (${this.reconnectCount}/${this.maxReconnectCount})`);
this.init(); // 重新初始化WebSocket连接
}, this.reconnectInterval);
} else {
console.error('达到最大重连次数,停止尝试');
this.triggerEvent('error', new Error('达到最大重连次数'));
this.reconnectTimer = null
this.close()
goto_login_fun('socket')
}
}
/**
* 启动心跳机制
*/
startHeartbeat() {
clearInterval(this.heartbeatTimer);
this.heartbeatTimer = setInterval(() => {
if (this.getStatus() === 1) {
this.send(this.heartbeatMsg); // 发送心跳消息
}
}, this.heartbeatInterval);
}
}
+92
View File
@@ -0,0 +1,92 @@
import request from '@/api/request'
const base_url = '/trapractice'
// 获取分类树接口
export const queryTraPartnerCatelogList = (data) => {
return request({
url: base_url + '/traPartnerCatelog/queryTraPartnerCatelogList',
method: 'post',
data
});
};
// 获取标签分类树 接口
export const getTraTagListApi = (data) => {
return request({
url: base_url + '/traPartnerTag/queryTraTagCatalogInfoList',
method: 'post',
data
});
};
// 全部模块
export const queryCrsInfoByCatalogApi = (data) => {
return request({
url: base_url + '/traPartnerInfo/queryTraPartnerInfoPaging',
method: 'post',
data
});
};
// id 查详情/trapractice/traPartnerInfo/queryTraPartnerInfoById
export const queryTraPartnerInfoById = (data) => {
return request({
url: base_url + '/traPartnerInfo/queryTraPartnerInfoById',
method: 'post',
data
});
};
// 陪练角色列表 traId
export const queryTraPartnerCharacterInfoList = (data) => {
return request({
url: base_url + '/traPartnerCharacterInfo/queryTraPartnerCharacterInfoList',
method: 'post',
data
});
};
// 陪练角色详情 + 对话详情 /traPartnerCharacterInfo/queryTraPartnerCharacterInfoById
export const queryTraPartnerCharacterInfoById = (data) => {
return request({
url: base_url + '/traPartnerCharacterInfo/queryTraPartnerCharacterInfoById',
method: 'post',
data
});
};
// 结束陪练 /trapractice/traPartnerChatReport/partnerChatReportTrigger?execId=
export const partnerChatReportTrigger = (data) => {
return request({
url: base_url + '/traPartnerChatReport/partnerChatReportTrigger',
method: 'post',
data
});
};
// 问答提示/trapractice/traPartnerChat/partnerChatPrompt
export const partnerChatPrompt = (data) => {
return request({
url: base_url + '/traPartnerChat/partnerChatPrompt',
method: 'post',
data
});
};
// 获取报告traPartnerChatReport/partnerChatReport
export const partnerChatReport = (data) => {
return request({
url: base_url + '/traPartnerChatReport/partnerChatReport',
method: 'post',
data
});
};
// 反馈
export const insertTraPartnerFeedback = (data) => {
return request({
url: base_url + '/traPartnerFeedback/insertTraPartnerFeedback',
method: 'post',
data
});
};
// 查询是否可以进入trapractice/traPartnerInfo/checkTraPartnerInfoById
export const checkTraPartnerInfoById = (data) => {
return request({
url: base_url + '/traPartnerInfo/checkTraPartnerInfoById',
method: 'post',
data
});
};
+206
View File
@@ -0,0 +1,206 @@
import request from '@/api/request'
import { getToken } from '@/common/common.js'
import { get_base_url } from '@/api/request'
const base_url = '/trastudy'
// 获取课程学习信息接口
export const getCourseStudyInfoApi = (data) => {
return request({
url: base_url + '/traStdyInfo/queryTraStdyProgressAndTeacher',
method: 'post',
toastErrors: true,
data
});
};
// 获取课程学习 知识点 接口
export const queryTraStdyBatchWaitParagraphInfoApi = (data) => {
return request({
url: base_url + '/traStdyInfo/queryTraStdyBatchWaitParagraphInfo',
method: 'post',
toastErrors: true,
data
});
};
// 重新开始 删除学习信息
export const deleteTraStdyInfoByIdApi = (data) => {
return request({
url: base_url + '/traStdyInfo/deleteTraStdyInfoById',
method: 'post',
toastErrors: true,
data
});
};
// 重新学习
export const afreshStudyCourseApi = (data) => {
return request({
url: base_url + '/traStdyInfo/afreshStudyCourse',
method: 'post',
toastErrors: true,
data
});
};
// 学习段落完成
export const studyParagraphOverApi = (data) => {
return request({
url: base_url + '/traStdyInfo/studyParagraphOver',
method: 'post',
toastErrors: true,
data
});
};
// 回答问题 /traStdyInfo/studyAnswerTextUp
export const studyAnswerTextUpApi = (data) => {
return request({
url: base_url + '/traStdyInfo/studyAnswerTextUp',
method: 'post',
toastErrors: true,
data
});
};
// 获取问题信息/traStdyInfo/queryTraStdyBatchWaitQuestionInfo
export const queryTraStdyBatchWaitQuestionInfoApi = (data) => {
return request({
url: base_url + '/traStdyInfo/queryTraStdyBatchWaitQuestionInfo',
method: 'post',
toastErrors: true,
data
});
};
// 上传语音翻译i文本 /traStdyInfo/audioTranscriptions
export const audioTranscriptionsApi = (data) => {
return request({
url: base_url + '/traStdyInfo/audioTranscriptions',
method: 'post',
toastErrors: true,
data
});
};
// 撤回消息
export const withdrawalAnswerApi = (data) => {
return request({
url: base_url + '/traStdyInfo/withdrawalAnswer',
method: 'post',
toastErrors: true,
data
});
};
// 异步完成回答
// 入参:
// qstLogId 必输项:true 类型:String
// 出参:
// qstLogId 问题日志ID String
// qnsId 问题ID String
// processStat 状态 String
// message 处理消息 String
export const waitMakeEvaluateApi = (data) => {
return request({
url: base_url + '/traStdyInfo/waitMakeEvaluate',
method: 'post',
toastErrors: true,
data
});
};
// 查询问题评价
// /traStdyInfo/queryQuestionEvalate
export const queryQuestionEvalateApi = (data) => {
return request({
url: base_url + '/traStdyInfo/queryQuestionEvalate',
method: 'post',
toastErrors: true,
data
});
};
// /trastudy/intgask/askAbout
// 流式获取文本
export const getAskAboutRawApi = (data) => {
// return request({
// baseUrl:'http://25.18.122.76:3030',
// url: '/api/v1/prediction/a402cf58-9f64-479b-a9f8-c527300a7905',
// isStream: true,
// header: {
// 'Accept':'text/event-stream',
// 'content-type': 'application/json; charset=UTF-8',
// 'Authorization': 'Bearer PjadGPFSHDg6F8ZRekROe-QxS2fPXlWRsmvP5Mp0vxA'
// },
// method: 'post',
// toastErrors: false,
// data
// });
return request({
url: base_url + '/intgask/askAbout',
header: {
'Accept':'text/event-stream',
'content-type': 'text/event-stream; charset=UTF-8',
},
responseType: 'arrayBuffer',
method: 'get',
isStream: true,
timeout: 30000,
toastErrors: false,
data
})
};
// 学习开始
export const studyStartApi = (data) => {
return request({
url: base_url + '/traStdychat/studyStart',
method: 'post',
toastErrors: true,
data
});
};
// 学习结束
export const studyEndApi = (data) => {
return request({
url: base_url + '/traStdychat/studyEnd',
method: 'post',
toastErrors: true,
data
});
};
// 学习 对话
export const textChartApi = (data) => {
return request({
url: base_url + '/traStdychat/textChart',
method: 'post',
toastErrors: true,
data
});
};
// 预览 对话
export const textChartPreviewApi = (data) => {
return request({
url: base_url + '/traStdyInfoPreview/previewCrsInfo',
method: 'post',
toastErrors: true,
data
});
};
export const previewAfreshStudyCourse = (data) => {
return request({
url: base_url + '/traStdyInfo/previewAfreshStudyCourse',
method: 'post',
toastErrors: true,
data
});
};
// 设置学习顾问音色/traStdyTeacher/setStudyTeacher
export const setStudyTeacherApi = (data) => {
return request({
url: base_url + '/traStdyTeacher/setStudyTeacher',
method: 'post',
toastErrors: true,
data
});
};
+11
View File
@@ -0,0 +1,11 @@
import request from '@/api/request'
export const prepareApi = (data) => {
return request({
url: '/trapractice/traPartnerVoiceCall/voiceCallPrepare' ,
method: 'post',
toastErrors: true,
data,
baseUrl: 'http://25.64.32.157:9603'
});
};
+13
View File
@@ -0,0 +1,13 @@
import request from '@/api/request'
const module_url = '/traexam'
// 1.查询考试列表分页
export const queryTraExamPapersPage = (data) => {
return request({
url: module_url + '/traExamPapers/queryTraExamPapersPage',
method: 'post',
data
});
};
+58
View File
@@ -0,0 +1,58 @@
import request from '@/api/request'
const base_url = '/trapractice'
/**练习*/
// 练习列表页
export const queryPracticeRecordPaging = (data) => {
return request({
url: base_url + '/traPartnerExecuteHis/queryTraPartnerExecutePaging',
method: 'post',
toastErrors: true,
data:{
traType:'01',
...data
}
});
};
// 练习每一项列表
export const queryPracticeHis = (data) => {
return request({
url: base_url + '/traPartnerExecuteHis/queryTraPartnerExecuteHisPaging',
method: 'post',
toastErrors: true,
data
});
};
// 练习详情 通过 exrId 查询
export const queryPracticeRecordByExrId = (data) => {
return request({
url: base_url + '/traCrsPractice/queryPracticeRecordByExrId',
method: 'post',
toastErrors: true,
data
});
};
/**考试*/
export const queryCrsExamRecordPaging = (data) => {
return request({
url: base_url + '/traPartnerExecuteHis/queryTraPartnerExecutePaging',
method: 'post',
toastErrors: true,
data:{
traType:'02',
...data
}
});
};
// 考试每一项列表
export const queryCrsExamHisByExamId = (data) => {
return request({
url: base_url + '/traPartnerExecuteHis/queryTraPartnerExecuteHisPaging',
method: 'post',
toastErrors: true,
data
});
};
+57
View File
@@ -0,0 +1,57 @@
import request from '@/api/request'
// 我的-统计信息
export const queryTraStdyStatisticsSelf = (data) => {
/**
userId 学员ID String
stdtTmLen 累计学习时长 double
accmTaskCnt 累计完成任务 int
accmPoint 累计获得积分 int
**/
return request({
url: '/traapp/traStdyInfo/queryTraStdyStatisticsSelf',
method: 'post',
data
});
};
// 更新性别
export const updateDfSysUserExtandInfoGender = (data) => {
/**
**/
return request({
url: '/traapp/dfSysUserInfo/updateDfSysUserExtandInfoGender ',
method: 'post',
data
});
};
// 更新照片
export const updateDfSysUserExtandInfoImageAddr = (data) => {
return request({
url: '/traapp/dfSysUserInfo/updateDfSysUserExtandInfoImageAddr',
method: 'post',
data
});
};
// 设置匿名 是Y 否N
export const updateTraUserAnony = (data) => {
return request({
url: '/traapp/traUserAnony/updateTraUserAnony',
method: 'post',
data
});
};
// 查询是否匿名 是Y 否N
export const queryTraUserAnonyByUserId = (data) => {
return request({
url: '/traapp/traUserAnony/queryTraUserAnonyByUserId',
method: 'post',
data
});
};

Some files were not shown because too many files have changed in this diff Show More