diff --git a/docs/docs/tutorials/multi_agent/agent_supervisor.ipynb b/docs/docs/tutorials/multi_agent/agent_supervisor.ipynb index 111d0b05e..930c0e989 100644 --- a/docs/docs/tutorials/multi_agent/agent_supervisor.ipynb +++ b/docs/docs/tutorials/multi_agent/agent_supervisor.ipynb @@ -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", diff --git a/docs/docs/tutorials/multi_agent/hierarchical_agent_teams.ipynb b/docs/docs/tutorials/multi_agent/hierarchical_agent_teams.ipynb index 39009e852..c46074831 100644 --- a/docs/docs/tutorials/multi_agent/hierarchical_agent_teams.ipynb +++ b/docs/docs/tutorials/multi_agent/hierarchical_agent_teams.ipynb @@ -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",