feat(model): 修改基础的 OpenAIBase 为异步调用
This commit is contained in:
parent
27b9e829d4
commit
138a0795bf
@ -20,6 +20,7 @@
|
||||
- 文件上传解析后,如何提示用户需要入库
|
||||
- 检索测试中,添加问答
|
||||
- 丰富当前智能体的 Prompt,最好支持从 markdown 解析
|
||||
- 模型的配置也传输到数据库(待定)
|
||||
|
||||
|
||||
### Bugs
|
||||
|
||||
@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import traceback
|
||||
import uuid
|
||||
|
||||
@ -10,7 +9,6 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from src.storage.postgres.models_business import User
|
||||
from server.routers.auth_router import get_admin_user
|
||||
from server.utils.auth_middleware import get_db, get_required_user
|
||||
from src import executor
|
||||
from src import config as conf
|
||||
from src.agents import agent_manager
|
||||
from src.models import select_model
|
||||
@ -130,11 +128,7 @@ async def call(query: str = Body(...), meta: dict = Body(None), current_user: Us
|
||||
model_spec=meta.get("model_spec") or meta.get("model"),
|
||||
)
|
||||
|
||||
async def call_async(query):
|
||||
loop = asyncio.get_event_loop()
|
||||
return await loop.run_in_executor(executor, model.call, query)
|
||||
|
||||
response = await call_async(query)
|
||||
response = await model.call(query)
|
||||
logger.debug({"query": query, "response": response.content})
|
||||
|
||||
return {"response": response.content, "request_id": meta["request_id"]}
|
||||
@ -409,7 +403,8 @@ async def chat_agent(
|
||||
async def get_chat_models(model_provider: str, current_user: User = Depends(get_admin_user)):
|
||||
"""获取指定模型提供商的模型列表(需要登录)"""
|
||||
model = select_model(model_provider=model_provider)
|
||||
return {"models": model.get_models()}
|
||||
models = await model.get_models()
|
||||
return {"models": models}
|
||||
|
||||
|
||||
@chat.post("/models/update")
|
||||
|
||||
@ -966,7 +966,7 @@ async def generate_sample_questions(
|
||||
# 选择模型并调用
|
||||
model = select_model()
|
||||
messages = [{"role": "system", "content": system_prompt}, {"role": "user", "content": user_message}]
|
||||
response = model.call(messages, stream=False)
|
||||
response = await model.call(messages, stream=False)
|
||||
|
||||
# 解析AI返回的JSON
|
||||
try:
|
||||
@ -1366,7 +1366,7 @@ async def generate_description(
|
||||
|
||||
try:
|
||||
model = select_model()
|
||||
response = await asyncio.to_thread(model.call, prompt)
|
||||
response = await model.call(prompt)
|
||||
description = response.content.strip()
|
||||
logger.debug(f"Generated description: {description}")
|
||||
return {"description": description, "status": "success"}
|
||||
|
||||
@ -7,7 +7,6 @@
|
||||
- 保存和加载思维导图配置
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
import textwrap
|
||||
@ -214,10 +213,10 @@ async def generate_mindmap(
|
||||
# 调用AI生成
|
||||
logger.info(f"开始生成思维导图,知识库: {db_name}, 文件数量: {len(files_info)}")
|
||||
|
||||
# 选择模型并调用(使用异步包装)
|
||||
# 选择模型并调用
|
||||
model = select_model()
|
||||
messages = [{"role": "system", "content": system_prompt}, {"role": "user", "content": user_message}]
|
||||
response = await asyncio.to_thread(model.call, messages, stream=False)
|
||||
response = await model.call(messages, stream=False)
|
||||
|
||||
# 解析AI返回的JSON
|
||||
try:
|
||||
|
||||
@ -1,7 +1,7 @@
|
||||
import os
|
||||
import traceback
|
||||
|
||||
from openai import OpenAI
|
||||
from openai import AsyncOpenAI
|
||||
from tenacity import before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential
|
||||
|
||||
from src import config
|
||||
@ -27,7 +27,7 @@ class OpenAIBase:
|
||||
def __init__(self, api_key, base_url, model_name, **kwargs):
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.client = OpenAI(api_key=api_key, base_url=base_url)
|
||||
self.client = AsyncOpenAI(api_key=api_key, base_url=base_url)
|
||||
self.model_name = model_name
|
||||
self.info = kwargs
|
||||
|
||||
@ -38,7 +38,7 @@ class OpenAIBase:
|
||||
before_sleep=before_sleep_log(logger, log_level="WARNING"),
|
||||
reraise=True,
|
||||
)
|
||||
def call(self, message, stream=False):
|
||||
async def call(self, message, stream=False):
|
||||
if isinstance(message, str):
|
||||
messages = [{"role": "user", "content": message}]
|
||||
else:
|
||||
@ -48,7 +48,7 @@ class OpenAIBase:
|
||||
if stream:
|
||||
response = self._stream_response(messages)
|
||||
else:
|
||||
response = self._get_response(messages)
|
||||
response = await self._get_response(messages)
|
||||
|
||||
except Exception as e:
|
||||
err = (
|
||||
@ -60,27 +60,27 @@ class OpenAIBase:
|
||||
|
||||
return response
|
||||
|
||||
def _stream_response(self, messages):
|
||||
response = self.client.chat.completions.create(
|
||||
async def _stream_response(self, messages):
|
||||
response = await self.client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=messages,
|
||||
stream=True,
|
||||
)
|
||||
for chunk in response:
|
||||
async for chunk in response:
|
||||
if len(chunk.choices) > 0:
|
||||
yield chunk.choices[0].delta
|
||||
|
||||
def _get_response(self, messages):
|
||||
response = self.client.chat.completions.create(
|
||||
async def _get_response(self, messages):
|
||||
response = await self.client.chat.completions.create(
|
||||
model=self.model_name,
|
||||
messages=messages,
|
||||
stream=False,
|
||||
)
|
||||
return response.choices[0].message
|
||||
|
||||
def get_models(self):
|
||||
async def get_models(self):
|
||||
try:
|
||||
return self.client.models.list(extra_query={"type": "text"})
|
||||
return await self.client.models.list(extra_query={"type": "text"})
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting models: {e}")
|
||||
return []
|
||||
@ -160,7 +160,7 @@ async def test_chat_model_status(provider: str, model_name: str) -> dict:
|
||||
test_messages = [{"role": "user", "content": "Say 1"}]
|
||||
|
||||
# 发送测试请求
|
||||
response = model.call(test_messages, stream=False)
|
||||
response = await model.call(test_messages, stream=False)
|
||||
logger.debug(f"Test chat model status response: {response}")
|
||||
|
||||
# 检查响应是否有效
|
||||
|
||||
@ -102,7 +102,7 @@ class ContentGuard:
|
||||
text_lower = text.lower()
|
||||
|
||||
prompt = PROMPT_TEMPLATE.format(content=text_lower)
|
||||
response = self.llm_model.call(prompt)
|
||||
response = await self.llm_model.call(prompt)
|
||||
logger.debug(f"LLM response: {response.content}")
|
||||
return True if "不合规" in response.content else False
|
||||
|
||||
|
||||
@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
@ -390,7 +389,7 @@ class EvaluationService:
|
||||
)
|
||||
|
||||
try:
|
||||
resp = await asyncio.to_thread(llm.call, prompt, False)
|
||||
resp = await llm.call(prompt, False)
|
||||
content = resp.content if resp else ""
|
||||
|
||||
import json_repair
|
||||
@ -606,8 +605,8 @@ class EvaluationService:
|
||||
"如果上下文中缺少相关信息,请回答“信息不足,无法回答”。\n\n"
|
||||
)
|
||||
|
||||
# 生成答案 - 使用 asyncio.to_thread 避免阻塞事件循环
|
||||
response = await asyncio.to_thread(llm.call, prompt, stream=False)
|
||||
# 生成答案
|
||||
response = await llm.call(prompt, stream=False)
|
||||
generated_answer = response.content if response else ""
|
||||
logger.debug(f"LLM 生成的答案长度: {len(generated_answer) if generated_answer else 0}")
|
||||
|
||||
@ -629,9 +628,8 @@ class EvaluationService:
|
||||
|
||||
if benchmark_row.has_gold_answers and question_data.get("gold_answer"):
|
||||
if judge_llm:
|
||||
# 评判过程包含 LLM 调用,使用 asyncio.to_thread 避免阻塞
|
||||
answer_scores = await asyncio.to_thread(
|
||||
EvaluationMetricsCalculator.calculate_answer_metrics,
|
||||
# 评判过程包含 LLM 调用
|
||||
answer_scores = await EvaluationMetricsCalculator.calculate_answer_metrics(
|
||||
query=question_data["query"],
|
||||
generated_answer=generated_answer,
|
||||
gold_answer=question_data["gold_answer"],
|
||||
|
||||
@ -45,7 +45,7 @@ class AnswerMetrics:
|
||||
"""答案评估指标计算"""
|
||||
|
||||
@staticmethod
|
||||
def judge_correctness(query: str, generated_answer: str, gold_answer: str, judge_llm: Any) -> dict[str, Any]:
|
||||
async def judge_correctness(query: str, generated_answer: str, gold_answer: str, judge_llm: Any) -> dict[str, Any]:
|
||||
"""
|
||||
使用LLM判断生成的答案是否正确
|
||||
"""
|
||||
@ -75,7 +75,7 @@ class AnswerMetrics:
|
||||
}}
|
||||
""")
|
||||
try:
|
||||
response = judge_llm.call(prompt, stream=False)
|
||||
response = await judge_llm.call(prompt, stream=False)
|
||||
content = response.content.strip()
|
||||
|
||||
# 尝试清理可能的 markdown 代码块
|
||||
@ -117,14 +117,14 @@ class EvaluationMetricsCalculator:
|
||||
return metrics
|
||||
|
||||
@staticmethod
|
||||
def calculate_answer_metrics(
|
||||
async def calculate_answer_metrics(
|
||||
query: str, generated_answer: str, gold_answer: str, judge_llm: Any = None
|
||||
) -> dict[str, Any]:
|
||||
"""计算答案指标 (LLM Judge)"""
|
||||
if not judge_llm:
|
||||
return {}
|
||||
|
||||
return AnswerMetrics.judge_correctness(query, generated_answer, gold_answer, judge_llm)
|
||||
return await AnswerMetrics.judge_correctness(query, generated_answer, gold_answer, judge_llm)
|
||||
|
||||
@staticmethod
|
||||
def calculate_overall_score(
|
||||
|
||||
Loading…
Reference in New Issue
Block a user