mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-12 20:57:52 +02:00
minimalistic tool registration
This commit is contained in:
@@ -43,8 +43,14 @@ def create_agent(
|
||||
|
||||
model = cast(BaseChatModel, init_chat_model(model))
|
||||
|
||||
# init tool node
|
||||
tool_node = tools if isinstance(tools, ToolNode) else ToolNode(tools=tools)
|
||||
if isinstance(tools, list):
|
||||
all_tools = [t for m in middleware for t in m.tools] + list(tools)
|
||||
tool_node = ToolNode(tools=all_tools)
|
||||
default_tools = tools
|
||||
else:
|
||||
# TODO: what do we do when middleware tools are specified, plus a ToolNode is used?
|
||||
tool_node = tools
|
||||
default_tools = list(tool_node.tools_by_name.values())
|
||||
|
||||
# validate middleware
|
||||
assert len({m.__class__.__name__ for m in middleware}) == len(middleware), (
|
||||
@@ -77,7 +83,7 @@ def create_agent(
|
||||
def model_request(state: AgentState) -> AgentState:
|
||||
request = state.model_request or ModelRequest(
|
||||
model=model,
|
||||
tools=list(tool_node.tools_by_name.values()),
|
||||
tools=default_tools,
|
||||
system_prompt=system_prompt,
|
||||
response_format=response_format,
|
||||
messages=state.messages,
|
||||
@@ -85,6 +91,7 @@ def create_agent(
|
||||
)
|
||||
|
||||
# prepare messages
|
||||
print(request.system_prompt)
|
||||
if request.system_prompt:
|
||||
messages = [SystemMessage(request.system_prompt)] + request.messages
|
||||
else:
|
||||
@@ -127,7 +134,7 @@ def create_agent(
|
||||
|
||||
default_model_request = ModelRequest(
|
||||
model=model,
|
||||
tools=list(tool_node.tools_by_name.values()),
|
||||
tools=default_tools,
|
||||
system_prompt=system_prompt,
|
||||
response_format=response_format,
|
||||
messages=state.messages,
|
||||
|
||||
@@ -15,8 +15,8 @@ class SwarmAgent:
|
||||
tools: list[BaseTool]
|
||||
|
||||
|
||||
class SwarmMiddleware(AgentMiddleware):
|
||||
"""Swarm middleware.
|
||||
class MultiAgentMiddleware(AgentMiddleware):
|
||||
"""Multi agent middleware (enabling swarm like behavior).
|
||||
|
||||
TODOs:
|
||||
* Support create_agent for handoffs
|
||||
@@ -59,6 +59,7 @@ class SwarmMiddleware(AgentMiddleware):
|
||||
request.tools = agent.tools
|
||||
|
||||
request.tools.extend(self.handoff_tools)
|
||||
|
||||
return request
|
||||
|
||||
def after_model(self, state) -> State | None:
|
||||
|
||||
@@ -2,6 +2,7 @@ from dataclasses import dataclass, field
|
||||
from typing import Annotated, Any, Dict, List, cast
|
||||
|
||||
from langchain_core.messages import AIMessage
|
||||
from typing_extensions import Annotated
|
||||
|
||||
from langgraph.agent.types import AgentJump, AgentMiddleware, AgentState, AgentUpdate
|
||||
|
||||
@@ -11,7 +12,15 @@ class ToolCallLimitMiddleware(AgentMiddleware):
|
||||
|
||||
@dataclass
|
||||
class State(AgentMiddleware.State):
|
||||
tool_call_count: dict[str, int] = field(default_factory=dict)
|
||||
important: Annotated[dict[str, int], Input, Output] = field(default_factory=dict)
|
||||
|
||||
@dataclass
|
||||
class InputState(AgentMiddleware.State):
|
||||
important: dict[str, int]
|
||||
|
||||
@dataclass
|
||||
class OutputState(AgentMiddleware.State):
|
||||
important: dict[str, int]
|
||||
|
||||
def __init__(self, tool_limits: dict[str, int]):
|
||||
self.tool_limits = tool_limits
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.ephemeral_value import EphemeralValue
|
||||
from langgraph.graph.message import Messages, add_messages
|
||||
from langgraph.runtime import Runtime
|
||||
|
||||
ResponseFormat = dict | type[BaseModel]
|
||||
JumpTo = Literal["tools", "model", "__end__"]
|
||||
@@ -39,6 +40,8 @@ class AgentMiddleware:
|
||||
class State(AgentState):
|
||||
pass
|
||||
|
||||
tools: list[BaseTool]
|
||||
|
||||
def before_model(self, state: State) -> AgentUpdate | AgentJump | None:
|
||||
pass
|
||||
|
||||
|
||||
Reference in New Issue
Block a user