Compare commits

...
Author SHA1 Message Date
Sydney Runkle 9aab70bd73 experimental stop when 2025-08-12 16:19:16 -04:00
Sydney Runkle aaea76b475 hacky solution for now 2025-08-11 19:44:02 -04:00
2 changed files with 82 additions and 126 deletions
@@ -12,6 +12,7 @@ from typing import (
cast, cast,
get_type_hints, get_type_hints,
) )
from operator import add
from warnings import warn from warnings import warn
from langchain_core.language_models import ( from langchain_core.language_models import (
@@ -46,7 +47,7 @@ from langgraph.managed import RemainingSteps
from langgraph.prebuilt.tool_node import ToolNode from langgraph.prebuilt.tool_node import ToolNode
from langgraph.runtime import Runtime from langgraph.runtime import Runtime
from langgraph.store.base import BaseStore 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.typing import ContextT
from langgraph.warnings import LangGraphDeprecatedSinceV10 from langgraph.warnings import LangGraphDeprecatedSinceV10
@@ -66,6 +67,8 @@ class AgentState(TypedDict):
remaining_steps: NotRequired[RemainingSteps] remaining_steps: NotRequired[RemainingSteps]
model_calls: Annotated[NotRequired[int], add]
class AgentStatePydantic(BaseModel): class AgentStatePydantic(BaseModel):
"""The state of the agent.""" """The state of the agent."""
@@ -191,29 +194,6 @@ def _should_bind_tools(
return False 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( def _validate_chat_history(
messages: Sequence[BaseMessage], messages: Sequence[BaseMessage],
) -> None: ) -> None:
@@ -245,6 +225,10 @@ def _validate_chat_history(
raise ValueError(error_message) raise ValueError(error_message)
class StepCountIs(BaseModel):
count: int
def create_react_agent( def create_react_agent(
model: Union[ model: Union[
str, str,
@@ -276,6 +260,7 @@ def create_react_agent(
debug: bool = False, debug: bool = False,
version: Literal["v1", "v2"] = "v2", version: Literal["v1", "v2"] = "v2",
name: Optional[str] = None, name: Optional[str] = None,
stop_when: Optional[Callable[[StateSchema], bool]] | StepCountIs = None,
**deprecated_kwargs: Any, **deprecated_kwargs: Any,
) -> CompiledStateGraph: ) -> CompiledStateGraph:
"""Creates an agent graph that calls tools in a loop until a stopping condition is met. """Creates an agent graph that calls tools in a loop until a stopping condition is met.
@@ -499,13 +484,19 @@ def create_react_agent(
else AgentState else AgentState
) )
structured_output_tools: list[type] = []
llm_builtin_tools: list[dict] = [] llm_builtin_tools: list[dict] = []
if isinstance(tools, ToolNode): if isinstance(tools, ToolNode):
tool_classes = list(tools.tools_by_name.values()) tool_classes = list(tools.tools_by_name.values())
tool_node = tools tool_node = tools
else: else:
llm_builtin_tools = [t for t in tools if isinstance(t, dict)] 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()) tool_classes = list(tool_node.tools_by_name.values())
is_dynamic_model = not isinstance(model, (str, Runnable)) and callable(model) is_dynamic_model = not isinstance(model, (str, Runnable)) and callable(model)
@@ -527,12 +518,19 @@ def create_react_agent(
model = cast(BaseChatModel, init_chat_model(model)) 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 ( if (
_should_bind_tools(model, tool_classes, num_builtin=len(llm_builtin_tools)) # type: ignore[arg-type] _should_bind_tools(model, tool_classes, num_builtin=len(llm_builtin_tools)) # type: ignore[arg-type]
and len(tool_classes + llm_builtin_tools) > 0 and len(tool_classes + llm_builtin_tools) > 0
): ):
model = cast(BaseChatModel, model).bind_tools( 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",
parallel_tool_calls=False,
) )
static_model: Optional[Runnable] = _get_prompt_runnable(prompt) | model # type: ignore[operator] static_model: Optional[Runnable] = _get_prompt_runnable(prompt) | model # type: ignore[operator]
@@ -542,7 +540,11 @@ def create_react_agent(
# If any of the tools are configured to return_directly after running, # If any of the tools are configured to return_directly after running,
# our graph needs to check if these were called # 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( def _resolve_model(
state: StateSchema, runtime: Runtime[ContextT] state: StateSchema, runtime: Runtime[ContextT]
@@ -608,7 +610,7 @@ def create_react_agent(
# Define the function that calls the model # Define the function that calls the model
def call_model( def call_model(
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
) -> StateSchema: ) -> dict[str, list[AIMessage]] | Command:
if is_async_dynamic_model: if is_async_dynamic_model:
msg = ( msg = (
"Async model callable provided but agent invoked synchronously. " "Async model callable provided but agent invoked synchronously. "
@@ -619,11 +621,25 @@ def create_react_agent(
model_input = _get_model_input_state(state) model_input = _get_model_input_state(state)
if stop_when is not None:
post_model_node = "post_model_hook" if post_model_hook is not None else END
if isinstance(stop_when, StepCountIs):
if (model_calls := _get_state_value(state, "model_calls", 0)) == (stop_when.count - 1):
# set tool_choice to structured output tool if response_format is provided
# though we don't currently expose support for that binding here.
...
elif model_calls == stop_when.count:
return Command(goto=post_model_node)
else:
if stop_when(state):
return Command(goto=post_model_node)
if is_dynamic_model: if is_dynamic_model:
# Resolve dynamic model at runtime and apply prompt # Resolve dynamic model at runtime and apply prompt
dynamic_model = _resolve_model(state, runtime) dynamic_model = _resolve_model(state, runtime)
response = cast(AIMessage, dynamic_model.invoke(model_input, config)) # type: ignore[arg-type] response = cast(AIMessage, dynamic_model.invoke(model_input, config)) # type: ignore[arg-type]
else: else:
response = cast(AIMessage, static_model.invoke(model_input, config)) # type: ignore[union-attr] response = cast(AIMessage, static_model.invoke(model_input, config)) # type: ignore[union-attr]
# add agent name to the AIMessage # add agent name to the AIMessage
@@ -638,12 +654,12 @@ def create_react_agent(
) )
] ]
} }
# We return a list, because this will get added to the existing list
return {"messages": [response]} return {"messages": [response], "model_calls": 1}
async def acall_model( async def acall_model(
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
) -> StateSchema: ) -> dict[str, list[AIMessage]]:
model_input = _get_model_input_state(state) model_input = _get_model_input_state(state)
if is_dynamic_model: if is_dynamic_model:
@@ -689,49 +705,6 @@ def create_react_agent(
else: else:
input_schema = state_schema 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: if not tool_calling_enabled:
# Define a new graph # Define a new graph
workflow = StateGraph(state_schema=state_schema, context_schema=context_schema) workflow = StateGraph(state_schema=state_schema, context_schema=context_schema)
@@ -753,19 +726,6 @@ def create_react_agent(
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type] workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
workflow.add_edge("agent", "post_model_hook") 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( return workflow.compile(
checkpointer=checkpointer, checkpointer=checkpointer,
store=store, store=store,
@@ -777,16 +737,14 @@ def create_react_agent(
# Define the function that determines whether to continue or not # Define the function that determines whether to continue or not
def should_continue(state: StateSchema) -> Union[str, list[Send]]: def should_continue(state: StateSchema) -> Union[str, list[Send]]:
post_model_node = "post_model_hook" if post_model_hook is not None else END
messages = _get_state_value(state, "messages") messages = _get_state_value(state, "messages")
last_message = messages[-1] last_message = messages[-1]
# If there is no function call, then we finish # If there is no function call, then we finish
if not isinstance(last_message, AIMessage) or not last_message.tool_calls: if not isinstance(last_message, AIMessage) or not last_message.tool_calls:
if post_model_hook is not None: return post_model_node
return "post_model_hook"
elif response_format is not None:
return "generate_structured_response"
else:
return END
# Otherwise if there is, we continue # Otherwise if there is, we continue
else: else:
if version == "v1": if version == "v1":
@@ -825,45 +783,19 @@ def create_react_agent(
# Set the entrypoint as `agent` # Set the entrypoint as `agent`
# This means that this node is the first one called # This means that this node is the first one called
workflow.set_entry_point(entrypoint) workflow.set_entry_point(entrypoint)
agent_paths = [] agent_paths = []
post_model_hook_paths = [entrypoint, "tools"]
# Add a post model hook node if post_model_hook is provided # Add a post model hook node if post_model_hook is provided
if post_model_hook is not None: if post_model_hook is not None:
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type] workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
agent_paths.append("post_model_hook") agent_paths.append("post_model_hook")
workflow.add_edge("agent", "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]]: def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]:
"""Route to the next node after post_model_hook. """Route to the next node after post_model_hook.
Routes to one of: Routes to one of:
* "tools": if there are pending tool calls without a corresponding message. * "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. * END: if no pending tool calls exist and no response_format is specified.
""" """
@@ -886,16 +818,17 @@ def create_react_agent(
return [Send("tools", [tool_call]) for tool_call in pending_tool_calls] return [Send("tools", [tool_call]) for tool_call in pending_tool_calls]
elif isinstance(messages[-1], ToolMessage): elif isinstance(messages[-1], ToolMessage):
return entrypoint return entrypoint
elif response_format is not None:
return "generate_structured_response"
else: else:
return END return END
workflow.add_conditional_edges( workflow.add_conditional_edges(
"post_model_hook", "post_model_hook",
post_model_hook_router, 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( workflow.add_conditional_edges(
"agent", "agent",
@@ -903,7 +836,7 @@ def create_react_agent(
path_map=agent_paths, 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")): for m in reversed(_get_state_value(state, "messages")):
if not isinstance(m, ToolMessage): if not isinstance(m, ToolMessage):
break break
+25 -2
View File
@@ -340,11 +340,19 @@ class ToolNode(RunnableCallable):
self.tools_by_name: dict[str, BaseTool] = {} self.tools_by_name: dict[str, BaseTool] = {}
self.tool_to_state_args: dict[str, dict[str, Optional[str]]] = {} self.tool_to_state_args: dict[str, dict[str, Optional[str]]] = {}
self.tool_to_store_arg: 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.handle_tool_errors = handle_tool_errors
self.messages_key = messages_key self.messages_key = messages_key
for tool_ in tools: 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_) tool_ = create_tool(tool_)
self.tools_by_name[tool_.name] = tool_ self.tools_by_name[tool_.name] = tool_
self.tool_to_state_args[tool_.name] = _get_state_args(tool_) self.tool_to_state_args[tool_.name] = _get_state_args(tool_)
self.tool_to_store_arg[tool_.name] = _get_store_arg(tool_) self.tool_to_store_arg[tool_.name] = _get_store_arg(tool_)
@@ -437,11 +445,26 @@ class ToolNode(RunnableCallable):
call: ToolCall, call: ToolCall,
input_type: Literal["list", "dict", "tool_calls"], input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig, config: RunnableConfig,
) -> ToolMessage: ) -> ToolMessage | Command:
"""Run a single tool call synchronously.""" """Run a single tool call synchronously."""
if invalid_tool_message := self._validate_tool_call(call): if invalid_tool_message := self._validate_tool_call(call):
return invalid_tool_message return invalid_tool_message
try: 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"}} call_args = {**call, **{"type": "tool_call"}}
response = self.tools_by_name[call["name"]].invoke(call_args, config) response = self.tools_by_name[call["name"]].invoke(call_args, config)