mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-18 13:45:44 +02:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
55880c9813 | ||
|
|
22af613437 | ||
|
|
a254978893 |
@@ -245,6 +245,435 @@ def _validate_chat_history(
|
||||
raise ValueError(error_message)
|
||||
|
||||
|
||||
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,
|
||||
pre_model_hook: Optional[RunnableLike] = None,
|
||||
post_model_hook: Optional[RunnableLike] = None,
|
||||
state_schema: Optional[StateSchemaType] = None,
|
||||
context_schema: Optional[Type[Any]] = None,
|
||||
version: Literal["v1", "v2"] = "v2",
|
||||
name: Optional[str] = None,
|
||||
):
|
||||
# Store all parameters
|
||||
self.model = model
|
||||
self.tools = tools
|
||||
self.prompt = prompt
|
||||
self.response_format = response_format
|
||||
self.pre_model_hook = pre_model_hook
|
||||
self.post_model_hook = post_model_hook
|
||||
self.state_schema = state_schema
|
||||
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_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
|
||||
}
|
||||
|
||||
# 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"
|
||||
)
|
||||
|
||||
self._final_state_schema = self.state_schema
|
||||
else:
|
||||
self._final_state_schema = (
|
||||
AgentStateWithStructuredResponse
|
||||
if self.response_format is not None
|
||||
else AgentState
|
||||
)
|
||||
|
||||
# Setup model
|
||||
model = self.model
|
||||
|
||||
# Convert string models
|
||||
if isinstance(model, str):
|
||||
try:
|
||||
from langchain.chat_models import init_chat_model # type: ignore[import-not-found]
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Please install langchain (`pip install langchain`) to use '<provider>:<model>' 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)
|
||||
)
|
||||
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]
|
||||
|
||||
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]] = (
|
||||
_get_state_value(state, "llm_input_messages")
|
||||
) or _get_state_value(state, "messages")
|
||||
error_msg: str = f"Expected input to call_model to have 'llm_input_messages' or 'messages' key, but got {state}"
|
||||
else:
|
||||
messages = _get_state_value(state, "messages")
|
||||
error_msg = f"Expected input to call_model to have 'messages' key, but got {state}"
|
||||
|
||||
if messages is None:
|
||||
raise ValueError(error_msg)
|
||||
|
||||
_validate_chat_history(messages)
|
||||
|
||||
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
|
||||
return state
|
||||
|
||||
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
|
||||
)
|
||||
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
|
||||
)
|
||||
)
|
||||
|
||||
def call_model(state: StateSchema, config: RunnableConfig) -> StateSchema:
|
||||
state = _get_model_input_state(state)
|
||||
response = cast(AIMessage, self._model_runnable.invoke(state, config)) # type: ignore[union-attr]
|
||||
response.name = self.name
|
||||
|
||||
if _are_more_steps_needed(state, response):
|
||||
return {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=response.id,
|
||||
content="Sorry, need more steps to process this request.",
|
||||
)
|
||||
]
|
||||
}
|
||||
return {"messages": [response]}
|
||||
|
||||
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.name = self.name
|
||||
|
||||
if _are_more_steps_needed(state, response):
|
||||
return {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=response.id,
|
||||
content="Sorry, need more steps to process this request.",
|
||||
)
|
||||
]
|
||||
}
|
||||
return {"messages": [response]}
|
||||
|
||||
# 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
|
||||
):
|
||||
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_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:
|
||||
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]
|
||||
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:
|
||||
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]
|
||||
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
|
||||
)
|
||||
|
||||
def create_model_router(self) -> Callable[[StateSchema], Union[str, list[Send]]]:
|
||||
"""Create routing function for model 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"
|
||||
elif self.response_format is not None:
|
||||
return "generate_structured_response"
|
||||
else:
|
||||
return END
|
||||
else:
|
||||
if self.version == "v1":
|
||||
return "tools"
|
||||
elif self.version == "v2":
|
||||
if self.post_model_hook is not None:
|
||||
return "post_model_hook"
|
||||
return [
|
||||
Send(
|
||||
"tools",
|
||||
ToolCallWithContext(
|
||||
__type="tool_call_with_context",
|
||||
tool_call=tool_call,
|
||||
state=state,
|
||||
),
|
||||
)
|
||||
for tool_call in last_message.tool_calls
|
||||
]
|
||||
|
||||
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):
|
||||
if not isinstance(m, ToolMessage):
|
||||
break
|
||||
if m.name in self._should_return_direct:
|
||||
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
|
||||
):
|
||||
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 _get_model_edges(self) -> list[str]:
|
||||
"""Get possible edge destinations from model node."""
|
||||
edges = []
|
||||
|
||||
# If post_model_hook exists, we don't add edges here - we use direct edge instead
|
||||
if not self.post_model_hook:
|
||||
if self._has_tools():
|
||||
edges.append("tools")
|
||||
if self.response_format:
|
||||
edges.append("generate_structured_response")
|
||||
if not self._has_tools() and not self.response_format:
|
||||
edges.append(END)
|
||||
|
||||
return edges
|
||||
|
||||
def _get_post_model_hook_edges(self) -> list[str]:
|
||||
"""Get possible edge destinations from post_model_hook node."""
|
||||
edges = [self._get_entry_point()]
|
||||
if self._has_tools():
|
||||
edges.append("tools")
|
||||
if self.response_format:
|
||||
edges.append("generate_structured_response")
|
||||
else:
|
||||
edges.append(END)
|
||||
return edges
|
||||
|
||||
def build(self) -> StateGraph:
|
||||
"""Build the agent workflow graph (uncompiled)."""
|
||||
# Create workflow
|
||||
workflow = StateGraph(
|
||||
state_schema=self._final_state_schema, # type: ignore[arg-type]
|
||||
context_schema=self.context_schema,
|
||||
)
|
||||
|
||||
# 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_edges = self._get_post_model_hook_edges()
|
||||
workflow.add_conditional_edges(
|
||||
"post_model_hook", self.post_model_hook_router, path_map=post_hook_edges
|
||||
) # type: ignore[arg-type]
|
||||
else:
|
||||
# Conditional edges from model node when no post_model_hook
|
||||
model_router = self.create_model_router()
|
||||
model_edges = self._get_model_edges()
|
||||
workflow.add_conditional_edges("agent", model_router, path_map=model_edges) # 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
|
||||
|
||||
|
||||
def create_react_agent(
|
||||
model: Union[str, LanguageModelLike],
|
||||
tools: Union[Sequence[Union[BaseTool, Callable, dict[str, Any]]], ToolNode],
|
||||
@@ -411,6 +840,7 @@ def create_react_agent(
|
||||
print(chunk)
|
||||
```
|
||||
"""
|
||||
# Handle deprecated config_schema parameter
|
||||
if (
|
||||
config_schema := deprecated_kwargs.pop("config_schema", MISSING)
|
||||
) is not MISSING:
|
||||
@@ -422,400 +852,29 @@ def create_react_agent(
|
||||
if context_schema is not None:
|
||||
context_schema = config_schema
|
||||
|
||||
# Validate version
|
||||
if version not in ("v1", "v2"):
|
||||
raise ValueError(
|
||||
f"Invalid version {version}. Supported versions are 'v1' and 'v2'."
|
||||
)
|
||||
|
||||
if state_schema is not None:
|
||||
required_keys = {"messages", "remaining_steps"}
|
||||
if response_format is not None:
|
||||
required_keys.add("structured_response")
|
||||
|
||||
schema_keys = set(get_type_hints(state_schema))
|
||||
if missing_keys := required_keys - set(schema_keys):
|
||||
raise ValueError(f"Missing required key(s) {missing_keys} in state_schema")
|
||||
|
||||
if state_schema is None:
|
||||
state_schema = (
|
||||
AgentStateWithStructuredResponse
|
||||
if response_format is not None
|
||||
else AgentState
|
||||
)
|
||||
|
||||
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)])
|
||||
tool_classes = list(tool_node.tools_by_name.values())
|
||||
|
||||
if isinstance(model, str):
|
||||
try:
|
||||
from langchain.chat_models import ( # type: ignore[import-not-found]
|
||||
init_chat_model,
|
||||
)
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Please install langchain (`pip install langchain`) to use '<provider>:<model>' string syntax for `model` parameter."
|
||||
)
|
||||
|
||||
model = cast(BaseChatModel, init_chat_model(model))
|
||||
|
||||
tool_calling_enabled = len(tool_classes) > 0
|
||||
|
||||
if (
|
||||
_should_bind_tools(model, tool_classes, num_builtin=len(llm_builtin_tools))
|
||||
and len(tool_classes + llm_builtin_tools) > 0
|
||||
):
|
||||
model = cast(BaseChatModel, model).bind_tools(tool_classes + llm_builtin_tools) # type: ignore[operator]
|
||||
|
||||
model_runnable = _get_prompt_runnable(prompt) | model
|
||||
|
||||
# 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}
|
||||
|
||||
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 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)
|
||||
)
|
||||
|
||||
def _get_model_input_state(state: StateSchema) -> StateSchema:
|
||||
if pre_model_hook is not None:
|
||||
messages = (
|
||||
_get_state_value(state, "llm_input_messages")
|
||||
) or _get_state_value(state, "messages")
|
||||
error_msg = f"Expected input to call_model to have 'llm_input_messages' or 'messages' key, but got {state}"
|
||||
else:
|
||||
messages = _get_state_value(state, "messages")
|
||||
error_msg = (
|
||||
f"Expected input to call_model to have 'messages' key, but got {state}"
|
||||
)
|
||||
|
||||
if messages is None:
|
||||
raise ValueError(error_msg)
|
||||
|
||||
_validate_chat_history(messages)
|
||||
# we're passing messages under `messages` key, as this is expected by the prompt
|
||||
if isinstance(state_schema, type) and issubclass(state_schema, BaseModel):
|
||||
state.messages = messages # type: ignore
|
||||
else:
|
||||
state["messages"] = messages # type: ignore
|
||||
|
||||
return state
|
||||
|
||||
# Define the function that calls the model
|
||||
def call_model(state: StateSchema, config: RunnableConfig) -> StateSchema:
|
||||
state = _get_model_input_state(state)
|
||||
response = cast(AIMessage, model_runnable.invoke(state, config))
|
||||
# add agent name to the AIMessage
|
||||
response.name = name
|
||||
|
||||
if _are_more_steps_needed(state, response):
|
||||
return {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=response.id,
|
||||
content="Sorry, need more steps to process this request.",
|
||||
)
|
||||
]
|
||||
}
|
||||
# We return a list, because this will get added to the existing list
|
||||
return {"messages": [response]}
|
||||
|
||||
async def acall_model(state: StateSchema, config: RunnableConfig) -> StateSchema:
|
||||
state = _get_model_input_state(state)
|
||||
response = cast(AIMessage, await model_runnable.ainvoke(state, config))
|
||||
# add agent name to the AIMessage
|
||||
response.name = name
|
||||
if _are_more_steps_needed(state, response):
|
||||
return {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
id=response.id,
|
||||
content="Sorry, need more steps to process this request.",
|
||||
)
|
||||
]
|
||||
}
|
||||
# We return a list, because this will get added to the existing list
|
||||
return {"messages": [response]}
|
||||
|
||||
input_schema: StateSchemaType
|
||||
if pre_model_hook is not None:
|
||||
# Dynamically create a schema that inherits from state_schema and adds 'llm_input_messages'
|
||||
if isinstance(state_schema, type) and issubclass(state_schema, BaseModel):
|
||||
# For Pydantic schemas
|
||||
from pydantic import create_model
|
||||
|
||||
input_schema = create_model(
|
||||
"CallModelInputSchema",
|
||||
llm_input_messages=(list[AnyMessage], ...),
|
||||
__base__=state_schema,
|
||||
)
|
||||
else:
|
||||
# For TypedDict schemas
|
||||
class CallModelInputSchema(state_schema): # type: ignore
|
||||
llm_input_messages: list[AnyMessage]
|
||||
|
||||
input_schema = CallModelInputSchema
|
||||
else:
|
||||
input_schema = state_schema
|
||||
|
||||
def generate_structured_response(
|
||||
state: StateSchema, 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)
|
||||
|
||||
model_with_structured_output = _get_model(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, 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)
|
||||
|
||||
model_with_structured_output = _get_model(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)
|
||||
workflow.add_node(
|
||||
"agent",
|
||||
RunnableCallable(call_model, acall_model),
|
||||
input_schema=input_schema,
|
||||
)
|
||||
if pre_model_hook is not None:
|
||||
workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type]
|
||||
workflow.add_edge("pre_model_hook", "agent")
|
||||
entrypoint = "pre_model_hook"
|
||||
else:
|
||||
entrypoint = "agent"
|
||||
|
||||
workflow.set_entry_point(entrypoint)
|
||||
|
||||
if post_model_hook is not None:
|
||||
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,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
debug=debug,
|
||||
name=name,
|
||||
)
|
||||
|
||||
# Define the function that determines whether to continue or not
|
||||
def should_continue(state: StateSchema) -> Union[str, list[Send]]:
|
||||
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
|
||||
# Otherwise if there is, we continue
|
||||
else:
|
||||
if version == "v1":
|
||||
return "tools"
|
||||
elif version == "v2":
|
||||
if post_model_hook is not None:
|
||||
return "post_model_hook"
|
||||
return [
|
||||
Send(
|
||||
"tools",
|
||||
ToolCallWithContext(
|
||||
__type="tool_call_with_context",
|
||||
tool_call=tool_call,
|
||||
state=state,
|
||||
),
|
||||
)
|
||||
for tool_call in last_message.tool_calls
|
||||
]
|
||||
|
||||
# Define a new graph
|
||||
workflow = StateGraph(
|
||||
state_schema=state_schema or AgentState, context_schema=context_schema
|
||||
# Build the graph using the internal builder
|
||||
builder = _AgentBuilder(
|
||||
model=model,
|
||||
tools=tools,
|
||||
prompt=prompt,
|
||||
response_format=response_format,
|
||||
pre_model_hook=pre_model_hook,
|
||||
post_model_hook=post_model_hook,
|
||||
state_schema=state_schema,
|
||||
context_schema=context_schema,
|
||||
version=version,
|
||||
name=name,
|
||||
)
|
||||
|
||||
# Define the two nodes we will cycle between
|
||||
workflow.add_node(
|
||||
"agent",
|
||||
RunnableCallable(call_model, acall_model),
|
||||
input_schema=input_schema,
|
||||
)
|
||||
workflow.add_node("tools", tool_node)
|
||||
workflow = builder.build()
|
||||
|
||||
# Optionally add a pre-model hook node that will be called
|
||||
# every time before the "agent" (LLM-calling node)
|
||||
if pre_model_hook is not None:
|
||||
workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type]
|
||||
workflow.add_edge("pre_model_hook", "agent")
|
||||
entrypoint = "pre_model_hook"
|
||||
else:
|
||||
entrypoint = "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.
|
||||
"""
|
||||
|
||||
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 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, # type: ignore[arg-type]
|
||||
path_map=post_model_hook_paths,
|
||||
)
|
||||
|
||||
workflow.add_conditional_edges(
|
||||
"agent",
|
||||
should_continue, # type: ignore[arg-type]
|
||||
path_map=agent_paths,
|
||||
)
|
||||
|
||||
def route_tool_responses(state: StateSchema) -> str:
|
||||
for m in reversed(_get_state_value(state, "messages")):
|
||||
if not isinstance(m, ToolMessage):
|
||||
break
|
||||
if m.name in should_return_direct:
|
||||
return END
|
||||
|
||||
# handle a case of parallel tool calls where
|
||||
# the tool w/ `return_direct` was executed in a different `Send`
|
||||
if isinstance(m, AIMessage) and m.tool_calls:
|
||||
if any(call["name"] in should_return_direct for call in m.tool_calls):
|
||||
return END
|
||||
|
||||
return entrypoint
|
||||
|
||||
if should_return_direct:
|
||||
workflow.add_conditional_edges(
|
||||
"tools", route_tool_responses, path_map=[entrypoint, END]
|
||||
)
|
||||
else:
|
||||
workflow.add_edge("tools", entrypoint)
|
||||
|
||||
# Finally, we compile it!
|
||||
# This compiles it into a LangChain Runnable,
|
||||
# meaning you can use it as you would any other runnable
|
||||
# Compile and return the graph
|
||||
return workflow.compile(
|
||||
checkpointer=checkpointer,
|
||||
store=store,
|
||||
|
||||
Reference in New Issue
Block a user