From aaea76b4753c0df843c7113f26b7eb1a3a804f55 Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Mon, 11 Aug 2025 19:44:02 -0400 Subject: [PATCH] hacky solution for now --- .../langgraph/prebuilt/chat_agent_executor.py | 146 ++++-------------- libs/prebuilt/langgraph/prebuilt/tool_node.py | 27 +++- 2 files changed, 53 insertions(+), 120 deletions(-) diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 8256bbc44..ce4fcfa28 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -46,7 +46,7 @@ from langgraph.managed import RemainingSteps from langgraph.prebuilt.tool_node import ToolNode from langgraph.runtime import Runtime from langgraph.store.base import BaseStore -from langgraph.types import Checkpointer, Send +from langgraph.types import Checkpointer, Command, Send from langgraph.typing import ContextT from langgraph.warnings import LangGraphDeprecatedSinceV10 @@ -191,29 +191,6 @@ def _should_bind_tools( return False -def _get_model(model: LanguageModelLike) -> BaseChatModel: - """Get the underlying model from a RunnableBinding or return the model itself.""" - if isinstance(model, RunnableSequence): - model = next( - ( - step - for step in model.steps - if isinstance(step, (RunnableBinding, BaseChatModel)) - ), - model, - ) - - if isinstance(model, RunnableBinding): - model = model.bound - - if not isinstance(model, BaseChatModel): - raise TypeError( - f"Expected `model` to be a ChatModel or RunnableBinding (e.g. model.bind_tools(...)), got {type(model)}" - ) - - return model - - def _validate_chat_history( messages: Sequence[BaseMessage], ) -> None: @@ -499,13 +476,19 @@ def create_react_agent( else AgentState ) + structured_output_tools: list[type] = [] llm_builtin_tools: list[dict] = [] if isinstance(tools, ToolNode): tool_classes = list(tools.tools_by_name.values()) tool_node = tools else: llm_builtin_tools = [t for t in tools if isinstance(t, dict)] - tool_node = ToolNode([t for t in tools if not isinstance(t, dict)]) + structured_output_tools = ( + [response_format] if response_format is not None else [] + ) + tool_node = ToolNode( + [t for t in [*tools, *structured_output_tools] if not isinstance(t, dict)] + ) tool_classes = list(tool_node.tools_by_name.values()) is_dynamic_model = not isinstance(model, (str, Runnable)) and callable(model) @@ -527,12 +510,18 @@ def create_react_agent( model = cast(BaseChatModel, init_chat_model(model)) + # Add structured output tool if response_format is provided + structured_output_tools = ( + [response_format] if response_format is not None else [] + ) + if ( _should_bind_tools(model, tool_classes, num_builtin=len(llm_builtin_tools)) # type: ignore[arg-type] and len(tool_classes + llm_builtin_tools) > 0 ): model = cast(BaseChatModel, model).bind_tools( - tool_classes + llm_builtin_tools # type: ignore[operator] + tool_classes + llm_builtin_tools + structured_output_tools, # type: ignore[operator], + tool_choice="any", ) static_model: Optional[Runnable] = _get_prompt_runnable(prompt) | model # type: ignore[operator] @@ -542,7 +531,11 @@ def create_react_agent( # If any of the tools are configured to return_directly after running, # our graph needs to check if these were called - should_return_direct = {t.name for t in tool_classes if t.return_direct} + should_return_direct = { + t.name for t in tool_classes if getattr(t, "return_direct", False) + } + if response_format is not None: + should_return_direct.add(response_format.__name__) def _resolve_model( state: StateSchema, runtime: Runtime[ContextT] @@ -608,7 +601,7 @@ def create_react_agent( # Define the function that calls the model def call_model( state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: + ) -> dict[str, list[AIMessage]]: if is_async_dynamic_model: msg = ( "Async model callable provided but agent invoked synchronously. " @@ -638,12 +631,12 @@ def create_react_agent( ) ] } - # We return a list, because this will get added to the existing list + return {"messages": [response]} async def acall_model( state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: + ) -> dict[str, list[AIMessage]]: model_input = _get_model_input_state(state) if is_dynamic_model: @@ -689,49 +682,6 @@ def create_react_agent( else: input_schema = state_schema - def generate_structured_response( - state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: - if is_async_dynamic_model: - msg = ( - "Async model callable provided but agent invoked synchronously. " - "Use agent.ainvoke() or agent.astream(), or provide a sync model callable." - ) - raise RuntimeError(msg) - - messages = _get_state_value(state, "messages") - structured_response_schema = response_format - if isinstance(response_format, tuple): - system_prompt, structured_response_schema = response_format - messages = [SystemMessage(content=system_prompt)] + list(messages) - - resolved_model = _resolve_model(state, runtime) - model_with_structured_output = _get_model( - resolved_model - ).with_structured_output( - cast(StructuredResponseSchema, structured_response_schema) - ) - response = model_with_structured_output.invoke(messages, config) - return {"structured_response": response} - - async def agenerate_structured_response( - state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: - messages = _get_state_value(state, "messages") - structured_response_schema = response_format - if isinstance(response_format, tuple): - system_prompt, structured_response_schema = response_format - messages = [SystemMessage(content=system_prompt)] + list(messages) - - resolved_model = await _aresolve_model(state, runtime) - model_with_structured_output = _get_model( - resolved_model - ).with_structured_output( - cast(StructuredResponseSchema, structured_response_schema) - ) - response = await model_with_structured_output.ainvoke(messages, config) - return {"structured_response": response} - if not tool_calling_enabled: # Define a new graph workflow = StateGraph(state_schema=state_schema, context_schema=context_schema) @@ -753,19 +703,6 @@ def create_react_agent( workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type] workflow.add_edge("agent", "post_model_hook") - if response_format is not None: - workflow.add_node( - "generate_structured_response", - RunnableCallable( - generate_structured_response, - agenerate_structured_response, - ), - ) - if post_model_hook is not None: - workflow.add_edge("post_model_hook", "generate_structured_response") - else: - workflow.add_edge("agent", "generate_structured_response") - return workflow.compile( checkpointer=checkpointer, store=store, @@ -783,8 +720,6 @@ def create_react_agent( if not isinstance(last_message, AIMessage) or not last_message.tool_calls: if post_model_hook is not None: return "post_model_hook" - elif response_format is not None: - return "generate_structured_response" else: return END # Otherwise if there is, we continue @@ -825,45 +760,19 @@ def create_react_agent( # Set the entrypoint as `agent` # This means that this node is the first one called workflow.set_entry_point(entrypoint) - agent_paths = [] - post_model_hook_paths = [entrypoint, "tools"] # Add a post model hook node if post_model_hook is provided if post_model_hook is not None: workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type] agent_paths.append("post_model_hook") workflow.add_edge("agent", "post_model_hook") - else: - agent_paths.append("tools") - - # Add a structured output node if response_format is provided - if response_format is not None: - workflow.add_node( - "generate_structured_response", - RunnableCallable( - generate_structured_response, - agenerate_structured_response, - ), - ) - if post_model_hook is not None: - post_model_hook_paths.append("generate_structured_response") - else: - agent_paths.append("generate_structured_response") - else: - if post_model_hook is not None: - post_model_hook_paths.append(END) - else: - agent_paths.append(END) - - if post_model_hook is not None: def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]: """Route to the next node after post_model_hook. Routes to one of: * "tools": if there are pending tool calls without a corresponding message. - * "generate_structured_response": if no pending tool calls exist and response_format is specified. * END: if no pending tool calls exist and no response_format is specified. """ @@ -886,16 +795,17 @@ def create_react_agent( return [Send("tools", [tool_call]) for tool_call in pending_tool_calls] elif isinstance(messages[-1], ToolMessage): return entrypoint - elif response_format is not None: - return "generate_structured_response" else: return END workflow.add_conditional_edges( "post_model_hook", post_model_hook_router, - path_map=post_model_hook_paths, + path_map=[entrypoint, "tools", END], ) + else: + agent_paths.append("tools") + agent_paths.append(END) workflow.add_conditional_edges( "agent", @@ -903,7 +813,7 @@ def create_react_agent( path_map=agent_paths, ) - def route_tool_responses(state: StateSchema) -> str: + def route_tool_responses(state: StateSchema) -> str | Command: for m in reversed(_get_state_value(state, "messages")): if not isinstance(m, ToolMessage): break diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index 252f88802..f2c8da11f 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_node.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_node.py @@ -340,11 +340,19 @@ class ToolNode(RunnableCallable): self.tools_by_name: dict[str, BaseTool] = {} self.tool_to_state_args: dict[str, dict[str, Optional[str]]] = {} self.tool_to_store_arg: dict[str, Optional[str]] = {} + self.structured_output_tools: list[str] = [] self.handle_tool_errors = handle_tool_errors self.messages_key = messages_key for tool_ in tools: - if not isinstance(tool_, BaseTool): + if issubclass(tool_, BaseModel): + self.tools_by_name[tool_.__name__] = tool_ + self.tool_to_state_args[tool_.__name__] = {} + self.tool_to_store_arg[tool_.__name__] = None + self.structured_output_tools.append(tool_.__name__) + continue + elif not isinstance(tool_, BaseTool): tool_ = create_tool(tool_) + self.tools_by_name[tool_.name] = tool_ self.tool_to_state_args[tool_.name] = _get_state_args(tool_) self.tool_to_store_arg[tool_.name] = _get_store_arg(tool_) @@ -437,11 +445,26 @@ class ToolNode(RunnableCallable): call: ToolCall, input_type: Literal["list", "dict", "tool_calls"], config: RunnableConfig, - ) -> ToolMessage: + ) -> ToolMessage | Command: """Run a single tool call synchronously.""" if invalid_tool_message := self._validate_tool_call(call): return invalid_tool_message try: + if call["name"] in self.structured_output_tools: + response_schema = self.tools_by_name[call["name"]] + return Command( + update={ + "messages": [ + ToolMessage( + content="structured output generated", + name="structured_output", + tool_call_id=call["id"], + status="success", + ), + ], + "structured_response": response_schema(**call["args"]), + } + ) call_args = {**call, **{"type": "tool_call"}} response = self.tools_by_name[call["name"]].invoke(call_args, config)