mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-08 18:57:52 +02:00
Merge pull request #1793 from langchain-ai/nc/21sep/detect-multi-subgraphs-in-node
Detect multiple subgraphs in single node
This commit is contained in:
@@ -10,6 +10,14 @@ from langgraph.types import Interrupt, Send # noqa: F401
|
||||
EMPTY_MAP: Mapping[str, Any] = MappingProxyType({})
|
||||
EMPTY_SEQ: tuple[str, ...] = tuple()
|
||||
|
||||
# --- Public constants ---
|
||||
TAG_HIDDEN = "langsmith:hidden"
|
||||
# tag to hide a node/edge from certain tracing/streaming environments
|
||||
START = "__start__"
|
||||
# the first (maybe virtual) node in graph-style Pregel
|
||||
END = "__end__"
|
||||
# the last (maybe virtual) node in graph-style Pregel
|
||||
|
||||
# --- Reserved write keys ---
|
||||
INPUT = "__input__"
|
||||
# for values passed as input to the graph
|
||||
@@ -23,10 +31,6 @@ SCHEDULED = "__scheduled__"
|
||||
# marker to signal node was scheduled (in distributed mode)
|
||||
TASKS = "__pregel_tasks"
|
||||
# for Send objects returned by nodes/edges, corresponds to PUSH below
|
||||
START = "__start__"
|
||||
# marker for the first (maybe virtual) node in graph-style Pregel
|
||||
END = "__end__"
|
||||
# marker for the last (maybe virtual) node in graph-style Pregel
|
||||
|
||||
# --- Reserved config.configurable keys ---
|
||||
CONFIG_KEY_SEND = "__pregel_send"
|
||||
@@ -43,8 +47,6 @@ CONFIG_KEY_STORE = "__pregel_store"
|
||||
# holds a `BaseStore` made available to managed values
|
||||
CONFIG_KEY_RESUMING = "__pregel_resuming"
|
||||
# holds a boolean indicating if subgraphs should resume from a previous checkpoint
|
||||
CONFIG_KEY_GRAPH_COUNT = "__pregel_graph_count"
|
||||
# holds the number of subgraphs executed in a given task, used to raise errors
|
||||
CONFIG_KEY_TASK_ID = "__pregel_task_id"
|
||||
# holds the task ID for the current task
|
||||
CONFIG_KEY_DEDUPE_TASKS = "__pregel_dedupe_tasks"
|
||||
@@ -68,14 +70,13 @@ PULL = "__pregel_pull"
|
||||
# denotes pull-style tasks, ie. those triggered by edges
|
||||
RUNTIME_PLACEHOLDER = "__pregel_runtime_placeholder__"
|
||||
# placeholder for managed values replaced at runtime
|
||||
TAG_HIDDEN = "langsmith:hidden"
|
||||
# tag to hide a node/edge from certain tracing/streaming environments
|
||||
NS_SEP = "|"
|
||||
# for checkpoint_ns, separates each level (ie. graph|subgraph|subsubgraph)
|
||||
NS_END = ":"
|
||||
# for checkpoint_ns, for each level, separates the namespace from the task_id
|
||||
|
||||
RESERVED = {
|
||||
TAG_HIDDEN,
|
||||
# reserved write keys
|
||||
INPUT,
|
||||
INTERRUPT,
|
||||
@@ -103,7 +104,6 @@ RESERVED = {
|
||||
PUSH,
|
||||
PULL,
|
||||
RUNTIME_PLACEHOLDER,
|
||||
TAG_HIDDEN,
|
||||
NS_SEP,
|
||||
NS_END,
|
||||
}
|
||||
|
||||
@@ -69,3 +69,13 @@ class CheckpointNotLatest(Exception):
|
||||
"""Raised when the checkpoint is not the latest version (for distributed mode)."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class MultipleSubgraphsError(Exception):
|
||||
"""Raised when multiple subgraphs are called inside the same node."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
_SEEN_CHECKPOINT_NS: set[str] = set()
|
||||
"""Used for subgraph detection."""
|
||||
|
||||
@@ -26,7 +26,6 @@ from langchain_core.runnables.graph import Node as DrawableNode
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.channels.ephemeral_value import EphemeralValue
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.constants import (
|
||||
END,
|
||||
NS_END,
|
||||
@@ -39,7 +38,7 @@ from langgraph.errors import InvalidUpdateError
|
||||
from langgraph.pregel import Channel, Pregel
|
||||
from langgraph.pregel.read import PregelNode
|
||||
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.types import All
|
||||
from langgraph.types import All, Checkpointer
|
||||
from langgraph.utils.runnable import RunnableCallable, coerce_to_runnable
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -406,7 +405,7 @@ class Graph:
|
||||
|
||||
def compile(
|
||||
self,
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None,
|
||||
checkpointer: Checkpointer = None,
|
||||
interrupt_before: Optional[Union[All, list[str]]] = None,
|
||||
interrupt_after: Optional[Union[All, list[str]]] = None,
|
||||
debug: bool = False,
|
||||
|
||||
@@ -32,7 +32,6 @@ from langgraph.channels.dynamic_barrier_value import DynamicBarrierValue, WaitFo
|
||||
from langgraph.channels.ephemeral_value import EphemeralValue
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.channels.named_barrier_value import NamedBarrierValue
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.constants import NS_END, NS_SEP, TAG_HIDDEN
|
||||
from langgraph.errors import InvalidUpdateError
|
||||
from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph, Send
|
||||
@@ -47,7 +46,7 @@ from langgraph.managed.base import (
|
||||
from langgraph.pregel.read import ChannelRead, PregelNode
|
||||
from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.types import All, RetryPolicy
|
||||
from langgraph.types import All, Checkpointer, RetryPolicy
|
||||
from langgraph.utils.fields import get_field_default
|
||||
from langgraph.utils.pydantic import create_model
|
||||
from langgraph.utils.runnable import coerce_to_runnable
|
||||
@@ -400,7 +399,7 @@ class StateGraph(Graph):
|
||||
|
||||
def compile(
|
||||
self,
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None,
|
||||
checkpointer: Checkpointer = None,
|
||||
*,
|
||||
store: Optional[BaseStore] = None,
|
||||
interrupt_before: Optional[Union[All, list[str]]] = None,
|
||||
@@ -413,7 +412,7 @@ class StateGraph(Graph):
|
||||
streamed, batched, and run asynchronously.
|
||||
|
||||
Args:
|
||||
checkpointer (Optional[BaseCheckpointSaver]): An optional checkpoint saver object.
|
||||
checkpointer (Checkpointer): An optional checkpoint saver object.
|
||||
This serves as a fully versioned "memory" for the graph, allowing
|
||||
the graph to be paused and resumed, and replayed from any point.
|
||||
interrupt_before (Optional[Sequence[str]]): An optional list of node names to interrupt before.
|
||||
|
||||
@@ -16,13 +16,13 @@ from langchain_core.runnables import Runnable, RunnableConfig, RunnableLambda
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from langgraph._api.deprecation import deprecated_parameter
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.graph.graph import CompiledGraph
|
||||
from langgraph.graph.message import add_messages
|
||||
from langgraph.managed import IsLastStep
|
||||
from langgraph.prebuilt.tool_executor import ToolExecutor
|
||||
from langgraph.prebuilt.tool_node import ToolNode
|
||||
from langgraph.types import Checkpointer
|
||||
|
||||
|
||||
# We create the AgentState that we will pass around
|
||||
@@ -132,7 +132,7 @@ def create_react_agent(
|
||||
state_schema: Optional[StateSchemaType] = None,
|
||||
messages_modifier: Optional[MessagesModifier] = None,
|
||||
state_modifier: Optional[StateModifier] = None,
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None,
|
||||
checkpointer: Checkpointer = None,
|
||||
interrupt_before: Optional[list[str]] = None,
|
||||
interrupt_after: Optional[list[str]] = None,
|
||||
debug: bool = False,
|
||||
|
||||
@@ -87,7 +87,7 @@ from langgraph.pregel.utils import get_new_channel_versions
|
||||
from langgraph.pregel.validate import validate_graph, validate_keys
|
||||
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.types import All, StateSnapshot, StreamMode
|
||||
from langgraph.types import All, Checkpointer, StateSnapshot, StreamMode
|
||||
from langgraph.utils.config import (
|
||||
ensure_config,
|
||||
merge_configs,
|
||||
@@ -197,7 +197,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
debug: bool
|
||||
"""Whether to print debug information during execution. Defaults to False."""
|
||||
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None
|
||||
checkpointer: Checkpointer = None
|
||||
"""Checkpointer used to save and load graph state. Defaults to None."""
|
||||
|
||||
store: Optional[BaseStore] = None
|
||||
@@ -281,7 +281,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
[spec for node in self.nodes.values() for spec in node.config_specs]
|
||||
+ (
|
||||
self.checkpointer.config_specs
|
||||
if self.checkpointer is not None
|
||||
if isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else []
|
||||
)
|
||||
+ (
|
||||
@@ -1059,6 +1059,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
Union[All, Sequence[str]],
|
||||
Optional[BaseCheckpointSaver],
|
||||
]:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
debug = debug if debug is not None else self.debug
|
||||
if output_keys is None:
|
||||
output_keys = self.stream_channels_asis
|
||||
@@ -1072,12 +1074,16 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
if CONFIG_KEY_TASK_ID in config.get("configurable", {}):
|
||||
# if being called as a node in another graph, always use values mode
|
||||
stream_mode = ["values"]
|
||||
if CONFIG_KEY_CHECKPOINTER in config.get("configurable", {}):
|
||||
checkpointer: Optional[BaseCheckpointSaver] = config["configurable"][
|
||||
CONFIG_KEY_CHECKPOINTER
|
||||
]
|
||||
if self.checkpointer is False:
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None
|
||||
elif CONFIG_KEY_CHECKPOINTER in config.get("configurable", {}):
|
||||
checkpointer = config["configurable"][CONFIG_KEY_CHECKPOINTER]
|
||||
else:
|
||||
checkpointer = self.checkpointer
|
||||
if checkpointer and not config.get("configurable"):
|
||||
raise ValueError(
|
||||
f"Checkpointer requires one or more of the following 'configurable' keys: {[s.id for s in checkpointer.config_specs]}"
|
||||
)
|
||||
return (
|
||||
debug,
|
||||
set(stream_mode),
|
||||
@@ -1193,12 +1199,6 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
run_id=config.get("run_id"),
|
||||
)
|
||||
try:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
if self.checkpointer and not config.get("configurable"):
|
||||
raise ValueError(
|
||||
f"Checkpointer requires one or more of the following 'configurable' keys: {[s.id for s in self.checkpointer.config_specs]}"
|
||||
)
|
||||
# assign defaults
|
||||
(
|
||||
debug,
|
||||
@@ -1414,12 +1414,6 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
None,
|
||||
)
|
||||
try:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
if self.checkpointer and not config.get("configurable"):
|
||||
raise ValueError(
|
||||
f"Checkpointer requires one or more of the following 'configurable' keys: {[s.id for s in self.checkpointer.config_specs]}"
|
||||
)
|
||||
# assign defaults
|
||||
(
|
||||
debug,
|
||||
|
||||
@@ -54,10 +54,12 @@ from langgraph.constants import (
|
||||
TASKS,
|
||||
)
|
||||
from langgraph.errors import (
|
||||
_SEEN_CHECKPOINT_NS,
|
||||
CheckpointNotLatest,
|
||||
EmptyInputError,
|
||||
GraphDelegate,
|
||||
GraphInterrupt,
|
||||
MultipleSubgraphsError,
|
||||
)
|
||||
from langgraph.managed.base import (
|
||||
ManagedValueMapping,
|
||||
@@ -195,6 +197,7 @@ class PregelLoop:
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
output_keys: Union[str, Sequence[str]],
|
||||
stream_keys: Union[str, Sequence[str]],
|
||||
check_subgraphs: bool = True,
|
||||
debug: bool = False,
|
||||
) -> None:
|
||||
self.stream = stream
|
||||
@@ -220,6 +223,11 @@ class PregelLoop:
|
||||
self.config = patch_configurable(
|
||||
self.config, {"checkpoint_ns": "", "checkpoint_id": None}
|
||||
)
|
||||
if check_subgraphs and self.is_nested and self.checkpointer is not None:
|
||||
if self.config["configurable"]["checkpoint_ns"] in _SEEN_CHECKPOINT_NS:
|
||||
raise MultipleSubgraphsError
|
||||
else:
|
||||
_SEEN_CHECKPOINT_NS.add(self.config["configurable"]["checkpoint_ns"])
|
||||
if (
|
||||
CONFIG_KEY_CHECKPOINT_MAP in self.config["configurable"]
|
||||
and self.config["configurable"].get("checkpoint_ns")
|
||||
@@ -634,6 +642,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
check_subgraphs: bool = True,
|
||||
debug: bool = False,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
@@ -646,6 +655,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
specs=specs,
|
||||
output_keys=output_keys,
|
||||
stream_keys=stream_keys,
|
||||
check_subgraphs=check_subgraphs,
|
||||
debug=debug,
|
||||
)
|
||||
self.stack = ExitStack()
|
||||
@@ -755,6 +765,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
check_subgraphs: bool = True,
|
||||
debug: bool = False,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
@@ -767,6 +778,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
specs=specs,
|
||||
output_keys=output_keys,
|
||||
stream_keys=stream_keys,
|
||||
check_subgraphs=check_subgraphs,
|
||||
debug=debug,
|
||||
)
|
||||
self.store = AsyncBatchedStore(self.store) if self.store else None
|
||||
|
||||
@@ -4,8 +4,8 @@ import random
|
||||
import time
|
||||
from typing import Optional, Sequence
|
||||
|
||||
from langgraph.constants import CONFIG_KEY_GRAPH_COUNT, CONFIG_KEY_RESUMING
|
||||
from langgraph.errors import GraphInterrupt
|
||||
from langgraph.constants import CONFIG_KEY_RESUMING
|
||||
from langgraph.errors import _SEEN_CHECKPOINT_NS, GraphInterrupt
|
||||
from langgraph.types import PregelExecutableTask, RetryPolicy
|
||||
from langgraph.utils.config import patch_configurable
|
||||
|
||||
@@ -70,9 +70,14 @@ def run_with_retry(
|
||||
exc_info=exc,
|
||||
)
|
||||
# signal subgraphs to resume (if available)
|
||||
config = patch_configurable(
|
||||
config, {CONFIG_KEY_RESUMING: True, CONFIG_KEY_GRAPH_COUNT: 0}
|
||||
)
|
||||
config = patch_configurable(config, {CONFIG_KEY_RESUMING: True})
|
||||
# clear checkpoint_ns seen (for subgraph detection)
|
||||
if checkpoint_ns := config["configurable"].get("checkpoint_ns"):
|
||||
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
|
||||
finally:
|
||||
# clear checkpoint_ns seen (for subgraph detection)
|
||||
if checkpoint_ns := config["configurable"].get("checkpoint_ns"):
|
||||
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
|
||||
|
||||
|
||||
async def arun_with_retry(
|
||||
@@ -138,6 +143,11 @@ async def arun_with_retry(
|
||||
exc_info=exc,
|
||||
)
|
||||
# signal subgraphs to resume (if available)
|
||||
config = patch_configurable(
|
||||
config, {CONFIG_KEY_RESUMING: True, CONFIG_KEY_GRAPH_COUNT: 0}
|
||||
)
|
||||
config = patch_configurable(config, {CONFIG_KEY_RESUMING: True})
|
||||
# clear checkpoint_ns seen (for subgraph detection)
|
||||
if checkpoint_ns := config["configurable"].get("checkpoint_ns"):
|
||||
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
|
||||
finally:
|
||||
# clear checkpoint_ns seen (for subgraph detection)
|
||||
if checkpoint_ns := config["configurable"].get("checkpoint_ns"):
|
||||
_SEEN_CHECKPOINT_NS.discard(checkpoint_ns)
|
||||
|
||||
@@ -1,12 +1,26 @@
|
||||
from collections import deque
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Literal, NamedTuple, Optional, Sequence, Type, Union
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
Sequence,
|
||||
Type,
|
||||
Union,
|
||||
)
|
||||
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
|
||||
from langgraph.checkpoint.base import CheckpointMetadata
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
|
||||
|
||||
All = Literal["*"]
|
||||
"""Special value to indicate that graph should interrupt on all nodes."""
|
||||
|
||||
Checkpointer = Union[None, Literal[False], BaseCheckpointSaver]
|
||||
"""Type of the checkpointer to use for a subgraph. False disables checkpointing,
|
||||
even if the parent graph has a checkpointer. None inherits checkpointer."""
|
||||
|
||||
StreamMode = Literal["values", "updates", "debug", "messages", "custom"]
|
||||
"""How the stream method should emit outputs.
|
||||
|
||||
@@ -53,7 +53,7 @@ from langgraph.checkpoint.base import (
|
||||
)
|
||||
from langgraph.checkpoint.memory import MemorySaver
|
||||
from langgraph.constants import ERROR, PULL, PUSH
|
||||
from langgraph.errors import InvalidUpdateError, NodeInterrupt
|
||||
from langgraph.errors import InvalidUpdateError, MultipleSubgraphsError, NodeInterrupt
|
||||
from langgraph.graph import END, Graph
|
||||
from langgraph.graph.graph import START
|
||||
from langgraph.graph.message import MessageGraph, add_messages
|
||||
@@ -1861,7 +1861,12 @@ def test_invoke_two_processes_two_in_join_two_out(mocker: MockerFixture) -> None
|
||||
assert [*executor.map(app.invoke, [2] * 100)] == [[13, 13]] * 100
|
||||
|
||||
|
||||
def test_invoke_join_then_call_other_pregel(mocker: MockerFixture) -> None:
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_invoke_join_then_call_other_pregel(
|
||||
mocker: MockerFixture, request: pytest.FixtureRequest, checkpointer_name: str
|
||||
) -> None:
|
||||
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
|
||||
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
add_10_each = mocker.Mock(side_effect=lambda x: [y + 10 for y in x])
|
||||
|
||||
@@ -1912,6 +1917,17 @@ def test_invoke_join_then_call_other_pregel(mocker: MockerFixture) -> None:
|
||||
with ThreadPoolExecutor() as executor:
|
||||
assert [*executor.map(app.invoke, [[2, 3]] * 10)] == [27] * 10
|
||||
|
||||
# add checkpointer
|
||||
app.checkpointer = checkpointer
|
||||
# subgraph is called twice in the same node, through .map(), so raises
|
||||
with pytest.raises(MultipleSubgraphsError):
|
||||
app.invoke([2, 3], {"configurable": {"thread_id": "1"}})
|
||||
|
||||
# set inner graph checkpointer NeverCheckpoint
|
||||
inner_app.checkpointer = False
|
||||
# subgraph still called twice, but checkpointing for inner graph is disabled
|
||||
assert app.invoke([2, 3], {"configurable": {"thread_id": "1"}}) == 27
|
||||
|
||||
|
||||
def test_invoke_two_processes_one_in_two_out(mocker: MockerFixture) -> None:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
|
||||
@@ -52,7 +52,7 @@ from langgraph.checkpoint.base import (
|
||||
)
|
||||
from langgraph.checkpoint.memory import MemorySaver
|
||||
from langgraph.constants import ERROR, PULL, PUSH
|
||||
from langgraph.errors import InvalidUpdateError, NodeInterrupt
|
||||
from langgraph.errors import InvalidUpdateError, MultipleSubgraphsError, NodeInterrupt
|
||||
from langgraph.graph import END, Graph, StateGraph
|
||||
from langgraph.graph.graph import START
|
||||
from langgraph.graph.message import MessageGraph, add_messages
|
||||
@@ -2080,7 +2080,10 @@ async def test_invoke_two_processes_two_in_join_two_out(mocker: MockerFixture) -
|
||||
]
|
||||
|
||||
|
||||
async def test_invoke_join_then_call_other_pregel(mocker: MockerFixture) -> None:
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_invoke_join_then_call_other_pregel(
|
||||
mocker: MockerFixture, checkpointer_name: str
|
||||
) -> None:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
add_10_each = mocker.Mock(side_effect=lambda x: [y + 10 for y in x])
|
||||
|
||||
@@ -2133,6 +2136,18 @@ async def test_invoke_join_then_call_other_pregel(mocker: MockerFixture) -> None
|
||||
27 for _ in range(10)
|
||||
]
|
||||
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
# add checkpointer
|
||||
app.checkpointer = checkpointer
|
||||
# subgraph is called twice in the same node, through .map(), so raises
|
||||
with pytest.raises(MultipleSubgraphsError):
|
||||
await app.ainvoke([2, 3], {"configurable": {"thread_id": "1"}})
|
||||
|
||||
# set inner graph checkpointer NeverCheckpoint
|
||||
inner_app.checkpointer = False
|
||||
# subgraph still called twice, but checkpointing for inner graph is disabled
|
||||
assert await app.ainvoke([2, 3], {"configurable": {"thread_id": "1"}}) == 27
|
||||
|
||||
|
||||
async def test_invoke_two_processes_one_in_two_out(mocker: MockerFixture) -> None:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
|
||||
@@ -55,6 +55,7 @@ def wait_for(
|
||||
raise ValueError(f"Callable did not return within {total_time}")
|
||||
|
||||
|
||||
@pytest.mark.skip("This test times out in CI")
|
||||
async def test_nested_tracing():
|
||||
lt_py_311 = sys.version_info < (3, 11)
|
||||
mock_client = _get_mock_client()
|
||||
|
||||
@@ -158,6 +158,7 @@ class AsyncKafkaOrchestrator(AbstractAsyncContextManager):
|
||||
specs=graph.channels,
|
||||
output_keys=graph.output_channels,
|
||||
stream_keys=graph.stream_channels,
|
||||
check_subgraphs=False,
|
||||
) as loop:
|
||||
if loop.tick(
|
||||
input_keys=graph.input_channels,
|
||||
@@ -347,6 +348,7 @@ class KafkaOrchestrator(AbstractContextManager):
|
||||
specs=graph.channels,
|
||||
output_keys=graph.output_channels,
|
||||
stream_keys=graph.stream_channels,
|
||||
check_subgraphs=False,
|
||||
) as loop:
|
||||
if loop.tick(
|
||||
input_keys=graph.input_channels,
|
||||
|
||||
Reference in New Issue
Block a user