完成请求推送,但组装mes还有问题
This commit is contained in:
@@ -4,9 +4,12 @@ LLM 客户端工具
|
||||
提供统一的 LLM 接口,支持多种模型提供商。
|
||||
使用 LangChain 的 ChatModel 抽象,简化不同厂商 API 的调用。
|
||||
"""
|
||||
from typing import Optional
|
||||
from typing import Optional, List, Dict, Any, AsyncGenerator
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage
|
||||
from langchain_core.callbacks import AsyncCallbackHandler
|
||||
from core.config import settings
|
||||
import time
|
||||
|
||||
|
||||
def get_llm(
|
||||
@@ -86,3 +89,249 @@ def get_creative_llm(provider: str = "openai") -> BaseChatModel:
|
||||
def get_streaming_llm(provider: str = "openai") -> BaseChatModel:
|
||||
"""获取支持流式输出的 LLM"""
|
||||
return get_llm(provider, streaming=True)
|
||||
|
||||
|
||||
class TokenUsageCallbackHandler(AsyncCallbackHandler):
|
||||
"""
|
||||
Token 使用回调处理器
|
||||
|
||||
用于捕获 LLM 调用的 token 使用情况
|
||||
"""
|
||||
def __init__(self):
|
||||
self.prompt_tokens = 0
|
||||
self.completion_tokens = 0
|
||||
self.total_tokens = 0
|
||||
self.response_content = ""
|
||||
|
||||
async def on_llm_start(self, serialized: Dict[str, Any], prompts: List[str] = None, **kwargs):
|
||||
"""LLM 开始时的回调"""
|
||||
pass
|
||||
|
||||
async def on_llm_end(self, response, **kwargs):
|
||||
"""LLM 结束时的回调,获取 token 统计"""
|
||||
try:
|
||||
# 从 response 中提取 token 信息
|
||||
if hasattr(response, 'llm_output') and response.llm_output:
|
||||
token_usage = response.llm_output.get('token_usage', {})
|
||||
self.prompt_tokens = token_usage.get('prompt_tokens', 0)
|
||||
self.completion_tokens = token_usage.get('completion_tokens', 0)
|
||||
self.total_tokens = token_usage.get('total_tokens', 0)
|
||||
except Exception as e:
|
||||
print(f"[TokenUsageCallback] 提取 token 信息失败: {e}")
|
||||
|
||||
async def on_llm_new_token(self, token: str, **kwargs):
|
||||
"""每个新 token 的回调(流式输出)"""
|
||||
self.response_content += token
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""
|
||||
LLM 客户端封装类
|
||||
|
||||
提供统一的异步接口,支持自定义API配置、流式输出和 token 统计
|
||||
"""
|
||||
|
||||
async def chat_completion(
|
||||
self,
|
||||
messages: List[BaseMessage],
|
||||
api_url: str,
|
||||
api_key: str,
|
||||
model: str = "gpt-3.5-turbo",
|
||||
temperature: float = 1.0,
|
||||
max_tokens: int = 500,
|
||||
request_timeout: int = 60,
|
||||
stream: bool = False,
|
||||
**kwargs
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
调用 LLM API 生成回复
|
||||
|
||||
Args:
|
||||
messages: LangChain 消息列表
|
||||
api_url: API 地址
|
||||
api_key: API 密钥
|
||||
model: 模型名称
|
||||
temperature: 温度参数
|
||||
max_tokens: 最大 token 数
|
||||
request_timeout: 请求超时时间(秒)
|
||||
stream: 是否启用流式输出
|
||||
**kwargs: 其他参数
|
||||
|
||||
Returns:
|
||||
OpenAI 格式的响应字典,包含 token 使用信息
|
||||
"""
|
||||
try:
|
||||
from langchain_openai import ChatOpenAI
|
||||
|
||||
# 创建回调处理器
|
||||
callback_handler = TokenUsageCallbackHandler()
|
||||
|
||||
# 创建自定义的 ChatOpenAI 实例
|
||||
llm = ChatOpenAI(
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
api_key=api_key,
|
||||
base_url=api_url if api_url else None,
|
||||
max_tokens=max_tokens,
|
||||
streaming=stream,
|
||||
callbacks=[callback_handler],
|
||||
request_timeout=request_timeout, # ✅ 设置超时时间
|
||||
**kwargs
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
if stream:
|
||||
# 流式模式
|
||||
full_content = ""
|
||||
async for chunk in llm.astream(messages):
|
||||
if hasattr(chunk, 'content'):
|
||||
full_content += chunk.content
|
||||
|
||||
duration = time.time() - start_time
|
||||
|
||||
return {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": full_content
|
||||
}
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": callback_handler.prompt_tokens,
|
||||
"completion_tokens": callback_handler.completion_tokens,
|
||||
"total_tokens": callback_handler.total_tokens
|
||||
},
|
||||
"duration": duration
|
||||
}
|
||||
else:
|
||||
# 非流式模式
|
||||
response = await llm.ainvoke(messages)
|
||||
duration = time.time() - start_time
|
||||
|
||||
return {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": response.content
|
||||
}
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": callback_handler.prompt_tokens,
|
||||
"completion_tokens": callback_handler.completion_tokens,
|
||||
"total_tokens": callback_handler.total_tokens
|
||||
},
|
||||
"duration": duration
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
print(f"[LLMClient] 调用失败: {e}")
|
||||
raise
|
||||
|
||||
async def stream_chat(
|
||||
self,
|
||||
messages: List[BaseMessage],
|
||||
api_url: str,
|
||||
api_key: str,
|
||||
model: str = "gpt-3.5-turbo",
|
||||
temperature: float = 1.0,
|
||||
max_tokens: int = 500,
|
||||
request_timeout: int = 60,
|
||||
**kwargs
|
||||
) -> AsyncGenerator[Dict[str, Any], None]:
|
||||
"""
|
||||
流式调用 LLM API
|
||||
|
||||
Args:
|
||||
messages: LangChain 消息列表
|
||||
api_url: API 地址
|
||||
api_key: API 密钥
|
||||
model: 模型名称
|
||||
temperature: 温度参数
|
||||
max_tokens: 最大 token 数
|
||||
request_timeout: 请求超时时间(秒)
|
||||
**kwargs: 其他参数
|
||||
|
||||
Yields:
|
||||
包含 token 片段的字典
|
||||
"""
|
||||
try:
|
||||
from langchain_openai import ChatOpenAI
|
||||
|
||||
print(f"\n[LLMClient] 🔧 创建 ChatOpenAI 实例")
|
||||
print(f" - Model: {model}")
|
||||
print(f" - API URL: {api_url[:50]}..." if len(api_url) > 50 else f" - API URL: {api_url}")
|
||||
print(f" - Temperature: {temperature}")
|
||||
print(f" - Max Tokens: {max_tokens}")
|
||||
print(f" - Request Timeout: {request_timeout}s")
|
||||
|
||||
# 创建回调处理器
|
||||
callback_handler = TokenUsageCallbackHandler()
|
||||
|
||||
# 创建自定义的 ChatOpenAI 实例
|
||||
llm = ChatOpenAI(
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
api_key=api_key,
|
||||
base_url=api_url if api_url else None,
|
||||
max_tokens=max_tokens,
|
||||
streaming=True,
|
||||
callbacks=[callback_handler],
|
||||
request_timeout=request_timeout, # ✅ 设置超时时间
|
||||
**kwargs
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
print(f"[LLMClient] 🚀 开始流式请求...")
|
||||
print(f" - Messages 数量: {len(messages)}")
|
||||
if messages:
|
||||
first_msg_role = getattr(messages[0], 'role', 'unknown')
|
||||
first_msg_preview = str(getattr(messages[0], 'content', ''))[:50]
|
||||
print(f" - 第一条消息: [{first_msg_role}] {first_msg_preview}...")
|
||||
|
||||
chunk_count = 0
|
||||
# 流式输出
|
||||
async for chunk in llm.astream(messages):
|
||||
if hasattr(chunk, 'content') and chunk.content:
|
||||
chunk_count += 1
|
||||
|
||||
# 第一个 chunk 时记录
|
||||
if chunk_count == 1:
|
||||
first_chunk_time = time.time()
|
||||
print(f"[LLMClient] ✨ 收到第一个 chunk (耗时: {first_chunk_time - start_time:.2f}s)")
|
||||
|
||||
yield {
|
||||
"type": "chunk",
|
||||
"content": chunk.content
|
||||
}
|
||||
|
||||
duration = time.time() - start_time
|
||||
|
||||
print(f"[LLMClient] ✅ 流式请求完成")
|
||||
print(f" - 总 Chunks: {chunk_count}")
|
||||
print(f" - 耗时: {duration:.2f}秒")
|
||||
print(f" - Prompt Tokens: {callback_handler.prompt_tokens}")
|
||||
print(f" - Completion Tokens: {callback_handler.completion_tokens}")
|
||||
print(f" - Total Tokens: {callback_handler.total_tokens}\n")
|
||||
|
||||
# 最后发送 token 使用信息
|
||||
yield {
|
||||
"type": "usage",
|
||||
"usage": {
|
||||
"prompt_tokens": callback_handler.prompt_tokens,
|
||||
"completion_tokens": callback_handler.completion_tokens,
|
||||
"total_tokens": callback_handler.total_tokens
|
||||
},
|
||||
"duration": duration
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n[LLMClient] ❌ 流式调用失败: {e}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
raise
|
||||
|
||||
Reference in New Issue
Block a user