x
This commit is contained in:
@@ -0,0 +1,173 @@
|
||||
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", "")))
|
||||
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()
|
||||
Reference in New Issue
Block a user