feat(model): 修改基础的 OpenAIBase 为异步调用

This commit is contained in:
Wenjie Zhang 2026-01-28 11:30:55 +08:00
parent 27b9e829d4
commit 138a0795bf
8 changed files with 30 additions and 37 deletions

View File

@ -20,6 +20,7 @@
- 文件上传解析后,如何提示用户需要入库
- 检索测试中,添加问答
- 丰富当前智能体的 Prompt最好支持从 markdown 解析
- 模型的配置也传输到数据库(待定)
### Bugs

View File

@ -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")

View File

@ -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"}

View File

@ -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:

View File

@ -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}")
# 检查响应是否有效

View File

@ -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

View File

@ -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"],

View File

@ -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(