docs: update state to include 'next' for supervisor notebooks (#3087)

This commit is contained in:
Vadym Barda
2025-01-17 14:32:15 -05:00
committed by GitHub
parent aa9d253978
commit 943dd28863
2 changed files with 27 additions and 19 deletions
@@ -83,7 +83,7 @@
},
{
"cell_type": "code",
"execution_count": 3,
"execution_count": 2,
"id": "f04c6778-403b-4b49-9b93-678e910d5cec",
"metadata": {},
"outputs": [],
@@ -126,7 +126,7 @@
},
{
"cell_type": "code",
"execution_count": 4,
"execution_count": 3,
"id": "df2bd80b-c477-4d74-8faa-1c0548622239",
"metadata": {},
"outputs": [],
@@ -162,7 +162,11 @@
"llm = ChatAnthropic(model=\"claude-3-5-sonnet-latest\")\n",
"\n",
"\n",
"def supervisor_node(state: MessagesState) -> Command[Literal[*members, \"__end__\"]]:\n",
"class State(MessagesState):\n",
" next: str\n",
"\n",
"\n",
"def supervisor_node(state: State) -> Command[Literal[*members, \"__end__\"]]:\n",
" messages = [\n",
" {\"role\": \"system\", \"content\": system_prompt},\n",
" ] + state[\"messages\"]\n",
@@ -171,7 +175,7 @@
" if goto == \"FINISH\":\n",
" goto = END\n",
"\n",
" return Command(goto=goto)"
" return Command(goto=goto, update={\"next\": goto})"
]
},
{
@@ -201,7 +205,7 @@
")\n",
"\n",
"\n",
"def research_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def research_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = research_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -217,7 +221,7 @@
"code_agent = create_react_agent(llm, tools=[python_repl_tool])\n",
"\n",
"\n",
"def code_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def code_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = code_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -229,7 +233,7 @@
" )\n",
"\n",
"\n",
"builder = StateGraph(MessagesState)\n",
"builder = StateGraph(State)\n",
"builder.add_edge(START, \"supervisor\")\n",
"builder.add_node(\"supervisor\", supervisor_node)\n",
"builder.add_node(\"researcher\", research_node)\n",
@@ -293,6 +293,10 @@
"from langchain_core.messages import HumanMessage, trim_messages\n",
"\n",
"\n",
"class State(MessagesState):\n",
" next: str\n",
"\n",
"\n",
"def make_supervisor_node(llm: BaseChatModel, members: list[str]) -> str:\n",
" options = [\"FINISH\"] + members\n",
" system_prompt = (\n",
@@ -308,7 +312,7 @@
"\n",
" next: Literal[*options]\n",
"\n",
" def supervisor_node(state: MessagesState) -> Command[Literal[*members, \"__end__\"]]:\n",
" def supervisor_node(state: State) -> Command[Literal[*members, \"__end__\"]]:\n",
" \"\"\"An LLM-based router.\"\"\"\n",
" messages = [\n",
" {\"role\": \"system\", \"content\": system_prompt},\n",
@@ -318,7 +322,7 @@
" if goto == \"FINISH\":\n",
" goto = END\n",
"\n",
" return Command(goto=goto)\n",
" return Command(goto=goto, update={\"next\": goto})\n",
"\n",
" return supervisor_node"
]
@@ -358,7 +362,7 @@
"search_agent = create_react_agent(llm, tools=[tavily_tool])\n",
"\n",
"\n",
"def search_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def search_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = search_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -374,7 +378,7 @@
"web_scraper_agent = create_react_agent(llm, tools=[scrape_webpages])\n",
"\n",
"\n",
"def web_scraper_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def web_scraper_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = web_scraper_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -410,7 +414,7 @@
},
"outputs": [],
"source": [
"research_builder = StateGraph(MessagesState)\n",
"research_builder = StateGraph(State)\n",
"research_builder.add_node(\"supervisor\", research_supervisor_node)\n",
"research_builder.add_node(\"search\", search_node)\n",
"research_builder.add_node(\"web_scraper\", web_scraper_node)\n",
@@ -528,7 +532,7 @@
")\n",
"\n",
"\n",
"def doc_writing_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def doc_writing_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = doc_writer_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -551,7 +555,7 @@
")\n",
"\n",
"\n",
"def note_taking_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def note_taking_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = note_taking_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -569,7 +573,7 @@
")\n",
"\n",
"\n",
"def chart_generating_node(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def chart_generating_node(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" result = chart_generating_agent.invoke(state)\n",
" return Command(\n",
" update={\n",
@@ -610,7 +614,7 @@
"outputs": [],
"source": [
"# Create the graph here\n",
"paper_writing_builder = StateGraph(MessagesState)\n",
"paper_writing_builder = StateGraph(State)\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",
@@ -730,7 +734,7 @@
},
"outputs": [],
"source": [
"def call_research_team(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def call_research_team(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" response = research_graph.invoke({\"messages\": state[\"messages\"][-1]})\n",
" return Command(\n",
" update={\n",
@@ -744,7 +748,7 @@
" )\n",
"\n",
"\n",
"def call_paper_writing_team(state: MessagesState) -> Command[Literal[\"supervisor\"]]:\n",
"def call_paper_writing_team(state: State) -> Command[Literal[\"supervisor\"]]:\n",
" response = paper_writing_graph.invoke({\"messages\": state[\"messages\"][-1]})\n",
" return Command(\n",
" update={\n",
@@ -759,7 +763,7 @@
"\n",
"\n",
"# Define the graph.\n",
"super_builder = StateGraph(MessagesState)\n",
"super_builder = StateGraph(State)\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",