diff --git a/libs/langgraph/langgraph/agent/__init__.py b/libs/langgraph/langgraph/agent/__init__.py new file mode 100644 index 000000000..a7b8c1d19 --- /dev/null +++ b/libs/langgraph/langgraph/agent/__init__.py @@ -0,0 +1,178 @@ +from collections.abc import Sequence +from typing import Callable, cast + +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.messages import AIMessage, SystemMessage +from langchain_core.tools import BaseTool + +from langgraph.agent.types import AgentInput, AgentMiddleware, AgentState, ModelRequest +from langgraph.constants import END, START +from langgraph.graph.state import StateGraph +from langgraph.prebuilt.tool_node import ToolNode + + +def create_agent( + *, + model: str | BaseChatModel, + tools: Sequence[BaseTool | Callable], + system_prompt: str, + middleware: Sequence[AgentMiddleware] = (), +) -> StateGraph[AgentState, None, AgentInput]: + # init chat model + 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 ':' string syntax for `model` parameter." + ) + + model = cast(BaseChatModel, init_chat_model(model)) + + # init tool node + tool_node = ToolNode(tools=tools) + + # validate middleware + assert len({m.__class__.__name__ for m in middleware}) == len(middleware), ( + "Please remove duplicate middleware instances." + ) # this is just to keep the node names simple, we can change if needed + middleware_w_before = [ + m + for m in middleware + if m.__class__.before_model is not AgentMiddleware.before_model + ] + middleware_w_after = [ + m + for m in middleware + if m.__class__.after_model is not AgentMiddleware.after_model + ] + + # create graph, add nodes + graph = StateGraph(AgentState, input_schema=AgentInput) + graph.add_node( + "model_request", + _make_model_request_node( + model=model, + tools=list(tool_node.tools_by_name.values()), + system_prompt=system_prompt, + middleware=middleware, + ), + ) + graph.add_node("tools", tool_node) + for m in middleware: + if m.__class__.before_model is not AgentMiddleware.before_model: + graph.add_node(f"{m.__class__.__name__}.before_model", m.before_model) + if m.__class__.after_model is not AgentMiddleware.after_model: + graph.add_node(f"{m.__class__.__name__}.after_model", m.after_model) + + # add start edge + first_node = ( + f"{middleware_w_before[0].__class__.__name__}.before_model" + if middleware_w_before + else "model_request" + ) + last_node = ( + f"{middleware_w_after[0].__class__.__name__}.after_model" + if middleware_w_after + else "model_request" + ) + graph.add_edge(START, first_node) + + # add cond edges + graph.add_conditional_edges( + "tools", + _make_tools_to_model_edge(tool_node, first_node), + [first_node, END], + ) + graph.add_conditional_edges(last_node, _make_model_to_tools_edge(), ["tools", END]) + + # add before model edges + if middleware_w_before: + for m1, m2 in zip(middleware_w_before, middleware_w_before[1:]): + graph.add_edge( + f"{m1.__class__.__name__}.before_model", + f"{m2.__class__.__name__}.before_model", + ) + graph.add_edge( + f"{middleware_w_before[-1].__class__.__name__}.before_model", + "model_request", + ) + + # add after model edges + if middleware_w_after: + graph.add_edge( + "model_request", f"{middleware_w_after[-1].__class__.__name__}.after_model" + ) + for idx in range(len(middleware_w_after) - 1, 0, -1): + m1 = middleware_w_after[idx] + m2 = middleware_w_after[idx - 1] + graph.add_edge( + f"{m1.__class__.__name__}.after_model", + f"{m2.__class__.__name__}.after_model", + ) + + return graph + + +def _make_model_request_node( + *, + system_prompt: str, + model: BaseChatModel, + tools: Sequence[BaseTool], + middleware: Sequence[AgentMiddleware], +) -> Callable[[AgentState], AgentState]: + def model_request(state: AgentState) -> AgentState: + # create request + request = ModelRequest( + model=model, + system_prompt=system_prompt, + messages=state.messages, + tool_choice=None, + tools=tools, + ) + # visit middleware in order + for mw in middleware: + request = mw.modify_model_request(request) + # prepare messages + if request.system_prompt: + messages = [SystemMessage(request.system_prompt)] + request.messages + else: + messages = request.messages + # call model + output = request.model.invoke( + messages, tools=request.tools, tool_choice=request.tool_choice + ) + return {"messages": output} + + return model_request + + +def _make_model_to_tools_edge() -> Callable[[AgentState], str | None]: + def model_to_tools(state: AgentState) -> str | None: + message = state.messages[-1] + if isinstance(message, AIMessage) and message.tool_calls: + return "tools" + + return END + + return model_to_tools + + +def _make_tools_to_model_edge( + tool_node: ToolNode, next_node: str +) -> Callable[[AgentState], str | None]: + def tools_to_model(state: AgentState) -> str | None: + ai_message = [m for m in state.messages if isinstance(m, AIMessage)][-1] + if all( + tool_node.tools_by_name[c["name"]].return_direct + for c in ai_message.tool_calls + if c["name"] in tool_node.tools_by_name + ): + return END + + return next_node + + return tools_to_model diff --git a/libs/langgraph/langgraph/agent/types.py b/libs/langgraph/langgraph/agent/types.py new file mode 100644 index 000000000..54ed7cc84 --- /dev/null +++ b/libs/langgraph/langgraph/agent/types.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Annotated, Any, Self + +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.messages import AnyMessage +from langchain_core.tools import BaseTool +from typing_extensions import TypedDict + +from langgraph.graph.message import Messages, add_messages + + +@dataclass +class ModelRequest: + model: BaseChatModel + system_prompt: str + messages: Sequence[AnyMessage] # excluding system prompt + tool_choice: Any + tools: Sequence[BaseTool] + + +class AgentMiddleware: + def __init__(self) -> None: + pass + + def __copy__(self) -> Self: + return self.__class__(**self.__dict__) + + def before_model(self, state: AgentState) -> AgentState | None: + pass + + def modify_model_request(self, request: ModelRequest) -> ModelRequest: + return request + + def after_model(self, state: AgentState) -> AgentState | None: + pass + + +class AgentInput(TypedDict): + messages: Messages + + +@dataclass +class AgentState: + messages: Annotated[list[AnyMessage], add_messages] diff --git a/libs/langgraph/tests/__snapshots__/test_agent.ambr b/libs/langgraph/tests/__snapshots__/test_agent.ambr new file mode 100644 index 000000000..8aa887354 --- /dev/null +++ b/libs/langgraph/tests/__snapshots__/test_agent.ambr @@ -0,0 +1,279 @@ +# serializer version: 1 +# name: test_create_agent_diagram + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + __end__([

__end__

]):::last + __start__ --> model_request; + model_request -.-> __end__; + model_request -.-> tools; + tools -.-> __end__; + tools -.-> model_request; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.1 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopOne_before_model(NoopOne.before_model) + __end__([

__end__

]):::last + NoopOne_before_model --> model_request; + __start__ --> NoopOne_before_model; + model_request -.-> __end__; + model_request -.-> tools; + tools -.-> NoopOne_before_model; + tools -.-> __end__; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.2 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopOne_before_model(NoopOne.before_model) + NoopTwo_before_model(NoopTwo.before_model) + __end__([

__end__

]):::last + NoopOne_before_model --> NoopTwo_before_model; + NoopTwo_before_model --> model_request; + __start__ --> NoopOne_before_model; + model_request -.-> __end__; + model_request -.-> tools; + tools -.-> NoopOne_before_model; + tools -.-> __end__; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.3 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopOne_before_model(NoopOne.before_model) + NoopTwo_before_model(NoopTwo.before_model) + NoopThree_before_model(NoopThree.before_model) + __end__([

__end__

]):::last + NoopOne_before_model --> NoopTwo_before_model; + NoopThree_before_model --> model_request; + NoopTwo_before_model --> NoopThree_before_model; + __start__ --> NoopOne_before_model; + model_request -.-> __end__; + model_request -.-> tools; + tools -.-> NoopOne_before_model; + tools -.-> __end__; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.4 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopFour_after_model(NoopFour.after_model) + __end__([

__end__

]):::last + NoopFour_after_model -.-> __end__; + NoopFour_after_model -.-> tools; + __start__ --> model_request; + model_request --> NoopFour_after_model; + tools -.-> __end__; + tools -.-> model_request; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.5 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopFour_after_model(NoopFour.after_model) + NoopFive_after_model(NoopFive.after_model) + __end__([

__end__

]):::last + NoopFive_after_model --> NoopFour_after_model; + NoopFour_after_model -.-> __end__; + NoopFour_after_model -.-> tools; + __start__ --> model_request; + model_request --> NoopFive_after_model; + tools -.-> __end__; + tools -.-> model_request; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.6 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopFour_after_model(NoopFour.after_model) + NoopFive_after_model(NoopFive.after_model) + NoopSix_after_model(NoopSix.after_model) + __end__([

__end__

]):::last + NoopFive_after_model --> NoopFour_after_model; + NoopFour_after_model -.-> __end__; + NoopFour_after_model -.-> tools; + NoopSix_after_model --> NoopFive_after_model; + __start__ --> model_request; + model_request --> NoopSix_after_model; + tools -.-> __end__; + tools -.-> model_request; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.7 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopSeven_before_model(NoopSeven.before_model) + NoopSeven_after_model(NoopSeven.after_model) + __end__([

__end__

]):::last + NoopSeven_after_model -.-> __end__; + NoopSeven_after_model -.-> tools; + NoopSeven_before_model --> model_request; + __start__ --> NoopSeven_before_model; + model_request --> NoopSeven_after_model; + tools -.-> NoopSeven_before_model; + tools -.-> __end__; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.8 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopSeven_before_model(NoopSeven.before_model) + NoopSeven_after_model(NoopSeven.after_model) + NoopEight_before_model(NoopEight.before_model) + NoopEight_after_model(NoopEight.after_model) + __end__([

__end__

]):::last + NoopEight_after_model --> NoopSeven_after_model; + NoopEight_before_model --> model_request; + NoopSeven_after_model -.-> __end__; + NoopSeven_after_model -.-> tools; + NoopSeven_before_model --> NoopEight_before_model; + __start__ --> NoopSeven_before_model; + model_request --> NoopEight_after_model; + tools -.-> NoopSeven_before_model; + tools -.-> __end__; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_create_agent_diagram.9 + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::first + model_request(model_request) + tools(tools) + NoopSeven_before_model(NoopSeven.before_model) + NoopSeven_after_model(NoopSeven.after_model) + NoopEight_before_model(NoopEight.before_model) + NoopEight_after_model(NoopEight.after_model) + NoopNine_before_model(NoopNine.before_model) + NoopNine_after_model(NoopNine.after_model) + __end__([

__end__

]):::last + NoopEight_after_model --> NoopSeven_after_model; + NoopEight_before_model --> NoopNine_before_model; + NoopNine_after_model --> NoopEight_after_model; + NoopNine_before_model --> model_request; + NoopSeven_after_model -.-> __end__; + NoopSeven_after_model -.-> tools; + NoopSeven_before_model --> NoopEight_before_model; + __start__ --> NoopSeven_before_model; + model_request --> NoopNine_after_model; + tools -.-> NoopSeven_before_model; + tools -.-> __end__; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- diff --git a/libs/langgraph/tests/test_agent.py b/libs/langgraph/tests/test_agent.py new file mode 100644 index 000000000..b75a282bd --- /dev/null +++ b/libs/langgraph/tests/test_agent.py @@ -0,0 +1,226 @@ +from langchain_core.messages import AIMessage, ToolCall +from syrupy import SnapshotAssertion + +from langgraph.agent import create_agent +from langgraph.agent.types import AgentMiddleware +from langgraph.checkpoint.base import BaseCheckpointSaver +from tests.fake_chat import FakeChatModel +from tests.messages import _AnyIdToolMessage + + +def test_create_agent_diagram( + snapshot: SnapshotAssertion, +): + class NoopOne(AgentMiddleware): + def before_model(self, state): + pass + + class NoopTwo(AgentMiddleware): + def before_model(self, state): + pass + + class NoopThree(AgentMiddleware): + def before_model(self, state): + pass + + class NoopFour(AgentMiddleware): + def after_model(self, state): + pass + + class NoopFive(AgentMiddleware): + def after_model(self, state): + pass + + class NoopSix(AgentMiddleware): + def after_model(self, state): + pass + + class NoopSeven(AgentMiddleware): + def before_model(self, state): + pass + + def after_model(self, state): + pass + + class NoopEight(AgentMiddleware): + def before_model(self, state): + pass + + def after_model(self, state): + pass + + class NoopNine(AgentMiddleware): + def before_model(self, state): + pass + + def after_model(self, state): + pass + + agent_zero = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + ) + + assert agent_zero.compile().get_graph().draw_mermaid() == snapshot + + agent_one = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopOne()], + ) + + assert agent_one.compile().get_graph().draw_mermaid() == snapshot + + agent_two = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopOne(), NoopTwo()], + ) + + assert agent_two.compile().get_graph().draw_mermaid() == snapshot + + agent_three = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopOne(), NoopTwo(), NoopThree()], + ) + + assert agent_three.compile().get_graph().draw_mermaid() == snapshot + + agent_four = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopFour()], + ) + + assert agent_four.compile().get_graph().draw_mermaid() == snapshot + + agent_five = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopFour(), NoopFive()], + ) + + assert agent_five.compile().get_graph().draw_mermaid() == snapshot + + agent_six = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopFour(), NoopFive(), NoopSix()], + ) + + assert agent_six.compile().get_graph().draw_mermaid() == snapshot + + agent_seven = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopSeven()], + ) + + assert agent_seven.compile().get_graph().draw_mermaid() == snapshot + + agent_eight = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopSeven(), NoopEight()], + ) + + assert agent_eight.compile().get_graph().draw_mermaid() == snapshot + + agent_nine = create_agent( + model=FakeChatModel(messages=[]), + tools=[], + system_prompt="You are a helpful assistant.", + middleware=[NoopSeven(), NoopEight(), NoopNine()], + ) + + assert agent_nine.compile().get_graph().draw_mermaid() == snapshot + + +def test_create_agent_invoke( + snapshot: SnapshotAssertion, + sync_checkpointer: BaseCheckpointSaver, +): + calls = [] + + class NoopSeven(AgentMiddleware): + def before_model(self, state): + calls.append("NoopSeven.before_model") + + def after_model(self, state): + calls.append("NoopSeven.after_model") + + class NoopEight(AgentMiddleware): + def before_model(self, state): + calls.append("NoopEight.before_model") + + def after_model(self, state): + calls.append("NoopEight.after_model") + + def my_tool(input: str) -> str: + """A great tool""" + calls.append("my_tool") + return input.upper() + + agent_one = create_agent( + model=FakeChatModel( + messages=[ + AIMessage( + "", + id="ai1", + tool_calls=[ToolCall(id="1", name="my_tool", args={"input": "yo"})], + ), + AIMessage(id="ai2", content="Hello, how can I assist you today?"), + ] + ), + tools=[my_tool], + system_prompt="You are a helpful assistant.", + middleware=[NoopSeven(), NoopEight()], + ).compile(checkpointer=sync_checkpointer) + + thread1 = {"configurable": {"thread_id": "1"}} + assert agent_one.invoke({"messages": []}, thread1) == { + "messages": [ + AIMessage( + content="", + additional_kwargs={}, + response_metadata={}, + id="ai1", + tool_calls=[ + { + "name": "my_tool", + "args": {"input": "yo"}, + "id": "1", + "type": "tool_call", + } + ], + ), + _AnyIdToolMessage(content="YO", name="my_tool", tool_call_id="1"), + AIMessage( + content="Hello, how can I assist you today?", + additional_kwargs={}, + response_metadata={}, + id="ai2", + ), + ] + } + assert calls == [ + "NoopSeven.before_model", + "NoopEight.before_model", + "NoopEight.after_model", + "NoopSeven.after_model", + "my_tool", + "NoopSeven.before_model", + "NoopEight.before_model", + "NoopEight.after_model", + "NoopSeven.after_model", + ]