优化了一大堆:(1)添加了一个适配 LangGraph 的 graph 的类,这样可以更加方便的调试 Agent 了。(2)增加了 Agent call 的报错体验,但是前端还没有加进去。(3)移除部分无关紧要的 logger

This commit is contained in:
Wenjie Zhang 2025-03-31 22:20:29 +08:00
parent f5302c4c4b
commit e43f1cd618
11 changed files with 127 additions and 103 deletions

View File

@ -1,19 +1,9 @@
from dataclasses import dataclass, field
from datetime import datetime, timezone
from langchain_core.tools import tool
from langchain_openai import ChatOpenAI
from src.agents.registry import Configuration
def get_default_requirements():
return ["TAVILY_API_KEY"]
@tool
def multiply(first_int: int, second_int: int) -> int:
"""Multiply two integers together."""
return first_int * second_int
@dataclass(kw_only=True)
class ChatbotConfiguration(Configuration):
"""Chatbot 的配置"""

View File

@ -12,26 +12,28 @@ from langchain_community.tools.tavily_search import TavilySearchResults
from src.agents.registry import State, BaseAgent
from src.agents.utils import load_chat_model
from src.agents.chatbot.configuration import ChatbotConfiguration, multiply
from src.agents.tools_factory import multiply, add, subtract, divide
from src.agents.chatbot.configuration import ChatbotConfiguration
class ChatbotAgent(BaseAgent):
name = "chatbot"
description = "A chatbot that can answer questions and help with tasks."
requirements = ["TAVILY_API_KEY", "ZHIPUAI_API_KEY"]
_graph_cache = None
config_schema = ChatbotConfiguration.to_dict()
config_schema = ChatbotConfiguration
def __init__(self, **kwargs):
super().__init__(**kwargs)
def _get_tools(self, config_schema: RunnableConfig):
"""根据配置获取工具"""
tools = [multiply, TavilySearchResults(max_results=10)]
tools = [multiply, add, subtract, divide, TavilySearchResults(max_results=10)]
return tools
def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
"""调用 llm 模型"""
config_schema = config or {}
conf = ChatbotConfiguration.from_runnable_config(config_schema)
conf = self.config_schema.from_runnable_config(config_schema)
model = load_chat_model(conf.model)
model_with_tools = model.bind_tools(self._get_tools(config_schema))
@ -40,9 +42,9 @@ class ChatbotAgent(BaseAgent):
)
return {"messages": [res]}
def get_graph(self, config_schema: RunnableConfig = None):
def get_graph(self, config_schema: RunnableConfig = None, **kwargs):
"""构建图"""
workflow = StateGraph(State, config_schema=ChatbotConfiguration)
workflow = StateGraph(State, config_schema=self.config_schema)
workflow.add_node("chatbot", self.llm_call)
workflow.add_node("tools", ToolNode(tools=self._get_tools(config_schema)))
workflow.add_edge(START, "chatbot")
@ -56,20 +58,6 @@ class ChatbotAgent(BaseAgent):
graph = workflow.compile(checkpointer=MemorySaver())
return graph
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None):
graph = self.get_graph(config_schema)
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema):
yield event["messages"]
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None):
graph = self.get_graph(config_schema)
conf = ChatbotConfiguration.from_runnable_config(config_schema)
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema):
msg_type = msg.type
return_keys =conf.return_keys
if not return_keys or msg_type in return_keys:
yield msg, metadata
def main():
agent = ChatbotAgent(ChatbotConfiguration())

View File

@ -17,60 +17,12 @@ from src.agents.react.configuration import ReActConfiguration, multiply
class ReActAgent(BaseAgent):
name = "react"
description = "A react agent that can answer questions and help with tasks."
_graph_cache = None
config_schema = ReActConfiguration.to_dict()
def __init__(self, **kwargs):
super().__init__(**kwargs)
def _get_tools(self, config_schema: RunnableConfig):
"""根据配置获取工具"""
tools = [multiply, TavilySearchResults(max_results=10)]
return tools
def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
"""调用 llm 模型"""
config_schema = config or {}
conf = ReActConfiguration.from_runnable_config(config_schema)
model = load_chat_model(conf.model)
model_with_tools = model.bind_tools(self._get_tools(config_schema))
res = model_with_tools.invoke(
[{"role": "system", "content": conf.system_prompt}, *state["messages"]]
)
return {"messages": [res]}
def get_graph(self, config_schema: RunnableConfig = None):
def get_graph(self, **kwargs):
"""构建图"""
workflow = StateGraph(State, config_schema=ReActConfiguration)
workflow.add_node("react", self.llm_call)
workflow.add_node("tools", ToolNode(tools=self._get_tools(config_schema)))
workflow.add_edge(START, "react")
workflow.add_conditional_edges(
"react",
tools_condition,
)
workflow.add_edge("tools", "react")
workflow.add_edge("react", END)
graph = workflow.compile(checkpointer=MemorySaver())
from .workflows import graph
return graph
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None):
graph = self.get_graph(config_schema)
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema):
yield event["messages"]
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None):
graph = self.get_graph(config_schema)
conf = ReActConfiguration.from_runnable_config(config_schema)
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema):
msg_type = msg.type
return_keys =conf.return_keys
if not return_keys or msg_type in return_keys:
yield msg, metadata
def main():
agent = ReActAgent(ReActConfiguration())

View File

@ -0,0 +1,36 @@
import os
from langchain_openai import ChatOpenAI
model = ChatOpenAI(model="glm-4-plus",
api_key=os.getenv("ZHIPUAI_API_KEY"),
base_url="https://open.bigmodel.cn/api/paas/v4/",
temperature=0)
# For this tutorial we will use custom tool that returns pre-defined values for weather in two cities (NYC & SF)
from typing import Literal
from langchain_core.tools import tool
@tool
def get_weather(city: Literal["nyc", "sf"]):
"""Use this to get weather information."""
if city == "nyc":
return "It might be cloudy in nyc"
elif city == "sf":
return "It's always sunny in sf"
else:
raise AssertionError("Unknown city")
tools = [get_weather]
# Define the graph
from langgraph.prebuilt import create_react_agent
graph = create_react_agent(model, tools=tools)

View File

@ -1,5 +1,7 @@
from __future__ import annotations
import os
from typing import Type, Annotated, Optional, TypedDict
from abc import abstractmethod
from dataclasses import dataclass, fields, field
@ -46,9 +48,40 @@ class Configuration(SimpleConfig):
class BaseAgent():
"""
定义一个基础 Agent 各类 graph 继承
"""
name: str = field(default="base_agent")
description: str = field(default="base_agent")
config_schema: Configuration = Configuration
requirements: list[str]
def __init__(self, **kwargs):
pass
self.check_requirements()
def check_requirements(self):
if not hasattr(self, "requirements") or not self.requirements:
return
for requirement in self.requirements:
if requirement not in os.environ:
raise ValueError(f"{requirement} is not set")
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
graph = self.get_graph(config_schema=config_schema, **kwargs)
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema):
yield event["messages"]
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
graph = self.get_graph(config_schema=config_schema, **kwargs)
conf = self.config_schema.from_runnable_config(config_schema)
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema):
msg_type = msg.type
return_keys =conf.return_keys
if not return_keys or msg_type in return_keys:
yield msg, metadata
@abstractmethod
def get_graph(self) -> CompiledStateGraph:
def get_graph(self, **kwargs) -> CompiledStateGraph:
pass

View File

@ -5,8 +5,7 @@ from typing import Any, Callable, Optional, Type, Union
from pydantic import BaseModel, Field
from langchain.agents import tool
from langchain_core.tools import Tool, BaseTool
from langchain_core.tools import tool, BaseTool
_TOOLS_REGISTRY = {}
@ -95,3 +94,24 @@ class BaseToolOutput:
else:
return str(self.data)
@tool
def multiply(first_int: int, second_int: int) -> int:
"""Multiply two integers together."""
return first_int * second_int
@tool
def add(first_int: int, second_int: int) -> int:
"""Add two integers together."""
return first_int + second_int
@tool
def subtract(first_int: int, second_int: int) -> int:
"""Subtract two integers."""
return first_int - second_int
@tool
def divide(first_int: int, second_int: int) -> int:
"""Divide two integers."""
return first_int / second_int

View File

@ -229,7 +229,7 @@ class Retriever:
return formatted_results
def format_query_results(self, results):
logger.debug(f"Graph Query Results: {results}")
# logger.debug(f"Graph Query Results: {results}")
formatted_results = {"nodes": [], "edges": []}
node_dict = {}

View File

@ -1,3 +1,4 @@
import os
import json
import asyncio
import traceback
@ -126,7 +127,7 @@ async def get_agent():
agents = [{
"name": agent.name,
"description": agent.description,
"config_schema": agent.config_schema
"config_schema": agent.config_schema.to_dict()
} for agent in agent_manager.agents.values()]
return {"agents": agents}
@ -138,19 +139,6 @@ def chat_agent(agent_name: str,
thread_id: str | None = Body(None)):
meta["server_model_name"] = agent_name
agent = agent_manager.get_runnable_agent(agent_name)
history_manager = HistoryManager(history)
messages = history_manager.get_history_with_msg(query, max_rounds=meta.get('history_round'))
history_manager.add_user(query) # 注意这里使用原始查询
runnable_config = {
"configurable": {
"thread_id": thread_id or str(uuid.uuid4()),
"use_web": meta.get("use_web", False),
"return_keys": []
}
}
def make_chunk(content=None, **kwargs):
return json.dumps({
@ -160,6 +148,23 @@ def chat_agent(agent_name: str,
**kwargs
}, 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')
history_manager = HistoryManager(history)
messages = history_manager.get_history_with_msg(query, max_rounds=meta.get('history_round'))
history_manager.add_user(query) # 注意这里使用原始查询
runnable_config = {
"configurable": {
"thread_id": thread_id or str(uuid.uuid4()),
"return_keys": []
}
}
def stream_messages():
content = ""
yield make_chunk(status="waiting")
@ -178,4 +183,4 @@ def chat_agent(agent_name: str,
yield make_chunk(status="finished", history=history_manager.update_ai(content))
return StreamingResponse(stream_messages(), media_type='application/json')
return StreamingResponse(stream_messages(), media_type='application/json')

View File

@ -72,7 +72,7 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod
@data.post("/add-by-chunks")
async def add_by_chunks(db_id: str = Body(...), file_chunks: dict = Body(...)):
logger.debug(f"Add chunks in {db_id}: {file_chunks}")
# logger.debug(f"Add chunks in {db_id}: {len(file_chunks)} chunks")
try:
loop = asyncio.get_event_loop()
await loop.run_in_executor(

View File

@ -50,7 +50,7 @@ async def route_index():
description=agent.description,
url=f"/agent/{agent.name}",
method="POST",
metadata=agent.config_schema,
metadata=agent.config_schema.to_dict(),
)
)

View File

@ -527,7 +527,7 @@ const updateExistingMessage = async (data, existingMsgIndex) => {
// loadingprocessing
if (msgInstance.status === 'loading') {
msgInstance.status = 'processing';
msgInstance.status = 'loading';
msgInstance.run_id = data.metadata?.run_id;
msgInstance.step = data.metadata?.langgraph_step;
}