docs: update multi-agent tutorials to use Command (#2643)

This commit is contained in:
Vadym Barda
2024-12-06 15:18:12 +00:00
committed by GitHub
parent 93e4c8cc1f
commit 5fa80e2a92
3 changed files with 159 additions and 229 deletions
File diff suppressed because one or more lines are too long
@@ -289,15 +289,10 @@
"from langchain_core.language_models.chat_models import BaseChatModel\n",
"\n",
"from langgraph.graph import StateGraph, MessagesState, START, END\n",
"from langgraph.types import Command\n",
"from langchain_core.messages import HumanMessage, trim_messages\n",
"\n",
"\n",
"# The agent state is the input to each node in the graph\n",
"class AgentState(MessagesState):\n",
" # The 'next' field indicates where to route to next\n",
" next: str\n",
"\n",
"\n",
"def make_supervisor_node(llm: BaseChatModel, members: list[str]) -> str:\n",
" options = [\"FINISH\"] + members\n",
" system_prompt = (\n",
@@ -313,17 +308,17 @@
"\n",
" next: Literal[*options]\n",
"\n",
" def supervisor_node(state: MessagesState) -> MessagesState:\n",
" def supervisor_node(state: MessagesState) -> Command[Literal[*members, \"__end__\"]]:\n",
" \"\"\"An LLM-based router.\"\"\"\n",
" messages = [\n",
" {\"role\": \"system\", \"content\": system_prompt},\n",
" ] + state[\"messages\"]\n",
" response = llm.with_structured_output(Router).invoke(messages)\n",
" next_ = response[\"next\"]\n",
" if next_ == \"FINISH\":\n",
" next_ = END\n",
" goto = response[\"next\"]\n",
" if goto == \"FINISH\":\n",
" goto = END\n",
"\n",
" return {\"next\": next_}\n",
" return Command(goto=goto)\n",
"\n",
" return supervisor_node"
]
@@ -363,25 +358,33 @@
"search_agent = create_react_agent(llm, tools=[tavily_tool])\n",
"\n",
"\n",
"def search_node(state: AgentState) -> AgentState:\n",
"def search_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" result = search_agent.invoke(state)\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"search\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"search\")\n",
" ]\n",
" },\n",
" # We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"web_scraper_agent = create_react_agent(llm, tools=[scrape_webpages])\n",
"\n",
"\n",
"def web_scraper_node(state: AgentState) -> AgentState:\n",
"def web_scraper_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" result = web_scraper_agent.invoke(state)\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"web_scraper\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"web_scraper\")\n",
" ]\n",
" },\n",
" # We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"research_supervisor_node = make_supervisor_node(llm, [\"search\", \"web_scraper\"])"
@@ -412,14 +415,7 @@
"research_builder.add_node(\"search\", search_node)\n",
"research_builder.add_node(\"web_scraper\", web_scraper_node)\n",
"\n",
"# Define the control flow\n",
"research_builder.add_edge(START, \"supervisor\")\n",
"# We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
"research_builder.add_edge(\"search\", \"supervisor\")\n",
"research_builder.add_edge(\"web_scraper\", \"supervisor\")\n",
"# Add the edges where routing applies\n",
"research_builder.add_conditional_edges(\"supervisor\", lambda state: state[\"next\"])\n",
"\n",
"research_graph = research_builder.compile()"
]
},
@@ -532,13 +528,17 @@
")\n",
"\n",
"\n",
"def doc_writing_node(state: AgentState) -> AgentState:\n",
"def doc_writing_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" result = doc_writer_agent.invoke(state)\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"doc_writer\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"doc_writer\")\n",
" ]\n",
" },\n",
" # We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"note_taking_agent = create_react_agent(\n",
@@ -551,13 +551,17 @@
")\n",
"\n",
"\n",
"def note_taking_node(state: AgentState) -> AgentState:\n",
"def note_taking_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" result = note_taking_agent.invoke(state)\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"note_taker\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"note_taker\")\n",
" ]\n",
" },\n",
" # We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"chart_generating_agent = create_react_agent(\n",
@@ -565,13 +569,19 @@
")\n",
"\n",
"\n",
"def chart_generating_node(state: AgentState) -> AgentState:\n",
"def chart_generating_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" result = chart_generating_agent.invoke(state)\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=result[\"messages\"][-1].content, name=\"chart_generator\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(\n",
" content=result[\"messages\"][-1].content, name=\"chart_generator\"\n",
" )\n",
" ]\n",
" },\n",
" # We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"doc_writing_supervisor_node = make_supervisor_node(\n",
@@ -600,21 +610,13 @@
"outputs": [],
"source": [
"# Create the graph here\n",
"paper_writing_builder = StateGraph(AgentState)\n",
"paper_writing_builder = StateGraph(MessagesState)\n",
"paper_writing_builder.add_node(\"supervisor\", doc_writing_supervisor_node)\n",
"paper_writing_builder.add_node(\"doc_writer\", doc_writing_node)\n",
"paper_writing_builder.add_node(\"note_taker\", note_taking_node)\n",
"paper_writing_builder.add_node(\"chart_generator\", chart_generating_node)\n",
"\n",
"# Define the control flow\n",
"paper_writing_builder.add_edge(START, \"supervisor\")\n",
"# We want our workers to ALWAYS \"report back\" to the supervisor when done\n",
"paper_writing_builder.add_edge(\"doc_writer\", \"supervisor\")\n",
"paper_writing_builder.add_edge(\"note_taker\", \"supervisor\")\n",
"paper_writing_builder.add_edge(\"chart_generator\", \"supervisor\")\n",
"# Add the edges where routing applies\n",
"paper_writing_builder.add_conditional_edges(\"supervisor\", lambda state: state[\"next\"])\n",
"\n",
"paper_writing_graph = paper_writing_builder.compile()"
]
},
@@ -728,37 +730,41 @@
},
"outputs": [],
"source": [
"def call_research_team(state: AgentState) -> AgentState:\n",
"def call_research_team(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" response = research_graph.invoke({\"messages\": state[\"messages\"][-1]})\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=response[\"messages\"][-1].content, name=\"research_team\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(\n",
" content=response[\"messages\"][-1].content, name=\"research_team\"\n",
" )\n",
" ]\n",
" },\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"def call_paper_writing_team(state: AgentState) -> AgentState:\n",
"def call_paper_writing_team(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
" response = paper_writing_graph.invoke({\"messages\": state[\"messages\"][-1]})\n",
" return {\n",
" \"messages\": [\n",
" HumanMessage(content=response[\"messages\"][-1].content, name=\"writing_team\")\n",
" ]\n",
" }\n",
" return Command(\n",
" update={\n",
" \"messages\": [\n",
" HumanMessage(\n",
" content=response[\"messages\"][-1].content, name=\"writing_team\"\n",
" )\n",
" ]\n",
" },\n",
" goto=\"supervisor\",\n",
" )\n",
"\n",
"\n",
"# Define the graph.\n",
"super_builder = StateGraph(AgentState)\n",
"super_builder = StateGraph(MessagesState)\n",
"super_builder.add_node(\"supervisor\", teams_supervisor_node)\n",
"super_builder.add_node(\"research_team\", call_research_team)\n",
"super_builder.add_node(\"writing_team\", call_paper_writing_team)\n",
"\n",
"# Define the control flow\n",
"super_builder.add_edge(START, \"supervisor\")\n",
"# We want our teams to ALWAYS \"report back\" to the top-level supervisor when done\n",
"super_builder.add_edge(\"research_team\", \"supervisor\")\n",
"super_builder.add_edge(\"writing_team\", \"supervisor\")\n",
"# Add the edges where routing applies\n",
"super_builder.add_conditional_edges(\"supervisor\", lambda state: state[\"next\"])\n",
"super_graph = super_builder.compile()"
]
},
File diff suppressed because one or more lines are too long