diff --git a/libs/langgraph/langgraph/func/__init__.py b/libs/langgraph/langgraph/func/__init__.py index 76dec76d2..73d308365 100644 --- a/libs/langgraph/langgraph/func/__init__.py +++ b/libs/langgraph/langgraph/func/__init__.py @@ -4,6 +4,7 @@ import concurrent.futures import functools import inspect import types +from collections.abc import Iterator from typing import ( Any, Awaitable, @@ -14,6 +15,8 @@ from typing import ( overload, ) +from langchain_core.runnables.base import Runnable +from langchain_core.runnables.graph import Graph, Node from typing_extensions import ParamSpec from langgraph.channels.ephemeral_value import EphemeralValue @@ -22,6 +25,7 @@ from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.constants import CONF, END, START, TAG_HIDDEN from langgraph.pregel import Pregel from langgraph.pregel.call import get_runnable_for_func +from langgraph.pregel.protocol import PregelProtocol from langgraph.pregel.read import PregelNode from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore @@ -96,9 +100,9 @@ def task( def _tick(__allargs__: tuple) -> T: return func(*__allargs__[0], **__allargs__[1]) - return functools.update_wrapper( - functools.partial(call, _tick, retry=retry), func - ) + wrapper = functools.partial(call, _tick, retry=retry) + object.__setattr__(wrapper, "_is_pregel_task", True) + return functools.update_wrapper(wrapper, func) if __func_or_none__ is not None: return decorator(__func_or_none__) @@ -174,6 +178,84 @@ def entrypoint( checkpointer=checkpointer, store=store, config_type=config_schema, + graph=_entrypoint_graph(bound), ) return _imp + + +def _find_children( + candidate: Runnable, parent: Node +) -> Iterator[tuple[Node, Union[Callable, PregelProtocol]]]: + from langchain_core.runnables.utils import get_function_nonlocals + + from langgraph.utils.runnable import ( + RunnableCallable, + RunnableLambda, + RunnableSeq, + RunnableSequence, + ) + + candidates: list[Runnable] = [candidate] + + for c in candidates: + print(c, type(c)) + if callable(c) and getattr(c, "_is_pregel_task", False) is True: + yield (parent, c) + elif isinstance(c, PregelProtocol): + yield (parent, c) + elif isinstance(c, RunnableSequence) or isinstance(c, RunnableSeq): + candidates.extend(c.steps) + elif isinstance(c, RunnableLambda): + candidates.extend(c.deps) + elif isinstance(c, RunnableCallable): + if c.func is not None: + candidates.extend( + nl.__self__ if hasattr(nl, "__self__") else nl + for nl in get_function_nonlocals(c.func) + ) + elif c.afunc is not None: + candidates.extend( + nl.__self__ if hasattr(nl, "__self__") else nl + for nl in get_function_nonlocals(c.afunc) + ) + + +def _entrypoint_graph(entrypoint: Runnable, xray: int = 0) -> Graph: + graph = Graph() + node = Node(f"__{entrypoint.name}", entrypoint.name, entrypoint, None) + graph.nodes[node.id] = node + candidates: list[tuple[Node, Union[Callable, PregelProtocol]]] = [ + *_find_children(entrypoint, node) + ] + seen: set[Callable] = set() + for parent, child in candidates: + if child in seen: + continue + else: + seen.add(child) + if callable(child): + node = Node(f"__{child.__name__}", child.__name__, child, None) + graph.nodes[node.id] = node + graph.add_edge(parent, node, conditional=True) + graph.add_edge(node, parent) + candidates.extend(_find_children(child, node)) + elif isinstance(child, Runnable): + if xray > 0: + graph = child.get_graph(xray=xray - 1 if xray else 0) + graph.trim_first_node() + graph.trim_last_node() + s, e = graph.extend(graph, prefix=child.name) + if s is None: + raise ValueError( + f"Could not extend subgraph '{child.name}' due to missing entrypoint" + ) + else: + graph.add_edge(parent, s, conditional=True) + if e is not None: + graph.add_edge(e, parent) + else: + node = graph.add_node(child, child.name) + graph.add_edge(parent, node, conditional=True) + graph.add_edge(node, parent) + return graph diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 877ddcf4e..fe6e56f5d 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -238,6 +238,8 @@ class Pregel(PregelProtocol): config: Optional[RunnableConfig] = None + graph: Optional[Graph] = None + name: str = "LangGraph" def __init__( @@ -260,6 +262,7 @@ class Pregel(PregelProtocol): retry_policy: Optional[RetryPolicy] = None, config_type: Optional[Type[Any]] = None, config: Optional[RunnableConfig] = None, + graph: Optional[Graph] = None, name: str = "LangGraph", ) -> None: self.nodes = nodes @@ -278,6 +281,7 @@ class Pregel(PregelProtocol): self.retry_policy = retry_policy self.config_type = config_type self.config = config + self.graph = graph self.name = name if auto_validate: self.validate() @@ -285,11 +289,17 @@ class Pregel(PregelProtocol): def get_graph( self, config: RunnableConfig | None = None, *, xray: int | bool = False ) -> Graph: + if self.graph is not None: + # TODO xray + return self.graph raise NotImplementedError async def aget_graph( self, config: RunnableConfig | None = None, *, xray: int | bool = False ) -> Graph: + if self.graph is not None: + # TODO xray + return self.graph raise NotImplementedError def copy(self, update: dict[str, Any] | None = None) -> Self: diff --git a/libs/langgraph/tests/__snapshots__/test_pregel.ambr b/libs/langgraph/tests/__snapshots__/test_pregel.ambr index 9e45447bf..67cb0f852 100644 --- a/libs/langgraph/tests/__snapshots__/test_pregel.ambr +++ b/libs/langgraph/tests/__snapshots__/test_pregel.ambr @@ -2456,6 +2456,37 @@ ''' # --- +# name: test_falsy_return_from_task[memory] + ''' + %%{init: {'flowchart': {'curve': 'linear'}}}%% + graph TD; + __graph([graph]):::last + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_imp_stream_order[memory] + ''' + %%{init: {'flowchart': {'curve': 'linear'}}}%% + graph TD; + __graph(graph) + __bar(bar) + __baz(baz) + __foo(foo) + __graph -.-> __bar; + __bar --> __graph; + __graph -.-> __baz; + __baz --> __graph; + __graph -.-> __foo; + __foo --> __graph; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- # name: test_in_one_fan_out_state_graph_waiting_edge[memory] ''' graph TD; @@ -3579,6 +3610,40 @@ ''' # --- +# name: test_interrupt_functional[memory] + ''' + %%{init: {'flowchart': {'curve': 'linear'}}}%% + graph TD; + __graph(graph) + __bar(bar) + __foo(foo) + __graph -.-> __bar; + __bar --> __graph; + __graph -.-> __foo; + __foo --> __graph; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- +# name: test_interrupt_task_functional[memory] + ''' + %%{init: {'flowchart': {'curve': 'linear'}}}%% + graph TD; + __graph(graph) + __bar(bar) + __foo(foo) + __graph -.-> __bar; + __bar --> __graph; + __graph -.-> __foo; + __foo --> __graph; + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- # name: test_message_graph.1 '{"title": "LangGraphOutput", "type": "array", "items": {"anyOf": [{"$ref": "#/definitions/AIMessage"}, {"$ref": "#/definitions/HumanMessage"}, {"$ref": "#/definitions/ChatMessage"}, {"$ref": "#/definitions/SystemMessage"}, {"$ref": "#/definitions/FunctionMessage"}, {"$ref": "#/definitions/ToolMessage"}]}, "definitions": {"ToolCall": {"title": "ToolCall", "type": "object", "properties": {"name": {"title": "Name", "type": "string"}, "args": {"title": "Args", "type": "object"}, "id": {"title": "Id", "type": "string"}, "type": {"title": "Type", "enum": ["tool_call"], "type": "string"}}, "required": ["name", "args", "id"]}, "InvalidToolCall": {"title": "InvalidToolCall", "type": "object", "properties": {"name": {"title": "Name", "type": "string"}, "args": {"title": "Args", "type": "string"}, "id": {"title": "Id", "type": "string"}, "error": {"title": "Error", "type": "string"}, "type": {"title": "Type", "enum": ["invalid_tool_call"], "type": "string"}}, "required": ["name", "args", "id", "error"]}, "UsageMetadata": {"title": "UsageMetadata", "type": "object", "properties": {"input_tokens": {"title": "Input Tokens", "type": "integer"}, "output_tokens": {"title": "Output Tokens", "type": "integer"}, "total_tokens": {"title": "Total Tokens", "type": "integer"}}, "required": ["input_tokens", "output_tokens", "total_tokens"]}, "AIMessage": {"title": "AIMessage", "description": "Message from an AI.\\n\\nAIMessage is returned from a chat model as a response to a prompt.\\n\\nThis message represents the output of the model and consists of both\\nthe raw output as returned by the model together standardized fields\\n(e.g., tool calls, usage metadata) added by the LangChain framework.", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "default": "ai", "enum": ["ai"], "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}, "example": {"title": "Example", "default": false, "type": "boolean"}, "tool_calls": {"title": "Tool Calls", "default": [], "type": "array", "items": {"$ref": "#/definitions/ToolCall"}}, "invalid_tool_calls": {"title": "Invalid Tool Calls", "default": [], "type": "array", "items": {"$ref": "#/definitions/InvalidToolCall"}}, "usage_metadata": {"$ref": "#/definitions/UsageMetadata"}}, "required": ["content"]}, "HumanMessage": {"title": "HumanMessage", "description": "Message from a human.\\n\\nHumanMessages are messages that are passed in from a human to the model.\\n\\nExample:\\n\\n .. code-block:: python\\n\\n from langchain_core.messages import HumanMessage, SystemMessage\\n\\n messages = [\\n SystemMessage(\\n content=\\"You are a helpful assistant! Your name is Bob.\\"\\n ),\\n HumanMessage(\\n content=\\"What is your name?\\"\\n )\\n ]\\n\\n # Instantiate a chat model and invoke it with the messages\\n model = ...\\n print(model.invoke(messages))", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "default": "human", "enum": ["human"], "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}, "example": {"title": "Example", "default": false, "type": "boolean"}}, "required": ["content"]}, "ChatMessage": {"title": "ChatMessage", "description": "Message that can be assigned an arbitrary speaker (i.e. role).", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "default": "chat", "enum": ["chat"], "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}, "role": {"title": "Role", "type": "string"}}, "required": ["content", "role"]}, "SystemMessage": {"title": "SystemMessage", "description": "Message for priming AI behavior.\\n\\nThe system message is usually passed in as the first of a sequence\\nof input messages.\\n\\nExample:\\n\\n .. code-block:: python\\n\\n from langchain_core.messages import HumanMessage, SystemMessage\\n\\n messages = [\\n SystemMessage(\\n content=\\"You are a helpful assistant! Your name is Bob.\\"\\n ),\\n HumanMessage(\\n content=\\"What is your name?\\"\\n )\\n ]\\n\\n # Define a chat model and invoke it with the messages\\n print(model.invoke(messages))", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "default": "system", "enum": ["system"], "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}}, "required": ["content"]}, "FunctionMessage": {"title": "FunctionMessage", "description": "Message for passing the result of executing a tool back to a model.\\n\\nFunctionMessage are an older version of the ToolMessage schema, and\\ndo not contain the tool_call_id field.\\n\\nThe tool_call_id field is used to associate the tool call request with the\\ntool call response. This is useful in situations where a chat model is able\\nto request multiple tool calls in parallel.", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "default": "function", "enum": ["function"], "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}}, "required": ["content", "name"]}, "ToolMessage": {"title": "ToolMessage", "description": "Message for passing the result of executing a tool back to a model.\\n\\nToolMessages contain the result of a tool invocation. Typically, the result\\nis encoded inside the `content` field.\\n\\nExample: A ToolMessage representing a result of 42 from a tool call with id\\n\\n .. code-block:: python\\n\\n from langchain_core.messages import ToolMessage\\n\\n ToolMessage(content=\'42\', tool_call_id=\'call_Jja7J89XsjrOLA5r!MEOW!SL\')\\n\\n\\nExample: A ToolMessage where only part of the tool output is sent to the model\\n and the full output is passed in to artifact.\\n\\n .. versionadded:: 0.2.17\\n\\n .. code-block:: python\\n\\n from langchain_core.messages import ToolMessage\\n\\n tool_output = {\\n \\"stdout\\": \\"From the graph we can see that the correlation between x and y is ...\\",\\n \\"stderr\\": None,\\n \\"artifacts\\": {\\"type\\": \\"image\\", \\"base64_data\\": \\"/9j/4gIcSU...\\"},\\n }\\n\\n ToolMessage(\\n content=tool_output[\\"stdout\\"],\\n artifact=tool_output,\\n tool_call_id=\'call_Jja7J89XsjrOLA5r!MEOW!SL\',\\n )\\n\\nThe tool_call_id field is used to associate the tool call request with the\\ntool call response. This is useful in situations where a chat model is able\\nto request multiple tool calls in parallel.", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "default": "tool", "enum": ["tool"], "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}, "tool_call_id": {"title": "Tool Call Id", "type": "string"}, "artifact": {"title": "Artifact"}, "status": {"title": "Status", "default": "success", "enum": ["success", "error"], "type": "string"}}, "required": ["content", "tool_call_id"]}}}' # --- @@ -4062,6 +4127,17 @@ ''' # --- +# name: test_multiple_interrupts_imperative[memory] + ''' + %%{init: {'flowchart': {'curve': 'linear'}}}%% + graph TD; + __graph([graph]):::last + classDef default fill:#f2f0ff,line-height:1.2 + classDef first fill-opacity:0 + classDef last fill:#bfb6fc + + ''' +# --- # name: test_multiple_sinks_subgraphs ''' %%{init: {'flowchart': {'curve': 'linear'}}}%% diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 430a98dd9..e941de088 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -1523,7 +1523,7 @@ def test_imp_task(request: pytest.FixtureRequest, checkpointer_name: str) -> Non @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_imp_stream_order( - request: pytest.FixtureRequest, checkpointer_name: str + request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion ) -> None: checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") @@ -1546,6 +1546,9 @@ def test_imp_stream_order( fut_baz = baz(fut_bar.result()) return fut_baz.result() + if checkpointer_name == "memory": + assert graph.get_graph().draw_mermaid() == snapshot + thread1 = {"configurable": {"thread_id": "1"}} assert [c for c in graph.stream({"a": "0"}, thread1)] == [ { @@ -4951,7 +4954,7 @@ def test_interrupt_loop(request: pytest.FixtureRequest, checkpointer_name: str): @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_interrupt_functional( - request: pytest.FixtureRequest, checkpointer_name: str + request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion ) -> None: checkpointer: BaseCheckpointSaver = request.getfixturevalue( f"checkpointer_{checkpointer_name}" @@ -4973,6 +4976,8 @@ def test_interrupt_functional( fut_bar = bar(bar_input) return fut_bar.result() + assert graph.get_graph().draw_mermaid() == snapshot + config = {"configurable": {"thread_id": "1"}} # First run, interrupted at bar graph.invoke({"a": ""}, config) @@ -4983,7 +4988,7 @@ def test_interrupt_functional( @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_interrupt_task_functional( - request: pytest.FixtureRequest, checkpointer_name: str + request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion ) -> None: checkpointer: BaseCheckpointSaver = request.getfixturevalue( f"checkpointer_{checkpointer_name}" @@ -5004,6 +5009,9 @@ def test_interrupt_task_functional( fut_bar = bar(fut_foo.result()) return fut_bar.result() + if checkpointer_name == "memory": + assert graph.get_graph().draw_mermaid() == snapshot + config = {"configurable": {"thread_id": "1"}} # First run, interrupted at bar graph.invoke({"a": ""}, config) @@ -5432,7 +5440,9 @@ def test_multiple_updates() -> None: @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) -def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_name: str): +def test_falsy_return_from_task( + request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion +): """Test with a falsy return from a task.""" checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") @@ -5446,6 +5456,9 @@ def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_nam falsy_task().result() interrupt("test") + if checkpointer_name == "memory": + assert graph.get_graph().draw_mermaid() == snapshot + configurable = {"configurable": {"thread_id": str(uuid.uuid4())}} graph.invoke({"a": 5}, configurable) graph.invoke(Command(resume="123"), configurable) @@ -5453,7 +5466,7 @@ def test_falsy_return_from_task(request: pytest.FixtureRequest, checkpointer_nam @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_multiple_interrupts_imperative( - request: pytest.FixtureRequest, checkpointer_name: str + request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion ): """Test multiple interrupts with an imperative API.""" checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") @@ -5478,6 +5491,9 @@ def test_multiple_interrupts_imperative( return {"values": values} + if checkpointer_name == "memory": + assert graph.get_graph().draw_mermaid() == snapshot + configurable = {"configurable": {"thread_id": str(uuid.uuid4())}} graph.invoke({}, configurable) graph.invoke(Command(resume="a"), configurable)