Merge pull request #147 from xerrors/agent-bug-fix

Agent bug fix
This commit is contained in:
Wenjie Zhang 2025-04-09 12:26:37 +08:00 committed by GitHub
commit dd40f90e49
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 104 additions and 40 deletions

View File

@ -171,32 +171,34 @@ def chat_agent(agent_name: str,
**kwargs **kwargs
}, ensure_ascii=False).encode('utf-8') + b"\n" }, ensure_ascii=False).encode('utf-8') + b"\n"
try:
agent = agent_manager.get_runnable_agent(agent_name)
except Exception as e:
logger.error(f"Error getting agent {agent_name}: {e}")
return StreamingResponse(make_chunk(message=f"Error getting agent {agent_name}: {e}", status="error"), media_type='application/json')
# 从config中获取history_round
history_round = config.get("history_round")
history_manager = HistoryManager(history)
messages = history_manager.get_history_with_msg(query, max_rounds=history_round)
history_manager.add_user(query)
# 如果没有thread_id则生成一个
if "thread_id" not in config or not config["thread_id"]:
config["thread_id"] = str(uuid.uuid4())
# 构造运行时配置
runnable_config = {
"configurable": {
**config
}
}
def stream_messages(): def stream_messages():
content = ""
# 代表服务端已经收到了请求
yield make_chunk(status="init", meta=meta) yield make_chunk(status="init", meta=meta)
try:
agent = agent_manager.get_runnable_agent(agent_name)
except Exception as e:
logger.error(f"Error getting agent {agent_name}: {e}, {traceback.format_exc()}")
yield make_chunk(message=f"Error getting agent {agent_name}: {e}", status="error")
return
# 从config中获取history_round
history_round = config.get("history_round")
history_manager = HistoryManager(history)
messages = history_manager.get_history_with_msg(query, max_rounds=history_round)
history_manager.add_user(query)
# 构造运行时配置如果没有thread_id则生成一个
if "thread_id" not in config or not config["thread_id"]:
config["thread_id"] = str(uuid.uuid4())
runnable_config = {"configurable": {**config}}
content = ""
try: try:
for msg, metadata in agent.stream_messages(messages, config_schema=runnable_config): for msg, metadata in agent.stream_messages(messages, config_schema=runnable_config):
if isinstance(msg, AIMessageChunk) and msg.content != "<tool_call>": if isinstance(msg, AIMessageChunk) and msg.content != "<tool_call>":
@ -214,7 +216,7 @@ def chat_agent(agent_name: str,
history=history_manager.update_ai(content), history=history_manager.update_ai(content),
meta=meta) meta=meta)
except Exception as e: except Exception as e:
logger.error(f"Error streaming messages: {e}") logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}")
yield make_chunk(message=f"Error streaming messages: {e}", status="error") yield make_chunk(message=f"Error streaming messages: {e}", status="error")
return StreamingResponse(stream_messages(), media_type='application/json') return StreamingResponse(stream_messages(), media_type='application/json')

View File

@ -3,6 +3,7 @@ from __future__ import annotations
import os import os
from typing import Type, Annotated, Optional, TypedDict from typing import Type, Annotated, Optional, TypedDict
from enum import Enum
from abc import abstractmethod from abc import abstractmethod
from dataclasses import dataclass, fields, field from dataclasses import dataclass, fields, field
@ -96,7 +97,7 @@ class BaseAgent():
return return
for requirement in self.requirements: for requirement in self.requirements:
if requirement not in os.environ: if requirement not in os.environ:
raise ValueError(f"{requirement} is not set") raise ValueError(f"没有配置{requirement} 环境变量,请在 src/.env 文件中配置,并重新启动服务")
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs): def stream_values(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
graph = self.get_graph(config_schema=config_schema, **kwargs) graph = self.get_graph(config_schema=config_schema, **kwargs)

View File

@ -47,7 +47,7 @@ class Config(SimpleConfig):
self.add_item("enable_reranker", default=False, des="是否开启重排序") self.add_item("enable_reranker", default=False, des="是否开启重排序")
self.add_item("enable_knowledge_base", default=False, des="是否开启知识库") self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱") self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱")
self.add_item("enable_web_search", default=False, des="是否开启网页搜索(需配置 TAVILY_API_KEY") self.add_item("enable_web_search", default=False, des="是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项")
# 模型配置 # 模型配置
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR ## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR
@ -150,9 +150,13 @@ class Config(SimpleConfig):
self.model_provider_status[provider] = all(conds_bool) self.model_provider_status[provider] = all(conds_bool)
# 检查web_search的环境变量 # 检查web_search的环境变量
if self.enable_web_search and not os.getenv("TAVILY_API_KEY"): # if self.enable_web_search and not os.getenv("TAVILY_API_KEY"):
logger.warning("TAVILY_API_KEY not set, web search will be disabled") # logger.warning("TAVILY_API_KEY not set, web search will be disabled")
self.enable_web_search = False # self.enable_web_search = False
# 2025.04.08 修改为不手动配置只要配置了TAVILY_API_KEY就默认开启web_search
if os.getenv("TAVILY_API_KEY"):
self.enable_web_search = True
self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v] self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v]
assert len(self.valuable_model_provider) > 0, f"No model provider available, please check your `.env` file. API_KEY_LIST: {conds}" assert len(self.valuable_model_provider) > 0, f"No model provider available, please check your `.env` file. API_KEY_LIST: {conds}"

View File

@ -1,18 +1,19 @@
import os import os
import requests
from openai import OpenAI from openai import OpenAI
from src.utils import logger, get_docker_safe_url from src.utils import logger, get_docker_safe_url
from langchain_openai import ChatOpenAI from langchain_openai import ChatOpenAI
class OpenAIBase(): class OpenAIBase():
def __init__(self, api_key, base_url, model_name): def __init__(self, api_key, base_url, model_name, **kwargs):
self.api_key = api_key self.api_key = api_key
self.base_url = base_url self.base_url = base_url
self.client = OpenAI(api_key=api_key, base_url=base_url) self.client = OpenAI(api_key=api_key, base_url=base_url)
self.model_name = model_name self.model_name = model_name
self.info = kwargs
self.chat_open_ai = ChatOpenAI(model=model_name, self.chat_open_ai = ChatOpenAI(model=model_name,
api_key=api_key, api_key=api_key,
base_url=base_url, base_url=base_url)
temperature=0.7)
def predict(self, message, stream=False): def predict(self, message, stream=False):
if isinstance(message, str): if isinstance(message, str):
@ -54,6 +55,49 @@ class OpenAIBase():
logger.error(f"Error getting models: {e}") logger.error(f"Error getting models: {e}")
return [] return []
def _get_model_by_model_url(self, model_url):
"""
Refs: https://docs.together.ai/reference/models-1
Return: [
{
"id": "meta-llama/Meta-Llama-3-70B-Instruct-Turbo",
"object": "model",
"created": 0,
"type": "chat",
"running": false,
"display_name": "Meta Llama 3 70B Instruct Turbo",
"organization": "Meta",
"link": "https://huggingface.co/meta-llama/Meta-Llama-3-70B-Instruct",
"license": "Llama-3 (Other)",
"context_length": 8192,
"config": {
"chat_template": "{% set loop_messages = messages %}{% for message in loop_messages %}{% set content = '<|start_header_id|>' + message['role'] + '<|end_header_id|>\n\n'+ message['content'] | trim + '<|eot_id|>' %}{% if loop.index0 == 0 %}{% set content = bos_token + content %}{% endif %}{{ content }}{% endfor %}{{ '<|start_header_id|>assistant<|end_header_id|>\n\n' }}",
"stop": [
"<|eot_id|>"
],
"bos_token": "<|begin_of_text|>",
"eos_token": "<|end_of_text|>"
},
"pricing": {
"hourly": 0,
"input": 0.88,
"output": 0.88,
"base": 0,
"finetune": 0
}
},
]
"""
headers = {
"accept": "application/json",
"authorization": f"Bearer {self.api_key}"
}
response = requests.get(model_url, headers=headers)
return response.json()
class OpenModel(OpenAIBase): class OpenModel(OpenAIBase):
def __init__(self, model_name=None): def __init__(self, model_name=None):

View File

@ -511,6 +511,18 @@ const handleStreamResponse = async (response) => {
} }
} catch (error) { } catch (error) {
console.error('流式处理出错:', error); console.error('流式处理出错:', error);
const lastMsg = messages.value[messages.value.length - 1];
if (lastMsg.role === 'assistant') {
lastMsg.status = 'error';
lastMsg.message = error.message;
} else {
messages.value.push({
role: 'assistant',
message: `发生错误: ${error.message}`,
status: 'error'
});
await scrollToBottom();
}
isProcessing.value = false; isProcessing.value = false;
} }
}; };

View File

@ -48,7 +48,7 @@
previewTheme="github" previewTheme="github"
:showCodeRowNumber="false" :showCodeRowNumber="false"
:modelValue="message.content" :modelValue="message.content"
:key="index" :key="message.id"
class="message-md"/> class="message-md"/>
<div v-if="message.isStoppedByUser" class="retry-hint"> <div v-if="message.isStoppedByUser" class="retry-hint">

View File

@ -6,7 +6,7 @@
</template> </template>
<template #actions> <template #actions>
<a-button :type="isNeedRestart ? 'primary' : 'default'" @click="sendRestart" :icon="h(ReloadOutlined)"> <a-button :type="isNeedRestart ? 'primary' : 'default'" @click="sendRestart" :icon="h(ReloadOutlined)">
{{ isNeedRestart ? '需要重启' : '重启服务' }} {{ isNeedRestart ? '需要刷新' : '重新加载' }}
</a-button> </a-button>
</template> </template>
</HeaderComponent> </HeaderComponent>
@ -63,10 +63,11 @@
</div> </div>
<div class="card"> <div class="card">
<span class="label">{{ items?.enable_web_search.des }}</span> <span class="label">{{ items?.enable_web_search.des }}</span>
<a-switch <!-- <a-switch
:checked="configStore.config.enable_web_search" :checked="configStore.config.enable_web_search"
@change="handleChange('enable_web_search', !configStore.config.enable_web_search)" @change="handleChange('enable_web_search', !configStore.config.enable_web_search)"
/> /> -->
<a-switch :checked="configStore.config.enable_web_search" />
</div> </div>
<div class="card"> <div class="card">
<span class="label">{{ items?.enable_reranker.des }}</span> <span class="label">{{ items?.enable_reranker.des }}</span>
@ -264,7 +265,7 @@
</a-checkbox-group> </a-checkbox-group>
</div> </div>
<div v-if="providerConfig.allModels.length === 0" class="modal-no-models"> <div v-if="providerConfig.allModels.length === 0" class="modal-no-models">
<a-alert v-if="!modelStatus[providerConfig.provider]" type="warning" message="请在 src/.env 中配置对应的 APIKEY" /> <a-alert v-if="!modelStatus[providerConfig.provider]" type="warning" message="请在 src/.env 中配置对应的 APIKEY,并重新启动服务" />
<a-alert v-else type="warning" message="该提供商暂未适配获取模型列表的方法,如果需要添加模型,请在 src/static/models.private.yml 中添加。" /> <a-alert v-else type="warning" message="该提供商暂未适配获取模型列表的方法,如果需要添加模型,请在 src/static/models.private.yml 中添加。" />
</div> </div>
</div> </div>
@ -387,11 +388,11 @@ const handleChange = (key, e) => {
if (!isNeedRestart.value) { if (!isNeedRestart.value) {
isNeedRestart.value = true isNeedRestart.value = true
notification.info({ notification.info({
message: '需要重启服务', message: '需要重新加载模型',
description: '请点击右下角按钮重启服务', description: '请点击右下角按钮重新加载模型',
placement: 'topLeft', placement: 'topLeft',
duration: 0, duration: 0,
btn: h(Button, { type: 'primary', onClick: sendRestart }, '立即重') btn: h(Button, { type: 'primary', onClick: sendRestart }, '立即重新加载')
}) })
} }
} }
@ -543,7 +544,7 @@ const openProviderConfig = (provider) => {
const saveProviderConfig = async () => { const saveProviderConfig = async () => {
if (!modelStatus.value[providerConfig.provider]) { if (!modelStatus.value[providerConfig.provider]) {
message.error('请在 src/.env 中配置对应的 APIKEY') message.error('请在 src/.env 中配置对应的 APIKEY,并重新启动服务')
return return
} }