diff --git a/docs/latest/changelog/roadmap.md b/docs/latest/changelog/roadmap.md index 2cf7375c..bf00ba1f 100644 --- a/docs/latest/changelog/roadmap.md +++ b/docs/latest/changelog/roadmap.md @@ -20,6 +20,7 @@ - 文件上传解析后,如何提示用户需要入库 - 检索测试中,添加问答 - 丰富当前智能体的 Prompt,最好支持从 markdown 解析 +- 模型的配置也传输到数据库(待定) ### Bugs diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index c56c273b..eb9343d5 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -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") diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 1677aad3..1c748356 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -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"} diff --git a/server/routers/mindmap_router.py b/server/routers/mindmap_router.py index f97d1dec..012bb8b1 100644 --- a/server/routers/mindmap_router.py +++ b/server/routers/mindmap_router.py @@ -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: diff --git a/src/models/chat.py b/src/models/chat.py index 4fc66f2d..27eda78a 100644 --- a/src/models/chat.py +++ b/src/models/chat.py @@ -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}") # 检查响应是否有效 diff --git a/src/plugins/guard.py b/src/plugins/guard.py index fb69887a..efbe19d6 100644 --- a/src/plugins/guard.py +++ b/src/plugins/guard.py @@ -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 diff --git a/src/services/evaluation_service.py b/src/services/evaluation_service.py index 1280fac7..e1b6968b 100644 --- a/src/services/evaluation_service.py +++ b/src/services/evaluation_service.py @@ -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"], diff --git a/src/utils/evaluation_metrics.py b/src/utils/evaluation_metrics.py index 19292fc4..d8b8f97a 100644 --- a/src/utils/evaluation_metrics.py +++ b/src/utils/evaluation_metrics.py @@ -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(