diff --git a/pyagentspec/src/pyagentspec/adapters/langgraph/_langgraphconverter.py b/pyagentspec/src/pyagentspec/adapters/langgraph/_langgraphconverter.py index a07c7365..fb54281f 100644 --- a/pyagentspec/src/pyagentspec/adapters/langgraph/_langgraphconverter.py +++ b/pyagentspec/src/pyagentspec/adapters/langgraph/_langgraphconverter.py @@ -197,9 +197,31 @@ def convert( converted_components: Optional[Dict[str, Any]] = None, checkpointer: Optional[Checkpointer] = None, config: Optional[RunnableConfig] = None, + middleware: Optional[List[Any]] = None, **kwargs: Any, ) -> Any: - """Convert the given PyAgentSpec component object into the corresponding LangGraph component""" + """Convert the given PyAgentSpec component object into the corresponding LangGraph component. + + Parameters + ---------- + agentspec_component: + The Agent Spec component to convert. + tool_registry: + Dictionary mapping tool names to LangGraph tool objects. + converted_components: + Optional cache of already-converted components (keyed by component id). + checkpointer: + Optional LangGraph checkpointer to wire into created graphs. + config: + Optional ``RunnableConfig`` to pass to created runnables/graphs. + middleware: + Optional list of LangChain agent middleware instances forwarded to + ``langchain_agents.create_agent(middleware=...)`` when compiling an Agent + Spec ``Agent`` into a ReAct graph. Order is preserved — index ``0`` is the + outermost middleware. When ``None`` or an empty list, the ``middleware`` + keyword is omitted entirely from the ``create_agent`` call. + """ + middleware_list: List[Any] = list(middleware or []) if converted_components is None: converted_components = {} if config is None: @@ -209,7 +231,12 @@ def convert( config = RunnableConfig({}) if agentspec_component.id not in converted_components: converted_components[agentspec_component.id] = self._convert( - agentspec_component, tool_registry, converted_components, checkpointer, config + agentspec_component, + tool_registry, + converted_components, + checkpointer, + config, + middleware_list, ) return converted_components[agentspec_component.id] @@ -220,6 +247,7 @@ def _convert( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> Any: if isinstance(agentspec_component, AgentSpecAgent): return self._agent_convert_to_langgraph( @@ -228,6 +256,7 @@ def _convert( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(agentspec_component, AgentSpecSwarm): return self._swarm_convert_to_langgraph( @@ -236,6 +265,7 @@ def _convert( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(agentspec_component, AgentSpecLlmConfig): return self._llm_convert_to_langgraph(agentspec_component, config=config) @@ -274,6 +304,7 @@ def _convert( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(agentspec_component, AgentSpecNode): return self._node_convert_to_langgraph( @@ -282,6 +313,7 @@ def _convert( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(agentspec_component, AgentSpecComponent): raise NotImplementedError( @@ -325,6 +357,7 @@ def _flow_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> CompiledStateGraph[Any, Any, Any]: graph_builder = StateGraph( @@ -340,6 +373,7 @@ def _flow_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) for node in flow.nodes } @@ -520,6 +554,7 @@ def _node_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> "NodeExecutor": if isinstance(node, AgentSpecStartNode): return self._start_node_convert_to_langgraph(node) @@ -548,6 +583,7 @@ def _node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(node, AgentSpecBranchingNode): return self._branching_node_convert_to_langgraph(node) @@ -560,6 +596,7 @@ def _node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(node, AgentSpecCatchExceptionNode): return self._catch_exception_node_convert_to_langgraph( @@ -568,6 +605,7 @@ def _node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) elif isinstance(node, AgentSpecInputMessageNode): return self._input_message_node_convert_to_langgraph(node) @@ -580,6 +618,7 @@ def _node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) else: raise NotImplementedError( @@ -609,6 +648,7 @@ def _map_node_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> "NodeExecutor": from pyagentspec.adapters.langgraph._node_execution import MapNodeExecutor @@ -618,6 +658,7 @@ def _map_node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) if not isinstance(subflow, CompiledStateGraph): raise TypeError("MapNodeExecutor can only be initialized with MapNode") @@ -631,6 +672,7 @@ def _flow_node_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> "NodeExecutor": from pyagentspec.adapters.langgraph._node_execution import FlowNodeExecutor @@ -640,6 +682,7 @@ def _flow_node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) if not isinstance(subflow, CompiledStateGraph): raise TypeError("FlowNodeExecutor can only initialize FlowNode") @@ -657,6 +700,7 @@ def _catch_exception_node_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> "NodeExecutor": from pyagentspec.adapters.langgraph._node_execution import CatchExceptionNodeExecutor @@ -666,6 +710,7 @@ def _catch_exception_node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) if not isinstance(subflow, CompiledStateGraph): raise TypeError( @@ -698,6 +743,7 @@ def _agent_node_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> "NodeExecutor": from pyagentspec.adapters.langgraph._node_execution import AgentNodeExecutor @@ -707,6 +753,7 @@ def _agent_node_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) def _llm_node_convert_to_langgraph( @@ -993,6 +1040,7 @@ def _swarm_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> CompiledStateGraph[Any, Any, Any]: if agentspec_component.handoff is AgentSpecHandoffMode.NEVER: # As of now, we cannot control what langgraph-swarm does internally in terms of conversation sharing. @@ -1023,6 +1071,7 @@ def _swarm_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) handoffs: dict[str, list[str]] = {agent_name: [] for agent_name in agents} for from_agent, to_agent in agentspec_component.relationships: @@ -1042,6 +1091,7 @@ def _swarm_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, additional_langgraph_tools=[ langgraph_swarm.create_handoff_tool(agent_name=to_agent_name) for to_agent_name in handoffs.get(agent.name, []) @@ -1069,6 +1119,7 @@ def _create_react_agent_with_given_info( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], additional_langgraph_tools: Optional[List[LangGraphTool]] = None, ) -> CompiledStateGraph[Any, Any, Any]: model = self.convert( @@ -1115,7 +1166,7 @@ def _create_react_agent_with_given_info( inputs=inputs, ) - compiled_graph = langchain_agents.create_agent( + create_agent_kwargs: Dict[str, Any] = dict( name=name, model=model, tools=langgraph_tools, @@ -1124,6 +1175,9 @@ def _create_react_agent_with_given_info( response_format=output_model, state_schema=state_schema, ) + if middleware: + create_agent_kwargs["middleware"] = middleware + compiled_graph = langchain_agents.create_agent(**create_agent_kwargs) # To enable flow execution traces monkey patch all the functions that invoke the compiled graph @@ -1209,6 +1263,7 @@ def _agent_convert_to_langgraph( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: List[Any], ) -> CompiledStateGraph[Any, Any, Any]: return self._create_react_agent_with_given_info( name=agentspec_component.name, @@ -1223,6 +1278,7 @@ def _agent_convert_to_langgraph( converted_components=converted_components, checkpointer=checkpointer, config=config, + middleware=middleware, ) def _llm_convert_to_langgraph( diff --git a/pyagentspec/src/pyagentspec/adapters/langgraph/_node_execution.py b/pyagentspec/src/pyagentspec/adapters/langgraph/_node_execution.py index dab0a0c7..f1999aef 100644 --- a/pyagentspec/src/pyagentspec/adapters/langgraph/_node_execution.py +++ b/pyagentspec/src/pyagentspec/adapters/langgraph/_node_execution.py @@ -487,6 +487,7 @@ def __init__( converted_components: Dict[str, Any], checkpointer: Optional[Checkpointer], config: RunnableConfig, + middleware: Optional[List[Any]] = None, ) -> None: super().__init__(node) if not isinstance(self.node, AgentSpecAgentNode): @@ -495,6 +496,7 @@ def __init__( self.checkpointer = checkpointer self.converted_components = converted_components self.config = config + self._middleware: List[Any] = list(middleware or []) self._agents_cache: Dict[str, CompiledStateGraph[Any, Any]] = {} def _create_react_agent_with_given_input_values( @@ -523,6 +525,7 @@ def _create_react_agent_with_given_input_values( converted_components=self.converted_components, checkpointer=self.checkpointer, config=self.config, + middleware=self._middleware, ) return self._agents_cache[system_prompt] diff --git a/pyagentspec/src/pyagentspec/adapters/langgraph/agentspecloader.py b/pyagentspec/src/pyagentspec/adapters/langgraph/agentspecloader.py index b0d768b3..00d9f239 100644 --- a/pyagentspec/src/pyagentspec/adapters/langgraph/agentspecloader.py +++ b/pyagentspec/src/pyagentspec/adapters/langgraph/agentspecloader.py @@ -39,6 +39,12 @@ class AgentSpecLoader(AdapterAgnosticAgentSpecLoader): enables features that require a checkpointer (e.g., client tools). config: Optional ``RunnableConfig`` to pass to created runnables/graphs. + middleware: + Optional list of LangChain agent middleware instances forwarded verbatim to + ``langchain_agents.create_agent(middleware=...)`` when compiling an Agent Spec + ``Agent`` into a ReAct graph. Order is preserved — index ``0`` is the outermost + middleware. When ``None`` or an empty list, the ``middleware`` keyword is + omitted entirely from the ``create_agent`` call. allowed_components: Optional iterable of Agent Spec component type names or Component classes allowed to be loaded. If omitted, all component types are allowed unless blocked. @@ -57,7 +63,7 @@ def __init__( plugins: Optional[List[ComponentDeserializationPlugin]] = None, checkpointer: Optional[Checkpointer] = None, config: Optional[RunnableConfig] = None, - *, + middleware: Optional[List[Any]] = None, allowed_components: Optional[ComponentPolicyInput] = None, blocked_components: Optional[ComponentPolicyInput] = None, ) -> None: @@ -69,6 +75,7 @@ def __init__( ) self.checkpointer = checkpointer self.config = config + self._middleware: List[Any] = list(middleware or []) @property def agentspec_to_runtime_converter(self) -> AgentSpecToLangGraphConverter: @@ -287,5 +294,6 @@ def load_component(self, agentspec_component: AgentSpecComponent) -> LangGraphRu tool_registry=self.tool_registry, checkpointer=self.checkpointer, config=self.config, + middleware=self._middleware, ), ) diff --git a/pyagentspec/tests/adapters/langgraph/test_middleware_parameter.py b/pyagentspec/tests/adapters/langgraph/test_middleware_parameter.py new file mode 100644 index 00000000..8f7381d9 --- /dev/null +++ b/pyagentspec/tests/adapters/langgraph/test_middleware_parameter.py @@ -0,0 +1,288 @@ +# Copyright © 2025, 2026 Oracle and/or its affiliates. +# +# This software is under the Apache License 2.0 +# (LICENSE-APACHE or http://www.apache.org/licenses/LICENSE-2.0) or Universal Permissive License +# (UPL) 1.0 (LICENSE-UPL or https://oss.oracle.com/licenses/upl), at your option. + +from typing import Any, Callable, Dict, List, Optional +from unittest.mock import patch + +import pytest + +from pyagentspec.agent import Agent +from pyagentspec.flows.edges import ControlFlowEdge +from pyagentspec.flows.flow import Flow +from pyagentspec.flows.nodes import AgentNode, EndNode, FlowNode, StartNode +from pyagentspec.llms import OpenAiCompatibleConfig +from pyagentspec.property import Property +from pyagentspec.tools import ClientTool + + +class _StopCreateAgent(Exception): + """Raised inside the ``create_agent`` spy to short-circuit graph compilation. + + The kwarg-capture tests only care about what ``create_agent`` is called + with; letting the real call proceed would require valid LangChain + middleware instances, which these tests deliberately do not construct. + """ + + +def _spy_create_agent(captured: Dict[str, Any]) -> Callable[..., Any]: + def spy(**kwargs: Any) -> Any: + captured.update(kwargs) + raise _StopCreateAgent() + + return spy + + +@pytest.fixture +def agent() -> Agent: + return Agent( + name="agent", + system_prompt="You are a helpful agent.", + llm_config=OpenAiCompatibleConfig(name="llm", model_id="fake", url="null"), + tools=[ + ClientTool( + name="ask_user", + description="Ask the user something", + inputs=[Property(title="question", json_schema={"type": "string"})], + outputs=[Property(title="answer", json_schema={})], + ) + ], + ) + + +@pytest.fixture +def agent_flow(agent: Agent) -> Flow: + start_node = StartNode(name="start") + agent_node = AgentNode(name="agent_node", agent=agent) + end_node = EndNode(name="end") + return Flow( + name="flow", + start_node=start_node, + nodes=[start_node, agent_node, end_node], + control_flow_connections=[ + ControlFlowEdge(name="start_to_agent", from_node=start_node, to_node=agent_node), + ControlFlowEdge(name="agent_to_end", from_node=agent_node, to_node=end_node), + ], + data_flow_connections=[], + ) + + +@pytest.fixture +def nested_agent_flow(agent_flow: Flow) -> Flow: + """An outer flow whose single ``FlowNode`` wraps ``agent_flow`` as a subflow. + + Exercises the recursive ``self.convert(subflow, ...)`` path in + ``_flow_node_convert_to_langgraph`` to ensure the ``middleware`` threaded into + ``convert`` still reaches an ``AgentNode`` nested inside a subflow. + """ + start_node = StartNode(name="outer_start") + flow_node = FlowNode(name="flow_node", subflow=agent_flow) + end_node = EndNode(name="outer_end") + return Flow( + name="outer_flow", + start_node=start_node, + nodes=[start_node, flow_node, end_node], + control_flow_connections=[ + ControlFlowEdge(name="start_to_flow", from_node=start_node, to_node=flow_node), + ControlFlowEdge(name="flow_to_end", from_node=flow_node, to_node=end_node), + ], + data_flow_connections=[], + ) + + +@pytest.fixture +def capture_create_agent_kwargs( + agent: Agent, +) -> Callable[[Optional[List[Any]]], Dict[str, Any]]: + """Return a callable that drives a conversion and returns the kwargs ``create_agent`` saw.""" + + def _capture(middleware: Optional[List[Any]]) -> Dict[str, Any]: + from langgraph.checkpoint.memory import MemorySaver + + from pyagentspec.adapters.langgraph._langgraphconverter import ( + AgentSpecToLangGraphConverter, + ) + from pyagentspec.adapters.langgraph._types import langchain_agents + + captured: Dict[str, Any] = {} + with patch.object( + langchain_agents, "create_agent", side_effect=_spy_create_agent(captured) + ): + loader = AgentSpecToLangGraphConverter() + with pytest.raises(_StopCreateAgent): + loader.convert( + agent, + tool_registry={}, + converted_components={agent.llm_config.id: object()}, + checkpointer=MemorySaver(), + middleware=middleware, + ) + return captured + + return _capture + + +@pytest.mark.parametrize( + "middleware", + [None, []], + ids=["none", "empty_list"], +) +def test_omits_middleware_kwarg_when_not_provided( + capture_create_agent_kwargs: Callable[[Optional[List[Any]]], Dict[str, Any]], + middleware: Optional[List[Any]], +) -> None: + """``AgentSpecLoader()`` without ``middleware`` (or with empty list) must not pass ``middleware=``.""" + captured = capture_create_agent_kwargs(middleware) + assert "middleware" not in captured + + +def test_middleware_forwarded_in_order( + capture_create_agent_kwargs: Callable[[Optional[List[Any]]], Dict[str, Any]], +) -> None: + """A non-empty list reaches ``create_agent`` in the original order.""" + a, b = object(), object() + captured = capture_create_agent_kwargs([a, b]) + assert captured.get("middleware") == [a, b] + + +def test_middleware_list_is_copied( + agent: Agent, +) -> None: + """Mutating the caller's list after construction must not leak into conversions.""" + from langgraph.checkpoint.memory import MemorySaver + + from pyagentspec.adapters.langgraph import AgentSpecLoader + from pyagentspec.adapters.langgraph._langgraphconverter import AgentSpecToLangGraphConverter + from pyagentspec.adapters.langgraph._types import langchain_agents + + a = object() + caller_list: List[Any] = [a] + loader = AgentSpecLoader(checkpointer=MemorySaver(), middleware=caller_list) + caller_list.append(object()) + caller_list[0] = object() + + captured: Dict[str, Any] = {} + with patch.object( + AgentSpecToLangGraphConverter, + "_llm_convert_to_langgraph", + return_value=object(), + ), patch.object(langchain_agents, "create_agent", side_effect=_spy_create_agent(captured)): + with pytest.raises(_StopCreateAgent): + loader.load_component(agent) + + assert captured.get("middleware") == [a] + + +def test_middleware_forwarded_through_flow_agent_node(agent_flow: Flow) -> None: + """Middleware must reach ``create_agent`` for agents inside flows.""" + from langgraph.checkpoint.memory import MemorySaver + + from pyagentspec.adapters.langgraph._langgraphconverter import ( + AgentSpecToLangGraphConverter, + ) + from pyagentspec.adapters.langgraph._types import langchain_agents + + agent_llm_id = agent_flow.nodes[1].agent.llm_config.id # type: ignore[union-attr] + captured: Dict[str, Any] = {} + sentinel = object() + checkpointer = MemorySaver() + compiled = AgentSpecToLangGraphConverter().convert( + agent_flow, + tool_registry={}, + converted_components={agent_llm_id: object()}, + checkpointer=checkpointer, + middleware=[sentinel], + ) + + with patch.object(langchain_agents, "create_agent", side_effect=_spy_create_agent(captured)): + # Triggering execution of the AgentNode lazily compiles the inner agent, + # which is where the middleware kwarg is forwarded. + with pytest.raises(_StopCreateAgent): + compiled.invoke( + {"inputs": {}, "messages": [{"role": "user", "content": ""}]}, + config={"configurable": {"thread_id": "flow-mw-regression"}}, + ) + + assert captured.get("middleware") == [sentinel] + + +def test_middleware_forwarded_through_nested_subflow_agent_node(nested_agent_flow: Flow) -> None: + """Regression: middleware must reach agents nested inside a ``FlowNode`` subflow. + + ``_flow_node_convert_to_langgraph`` recurses via ``self.convert(subflow, ...)``, + so the ``middleware`` threaded into ``convert`` must be forwarded through the + subflow conversion to reach an ``AgentNode`` one (or more) subflows deep. + """ + from langgraph.checkpoint.memory import MemorySaver + + from pyagentspec.adapters.langgraph._langgraphconverter import ( + AgentSpecToLangGraphConverter, + ) + from pyagentspec.adapters.langgraph._types import langchain_agents + + nested_agent = nested_agent_flow.nodes[1].subflow.nodes[1].agent # type: ignore[union-attr] + captured: Dict[str, Any] = {} + sentinel = object() + compiled = AgentSpecToLangGraphConverter().convert( + nested_agent_flow, + tool_registry={}, + converted_components={nested_agent.llm_config.id: object()}, + checkpointer=MemorySaver(), + middleware=[sentinel], + ) + + with patch.object(langchain_agents, "create_agent", side_effect=_spy_create_agent(captured)): + with pytest.raises(_StopCreateAgent): + compiled.invoke( + {"inputs": {}, "messages": [{"role": "user", "content": ""}]}, + config={"configurable": {"thread_id": "nested-flow-mw-regression"}}, + ) + + assert captured.get("middleware") == [sentinel] + + +def test_middleware_hook_runs_for_flow_agent_node(agent_flow: Flow) -> None: + """Execution: a middleware instance threaded through a flow's AgentNode is actually invoked. + + A real ``AgentMiddleware`` subclass records each ``before_agent`` call; the + LLM is injected via ``converted_components`` so we never hit the network + and the agent finishes after a single ``AIMessage``. + """ + from langchain.agents.middleware import AgentMiddleware + from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel + from langchain_core.messages import AIMessage + from langgraph.checkpoint.memory import MemorySaver + + from pyagentspec.adapters.langgraph._langgraphconverter import AgentSpecToLangGraphConverter + + class _FakeModel(FakeMessagesListChatModel): + def bind_tools(self, tools: Any, **kwargs: Any) -> Any: + return self + + fake_model = _FakeModel(responses=[AIMessage(content="Done")]) + + calls: List[str] = [] + + class _RecordingMiddleware(AgentMiddleware): + def before_agent(self, state: Any, runtime: Any) -> None: # type: ignore[override] + calls.append("before_agent") + + middleware_instance = _RecordingMiddleware() + agent_llm_id = agent_flow.nodes[1].agent.llm_config.id # type: ignore[union-attr] + checkpointer = MemorySaver() + + compiled = AgentSpecToLangGraphConverter().convert( + agent_flow, + tool_registry={}, + converted_components={agent_llm_id: fake_model}, + checkpointer=checkpointer, + middleware=[middleware_instance], + ) + compiled.invoke( + {"inputs": {}, "messages": [{"role": "user", "content": "hi"}]}, + config={"configurable": {"thread_id": "flow-mw-execution"}}, + ) + + assert calls == ["before_agent"]