diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 193512cc9..e05b418b5 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -247,14 +247,16 @@ def _validate_chat_history( class _AgentBuilder: """Internal builder class for constructing React agents with intuitive method-to-node mapping.""" - + def __init__( self, model: Union[str, LanguageModelLike], tools: Union[Sequence[Union[BaseTool, Callable, dict[str, Any]]], ToolNode], *, prompt: Optional[Prompt] = None, - response_format: Optional[Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]]] = None, + response_format: Optional[ + Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]] + ] = None, pre_model_hook: Optional[RunnableLike] = None, post_model_hook: Optional[RunnableLike] = None, state_schema: Optional[StateSchemaType] = None, @@ -273,28 +275,34 @@ class _AgentBuilder: self.context_schema = context_schema self.version = version self.name = name - + # Setup tools if isinstance(self.tools, ToolNode): self._tool_classes = list(self.tools.tools_by_name.values()) self._tool_node = self.tools else: self._llm_builtin_tools = [t for t in self.tools if isinstance(t, dict)] - self._tool_node = ToolNode([t for t in self.tools if not isinstance(t, dict)]) + self._tool_node = ToolNode( + [t for t in self.tools if not isinstance(t, dict)] + ) self._tool_classes = list(self._tool_node.tools_by_name.values()) - - self._should_return_direct: set[str] = {t.name for t in self._tool_classes if t.return_direct} + + self._should_return_direct: set[str] = { + t.name for t in self._tool_classes if t.return_direct + } # Setup state schema if self.state_schema is not None: required_keys = {"messages", "remaining_steps"} if self.response_format is not None: required_keys.add("structured_response") - + schema_keys = set(get_type_hints(self.state_schema)) if missing_keys := required_keys - schema_keys: - raise ValueError(f"Missing required key(s) {missing_keys} in state_schema") - + raise ValueError( + f"Missing required key(s) {missing_keys} in state_schema" + ) + self._final_state_schema = self.state_schema else: self._final_state_schema = ( @@ -305,7 +313,7 @@ class _AgentBuilder: # Setup model model = self.model - + # Convert string models if isinstance(model, str): try: @@ -315,19 +323,23 @@ class _AgentBuilder: "Please install langchain (`pip install langchain`) to use ':' string syntax for `model` parameter." ) model = cast(BaseChatModel, init_chat_model(model)) - + # Bind tools if needed if ( - _should_bind_tools(model, self._tool_classes, num_builtin=len(self._llm_builtin_tools)) + _should_bind_tools( + model, self._tool_classes, num_builtin=len(self._llm_builtin_tools) + ) and len(self._tool_classes + self._llm_builtin_tools) > 0 ): - model = cast(BaseChatModel, model).bind_tools(self._tool_classes + self._llm_builtin_tools) # type: ignore[operator] - + model = cast(BaseChatModel, model).bind_tools( + self._tool_classes + self._llm_builtin_tools + ) # type: ignore[operator] + self._model_runnable = _get_prompt_runnable(self.prompt) | model def create_model_node(self) -> RunnableCallable: """Create the 'agent' node that calls the LLM.""" - + def _get_model_input_state(state: StateSchema) -> StateSchema: if self.pre_model_hook is not None: messages: Optional[Sequence[BaseMessage]] = ( @@ -342,8 +354,10 @@ class _AgentBuilder: raise ValueError(error_msg) _validate_chat_history(messages) - - if isinstance(self._final_state_schema, type) and issubclass(self._final_state_schema, BaseModel): + + if isinstance(self._final_state_schema, type) and issubclass( + self._final_state_schema, BaseModel + ): state.messages = messages # type: ignore else: state["messages"] = messages # type: ignore @@ -352,15 +366,27 @@ class _AgentBuilder: def _are_more_steps_needed(state: StateSchema, response: BaseMessage) -> bool: has_tool_calls = isinstance(response, AIMessage) and response.tool_calls all_tools_return_direct = ( - all(call["name"] in self._should_return_direct for call in response.tool_calls) - if isinstance(response, AIMessage) else False + all( + call["name"] in self._should_return_direct + for call in response.tool_calls + ) + if isinstance(response, AIMessage) + else False ) remaining_steps = _get_state_value(state, "remaining_steps", None) is_last_step = _get_state_value(state, "is_last_step", False) return ( (remaining_steps is None and is_last_step and has_tool_calls) - or (remaining_steps is not None and remaining_steps < 1 and all_tools_return_direct) - or (remaining_steps is not None and remaining_steps < 2 and has_tool_calls) + or ( + remaining_steps is not None + and remaining_steps < 1 + and all_tools_return_direct + ) + or ( + remaining_steps is not None + and remaining_steps < 2 + and has_tool_calls + ) ) def call_model(state: StateSchema, config: RunnableConfig) -> StateSchema: @@ -379,11 +405,15 @@ class _AgentBuilder: } return {"messages": [response]} - async def acall_model(state: StateSchema, config: RunnableConfig) -> StateSchema: + async def acall_model( + state: StateSchema, config: RunnableConfig + ) -> StateSchema: state = _get_model_input_state(state) - response = cast(AIMessage, await self._model_runnable.ainvoke(state, config)) # type: ignore[union-attr] + response = cast( + AIMessage, await self._model_runnable.ainvoke(state, config) + ) # type: ignore[union-attr] response.name = self.name - + if _are_more_steps_needed(state, response): return { "messages": [ @@ -398,71 +428,75 @@ class _AgentBuilder: # Determine input schema input_schema = self._final_state_schema if self.pre_model_hook is not None: - if isinstance(self._final_state_schema, type) and issubclass(self._final_state_schema, BaseModel): + if isinstance(self._final_state_schema, type) and issubclass( + self._final_state_schema, BaseModel + ): from pydantic import create_model + input_schema = create_model( "CallModelInputSchema", llm_input_messages=(list[AnyMessage], ...), __base__=self._final_state_schema, ) else: + class CallModelInputSchema(self._final_state_schema): # type: ignore llm_input_messages: list[AnyMessage] + input_schema = CallModelInputSchema return RunnableCallable(call_model, acall_model, input_schema=input_schema) - - def create_tools_node(self) -> ToolNode: - """Create the 'tools' node that executes tools.""" - return self._tool_node - - def create_pre_model_hook_node(self) -> Optional[RunnableLike]: - """Create the 'pre_model_hook' node if configured.""" - return self.pre_model_hook - - def create_post_model_hook_node(self) -> Optional[RunnableLike]: - """Create the 'post_model_hook' node if configured.""" - return self.post_model_hook - + def create_structured_response_node(self) -> Optional[RunnableCallable]: """Create the 'generate_structured_response' node if configured.""" if self.response_format is None: return None - def generate_structured_response(state: StateSchema, config: RunnableConfig) -> StateSchema: + def generate_structured_response( + state: StateSchema, config: RunnableConfig + ) -> StateSchema: messages = _get_state_value(state, "messages") structured_response_schema = self.response_format if isinstance(self.response_format, tuple): system_prompt, structured_response_schema = self.response_format messages = [SystemMessage(content=system_prompt)] + list(messages) - model_with_structured_output = _get_model(self._model_runnable).with_structured_output( # type: ignore[arg-type] + model_with_structured_output = _get_model( + self._model_runnable + ).with_structured_output( # type: ignore[arg-type] cast(StructuredResponseSchema, structured_response_schema) ) response = model_with_structured_output.invoke(messages, config) return {"structured_response": response} - async def agenerate_structured_response(state: StateSchema, config: RunnableConfig) -> StateSchema: + async def agenerate_structured_response( + state: StateSchema, config: RunnableConfig + ) -> StateSchema: messages = _get_state_value(state, "messages") structured_response_schema = self.response_format if isinstance(self.response_format, tuple): system_prompt, structured_response_schema = self.response_format messages = [SystemMessage(content=system_prompt)] + list(messages) - model_with_structured_output = _get_model(self._model_runnable).with_structured_output( # type: ignore[arg-type] + model_with_structured_output = _get_model( + self._model_runnable + ).with_structured_output( # type: ignore[arg-type] cast(StructuredResponseSchema, structured_response_schema) ) response = await model_with_structured_output.ainvoke(messages, config) return {"structured_response": response} - return RunnableCallable(generate_structured_response, agenerate_structured_response) + return RunnableCallable( + generate_structured_response, agenerate_structured_response + ) + + def create_model_router(self) -> Callable[[StateSchema], Union[str, list[Send]]]: + """Create routing function for model node conditional edges.""" - def create_agent_router(self) -> Callable[[StateSchema], Union[str, list[Send]]]: - """Create routing function for agent node conditional edges.""" def should_continue(state: StateSchema) -> Union[str, list[Send]]: messages = _get_state_value(state, "messages") last_message = messages[-1] - + if not isinstance(last_message, AIMessage) or not last_message.tool_calls: if self.post_model_hook is not None: return "post_model_hook" @@ -487,46 +521,44 @@ class _AgentBuilder: ) for tool_call in last_message.tool_calls ] - return should_continue - - def create_post_model_hook_router(self) -> Optional[Callable[[StateSchema], Union[str, list[Send]]]]: - """Create routing function for post_model_hook node conditional edges.""" - if self.post_model_hook is None: - return None - - def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]: - messages = _get_state_value(state, "messages") - tool_messages = [m.tool_call_id for m in messages if isinstance(m, ToolMessage)] - last_ai_message = next(m for m in reversed(messages) if isinstance(m, AIMessage)) - pending_tool_calls = [ - c for c in last_ai_message.tool_calls if c["id"] not in tool_messages - ] - if pending_tool_calls: - return [ - Send( - "tools", - ToolCallWithContext( - __type="tool_call_with_context", - tool_call=tool_call, - state=state, - ), - ) - for tool_call in pending_tool_calls - ] - elif isinstance(messages[-1], ToolMessage): - return self._get_entry_point() - elif self.response_format is not None: - return "generate_structured_response" - else: - return END - return post_model_hook_router - + return should_continue + + def post_model_hook_router(self, state: StateSchema) -> Union[str, list[Send]]: + """Route to the next node after post_model_hook.""" + messages = _get_state_value(state, "messages") + tool_messages = [m.tool_call_id for m in messages if isinstance(m, ToolMessage)] + last_ai_message = next( + m for m in reversed(messages) if isinstance(m, AIMessage) + ) + pending_tool_calls = [ + c for c in last_ai_message.tool_calls if c["id"] not in tool_messages + ] + + if pending_tool_calls: + return [ + Send( + "tools", + ToolCallWithContext( + __type="tool_call_with_context", + tool_call=tool_call, + state=state, + ), + ) + for tool_call in pending_tool_calls + ] + elif isinstance(messages[-1], ToolMessage): + return self._get_entry_point() + elif self.response_format is not None: + return "generate_structured_response" + else: + return END + def create_tools_router(self) -> Optional[Callable[[StateSchema], str]]: """Create routing function for tools node conditional edges.""" if not self._should_return_direct: return None - + def route_tool_responses(state: StateSchema) -> str: messages = _get_state_value(state, "messages") for m in reversed(messages): @@ -536,75 +568,27 @@ class _AgentBuilder: return END if isinstance(m, AIMessage) and m.tool_calls: - if any(call["name"] in self._should_return_direct for call in m.tool_calls): + if any( + call["name"] in self._should_return_direct for call in m.tool_calls + ): return END return self._get_entry_point() + return route_tool_responses def _get_entry_point(self) -> str: """Get the workflow entry point.""" return "pre_model_hook" if self.pre_model_hook else "agent" - + def _has_tools(self) -> bool: """Check if agent has tools enabled.""" return len(self._tool_classes) > 0 - - def _add_nodes(self, workflow: StateGraph) -> None: - """Add all nodes to the workflow.""" - # Always add agent node - workflow.add_node("agent", self.create_model_node()) - - # Add tools node if needed - if self._has_tools(): - workflow.add_node("tools", self.create_tools_node()) - - # Add hook nodes if configured - if self.pre_model_hook: - workflow.add_node("pre_model_hook", self.create_pre_model_hook_node()) # type: ignore[arg-type] - if self.post_model_hook: - workflow.add_node("post_model_hook", self.create_post_model_hook_node()) # type: ignore[arg-type] - - # Add structured response node if configured - structured_node = self.create_structured_response_node() - if structured_node: - workflow.add_node("generate_structured_response", structured_node) - - def _add_edges(self, workflow: StateGraph) -> None: - """Add all edges to the workflow.""" - entry_point = self._get_entry_point() - workflow.set_entry_point(entry_point) - - # Pre-model hook edge - if self.pre_model_hook: - workflow.add_edge("pre_model_hook", "agent") - - # Agent edges - if self.post_model_hook: - # Direct edge from agent to post_model_hook when post_model_hook exists - workflow.add_edge("agent", "post_model_hook") - # Post-model hook conditional edges - post_hook_router = self.create_post_model_hook_router() - post_hook_paths = self._get_post_model_hook_paths() - workflow.add_conditional_edges("post_model_hook", post_hook_router, path_map=post_hook_paths) # type: ignore[arg-type] - else: - # Conditional edges from agent when no post_model_hook - agent_router = self.create_agent_router() - agent_paths = self._get_agent_paths() - workflow.add_conditional_edges("agent", agent_router, path_map=agent_paths) # type: ignore[arg-type] - - # Tools edges - if self._has_tools(): - tools_router = self.create_tools_router() - if tools_router: - workflow.add_conditional_edges("tools", tools_router, path_map=[entry_point, END]) - else: - workflow.add_edge("tools", entry_point) - - def _get_agent_paths(self) -> list[str]: - """Get possible paths from agent node.""" + + def _get_model_paths(self) -> list[str]: + """Get possible paths from model node.""" paths = [] - + # If post_model_hook exists, we don't add paths here - we use direct edge instead if not self.post_model_hook: if self._has_tools(): @@ -613,9 +597,9 @@ class _AgentBuilder: paths.append("generate_structured_response") if not self._has_tools() and not self.response_format: paths.append(END) - + return paths - + def _get_post_model_hook_paths(self) -> list[str]: """Get possible paths from post_model_hook node.""" paths = [self._get_entry_point()] @@ -632,13 +616,61 @@ class _AgentBuilder: # Create workflow workflow = StateGraph( state_schema=self._final_state_schema, # type: ignore[arg-type] - context_schema=self.context_schema + context_schema=self.context_schema, ) - - # Add nodes and edges - self._add_nodes(workflow) - self._add_edges(workflow) - + + # Add nodes + # Always add model node (named 'agent' for backwards compatibility) + workflow.add_node("agent", self.create_model_node()) + + # Add tools node if needed + if self._has_tools(): + workflow.add_node("tools", self._tool_node) + + # Add hook nodes if configured + if self.pre_model_hook: + workflow.add_node("pre_model_hook", self.pre_model_hook) # type: ignore[arg-type] + if self.post_model_hook: + workflow.add_node("post_model_hook", self.post_model_hook) # type: ignore[arg-type] + + # Add structured response node if configured + structured_node = self.create_structured_response_node() + if structured_node: + workflow.add_node("generate_structured_response", structured_node) + + # Add edges + entry_point = self._get_entry_point() + workflow.set_entry_point(entry_point) + + # Pre-model hook edge + if self.pre_model_hook: + workflow.add_edge("pre_model_hook", "agent") + + # Model node edges + if self.post_model_hook: + # Direct edge from model node to post_model_hook when post_model_hook exists + workflow.add_edge("agent", "post_model_hook") + # Post-model hook conditional edges + post_hook_paths = self._get_post_model_hook_paths() + workflow.add_conditional_edges( + "post_model_hook", self.post_model_hook_router, path_map=post_hook_paths + ) # type: ignore[arg-type] + else: + # Conditional edges from model node when no post_model_hook + model_router = self.create_model_router() + model_paths = self._get_model_paths() + workflow.add_conditional_edges("agent", model_router, path_map=model_paths) # type: ignore[arg-type] + + # Tools edges + if self._has_tools(): + tools_router = self.create_tools_router() + if tools_router: + workflow.add_conditional_edges( + "tools", tools_router, path_map=[entry_point, END] + ) + else: + workflow.add_edge("tools", entry_point) + return workflow @@ -839,9 +871,9 @@ def create_react_agent( version=version, name=name, ) - + workflow = builder.build() - + # Compile and return the graph return workflow.compile( checkpointer=checkpointer,