refactor: 重构智能体配置管理,移除不必要的文件操作和API接口

This commit is contained in:
Wenjie Zhang 2026-03-18 22:55:52 +08:00
parent 1ae2974fad
commit 3c75b0acb3
9 changed files with 62 additions and 248 deletions

View File

@ -30,28 +30,35 @@ class ChatbotAgent(BaseAgent):
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
async def get_graph(self, **kwargs): async def _build_middlewares(self, context):
"""构建图""" """构建中间件列表"""
context = self.context_schema()
all_mcp_tools = ( all_mcp_tools = (
await get_tools_from_all_servers() await get_tools_from_all_servers()
) # 因为异步加载,无法放在 RuntimeConfigMiddleware 的 __init__ 中 ) # 因为异步加载,无法放在 RuntimeConfigMiddleware 的 __init__ 中
middlewares = [
save_attachments_to_fs, # 附件注入提示词
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
KnowledgeBaseMiddleware(), # 知识库工具
RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/MCP/提示词)
SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活)
ModelRetryMiddleware(), # 模型重试中间件
TodoListMiddleware(),
PatchToolCallsMiddleware(),
]
return middlewares
async def get_graph(self, context=None, **kwargs):
context = context or self.context_schema() # 获取上下文配置
# 使用 create_agent 创建智能体 # 使用 create_agent 创建智能体
# 注意tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置 # 注意tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置
graph = create_agent( graph = create_agent(
model=load_chat_model(fully_specified_name=context.model), model=load_chat_model(fully_specified_name=context.model),
system_prompt=context.system_prompt, system_prompt=context.system_prompt,
middleware=[ middleware=await self._build_middlewares(context),
save_attachments_to_fs, # 附件注入提示词
FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端
KnowledgeBaseMiddleware(), # 知识库工具
RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/MCP/提示词)
SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活)
ModelRetryMiddleware(), # 模型重试中间件
TodoListMiddleware(),
PatchToolCallsMiddleware(),
],
checkpointer=await self._get_checkpointer(), checkpointer=await self._get_checkpointer(),
) )

View File

@ -56,10 +56,9 @@ class DeepAgent(BaseAgent):
logger.warning("No search tools configured, DeepAgent will work without web search") logger.warning("No search tools configured, DeepAgent will work without web search")
return tools return tools
async def get_graph(self, **kwargs): async def get_graph(self, context=None, **kwargs):
"""构建 Deep Agent 的图"""
# 获取上下文配置 context = context or self.context_schema() # 获取上下文配置
context = self.context_schema.from_file(module_name=self.module_name)
model = load_chat_model(context.model) model = load_chat_model(context.model)
sub_model = load_chat_model(context.subagents_model) sub_model = load_chat_model(context.subagents_model)

View File

@ -69,9 +69,10 @@ class SqlReporterAgent(BaseAgent):
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
async def get_graph(self, **kwargs): async def get_graph(self, context=None, **kwargs):
"""构建图"""
context = self.context_schema.from_file(module_name=self.module_name) context = context or self.context_schema() # 获取上下文配置
all_mcp_tools = await get_tools_from_all_servers() all_mcp_tools = await get_tools_from_all_servers()
graph = create_agent( graph = create_agent(

View File

@ -65,26 +65,20 @@ class BaseAgent:
} }
async def get_config(self): async def get_config(self):
return self.context_schema.from_file(module_name=self.module_name) return self.context_schema()
async def stream_values(self, messages: list[str], input_context=None, **kwargs): async def stream_values(self, messages: list[str], input_context=None, **kwargs):
graph = await self.get_graph()
context = self.context_schema() context = self.context_schema()
agent_config = (input_context or {}).get("agent_config") context.update_from_dict(input_context or {})
if isinstance(agent_config, dict): graph = await self.get_graph(context=context)
context.update(agent_config)
context.update(input_context or {})
for event in graph.astream({"messages": messages}, stream_mode="values", context=context): for event in graph.astream({"messages": messages}, stream_mode="values", context=context):
yield event["messages"] yield event["messages"]
async def stream_messages(self, messages: list[str], input_context=None, **kwargs): async def stream_messages(self, messages: list[str], input_context=None, **kwargs):
graph = await self.get_graph()
context = self.context_schema() context = self.context_schema()
agent_config = (input_context or {}).get("agent_config") context.update_from_dict(input_context or {})
if isinstance(agent_config, dict): graph = await self.get_graph(context=context)
context.update(agent_config) logger.debug(f"stream_messages: {context=}")
context.update(input_context or {})
logger.debug(f"stream_messages: {context}")
# 构建配置LangGraph 会自动从 checkpointer 恢复 state # 构建配置LangGraph 会自动从 checkpointer 恢复 state
input_config = { input_config = {
@ -101,12 +95,9 @@ class BaseAgent:
yield msg, metadata yield msg, metadata
async def invoke_messages(self, messages: list[str], input_context=None, **kwargs): async def invoke_messages(self, messages: list[str], input_context=None, **kwargs):
graph = await self.get_graph()
context = self.context_schema() context = self.context_schema()
agent_config = (input_context or {}).get("agent_config") context.update_from_dict(input_context or {})
if isinstance(agent_config, dict): graph = await self.get_graph(context=context)
context.update(agent_config)
context.update(input_context or {})
logger.debug(f"invoke_messages: {context}") logger.debug(f"invoke_messages: {context}")
# 构建配置 # 构建配置

View File

@ -1,16 +1,11 @@
"""Define the configurable parameters for the agent.""" """Define the configurable parameters for the agent."""
import os
import uuid import uuid
from dataclasses import MISSING, dataclass, field, fields from dataclasses import MISSING, dataclass, field, fields
from pathlib import Path
from typing import Annotated, get_args, get_origin from typing import Annotated, get_args, get_origin
import yaml
from yuxi import config as sys_config from yuxi import config as sys_config
from yuxi.services.mcp_service import get_mcp_server_names from yuxi.services.mcp_service import get_mcp_server_names
from yuxi.utils import logger
@dataclass(kw_only=True) @dataclass(kw_only=True)
@ -20,8 +15,7 @@ class BaseContext:
配置优先级: 配置优先级:
1. 运行时配置(RunnableConfig)最高优先级直接从函数参数传入 1. 运行时配置(RunnableConfig)最高优先级直接从函数参数传入
2. 文件配置(config.private.yaml)中等优先级从文件加载 2. 类默认配置最低优先级类中定义的默认值
3. 类默认配置最低优先级类中定义的默认值
""" """
def update(self, data: dict): def update(self, data: dict):
@ -40,10 +34,11 @@ class BaseContext:
metadata={"name": "用户ID", "configurable": False, "description": "用来唯一标识一个用户"}, metadata={"name": "用户ID", "configurable": False, "description": "用来唯一标识一个用户"},
) )
department_id: int | None = field( # 不需要了,使用 user_id 判断
default=None, # department_id: int | None = field(
metadata={"name": "部门ID", "configurable": False, "description": "用来唯一标识一个部门"}, # default=None,
) # metadata={"name": "部门ID", "configurable": False, "description": "用来唯一标识一个部门"},
# )
system_prompt: Annotated[str, {"__template_metadata__": {"kind": "prompt"}}] = field( system_prompt: Annotated[str, {"__template_metadata__": {"kind": "prompt"}}] = field(
default="You are a helpful assistant.", default="You are a helpful assistant.",
@ -99,50 +94,6 @@ class BaseContext:
}, },
) )
@classmethod
def from_file(cls, module_name: str, input_context: dict = None) -> "BaseContext":
"""Load configuration from a YAML file. 用于持久化配置"""
# 从文件加载配置
context = cls()
config_file_path = Path(sys_config.save_dir) / "agents" / module_name / "config.yaml"
if module_name is not None and os.path.exists(config_file_path):
file_config = {}
try:
with open(config_file_path, encoding="utf-8") as f:
file_config = yaml.safe_load(f) or {}
except Exception as e:
logger.error(f"加载智能体配置文件出错: {e}")
context.update(file_config)
if input_context:
context.update(input_context)
return context
@classmethod
def save_to_file(cls, config: dict, module_name: str) -> bool:
"""Save configuration to a YAML file 用于持久化配置"""
configurable_items = cls.get_configurable_items()
configurable_config = {}
for k, v in config.items():
if k in configurable_items:
configurable_config[k] = v
try:
config_file_path = Path(sys_config.save_dir) / "agents" / module_name / "config.yaml"
# 确保目录存在
os.makedirs(os.path.dirname(config_file_path), exist_ok=True)
with open(config_file_path, "w", encoding="utf-8") as f:
yaml.dump(configurable_config, f, indent=2, allow_unicode=True)
return True
except Exception as e:
logger.error(f"保存智能体配置文件出错: {e}")
return False
@classmethod @classmethod
def get_configurable_items(cls): def get_configurable_items(cls):
"""实现一个可配置的参数列表,在 UI 上配置时使用""" """实现一个可配置的参数列表,在 UI 上配置时使用"""
@ -210,3 +161,9 @@ class BaseContext:
if isinstance(metadata, dict) and "__template_metadata__" in metadata: if isinstance(metadata, dict) and "__template_metadata__" in metadata:
return metadata["__template_metadata__"] return metadata["__template_metadata__"]
return {} return {}
def update_from_dict(self, data: dict):
"""从字典更新配置字段"""
for key, value in data.items():
if hasattr(self, key):
setattr(self, key, value)

View File

@ -72,17 +72,6 @@ class RuntimeConfigMiddleware(AgentMiddleware):
"将忽略 extra_tools 并不会应用任何工具覆盖。" "将忽略 extra_tools 并不会应用任何工具覆盖。"
) )
logger.debug(
f"Initialized RuntimeConfigMiddleware with custom field names: model={model_context_name}, "
f"system_prompt={system_prompt_context_name}, tools={tools_context_name}, "
f"knowledges={knowledges_context_name}, mcps={mcps_context_name}"
)
async def abefore_agent(self, state, runtime) -> dict[str, Any] | None:
# abefore_agent 在 RuntimeConfigMiddleware 中暂无额外逻辑
# Skills 相关逻辑已移至 SkillsMiddleware
return None
async def awrap_model_call( async def awrap_model_call(
self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
) -> ModelResponse: ) -> ModelResponse:

View File

@ -14,6 +14,7 @@ from yuxi.plugins.guard import content_guard
from yuxi.repositories.agent_config_repository import AgentConfigRepository from yuxi.repositories.agent_config_repository import AgentConfigRepository
from yuxi.repositories.conversation_repository import ConversationRepository from yuxi.repositories.conversation_repository import ConversationRepository
from yuxi.storage.postgres.manager import pg_manager from yuxi.storage.postgres.manager import pg_manager
from yuxi.storage.postgres.models_business import User
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger
from yuxi.utils.question_utils import ( from yuxi.utils.question_utils import (
normalize_options as _normalize_interrupt_options, normalize_options as _normalize_interrupt_options,
@ -297,28 +298,23 @@ def _ensure_full_msg(full_msg: AIMessage | None, accumulated_content: list[str])
return full_msg return full_msg
async def _resolve_agent_config( async def _resolve_agent_config(db, agent_id: str, user: User, agent_config_id):
db, agent_id: str, department_id, user_id: str, agent_config_id: int | str | None """解析 agent_config返回 agent_config"""
) -> tuple: department_id = user.department_id
"""解析 agent_config返回 (config_item, agent_config_id)"""
config_repo = AgentConfigRepository(db) agent_config_repo = AgentConfigRepository(db)
config_item = None config_item = None
if agent_config_id is not None: if agent_config_id is not None:
try: config_item = await agent_config_repo.get_by_id(config_id=int(agent_config_id))
config_item = await config_repo.get_by_id(int(agent_config_id))
except Exception:
logger.warning(f"Failed to fetch agent config {agent_config_id}: {traceback.format_exc()}")
config_item = None
if config_item is not None and (config_item.department_id != department_id or config_item.agent_id != agent_id): if config_item is not None and (config_item.department_id != department_id or config_item.agent_id != agent_id):
config_item = None config_item = None
if config_item is None: if config_item is None:
config_item = await config_repo.get_or_create_default( config_item = await agent_config_repo.get_or_create_default(
department_id=department_id, agent_id=agent_id, created_by=user_id department_id=department_id, agent_id=agent_id, created_by=str(user.id)
) )
agent_config_id = config_item.id
return config_item, agent_config_id return (config_item.config_json or {}).get("context", {})
async def check_and_handle_interrupts( async def check_and_handle_interrupts(
@ -408,26 +404,14 @@ async def stream_agent_chat(
messages = [human_message] messages = [human_message]
user_id = str(current_user.id) user_id = str(current_user.id)
department_id = current_user.department_id
if not department_id:
yield make_chunk(status="error", error_type="no_department", error_message="当前用户未绑定部门", meta=meta)
return
agent_config_id = config.get("agent_config_id") agent_config_id = config.get("agent_config_id")
config_item, agent_config_id = await _resolve_agent_config(db, agent_id, department_id, user_id, agent_config_id) agent_config = await _resolve_agent_config(db, agent_id, current_user, agent_config_id)
if not (thread_id := config.get("thread_id")): if not (thread_id := config.get("thread_id")):
thread_id = str(uuid.uuid4()) thread_id = str(uuid.uuid4())
logger.warning(f"No thread_id provided, generated new thread_id: {thread_id}") logger.warning(f"No thread_id provided, generated new thread_id: {thread_id}")
agent_config = (config_item.config_json or {}).get("context", {}) input_context = agent_config | {"user_id": user_id, "thread_id": thread_id}
input_context = {
"user_id": user_id,
"thread_id": thread_id,
"department_id": department_id,
"agent_config_id": agent_config_id,
"agent_config": agent_config,
}
full_msg = None full_msg = None
accumulated_content: list[str] = [] accumulated_content: list[str] = []
@ -595,28 +579,12 @@ async def stream_agent_resume(
graph = await agent.get_graph() graph = await agent.get_graph()
user_id = str(current_user.id) user_id = str(current_user.id)
department_id = current_user.department_id
if not department_id:
yield make_resume_chunk(
status="error", error_type="no_department", error_message="当前用户未绑定部门", meta=meta
)
return
agent_config_id = (config or {}).get("agent_config_id") agent_config_id = (config or {}).get("agent_config_id")
config_item, agent_config_id = await _resolve_agent_config(db, agent_id, department_id, user_id, agent_config_id) agent_config = await _resolve_agent_config(db, agent_id, current_user, agent_config_id)
input_context = {
"user_id": user_id,
"thread_id": thread_id,
"department_id": department_id,
"agent_config_id": agent_config_id,
"agent_config": (config_item.config_json or {}).get("context", config_item.config_json or {}),
}
context = agent.context_schema() context = agent.context_schema()
agent_config = input_context.get("agent_config") context.update(agent_config or {})
if isinstance(agent_config, dict): context.update({"user_id": user_id, "thread_id": thread_id})
context.update(agent_config)
context.update(input_context)
stream_source = graph.astream( stream_source = graph.astream(
resume_command, resume_command,

View File

@ -593,64 +593,6 @@ async def resume_agent_chat(
) )
@chat.post("/agent/{agent_id}/config")
async def save_agent_config(
agent_id: str,
config: dict = Body(...),
reload_graph: bool = Query(True),
current_user: User = Depends(get_required_user),
):
"""保存智能体配置到YAML文件需要登录"""
try:
# 获取Agent实例和配置类
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
# === 校验知识库权限 ===
from yuxi import knowledge_base
if "knowledges" in config and config["knowledges"]:
# 获取用户有权访问的知识库名称
try:
accessible_databases = await knowledge_base.get_databases_by_user_id(current_user.user_id)
accessible_kb_names = {
db.get("name") for db in accessible_databases.get("databases", []) if db.get("name")
}
except Exception as db_error:
logger.warning(f"获取知识库列表失败: {db_error}")
# 如果获取失败superadmin 可以访问所有,非 superadmin 无法访问任何
if current_user.role != "superadmin":
raise HTTPException(status_code=500, detail="无法获取知识库列表")
# 回退:获取所有数据库名称
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
kb_repo = KnowledgeBaseRepository()
rows = await kb_repo.get_all()
accessible_kb_names = {row.name for row in rows if row.name}
# 检查配置中的知识库是否都可用
invalid_kbs = [kb for kb in config["knowledges"] if kb not in accessible_kb_names]
if invalid_kbs:
raise HTTPException(status_code=403, detail=f"无权访问以下知识库: {', '.join(invalid_kbs)}")
# === 校验结束 ===
# 使用配置类的save_to_file方法保存配置
result = agent.context_schema.save_to_file(config, agent.module_name)
if result:
if reload_graph:
agent_manager.get_agent(agent_id, reload_graph=True)
return {"success": True, "message": f"智能体 {agent.name} 配置已保存"}
else:
raise HTTPException(status_code=500, detail="保存智能体配置失败")
except HTTPException:
raise
except Exception as e:
logger.error(f"保存智能体配置出错: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"保存智能体配置出错: {str(e)}")
@chat.get("/agent/{agent_id}/history") @chat.get("/agent/{agent_id}/history")
async def get_agent_history( async def get_agent_history(
agent_id: str, thread_id: str, current_user: User = Depends(get_required_user), db: AsyncSession = Depends(get_db) agent_id: str, thread_id: str, current_user: User = Depends(get_required_user), db: AsyncSession = Depends(get_db)
@ -691,23 +633,6 @@ async def get_agent_state(
raise HTTPException(status_code=500, detail=f"获取AgentState出错: {str(e)}") raise HTTPException(status_code=500, detail=f"获取AgentState出错: {str(e)}")
@chat.get("/agent/{agent_id}/config")
async def get_agent_config(agent_id: str, current_user: User = Depends(get_required_user)):
"""从YAML文件加载智能体配置需要登录"""
try:
# 检查智能体是否存在
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
config = await agent.get_config()
logger.debug(f"config: {config}, ContextClass: {agent.context_schema=}")
return {"success": True, "config": config}
except Exception as e:
logger.error(f"加载智能体配置出错: {e}, {traceback.format_exc()}")
raise HTTPException(status_code=500, detail=f"加载智能体配置出错: {str(e)}")
# ==================== 线程管理 API ==================== # ==================== 线程管理 API ====================

View File

@ -3,7 +3,6 @@ import {
apiPost, apiPost,
apiDelete, apiDelete,
apiPut, apiPut,
apiAdminGet,
apiAdminPost, apiAdminPost,
apiAdminDelete, apiAdminDelete,
apiRequest apiRequest
@ -123,28 +122,6 @@ export const agentApi = {
updateProviderModels: (provider, models) => updateProviderModels: (provider, models) =>
apiPost(`/api/chat/models/update?model_provider=${provider}`, models), apiPost(`/api/chat/models/update?model_provider=${provider}`, models),
/**
* 获取智能体配置
* @param {string} agentName - 智能体名称
* @returns {Promise} - 智能体配置
*/
getAgentConfig: async (agentName) => {
return apiAdminGet(`/api/chat/agent/${agentName}/config`)
},
/**
* 保存智能体配置
* @param {string} agentName - 智能体名称
* @param {Object} config - 配置对象
* @param {Object} options - 额外参数 (e.g., { reload_graph: true })
* @returns {Promise} - 保存结果
*/
saveAgentConfig: async (agentName, config, options = {}) => {
const queryParams = new URLSearchParams(options).toString()
const url = `/api/chat/agent/${agentName}/config` + (queryParams ? `?${queryParams}` : '')
return apiAdminPost(url, config)
},
getAgentConfigs: (agentId) => apiGet(`/api/chat/agent/${agentId}/configs`), getAgentConfigs: (agentId) => apiGet(`/api/chat/agent/${agentId}/configs`),
getAgentConfigProfile: (agentId, configId) => getAgentConfigProfile: (agentId, configId) =>