mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-18 21:55:46 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9aab70bd73 | ||
|
|
aaea76b475 |
@@ -12,6 +12,7 @@ from typing import (
|
||||
cast,
|
||||
get_type_hints,
|
||||
)
|
||||
from operator import add
|
||||
from warnings import warn
|
||||
|
||||
from langchain_core.language_models import (
|
||||
@@ -46,7 +47,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
|
||||
|
||||
@@ -66,6 +67,8 @@ class AgentState(TypedDict):
|
||||
|
||||
remaining_steps: NotRequired[RemainingSteps]
|
||||
|
||||
model_calls: Annotated[NotRequired[int], add]
|
||||
|
||||
|
||||
class AgentStatePydantic(BaseModel):
|
||||
"""The state of the agent."""
|
||||
@@ -191,29 +194,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:
|
||||
@@ -245,6 +225,10 @@ def _validate_chat_history(
|
||||
raise ValueError(error_message)
|
||||
|
||||
|
||||
class StepCountIs(BaseModel):
|
||||
count: int
|
||||
|
||||
|
||||
def create_react_agent(
|
||||
model: Union[
|
||||
str,
|
||||
@@ -276,6 +260,7 @@ def create_react_agent(
|
||||
debug: bool = False,
|
||||
version: Literal["v1", "v2"] = "v2",
|
||||
name: Optional[str] = None,
|
||||
stop_when: Optional[Callable[[StateSchema], bool]] | StepCountIs = None,
|
||||
**deprecated_kwargs: Any,
|
||||
) -> CompiledStateGraph:
|
||||
"""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
|
||||
)
|
||||
|
||||
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 +518,19 @@ 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",
|
||||
parallel_tool_calls=False,
|
||||
)
|
||||
|
||||
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,
|
||||
# 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 +610,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]] | Command:
|
||||
if is_async_dynamic_model:
|
||||
msg = (
|
||||
"Async model callable provided but agent invoked synchronously. "
|
||||
@@ -619,11 +621,25 @@ def create_react_agent(
|
||||
|
||||
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:
|
||||
# Resolve dynamic model at runtime and apply prompt
|
||||
dynamic_model = _resolve_model(state, runtime)
|
||||
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]
|
||||
|
||||
# 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(
|
||||
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 +705,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 +726,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,
|
||||
@@ -777,16 +737,14 @@ def create_react_agent(
|
||||
|
||||
# Define the function that determines whether to continue or not
|
||||
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")
|
||||
last_message = messages[-1]
|
||||
# If there is no function call, then we finish
|
||||
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
|
||||
return post_model_node
|
||||
# Otherwise if there is, we continue
|
||||
else:
|
||||
if version == "v1":
|
||||
@@ -825,45 +783,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 +818,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 +836,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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user