mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 05:07:51 +02:00
Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
33289349a9 | ||
|
|
727a9a1f6c | ||
|
|
dec15d13dc | ||
|
|
f4a4c4b7b1 | ||
|
|
f73983e2fa | ||
|
|
9023ad9931 | ||
|
|
353b0d8fb4 | ||
|
|
2b9c63645e | ||
|
|
54d0c94f54 | ||
|
|
7d7bc0e42b | ||
|
|
bbc35a30f8 |
@@ -6,6 +6,7 @@ import typing
|
|||||||
import warnings
|
import warnings
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from collections.abc import Awaitable, Callable, Hashable, Sequence
|
from collections.abc import Awaitable, Callable, Hashable, Sequence
|
||||||
|
from dataclasses import is_dataclass
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from inspect import isclass, isfunction, ismethod, signature
|
from inspect import isclass, isfunction, ismethod, signature
|
||||||
from types import FunctionType
|
from types import FunctionType
|
||||||
@@ -14,6 +15,7 @@ from typing import (
|
|||||||
Any,
|
Any,
|
||||||
Generic,
|
Generic,
|
||||||
Literal,
|
Literal,
|
||||||
|
TypeVar,
|
||||||
Union,
|
Union,
|
||||||
cast,
|
cast,
|
||||||
get_args,
|
get_args,
|
||||||
@@ -1164,6 +1166,20 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
|
|||||||
for key, node in self.nodes.items():
|
for key, node in self.nodes.items():
|
||||||
compiled.attach_node(key, node)
|
compiled.attach_node(key, node)
|
||||||
|
|
||||||
|
# Record output/state mappers for v2 stream coercion (pydantic/dataclass only)
|
||||||
|
compiled._output_mapper = _pick_mapper(
|
||||||
|
list(output_channels)
|
||||||
|
if isinstance(output_channels, list)
|
||||||
|
else [output_channels],
|
||||||
|
self.output_schema,
|
||||||
|
)
|
||||||
|
compiled._state_mapper = _pick_mapper(
|
||||||
|
list(stream_channels)
|
||||||
|
if isinstance(stream_channels, list)
|
||||||
|
else [stream_channels],
|
||||||
|
self.state_schema,
|
||||||
|
)
|
||||||
|
|
||||||
for start, end in self.edges:
|
for start, end in self.edges:
|
||||||
compiled.attach_edge(start, end)
|
compiled.attach_edge(start, end)
|
||||||
|
|
||||||
@@ -1183,6 +1199,8 @@ class CompiledStateGraph(
|
|||||||
):
|
):
|
||||||
builder: StateGraph[StateT, ContextT, InputT, OutputT]
|
builder: StateGraph[StateT, ContextT, InputT, OutputT]
|
||||||
schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None]
|
schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None]
|
||||||
|
_output_mapper: Callable[[Any], Any] | None
|
||||||
|
_state_mapper: Callable[[Any], Any] | None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -1504,12 +1522,15 @@ def _pick_mapper(
|
|||||||
) -> Callable[[Any], Any] | None:
|
) -> Callable[[Any], Any] | None:
|
||||||
if state_keys == ["__root__"]:
|
if state_keys == ["__root__"]:
|
||||||
return None
|
return None
|
||||||
if isclass(schema) and issubclass(schema, dict):
|
if isclass(schema) and (issubclass(schema, BaseModel) or is_dataclass(schema)):
|
||||||
return None
|
return partial(_coerce_state, schema)
|
||||||
return partial(_coerce_state, schema)
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _coerce_state(schema: type[Any], input: dict[str, Any]) -> dict[str, Any]:
|
_S = TypeVar("_S")
|
||||||
|
|
||||||
|
|
||||||
|
def _coerce_state(schema: type[_S], input: dict[str, Any]) -> _S:
|
||||||
return schema(**input)
|
return schema(**input)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ from uuid import UUID
|
|||||||
|
|
||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
from langgraph.checkpoint.base import CheckpointMetadata, PendingWrite
|
from langgraph.checkpoint.base import CheckpointMetadata, PendingWrite
|
||||||
from typing_extensions import TypedDict
|
|
||||||
|
|
||||||
from langgraph._internal._config import patch_checkpoint_map
|
from langgraph._internal._config import patch_checkpoint_map
|
||||||
from langgraph._internal._constants import (
|
from langgraph._internal._constants import (
|
||||||
@@ -23,42 +22,14 @@ from langgraph._internal._typing import MISSING
|
|||||||
from langgraph.channels.base import BaseChannel
|
from langgraph.channels.base import BaseChannel
|
||||||
from langgraph.constants import TAG_HIDDEN
|
from langgraph.constants import TAG_HIDDEN
|
||||||
from langgraph.pregel._io import read_channels
|
from langgraph.pregel._io import read_channels
|
||||||
from langgraph.types import PregelExecutableTask, PregelTask, StateSnapshot
|
from langgraph.types import (
|
||||||
|
CheckpointPayload,
|
||||||
__all__ = ("TaskPayload", "TaskResultPayload", "CheckpointTask", "CheckpointPayload")
|
PregelExecutableTask,
|
||||||
|
PregelTask,
|
||||||
|
StateSnapshot,
|
||||||
class TaskPayload(TypedDict):
|
TaskPayload,
|
||||||
id: str
|
TaskResultPayload,
|
||||||
name: str
|
)
|
||||||
input: Any
|
|
||||||
triggers: list[str]
|
|
||||||
|
|
||||||
|
|
||||||
class TaskResultPayload(TypedDict):
|
|
||||||
id: str
|
|
||||||
name: str
|
|
||||||
error: str | None
|
|
||||||
interrupts: list[dict]
|
|
||||||
result: dict[str, Any]
|
|
||||||
|
|
||||||
|
|
||||||
class CheckpointTask(TypedDict):
|
|
||||||
id: str
|
|
||||||
name: str
|
|
||||||
error: str | None
|
|
||||||
interrupts: list[dict]
|
|
||||||
state: StateSnapshot | RunnableConfig | None
|
|
||||||
|
|
||||||
|
|
||||||
class CheckpointPayload(TypedDict):
|
|
||||||
config: RunnableConfig | None
|
|
||||||
metadata: CheckpointMetadata
|
|
||||||
values: dict[str, Any]
|
|
||||||
next: list[str]
|
|
||||||
parent_config: RunnableConfig | None
|
|
||||||
tasks: list[CheckpointTask]
|
|
||||||
|
|
||||||
|
|
||||||
TASK_NAMESPACE = UUID("6ba7b831-9dad-11d1-80b4-00c04fd430c8")
|
TASK_NAMESPACE = UUID("6ba7b831-9dad-11d1-80b4-00c04fd430c8")
|
||||||
|
|
||||||
|
|||||||
@@ -22,8 +22,10 @@ from inspect import isclass
|
|||||||
from typing import (
|
from typing import (
|
||||||
Any,
|
Any,
|
||||||
Generic,
|
Generic,
|
||||||
|
Literal,
|
||||||
cast,
|
cast,
|
||||||
get_type_hints,
|
get_type_hints,
|
||||||
|
overload,
|
||||||
)
|
)
|
||||||
from uuid import UUID, uuid5
|
from uuid import UUID, uuid5
|
||||||
|
|
||||||
@@ -121,7 +123,10 @@ from langgraph.pregel._checkpoint import (
|
|||||||
)
|
)
|
||||||
from langgraph.pregel._draw import draw_graph
|
from langgraph.pregel._draw import draw_graph
|
||||||
from langgraph.pregel._io import map_input, read_channels
|
from langgraph.pregel._io import map_input, read_channels
|
||||||
from langgraph.pregel._loop import AsyncPregelLoop, SyncPregelLoop
|
from langgraph.pregel._loop import (
|
||||||
|
AsyncPregelLoop,
|
||||||
|
SyncPregelLoop,
|
||||||
|
)
|
||||||
from langgraph.pregel._messages import StreamMessagesHandler
|
from langgraph.pregel._messages import StreamMessagesHandler
|
||||||
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
|
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
|
||||||
from langgraph.pregel._retry import RetryPolicy
|
from langgraph.pregel._retry import RetryPolicy
|
||||||
@@ -138,11 +143,13 @@ from langgraph.types import (
|
|||||||
Checkpointer,
|
Checkpointer,
|
||||||
Command,
|
Command,
|
||||||
Durability,
|
Durability,
|
||||||
|
GraphOutput,
|
||||||
Interrupt,
|
Interrupt,
|
||||||
Send,
|
Send,
|
||||||
StateSnapshot,
|
StateSnapshot,
|
||||||
StateUpdate,
|
StateUpdate,
|
||||||
StreamMode,
|
StreamMode,
|
||||||
|
StreamPart,
|
||||||
ensure_valid_checkpointer,
|
ensure_valid_checkpointer,
|
||||||
)
|
)
|
||||||
from langgraph.typing import ContextT, InputT, OutputT, StateT
|
from langgraph.typing import ContextT, InputT, OutputT, StateT
|
||||||
@@ -993,6 +1000,11 @@ class Pregel(
|
|||||||
for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse):
|
for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse):
|
||||||
yield name, node
|
yield name, node
|
||||||
|
|
||||||
|
# Mappers for v2 stream coercion (pydantic/dataclass).
|
||||||
|
# Set by CompiledStateGraph; None for base Pregel.
|
||||||
|
_output_mapper: Callable[[Any], Any] | None = None
|
||||||
|
_state_mapper: Callable[[Any], Any] | None = None
|
||||||
|
|
||||||
def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:
|
def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:
|
||||||
"""Migrate a saved checkpoint to new channel layout."""
|
"""Migrate a saved checkpoint to new channel layout."""
|
||||||
if checkpoint["v"] < 4 and checkpoint.get("pending_sends"):
|
if checkpoint["v"] < 4 and checkpoint.get("pending_sends"):
|
||||||
@@ -2427,6 +2439,7 @@ class Pregel(
|
|||||||
durability,
|
durability,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@overload
|
||||||
def stream(
|
def stream(
|
||||||
self,
|
self,
|
||||||
input: InputT | Command | None,
|
input: InputT | Command | None,
|
||||||
@@ -2441,6 +2454,44 @@ class Pregel(
|
|||||||
durability: Durability | None = None,
|
durability: Durability | None = None,
|
||||||
subgraphs: bool = False,
|
subgraphs: bool = False,
|
||||||
debug: bool | None = None,
|
debug: bool | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Unpack[DeprecatedKwargs],
|
||||||
|
) -> Iterator[dict[str, Any] | Any]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
debug: bool | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Unpack[DeprecatedKwargs],
|
||||||
|
) -> Iterator[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
debug: bool | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Unpack[DeprecatedKwargs],
|
**kwargs: Unpack[DeprecatedKwargs],
|
||||||
) -> Iterator[dict[str, Any] | Any]:
|
) -> Iterator[dict[str, Any] | Any]:
|
||||||
"""Stream graph steps for a single input.
|
"""Stream graph steps for a single input.
|
||||||
@@ -2602,6 +2653,10 @@ class Pregel(
|
|||||||
runtime = parent_runtime.merge(runtime)
|
runtime = parent_runtime.merge(runtime)
|
||||||
config[CONF][CONFIG_KEY_RUNTIME] = runtime
|
config[CONF][CONFIG_KEY_RUNTIME] = runtime
|
||||||
|
|
||||||
|
# resolve mappers for v2 stream coercion
|
||||||
|
_output_mapper = self._output_mapper if stream_version == "v2" else None
|
||||||
|
_state_mapper = self._state_mapper if stream_version == "v2" else None
|
||||||
|
|
||||||
with SyncPregelLoop(
|
with SyncPregelLoop(
|
||||||
input,
|
input,
|
||||||
stream=StreamProtocol(stream.put, stream_modes),
|
stream=StreamProtocol(stream.put, stream_modes),
|
||||||
@@ -2674,7 +2729,14 @@ class Pregel(
|
|||||||
):
|
):
|
||||||
# emit output
|
# emit output
|
||||||
yield from _output(
|
yield from _output(
|
||||||
stream_mode, print_mode, subgraphs, stream.get, queue.Empty
|
stream_mode,
|
||||||
|
print_mode,
|
||||||
|
subgraphs,
|
||||||
|
stream.get,
|
||||||
|
queue.Empty,
|
||||||
|
stream_version,
|
||||||
|
_output_mapper,
|
||||||
|
_state_mapper,
|
||||||
)
|
)
|
||||||
loop.after_tick()
|
loop.after_tick()
|
||||||
# wait for checkpoint
|
# wait for checkpoint
|
||||||
@@ -2682,7 +2744,14 @@ class Pregel(
|
|||||||
loop._put_checkpoint_fut.result()
|
loop._put_checkpoint_fut.result()
|
||||||
# emit output
|
# emit output
|
||||||
yield from _output(
|
yield from _output(
|
||||||
stream_mode, print_mode, subgraphs, stream.get, queue.Empty
|
stream_mode,
|
||||||
|
print_mode,
|
||||||
|
subgraphs,
|
||||||
|
stream.get,
|
||||||
|
queue.Empty,
|
||||||
|
stream_version,
|
||||||
|
_output_mapper,
|
||||||
|
_state_mapper,
|
||||||
)
|
)
|
||||||
# handle exit
|
# handle exit
|
||||||
if loop.status == "out_of_steps":
|
if loop.status == "out_of_steps":
|
||||||
@@ -2701,6 +2770,44 @@ class Pregel(
|
|||||||
run_manager.on_chain_error(e)
|
run_manager.on_chain_error(e)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def astream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
debug: bool | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Unpack[DeprecatedKwargs],
|
||||||
|
) -> AsyncIterator[dict[str, Any] | Any]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def astream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
debug: bool | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Unpack[DeprecatedKwargs],
|
||||||
|
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
async def astream(
|
async def astream(
|
||||||
self,
|
self,
|
||||||
input: InputT | Command | None,
|
input: InputT | Command | None,
|
||||||
@@ -2715,6 +2822,7 @@ class Pregel(
|
|||||||
durability: Durability | None = None,
|
durability: Durability | None = None,
|
||||||
subgraphs: bool = False,
|
subgraphs: bool = False,
|
||||||
debug: bool | None = None,
|
debug: bool | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Unpack[DeprecatedKwargs],
|
**kwargs: Unpack[DeprecatedKwargs],
|
||||||
) -> AsyncIterator[dict[str, Any] | Any]:
|
) -> AsyncIterator[dict[str, Any] | Any]:
|
||||||
"""Asynchronously stream graph steps for a single input.
|
"""Asynchronously stream graph steps for a single input.
|
||||||
@@ -2911,6 +3019,10 @@ class Pregel(
|
|||||||
runtime = parent_runtime.merge(runtime)
|
runtime = parent_runtime.merge(runtime)
|
||||||
config[CONF][CONFIG_KEY_RUNTIME] = runtime
|
config[CONF][CONFIG_KEY_RUNTIME] = runtime
|
||||||
|
|
||||||
|
# resolve mappers for v2 stream coercion
|
||||||
|
_output_mapper = self._output_mapper if stream_version == "v2" else None
|
||||||
|
_state_mapper = self._state_mapper if stream_version == "v2" else None
|
||||||
|
|
||||||
async with AsyncPregelLoop(
|
async with AsyncPregelLoop(
|
||||||
input,
|
input,
|
||||||
stream=StreamProtocol(stream.put_nowait, stream_modes),
|
stream=StreamProtocol(stream.put_nowait, stream_modes),
|
||||||
@@ -3007,6 +3119,9 @@ class Pregel(
|
|||||||
subgraphs,
|
subgraphs,
|
||||||
stream.get_nowait,
|
stream.get_nowait,
|
||||||
asyncio.QueueEmpty,
|
asyncio.QueueEmpty,
|
||||||
|
stream_version,
|
||||||
|
_output_mapper,
|
||||||
|
_state_mapper,
|
||||||
):
|
):
|
||||||
yield o
|
yield o
|
||||||
loop.after_tick()
|
loop.after_tick()
|
||||||
@@ -3025,6 +3140,9 @@ class Pregel(
|
|||||||
subgraphs,
|
subgraphs,
|
||||||
stream.get_nowait,
|
stream.get_nowait,
|
||||||
asyncio.QueueEmpty,
|
asyncio.QueueEmpty,
|
||||||
|
stream_version,
|
||||||
|
_output_mapper,
|
||||||
|
_state_mapper,
|
||||||
):
|
):
|
||||||
yield o
|
yield o
|
||||||
# handle exit
|
# handle exit
|
||||||
@@ -3044,6 +3162,7 @@ class Pregel(
|
|||||||
await asyncio.shield(run_manager.on_chain_error(e))
|
await asyncio.shield(run_manager.on_chain_error(e))
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
@overload
|
||||||
def invoke(
|
def invoke(
|
||||||
self,
|
self,
|
||||||
input: InputT | Command | None,
|
input: InputT | Command | None,
|
||||||
@@ -3056,6 +3175,57 @@ class Pregel(
|
|||||||
interrupt_before: All | Sequence[str] | None = None,
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
durability: Durability | None = None,
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> dict[str, Any] | Any: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: Literal["values"] = ...,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> GraphOutput[OutputT]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> list[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode = "values",
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> dict[str, Any] | Any:
|
) -> dict[str, Any] | Any:
|
||||||
"""Run the graph with a single input and config.
|
"""Run the graph with a single input and config.
|
||||||
@@ -3079,6 +3249,9 @@ class Pregel(
|
|||||||
- `"sync"`: Changes are persisted synchronously before the next step starts.
|
- `"sync"`: Changes are persisted synchronously before the next step starts.
|
||||||
- `"async"`: Changes are persisted asynchronously while the next step executes.
|
- `"async"`: Changes are persisted asynchronously while the next step executes.
|
||||||
- `"exit"`: Changes are persisted only when the graph exits.
|
- `"exit"`: Changes are persisted only when the graph exits.
|
||||||
|
stream_version: The streaming format version. `"v1"` (default) returns the
|
||||||
|
traditional format, `"v2"` returns `StreamPart` typed dicts when
|
||||||
|
`stream_mode` is not `"values"`.
|
||||||
**kwargs: Additional keyword arguments to pass to the graph run.
|
**kwargs: Additional keyword arguments to pass to the graph run.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -3091,39 +3264,64 @@ class Pregel(
|
|||||||
chunks: list[dict[str, Any] | Any] = []
|
chunks: list[dict[str, Any] | Any] = []
|
||||||
interrupts: list[Interrupt] = []
|
interrupts: list[Interrupt] = []
|
||||||
|
|
||||||
for chunk in self.stream(
|
if stream_version == "v2":
|
||||||
input,
|
# v2: values stream parts carry interrupts directly
|
||||||
config,
|
for chunk in self.stream(
|
||||||
context=context,
|
input,
|
||||||
stream_mode=["updates", "values"]
|
config,
|
||||||
if stream_mode == "values"
|
context=context,
|
||||||
else stream_mode,
|
stream_mode="values" if stream_mode == "values" else stream_mode,
|
||||||
print_mode=print_mode,
|
print_mode=print_mode,
|
||||||
output_keys=output_keys,
|
output_keys=output_keys,
|
||||||
interrupt_before=interrupt_before,
|
interrupt_before=interrupt_before,
|
||||||
interrupt_after=interrupt_after,
|
interrupt_after=interrupt_after,
|
||||||
durability=durability,
|
durability=durability,
|
||||||
**kwargs,
|
stream_version=stream_version,
|
||||||
):
|
**kwargs,
|
||||||
if stream_mode == "values":
|
):
|
||||||
if len(chunk) == 2:
|
if stream_mode == "values":
|
||||||
mode, payload = cast(tuple[StreamMode, Any], chunk)
|
latest = chunk["data"]
|
||||||
|
if chunk_ints := chunk.get("interrupts", ()):
|
||||||
|
interrupts.extend(chunk_ints) # type: ignore[arg-type]
|
||||||
else:
|
else:
|
||||||
_, mode, payload = cast(
|
chunks.append(chunk)
|
||||||
tuple[tuple[str, ...], StreamMode, Any], chunk
|
else:
|
||||||
)
|
# v1: collect interrupts from updates stream
|
||||||
if (
|
for chunk in self.stream(
|
||||||
mode == "updates"
|
input,
|
||||||
and isinstance(payload, dict)
|
config,
|
||||||
and (ints := payload.get(INTERRUPT)) is not None
|
context=context,
|
||||||
):
|
stream_mode=(
|
||||||
interrupts.extend(ints)
|
["updates", "values"] if stream_mode == "values" else stream_mode
|
||||||
elif mode == "values":
|
),
|
||||||
latest = payload
|
print_mode=print_mode,
|
||||||
else:
|
output_keys=output_keys,
|
||||||
chunks.append(chunk)
|
interrupt_before=interrupt_before,
|
||||||
|
interrupt_after=interrupt_after,
|
||||||
|
durability=durability,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
if stream_mode == "values":
|
||||||
|
if len(chunk) == 2:
|
||||||
|
mode, payload = cast(tuple[StreamMode, Any], chunk)
|
||||||
|
else:
|
||||||
|
_, mode, payload = cast(
|
||||||
|
tuple[tuple[str, ...], StreamMode, Any], chunk
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
mode == "updates"
|
||||||
|
and isinstance(payload, dict)
|
||||||
|
and (ints := payload.get(INTERRUPT)) is not None
|
||||||
|
):
|
||||||
|
interrupts.extend(ints)
|
||||||
|
elif mode == "values":
|
||||||
|
latest = payload
|
||||||
|
else:
|
||||||
|
chunks.append(chunk)
|
||||||
|
|
||||||
if stream_mode == "values":
|
if stream_mode == "values":
|
||||||
|
if stream_version == "v2":
|
||||||
|
return GraphOutput(value=latest, interrupts=tuple(interrupts))
|
||||||
if interrupts:
|
if interrupts:
|
||||||
return (
|
return (
|
||||||
{**latest, INTERRUPT: interrupts}
|
{**latest, INTERRUPT: interrupts}
|
||||||
@@ -3134,6 +3332,57 @@ class Pregel(
|
|||||||
else:
|
else:
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode = "values",
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> dict[str, Any] | Any: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: Literal["values"] = ...,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> GraphOutput[OutputT]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode,
|
||||||
|
print_mode: StreamMode | Sequence[StreamMode] = (),
|
||||||
|
output_keys: str | Sequence[str] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> list[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
async def ainvoke(
|
async def ainvoke(
|
||||||
self,
|
self,
|
||||||
input: InputT | Command | None,
|
input: InputT | Command | None,
|
||||||
@@ -3146,6 +3395,7 @@ class Pregel(
|
|||||||
interrupt_before: All | Sequence[str] | None = None,
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
durability: Durability | None = None,
|
durability: Durability | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> dict[str, Any] | Any:
|
) -> dict[str, Any] | Any:
|
||||||
"""Asynchronously run the graph with a single input and config.
|
"""Asynchronously run the graph with a single input and config.
|
||||||
@@ -3169,6 +3419,9 @@ class Pregel(
|
|||||||
- `"sync"`: Changes are persisted synchronously before the next step starts.
|
- `"sync"`: Changes are persisted synchronously before the next step starts.
|
||||||
- `"async"`: Changes are persisted asynchronously while the next step executes.
|
- `"async"`: Changes are persisted asynchronously while the next step executes.
|
||||||
- `"exit"`: Changes are persisted only when the graph exits.
|
- `"exit"`: Changes are persisted only when the graph exits.
|
||||||
|
stream_version: The streaming format version. `"v1"` (default) returns the
|
||||||
|
traditional format, `"v2"` returns `StreamPart` typed dicts when
|
||||||
|
`stream_mode` is not `"values"`.
|
||||||
**kwargs: Additional keyword arguments to pass to the graph run.
|
**kwargs: Additional keyword arguments to pass to the graph run.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -3181,39 +3434,64 @@ class Pregel(
|
|||||||
chunks: list[dict[str, Any] | Any] = []
|
chunks: list[dict[str, Any] | Any] = []
|
||||||
interrupts: list[Interrupt] = []
|
interrupts: list[Interrupt] = []
|
||||||
|
|
||||||
async for chunk in self.astream(
|
if stream_version == "v2":
|
||||||
input,
|
# v2: values stream parts carry interrupts directly
|
||||||
config,
|
async for chunk in self.astream(
|
||||||
context=context,
|
input,
|
||||||
stream_mode=["updates", "values"]
|
config,
|
||||||
if stream_mode == "values"
|
context=context,
|
||||||
else stream_mode,
|
stream_mode="values" if stream_mode == "values" else stream_mode,
|
||||||
print_mode=print_mode,
|
print_mode=print_mode,
|
||||||
output_keys=output_keys,
|
output_keys=output_keys,
|
||||||
interrupt_before=interrupt_before,
|
interrupt_before=interrupt_before,
|
||||||
interrupt_after=interrupt_after,
|
interrupt_after=interrupt_after,
|
||||||
durability=durability,
|
durability=durability,
|
||||||
**kwargs,
|
stream_version=stream_version,
|
||||||
):
|
**kwargs,
|
||||||
if stream_mode == "values":
|
):
|
||||||
if len(chunk) == 2:
|
if stream_mode == "values":
|
||||||
mode, payload = cast(tuple[StreamMode, Any], chunk)
|
latest = chunk["data"]
|
||||||
|
if chunk_ints := chunk.get("interrupts", ()):
|
||||||
|
interrupts.extend(chunk_ints) # type: ignore[arg-type]
|
||||||
else:
|
else:
|
||||||
_, mode, payload = cast(
|
chunks.append(chunk)
|
||||||
tuple[tuple[str, ...], StreamMode, Any], chunk
|
else:
|
||||||
)
|
# v1: collect interrupts from updates stream
|
||||||
if (
|
async for chunk in self.astream(
|
||||||
mode == "updates"
|
input,
|
||||||
and isinstance(payload, dict)
|
config,
|
||||||
and (ints := payload.get(INTERRUPT)) is not None
|
context=context,
|
||||||
):
|
stream_mode=(
|
||||||
interrupts.extend(ints)
|
["updates", "values"] if stream_mode == "values" else stream_mode
|
||||||
elif mode == "values":
|
),
|
||||||
latest = payload
|
print_mode=print_mode,
|
||||||
else:
|
output_keys=output_keys,
|
||||||
chunks.append(chunk)
|
interrupt_before=interrupt_before,
|
||||||
|
interrupt_after=interrupt_after,
|
||||||
|
durability=durability,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
if stream_mode == "values":
|
||||||
|
if len(chunk) == 2:
|
||||||
|
mode, payload = cast(tuple[StreamMode, Any], chunk)
|
||||||
|
else:
|
||||||
|
_, mode, payload = cast(
|
||||||
|
tuple[tuple[str, ...], StreamMode, Any], chunk
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
mode == "updates"
|
||||||
|
and isinstance(payload, dict)
|
||||||
|
and (ints := payload.get(INTERRUPT)) is not None
|
||||||
|
):
|
||||||
|
interrupts.extend(ints)
|
||||||
|
elif mode == "values":
|
||||||
|
latest = payload
|
||||||
|
else:
|
||||||
|
chunks.append(chunk)
|
||||||
|
|
||||||
if stream_mode == "values":
|
if stream_mode == "values":
|
||||||
|
if stream_version == "v2":
|
||||||
|
return GraphOutput(value=latest, interrupts=tuple(interrupts))
|
||||||
if interrupts:
|
if interrupts:
|
||||||
return (
|
return (
|
||||||
{**latest, INTERRUPT: interrupts}
|
{**latest, INTERRUPT: interrupts}
|
||||||
@@ -3278,6 +3556,9 @@ def _output(
|
|||||||
stream_subgraphs: bool,
|
stream_subgraphs: bool,
|
||||||
getter: Callable[[], tuple[tuple[str, ...], str, Any]],
|
getter: Callable[[], tuple[tuple[str, ...], str, Any]],
|
||||||
empty_exc: type[Exception],
|
empty_exc: type[Exception],
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
|
output_mapper: Callable[[Any], Any] | None = None,
|
||||||
|
state_mapper: Callable[[Any], Any] | None = None,
|
||||||
) -> Iterator:
|
) -> Iterator:
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
@@ -3305,7 +3586,23 @@ def _output(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if mode in stream_mode:
|
if mode in stream_mode:
|
||||||
if stream_subgraphs and isinstance(stream_mode, list):
|
if stream_version == "v2":
|
||||||
|
if mode == "values":
|
||||||
|
# pop __interrupt__ into typed field, coerce data
|
||||||
|
ints: tuple[Interrupt, ...] = ()
|
||||||
|
if isinstance(payload, dict):
|
||||||
|
ints = payload.pop(INTERRUPT, ())
|
||||||
|
if output_mapper:
|
||||||
|
payload = output_mapper(payload)
|
||||||
|
yield {"type": mode, "ns": ns, "data": payload, "interrupts": ints}
|
||||||
|
elif mode in ("checkpoints", "debug"):
|
||||||
|
# coerce state values in checkpoint/debug payloads
|
||||||
|
if state_mapper:
|
||||||
|
_coerce_checkpoint_values(payload, state_mapper)
|
||||||
|
yield {"type": mode, "ns": ns, "data": payload}
|
||||||
|
else:
|
||||||
|
yield {"type": mode, "ns": ns, "data": payload}
|
||||||
|
elif stream_subgraphs and isinstance(stream_mode, list):
|
||||||
yield (ns, mode, payload)
|
yield (ns, mode, payload)
|
||||||
elif isinstance(stream_mode, list):
|
elif isinstance(stream_mode, list):
|
||||||
yield (mode, payload)
|
yield (mode, payload)
|
||||||
@@ -3315,6 +3612,31 @@ def _output(
|
|||||||
yield payload
|
yield payload
|
||||||
|
|
||||||
|
|
||||||
|
def _coerce_checkpoint_values(payload: Any, mapper: Callable[[Any], Any]) -> None:
|
||||||
|
"""Coerce `values` dicts inside checkpoint or debug payloads in-place.
|
||||||
|
|
||||||
|
Skips the initial checkpoint (where next contains ``__start__``) because
|
||||||
|
not all channels are populated yet and coercion would fail.
|
||||||
|
"""
|
||||||
|
_START = "__start__"
|
||||||
|
# debug wrapper: {"type": "checkpoint", "payload": {"values": dict, ...}}
|
||||||
|
if (
|
||||||
|
isinstance(payload, dict)
|
||||||
|
and payload.get("type") == "checkpoint"
|
||||||
|
and isinstance(payload.get("payload"), dict)
|
||||||
|
and isinstance(payload["payload"].get("values"), dict)
|
||||||
|
and _START not in payload["payload"].get("next", ())
|
||||||
|
):
|
||||||
|
payload["payload"]["values"] = mapper(payload["payload"]["values"])
|
||||||
|
# direct checkpoint payload: {"values": dict, ...}
|
||||||
|
elif (
|
||||||
|
isinstance(payload, dict)
|
||||||
|
and isinstance(payload.get("values"), dict)
|
||||||
|
and _START not in payload.get("next", ())
|
||||||
|
):
|
||||||
|
payload["values"] = mapper(payload["values"])
|
||||||
|
|
||||||
|
|
||||||
def _coerce_context(
|
def _coerce_context(
|
||||||
context_schema: type[ContextT] | None, context: Any
|
context_schema: type[ContextT] | None, context: Any
|
||||||
) -> ContextT | None:
|
) -> ContextT | None:
|
||||||
|
|||||||
@@ -2,13 +2,21 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from collections.abc import AsyncIterator, Callable, Iterator, Sequence
|
from collections.abc import AsyncIterator, Callable, Iterator, Sequence
|
||||||
from typing import Any, Generic, cast
|
from typing import Any, Generic, Literal, cast, overload
|
||||||
|
|
||||||
from langchain_core.runnables import Runnable, RunnableConfig
|
from langchain_core.runnables import Runnable, RunnableConfig
|
||||||
from langchain_core.runnables.graph import Graph as DrawableGraph
|
from langchain_core.runnables.graph import Graph as DrawableGraph
|
||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from langgraph.types import All, Command, StateSnapshot, StateUpdate, StreamMode
|
from langgraph.types import (
|
||||||
|
All,
|
||||||
|
Command,
|
||||||
|
GraphOutput,
|
||||||
|
StateSnapshot,
|
||||||
|
StateUpdate,
|
||||||
|
StreamMode,
|
||||||
|
StreamPart,
|
||||||
|
)
|
||||||
from langgraph.typing import ContextT, InputT, OutputT, StateT
|
from langgraph.typing import ContextT, InputT, OutputT, StateT
|
||||||
|
|
||||||
__all__ = ("PregelProtocol", "StreamProtocol")
|
__all__ = ("PregelProtocol", "StreamProtocol")
|
||||||
@@ -96,6 +104,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
|||||||
as_node: str | None = None,
|
as_node: str | None = None,
|
||||||
) -> RunnableConfig: ...
|
) -> RunnableConfig: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def stream(
|
def stream(
|
||||||
self,
|
self,
|
||||||
@@ -107,8 +116,68 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
|||||||
interrupt_before: All | Sequence[str] | None = None,
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
subgraphs: bool = False,
|
subgraphs: bool = False,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
) -> Iterator[dict[str, Any] | Any]: ...
|
) -> Iterator[dict[str, Any] | Any]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
) -> Iterator[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
|
) -> Iterator[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
def astream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
) -> AsyncIterator[dict[str, Any] | Any]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
def astream(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def astream(
|
def astream(
|
||||||
self,
|
self,
|
||||||
@@ -120,7 +189,34 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
|||||||
interrupt_before: All | Sequence[str] | None = None,
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
subgraphs: bool = False,
|
subgraphs: bool = False,
|
||||||
) -> AsyncIterator[dict[str, Any] | Any]: ...
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
|
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
) -> dict[str, Any] | Any: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
) -> GraphOutput[OutputT]: ...
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def invoke(
|
def invoke(
|
||||||
@@ -131,8 +227,35 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
|||||||
context: ContextT | None = None,
|
context: ContextT | None = None,
|
||||||
interrupt_before: All | Sequence[str] | None = None,
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
|
) -> GraphOutput[OutputT]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
) -> dict[str, Any] | Any: ...
|
) -> dict[str, Any] | Any: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
@abstractmethod
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: InputT | Command | None,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
context: ContextT | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
) -> GraphOutput[OutputT]: ...
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def ainvoke(
|
async def ainvoke(
|
||||||
self,
|
self,
|
||||||
@@ -142,7 +265,8 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
|
|||||||
context: ContextT | None = None,
|
context: ContextT | None = None,
|
||||||
interrupt_before: All | Sequence[str] | None = None,
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
) -> dict[str, Any] | Any: ...
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
|
) -> GraphOutput[OutputT]: ...
|
||||||
|
|
||||||
|
|
||||||
StreamChunk = tuple[tuple[str, ...], str, Any]
|
StreamChunk = tuple[tuple[str, ...], str, Any]
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from typing import (
|
|||||||
Any,
|
Any,
|
||||||
Literal,
|
Literal,
|
||||||
cast,
|
cast,
|
||||||
|
overload,
|
||||||
)
|
)
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
@@ -57,10 +58,12 @@ from langgraph.pregel.protocol import PregelProtocol, StreamProtocol
|
|||||||
from langgraph.types import (
|
from langgraph.types import (
|
||||||
All,
|
All,
|
||||||
Command,
|
Command,
|
||||||
|
GraphOutput,
|
||||||
Interrupt,
|
Interrupt,
|
||||||
PregelTask,
|
PregelTask,
|
||||||
StateSnapshot,
|
StateSnapshot,
|
||||||
StreamMode,
|
StreamMode,
|
||||||
|
StreamPart,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -682,6 +685,7 @@ class RemoteGraph(PregelProtocol):
|
|||||||
updated_stream_modes.remove("events")
|
updated_stream_modes.remove("events")
|
||||||
return (updated_stream_modes, requested_stream_modes, req_single, stream)
|
return (updated_stream_modes, requested_stream_modes, req_single, stream)
|
||||||
|
|
||||||
|
@overload
|
||||||
def stream(
|
def stream(
|
||||||
self,
|
self,
|
||||||
input: dict[str, Any] | Any,
|
input: dict[str, Any] | Any,
|
||||||
@@ -693,6 +697,38 @@ class RemoteGraph(PregelProtocol):
|
|||||||
subgraphs: bool = False,
|
subgraphs: bool = False,
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Iterator[dict[str, Any] | Any]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Iterator[StreamPart]: ...
|
||||||
|
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> Iterator[dict[str, Any] | Any]:
|
) -> Iterator[dict[str, Any] | Any]:
|
||||||
"""Create a run and stream the results.
|
"""Create a run and stream the results.
|
||||||
@@ -774,10 +810,12 @@ class RemoteGraph(PregelProtocol):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if chunk.event.startswith("messages"):
|
if chunk.event.startswith("messages"):
|
||||||
chunk = chunk._replace(data=tuple(chunk.data)) # type: ignore
|
chunk = chunk._replace(data=tuple(chunk.data))
|
||||||
|
|
||||||
# emit chunk
|
# emit chunk
|
||||||
if subgraphs:
|
if stream_version == "v2":
|
||||||
|
yield {"type": mode, "ns": ns, "data": chunk.data}
|
||||||
|
elif subgraphs:
|
||||||
if NS_SEP in chunk.event:
|
if NS_SEP in chunk.event:
|
||||||
mode, ns_ = chunk.event.split(NS_SEP, 1)
|
mode, ns_ = chunk.event.split(NS_SEP, 1)
|
||||||
ns = tuple(ns_.split(NS_SEP))
|
ns = tuple(ns_.split(NS_SEP))
|
||||||
@@ -792,6 +830,38 @@ class RemoteGraph(PregelProtocol):
|
|||||||
else:
|
else:
|
||||||
yield chunk
|
yield chunk
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def astream(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> AsyncIterator[dict[str, Any] | Any]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def astream(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
stream_mode: StreamMode | list[StreamMode] | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
subgraphs: bool = False,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> AsyncIterator[StreamPart]: ...
|
||||||
|
|
||||||
async def astream(
|
async def astream(
|
||||||
self,
|
self,
|
||||||
input: dict[str, Any] | Any,
|
input: dict[str, Any] | Any,
|
||||||
@@ -803,6 +873,7 @@ class RemoteGraph(PregelProtocol):
|
|||||||
subgraphs: bool = False,
|
subgraphs: bool = False,
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> AsyncIterator[dict[str, Any] | Any]:
|
) -> AsyncIterator[dict[str, Any] | Any]:
|
||||||
"""Create a run and stream the results.
|
"""Create a run and stream the results.
|
||||||
@@ -884,10 +955,12 @@ class RemoteGraph(PregelProtocol):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if chunk.event.startswith("messages"):
|
if chunk.event.startswith("messages"):
|
||||||
chunk = chunk._replace(data=tuple(chunk.data)) # type: ignore
|
chunk = chunk._replace(data=tuple(chunk.data))
|
||||||
|
|
||||||
# emit chunk
|
# emit chunk
|
||||||
if subgraphs:
|
if stream_version == "v2":
|
||||||
|
yield {"type": mode, "ns": ns, "data": chunk.data}
|
||||||
|
elif subgraphs:
|
||||||
if NS_SEP in chunk.event:
|
if NS_SEP in chunk.event:
|
||||||
mode, ns_ = chunk.event.split(NS_SEP, 1)
|
mode, ns_ = chunk.event.split(NS_SEP, 1)
|
||||||
ns = tuple(ns_.split(NS_SEP))
|
ns = tuple(ns_.split(NS_SEP))
|
||||||
@@ -918,6 +991,7 @@ class RemoteGraph(PregelProtocol):
|
|||||||
) -> AsyncIterator[dict[str, Any]]:
|
) -> AsyncIterator[dict[str, Any]]:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@overload
|
||||||
def invoke(
|
def invoke(
|
||||||
self,
|
self,
|
||||||
input: dict[str, Any] | Any,
|
input: dict[str, Any] | Any,
|
||||||
@@ -927,6 +1001,34 @@ class RemoteGraph(PregelProtocol):
|
|||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> dict[str, Any] | Any: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> GraphOutput[dict[str, Any]]: ...
|
||||||
|
|
||||||
|
def invoke(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> dict[str, Any] | Any:
|
) -> dict[str, Any] | Any:
|
||||||
"""Create a run, wait until it finishes and return the final state.
|
"""Create a run, wait until it finishes and return the final state.
|
||||||
@@ -937,12 +1039,14 @@ class RemoteGraph(PregelProtocol):
|
|||||||
interrupt_before: Interrupt the graph before these nodes.
|
interrupt_before: Interrupt the graph before these nodes.
|
||||||
interrupt_after: Interrupt the graph after these nodes.
|
interrupt_after: Interrupt the graph after these nodes.
|
||||||
headers: Additional headers to pass to the request.
|
headers: Additional headers to pass to the request.
|
||||||
|
stream_version: The streaming format version. `"v1"` (default) returns the
|
||||||
|
traditional format, `"v2"` returns `StreamPart` typed dicts.
|
||||||
**kwargs: Additional params to pass to RemoteGraph.stream.
|
**kwargs: Additional params to pass to RemoteGraph.stream.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The output of the graph.
|
The output of the graph.
|
||||||
"""
|
"""
|
||||||
for chunk in self.stream(
|
for chunk in self.stream( # type: ignore[misc, call-overload]
|
||||||
input,
|
input,
|
||||||
config=config,
|
config=config,
|
||||||
interrupt_before=interrupt_before,
|
interrupt_before=interrupt_before,
|
||||||
@@ -950,15 +1054,49 @@ class RemoteGraph(PregelProtocol):
|
|||||||
headers=headers,
|
headers=headers,
|
||||||
stream_mode="values",
|
stream_mode="values",
|
||||||
params=params,
|
params=params,
|
||||||
|
stream_version=stream_version,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
pass
|
pass
|
||||||
try:
|
try:
|
||||||
|
if stream_version == "v2":
|
||||||
|
return GraphOutput(
|
||||||
|
value=chunk["data"],
|
||||||
|
interrupts=tuple(chunk.get("interrupts", ())),
|
||||||
|
)
|
||||||
return chunk
|
return chunk
|
||||||
except UnboundLocalError:
|
except UnboundLocalError:
|
||||||
logger.warning("No events received from remote graph")
|
logger.warning("No events received from remote graph")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> dict[str, Any] | Any: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
async def ainvoke(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None = None,
|
||||||
|
*,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
headers: dict[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v2"] = ...,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> GraphOutput[dict[str, Any]]: ...
|
||||||
|
|
||||||
async def ainvoke(
|
async def ainvoke(
|
||||||
self,
|
self,
|
||||||
input: dict[str, Any] | Any,
|
input: dict[str, Any] | Any,
|
||||||
@@ -968,6 +1106,7 @@ class RemoteGraph(PregelProtocol):
|
|||||||
interrupt_after: All | Sequence[str] | None = None,
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
headers: dict[str, str] | None = None,
|
headers: dict[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
|
stream_version: Literal["v1", "v2"] = "v2",
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> dict[str, Any] | Any:
|
) -> dict[str, Any] | Any:
|
||||||
"""Create a run, wait until it finishes and return the final state.
|
"""Create a run, wait until it finishes and return the final state.
|
||||||
@@ -978,12 +1117,14 @@ class RemoteGraph(PregelProtocol):
|
|||||||
interrupt_before: Interrupt the graph before these nodes.
|
interrupt_before: Interrupt the graph before these nodes.
|
||||||
interrupt_after: Interrupt the graph after these nodes.
|
interrupt_after: Interrupt the graph after these nodes.
|
||||||
headers: Additional headers to pass to the request.
|
headers: Additional headers to pass to the request.
|
||||||
|
stream_version: The streaming format version. `"v1"` (default) returns the
|
||||||
|
traditional format, `"v2"` returns `StreamPart` typed dicts.
|
||||||
**kwargs: Additional params to pass to RemoteGraph.astream.
|
**kwargs: Additional params to pass to RemoteGraph.astream.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The output of the graph.
|
The output of the graph.
|
||||||
"""
|
"""
|
||||||
async for chunk in self.astream(
|
async for chunk in self.astream( # type: ignore[misc, call-overload]
|
||||||
input,
|
input,
|
||||||
config=config,
|
config=config,
|
||||||
interrupt_before=interrupt_before,
|
interrupt_before=interrupt_before,
|
||||||
@@ -991,10 +1132,16 @@ class RemoteGraph(PregelProtocol):
|
|||||||
headers=headers,
|
headers=headers,
|
||||||
stream_mode="values",
|
stream_mode="values",
|
||||||
params=params,
|
params=params,
|
||||||
|
stream_version=stream_version,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
pass
|
pass
|
||||||
try:
|
try:
|
||||||
|
if stream_version == "v2":
|
||||||
|
return GraphOutput(
|
||||||
|
value=chunk["data"],
|
||||||
|
interrupts=tuple(chunk.get("interrupts", ())),
|
||||||
|
)
|
||||||
return chunk
|
return chunk
|
||||||
except UnboundLocalError:
|
except UnboundLocalError:
|
||||||
logger.warning("No events received from remote graph")
|
logger.warning("No events received from remote graph")
|
||||||
|
|||||||
@@ -16,17 +16,26 @@ from typing import (
|
|||||||
)
|
)
|
||||||
from warnings import warn
|
from warnings import warn
|
||||||
|
|
||||||
|
from langchain_core.messages import AnyMessage
|
||||||
from langchain_core.runnables import Runnable, RunnableConfig
|
from langchain_core.runnables import Runnable, RunnableConfig
|
||||||
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
|
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
|
||||||
from typing_extensions import Unpack, deprecated
|
from typing_extensions import NotRequired, TypeAliasType, TypedDict, Unpack, deprecated
|
||||||
from xxhash import xxh3_128_hexdigest
|
from xxhash import xxh3_128_hexdigest
|
||||||
|
|
||||||
from langgraph._internal._cache import default_cache_key
|
from langgraph._internal._cache import default_cache_key
|
||||||
|
from langgraph._internal._constants import INTERRUPT as _INTERRUPT_KEY
|
||||||
from langgraph._internal._fields import get_cached_annotated_keys, get_update_as_tuples
|
from langgraph._internal._fields import get_cached_annotated_keys, get_update_as_tuples
|
||||||
from langgraph._internal._retry import default_retry_on
|
from langgraph._internal._retry import default_retry_on
|
||||||
from langgraph._internal._typing import MISSING, DeprecatedKwargs
|
from langgraph._internal._typing import MISSING, DeprecatedKwargs
|
||||||
from langgraph.warnings import LangGraphDeprecatedSinceV10
|
from langgraph.warnings import LangGraphDeprecatedSinceV10
|
||||||
|
|
||||||
|
# Local TypeVars for generic stream TypedDicts.
|
||||||
|
# We use separate TypeVars here (rather than importing from langgraph.typing)
|
||||||
|
# because the typing module TypeVars have defaults that cause mypy issues
|
||||||
|
# when used in standalone type aliases.
|
||||||
|
StateT = TypeVar("StateT")
|
||||||
|
OutputT = TypeVar("OutputT")
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.pregel.protocol import PregelProtocol
|
from langgraph.pregel.protocol import PregelProtocol
|
||||||
|
|
||||||
@@ -44,6 +53,19 @@ __all__ = (
|
|||||||
"Checkpointer",
|
"Checkpointer",
|
||||||
"StreamMode",
|
"StreamMode",
|
||||||
"StreamWriter",
|
"StreamWriter",
|
||||||
|
"StreamPart",
|
||||||
|
"ValuesStreamPart",
|
||||||
|
"UpdatesStreamPart",
|
||||||
|
"MessagesStreamPart",
|
||||||
|
"CustomStreamPart",
|
||||||
|
"CheckpointStreamPart",
|
||||||
|
"TasksStreamPart",
|
||||||
|
"DebugStreamPart",
|
||||||
|
"TaskPayload",
|
||||||
|
"TaskResultPayload",
|
||||||
|
"CheckpointTask",
|
||||||
|
"CheckpointPayload",
|
||||||
|
"DebugPayload",
|
||||||
"RetryPolicy",
|
"RetryPolicy",
|
||||||
"CachePolicy",
|
"CachePolicy",
|
||||||
"Interrupt",
|
"Interrupt",
|
||||||
@@ -56,6 +78,7 @@ __all__ = (
|
|||||||
"Durability",
|
"Durability",
|
||||||
"interrupt",
|
"interrupt",
|
||||||
"Overwrite",
|
"Overwrite",
|
||||||
|
"GraphOutput",
|
||||||
"ensure_valid_checkpointer",
|
"ensure_valid_checkpointer",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -113,6 +136,221 @@ StreamWriter = Callable[[Any], None]
|
|||||||
Always injected into nodes if requested as a keyword argument, but it's a no-op
|
Always injected into nodes if requested as a keyword argument, but it's a no-op
|
||||||
when not using `stream_mode="custom"`."""
|
when not using `stream_mode="custom"`."""
|
||||||
|
|
||||||
|
|
||||||
|
class TaskPayload(TypedDict):
|
||||||
|
"""Payload for a task start event."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
input: Any
|
||||||
|
triggers: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
class TaskResultPayload(TypedDict):
|
||||||
|
"""Payload for a task result event."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
error: str | None
|
||||||
|
interrupts: list[dict]
|
||||||
|
result: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointTask(TypedDict):
|
||||||
|
"""A task entry within a `CheckpointPayload`.
|
||||||
|
|
||||||
|
The keys present depend on the task's state:
|
||||||
|
|
||||||
|
- **Error:** `id`, `name`, `error`, `state`
|
||||||
|
- **Has result:** `id`, `name`, `result`, `interrupts`, `state`
|
||||||
|
- **Pending:** `id`, `name`, `interrupts`, `state`
|
||||||
|
"""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
error: NotRequired[str]
|
||||||
|
result: NotRequired[Any]
|
||||||
|
interrupts: NotRequired[list[dict]]
|
||||||
|
state: StateSnapshot | RunnableConfig | None
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointPayload(TypedDict, Generic[StateT]):
|
||||||
|
"""Payload for a checkpoint event."""
|
||||||
|
|
||||||
|
config: RunnableConfig | None
|
||||||
|
metadata: CheckpointMetadata
|
||||||
|
values: StateT
|
||||||
|
next: list[str]
|
||||||
|
parent_config: RunnableConfig | None
|
||||||
|
tasks: list[CheckpointTask]
|
||||||
|
|
||||||
|
|
||||||
|
class _DebugCheckpointPayload(TypedDict, Generic[StateT]):
|
||||||
|
step: int
|
||||||
|
timestamp: str
|
||||||
|
type: Literal["checkpoint"]
|
||||||
|
payload: CheckpointPayload[StateT]
|
||||||
|
|
||||||
|
|
||||||
|
class _DebugTaskPayload(TypedDict):
|
||||||
|
step: int
|
||||||
|
timestamp: str
|
||||||
|
type: Literal["task"]
|
||||||
|
payload: TaskPayload
|
||||||
|
|
||||||
|
|
||||||
|
class _DebugTaskResultPayload(TypedDict):
|
||||||
|
step: int
|
||||||
|
timestamp: str
|
||||||
|
type: Literal["task_result"]
|
||||||
|
payload: TaskResultPayload
|
||||||
|
|
||||||
|
|
||||||
|
DebugPayload = TypeAliasType(
|
||||||
|
"DebugPayload",
|
||||||
|
_DebugCheckpointPayload[StateT] | _DebugTaskPayload | _DebugTaskResultPayload,
|
||||||
|
type_params=(StateT,),
|
||||||
|
)
|
||||||
|
"""Wrapper payload for debug events. Discriminate on `type`."""
|
||||||
|
|
||||||
|
|
||||||
|
class ValuesStreamPart(TypedDict, Generic[OutputT]):
|
||||||
|
"""Stream part emitted for `stream_mode="values"`.
|
||||||
|
|
||||||
|
`data` contains the full state after each step, as returned by `read_channels()`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
type: Literal["values"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: OutputT
|
||||||
|
interrupts: tuple[Interrupt, ...]
|
||||||
|
|
||||||
|
|
||||||
|
class UpdatesStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for `stream_mode="updates"`.
|
||||||
|
|
||||||
|
`data` maps node names to their outputs. May also contain
|
||||||
|
`__interrupt__` (tuple of `Interrupt` dicts) and `__metadata__` keys.
|
||||||
|
"""
|
||||||
|
|
||||||
|
type: Literal["updates"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class MessagesStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for `stream_mode="messages"`.
|
||||||
|
|
||||||
|
`data` is a 2-tuple of `(message, metadata)` where `message` is a
|
||||||
|
`BaseMessage` (e.g. `AIMessageChunk`) and `metadata` is a dict containing
|
||||||
|
keys like `langgraph_step`, `langgraph_node`, `langgraph_triggers`, etc.
|
||||||
|
"""
|
||||||
|
|
||||||
|
type: Literal["messages"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: tuple[AnyMessage, dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
class CustomStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for `stream_mode="custom"`.
|
||||||
|
|
||||||
|
`data` is whatever value was passed to `StreamWriter` inside a node.
|
||||||
|
"""
|
||||||
|
|
||||||
|
type: Literal["custom"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: Any
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointStreamPart(TypedDict, Generic[StateT]):
|
||||||
|
"""Stream part emitted for `stream_mode="checkpoints"`."""
|
||||||
|
|
||||||
|
type: Literal["checkpoints"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: CheckpointPayload[StateT]
|
||||||
|
|
||||||
|
|
||||||
|
class TasksStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for `stream_mode="tasks"`.
|
||||||
|
|
||||||
|
For task start events, `data` is a `TaskPayload` with `id`, `name`,
|
||||||
|
`input`, and `triggers` keys.
|
||||||
|
|
||||||
|
For task result events, `data` is a `TaskResultPayload` with `id`,
|
||||||
|
`name`, `error`, `interrupts`, and `result` keys.
|
||||||
|
"""
|
||||||
|
|
||||||
|
type: Literal["tasks"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: TaskPayload | TaskResultPayload
|
||||||
|
|
||||||
|
|
||||||
|
class DebugStreamPart(TypedDict, Generic[StateT]):
|
||||||
|
"""Stream part emitted for `stream_mode="debug"`."""
|
||||||
|
|
||||||
|
type: Literal["debug"]
|
||||||
|
ns: tuple[str, ...]
|
||||||
|
data: DebugPayload[StateT]
|
||||||
|
|
||||||
|
|
||||||
|
StreamPart = TypeAliasType(
|
||||||
|
"StreamPart",
|
||||||
|
ValuesStreamPart[OutputT]
|
||||||
|
| UpdatesStreamPart
|
||||||
|
| MessagesStreamPart
|
||||||
|
| CustomStreamPart
|
||||||
|
| CheckpointStreamPart[StateT]
|
||||||
|
| TasksStreamPart
|
||||||
|
| DebugStreamPart[StateT],
|
||||||
|
type_params=(OutputT, StateT),
|
||||||
|
)
|
||||||
|
"""A discriminated union of all v2 stream part types.
|
||||||
|
|
||||||
|
Use `part["type"]` to narrow the type:
|
||||||
|
|
||||||
|
```python
|
||||||
|
async for part in graph.astream(input, stream_version="v2"):
|
||||||
|
if part["type"] == "values":
|
||||||
|
part["data"] # OutputT — full state (pydantic/dataclass/dict)
|
||||||
|
elif part["type"] == "messages":
|
||||||
|
part["data"] # tuple[BaseMessage, dict] — (message, metadata)
|
||||||
|
elif part["type"] == "custom":
|
||||||
|
part["data"] # Any — user-defined
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class GraphOutput(Generic[OutputT]):
|
||||||
|
"""Typed container returned by `invoke()` / `ainvoke()` with `stream_version="v2"`.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
value: The final output of the graph (dict, Pydantic model, dataclass, etc.).
|
||||||
|
interrupts: Any interrupts that occurred during execution.
|
||||||
|
"""
|
||||||
|
|
||||||
|
value: OutputT
|
||||||
|
interrupts: tuple[Interrupt, ...] = ()
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> Any:
|
||||||
|
"""Backward compat: `result['__interrupt__']` and dict-key access."""
|
||||||
|
if key == _INTERRUPT_KEY:
|
||||||
|
return self.interrupts
|
||||||
|
if isinstance(self.value, dict):
|
||||||
|
return self.value[key]
|
||||||
|
try:
|
||||||
|
return getattr(self.value, key)
|
||||||
|
except AttributeError:
|
||||||
|
raise KeyError(key)
|
||||||
|
|
||||||
|
def __contains__(self, key: object) -> bool:
|
||||||
|
if key == _INTERRUPT_KEY:
|
||||||
|
return bool(self.interrupts)
|
||||||
|
if isinstance(self.value, dict):
|
||||||
|
return key in self.value
|
||||||
|
return isinstance(key, str) and hasattr(self.value, key)
|
||||||
|
|
||||||
|
|
||||||
_DC_KWARGS = {"kw_only": True, "slots": True, "frozen": True}
|
_DC_KWARGS = {"kw_only": True, "slots": True, "frozen": True}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -5,12 +5,15 @@ from __future__ import annotations
|
|||||||
import builtins
|
import builtins
|
||||||
import warnings
|
import warnings
|
||||||
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
|
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
|
||||||
from typing import Any, overload
|
from typing import Any, Literal, overload
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from langgraph_sdk._async.http import HttpClient
|
from langgraph_sdk._async.http import HttpClient
|
||||||
from langgraph_sdk._shared.utilities import _get_run_metadata_from_response
|
from langgraph_sdk._shared.utilities import (
|
||||||
|
_get_run_metadata_from_response,
|
||||||
|
_sse_to_v2_dict,
|
||||||
|
)
|
||||||
from langgraph_sdk.schema import (
|
from langgraph_sdk.schema import (
|
||||||
All,
|
All,
|
||||||
BulkCancelRunsStatus,
|
BulkCancelRunsStatus,
|
||||||
@@ -33,9 +36,21 @@ from langgraph_sdk.schema import (
|
|||||||
RunStatus,
|
RunStatus,
|
||||||
StreamMode,
|
StreamMode,
|
||||||
StreamPart,
|
StreamPart,
|
||||||
|
StreamPartV2,
|
||||||
|
StreamVersion,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _wrap_stream_v2(
|
||||||
|
raw: AsyncIterator[StreamPart],
|
||||||
|
) -> AsyncIterator[StreamPartV2]:
|
||||||
|
"""Wrap a raw SSE stream, converting each event to a v2 dict."""
|
||||||
|
async for part in raw:
|
||||||
|
v2 = _sse_to_v2_dict(part.event, part.data)
|
||||||
|
if v2 is not None:
|
||||||
|
yield v2
|
||||||
|
|
||||||
|
|
||||||
class RunsClient:
|
class RunsClient:
|
||||||
"""Client for managing runs in LangGraph.
|
"""Client for managing runs in LangGraph.
|
||||||
|
|
||||||
@@ -81,6 +96,66 @@ class RunsClient:
|
|||||||
headers: Mapping[str, str] | None = None,
|
headers: Mapping[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
) -> AsyncIterator[StreamPart]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
thread_id: str,
|
||||||
|
assistant_id: str,
|
||||||
|
*,
|
||||||
|
input: Input | None = None,
|
||||||
|
command: Command | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] = "values",
|
||||||
|
stream_subgraphs: bool = False,
|
||||||
|
stream_resumable: bool = False,
|
||||||
|
metadata: Mapping[str, Any] | None = None,
|
||||||
|
config: Config | None = None,
|
||||||
|
context: Context | None = None,
|
||||||
|
checkpoint: Checkpoint | None = None,
|
||||||
|
checkpoint_id: str | None = None,
|
||||||
|
checkpoint_during: bool | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
feedback_keys: Sequence[str] | None = None,
|
||||||
|
on_disconnect: DisconnectMode | None = None,
|
||||||
|
webhook: str | None = None,
|
||||||
|
multitask_strategy: MultitaskStrategy | None = None,
|
||||||
|
if_not_exists: IfNotExists | None = None,
|
||||||
|
after_seconds: int | None = None,
|
||||||
|
headers: Mapping[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
|
stream_version: Literal["v2"] = "v2",
|
||||||
|
) -> AsyncIterator[StreamPartV2]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
thread_id: None,
|
||||||
|
assistant_id: str,
|
||||||
|
*,
|
||||||
|
input: Input | None = None,
|
||||||
|
command: Command | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] = "values",
|
||||||
|
stream_subgraphs: bool = False,
|
||||||
|
stream_resumable: bool = False,
|
||||||
|
metadata: Mapping[str, Any] | None = None,
|
||||||
|
config: Config | None = None,
|
||||||
|
checkpoint_during: bool | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
feedback_keys: Sequence[str] | None = None,
|
||||||
|
on_disconnect: DisconnectMode | None = None,
|
||||||
|
on_completion: OnCompletionBehavior | None = None,
|
||||||
|
if_not_exists: IfNotExists | None = None,
|
||||||
|
webhook: str | None = None,
|
||||||
|
after_seconds: int | None = None,
|
||||||
|
headers: Mapping[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
) -> AsyncIterator[StreamPart]: ...
|
) -> AsyncIterator[StreamPart]: ...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -108,7 +183,8 @@ class RunsClient:
|
|||||||
headers: Mapping[str, str] | None = None,
|
headers: Mapping[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
) -> AsyncIterator[StreamPart]: ...
|
stream_version: Literal["v2"] = "v2",
|
||||||
|
) -> AsyncIterator[StreamPartV2]: ...
|
||||||
|
|
||||||
def stream(
|
def stream(
|
||||||
self,
|
self,
|
||||||
@@ -139,7 +215,8 @@ class RunsClient:
|
|||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
durability: Durability | None = None,
|
durability: Durability | None = None,
|
||||||
) -> AsyncIterator[StreamPart]:
|
stream_version: StreamVersion = "v2",
|
||||||
|
) -> AsyncIterator[StreamPart | StreamPartV2]:
|
||||||
"""Create a run and stream the results.
|
"""Create a run and stream the results.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -180,6 +257,8 @@ class RunsClient:
|
|||||||
"async" means checkpoints are persisted async while next graph step executes, replaces checkpoint_during=True
|
"async" means checkpoints are persisted async while next graph step executes, replaces checkpoint_during=True
|
||||||
"sync" means checkpoints are persisted sync after graph step executes, replaces checkpoint_during=False
|
"sync" means checkpoints are persisted sync after graph step executes, replaces checkpoint_during=False
|
||||||
"exit" means checkpoints are only persisted when the run exits, does not save intermediate steps
|
"exit" means checkpoints are only persisted when the run exits, does not save intermediate steps
|
||||||
|
stream_version: Stream format version. "v1" (default) returns raw SSE StreamPart
|
||||||
|
NamedTuples. "v2" returns typed dicts with `type`, `ns`, and `data` keys.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Asynchronous iterator of stream results.
|
Asynchronous iterator of stream results.
|
||||||
@@ -222,7 +301,7 @@ class RunsClient:
|
|||||||
stacklevel=2,
|
stacklevel=2,
|
||||||
)
|
)
|
||||||
|
|
||||||
payload = {
|
payload: dict[str, Any] = {
|
||||||
"input": input,
|
"input": input,
|
||||||
"command": (
|
"command": (
|
||||||
{k: v for k, v in command.items() if v is not None} if command else None
|
{k: v for k, v in command.items() if v is not None} if command else None
|
||||||
@@ -259,7 +338,7 @@ class RunsClient:
|
|||||||
if on_run_created and (metadata := _get_run_metadata_from_response(res)):
|
if on_run_created and (metadata := _get_run_metadata_from_response(res)):
|
||||||
on_run_created(metadata)
|
on_run_created(metadata)
|
||||||
|
|
||||||
return self.http.stream(
|
raw = self.http.stream(
|
||||||
endpoint,
|
endpoint,
|
||||||
"POST",
|
"POST",
|
||||||
json={k: v for k, v in payload.items() if v is not None},
|
json={k: v for k, v in payload.items() if v is not None},
|
||||||
@@ -267,6 +346,9 @@ class RunsClient:
|
|||||||
headers=headers,
|
headers=headers,
|
||||||
on_response=on_response if on_run_created else None,
|
on_response=on_response if on_run_created else None,
|
||||||
)
|
)
|
||||||
|
if stream_version == "v2":
|
||||||
|
return _wrap_stream_v2(raw)
|
||||||
|
return raw
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
async def create(
|
async def create(
|
||||||
|
|||||||
@@ -107,6 +107,19 @@ def _get_run_metadata_from_response(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _sse_to_v2_dict(event: str, data: Any) -> dict[str, Any] | None:
|
||||||
|
"""Convert an SSE event+data pair into a v2 stream part dict.
|
||||||
|
|
||||||
|
Returns None for ``end`` events (signals end of stream).
|
||||||
|
"""
|
||||||
|
if event == "end":
|
||||||
|
return None
|
||||||
|
parts = event.split("|")
|
||||||
|
event_type = parts[0]
|
||||||
|
ns = parts[1:] if len(parts) > 1 else []
|
||||||
|
return {"type": event_type, "ns": ns, "data": data}
|
||||||
|
|
||||||
|
|
||||||
def _provided_vals(d: Mapping[str, Any]) -> dict[str, Any]:
|
def _provided_vals(d: Mapping[str, Any]) -> dict[str, Any]:
|
||||||
return {k: v for k, v in d.items() if v is not None}
|
return {k: v for k, v in d.items() if v is not None}
|
||||||
|
|
||||||
|
|||||||
@@ -5,11 +5,14 @@ from __future__ import annotations
|
|||||||
import builtins
|
import builtins
|
||||||
import warnings
|
import warnings
|
||||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||||
from typing import Any, overload
|
from typing import Any, Literal, overload
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from langgraph_sdk._shared.utilities import _get_run_metadata_from_response
|
from langgraph_sdk._shared.utilities import (
|
||||||
|
_get_run_metadata_from_response,
|
||||||
|
_sse_to_v2_dict,
|
||||||
|
)
|
||||||
from langgraph_sdk._sync.http import SyncHttpClient
|
from langgraph_sdk._sync.http import SyncHttpClient
|
||||||
from langgraph_sdk.schema import (
|
from langgraph_sdk.schema import (
|
||||||
All,
|
All,
|
||||||
@@ -33,9 +36,21 @@ from langgraph_sdk.schema import (
|
|||||||
RunStatus,
|
RunStatus,
|
||||||
StreamMode,
|
StreamMode,
|
||||||
StreamPart,
|
StreamPart,
|
||||||
|
StreamPartV2,
|
||||||
|
StreamVersion,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_stream_v2_sync(
|
||||||
|
raw: Iterator[StreamPart],
|
||||||
|
) -> Iterator[StreamPartV2]:
|
||||||
|
"""Wrap a raw SSE stream, converting each event to a v2 dict."""
|
||||||
|
for part in raw:
|
||||||
|
v2 = _sse_to_v2_dict(part.event, part.data)
|
||||||
|
if v2 is not None:
|
||||||
|
yield v2
|
||||||
|
|
||||||
|
|
||||||
class SyncRunsClient:
|
class SyncRunsClient:
|
||||||
"""Synchronous client for managing runs in LangGraph.
|
"""Synchronous client for managing runs in LangGraph.
|
||||||
|
|
||||||
@@ -80,6 +95,66 @@ class SyncRunsClient:
|
|||||||
headers: Mapping[str, str] | None = None,
|
headers: Mapping[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
|
) -> Iterator[StreamPart]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
thread_id: str,
|
||||||
|
assistant_id: str,
|
||||||
|
*,
|
||||||
|
input: Input | None = None,
|
||||||
|
command: Command | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] = "values",
|
||||||
|
stream_subgraphs: bool = False,
|
||||||
|
metadata: Mapping[str, Any] | None = None,
|
||||||
|
config: Config | None = None,
|
||||||
|
context: Context | None = None,
|
||||||
|
checkpoint: Checkpoint | None = None,
|
||||||
|
checkpoint_id: str | None = None,
|
||||||
|
checkpoint_during: bool | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
feedback_keys: Sequence[str] | None = None,
|
||||||
|
on_disconnect: DisconnectMode | None = None,
|
||||||
|
webhook: str | None = None,
|
||||||
|
multitask_strategy: MultitaskStrategy | None = None,
|
||||||
|
if_not_exists: IfNotExists | None = None,
|
||||||
|
after_seconds: int | None = None,
|
||||||
|
headers: Mapping[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
|
stream_version: Literal["v2"] = "v2",
|
||||||
|
) -> Iterator[StreamPartV2]: ...
|
||||||
|
|
||||||
|
@overload
|
||||||
|
def stream(
|
||||||
|
self,
|
||||||
|
thread_id: None,
|
||||||
|
assistant_id: str,
|
||||||
|
*,
|
||||||
|
input: Input | None = None,
|
||||||
|
command: Command | None = None,
|
||||||
|
stream_mode: StreamMode | Sequence[StreamMode] = "values",
|
||||||
|
stream_subgraphs: bool = False,
|
||||||
|
stream_resumable: bool = False,
|
||||||
|
metadata: Mapping[str, Any] | None = None,
|
||||||
|
config: Config | None = None,
|
||||||
|
context: Context | None = None,
|
||||||
|
checkpoint_during: bool | None = None,
|
||||||
|
interrupt_before: All | Sequence[str] | None = None,
|
||||||
|
interrupt_after: All | Sequence[str] | None = None,
|
||||||
|
feedback_keys: Sequence[str] | None = None,
|
||||||
|
on_disconnect: DisconnectMode | None = None,
|
||||||
|
on_completion: OnCompletionBehavior | None = None,
|
||||||
|
if_not_exists: IfNotExists | None = None,
|
||||||
|
webhook: str | None = None,
|
||||||
|
after_seconds: int | None = None,
|
||||||
|
headers: Mapping[str, str] | None = None,
|
||||||
|
params: QueryParamTypes | None = None,
|
||||||
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
|
stream_version: Literal["v1"],
|
||||||
) -> Iterator[StreamPart]: ...
|
) -> Iterator[StreamPart]: ...
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
@@ -108,7 +183,8 @@ class SyncRunsClient:
|
|||||||
headers: Mapping[str, str] | None = None,
|
headers: Mapping[str, str] | None = None,
|
||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
) -> Iterator[StreamPart]: ...
|
stream_version: Literal["v2"] = "v2",
|
||||||
|
) -> Iterator[StreamPartV2]: ...
|
||||||
|
|
||||||
def stream(
|
def stream(
|
||||||
self,
|
self,
|
||||||
@@ -139,7 +215,8 @@ class SyncRunsClient:
|
|||||||
params: QueryParamTypes | None = None,
|
params: QueryParamTypes | None = None,
|
||||||
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
|
||||||
durability: Durability | None = None,
|
durability: Durability | None = None,
|
||||||
) -> Iterator[StreamPart]:
|
stream_version: StreamVersion = "v2",
|
||||||
|
) -> Iterator[StreamPart | StreamPartV2]:
|
||||||
"""Create a run and stream the results.
|
"""Create a run and stream the results.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -179,7 +256,8 @@ class SyncRunsClient:
|
|||||||
"async" means checkpoints are persisted async while next graph step executes, replaces checkpoint_during=True
|
"async" means checkpoints are persisted async while next graph step executes, replaces checkpoint_during=True
|
||||||
"sync" means checkpoints are persisted sync after graph step executes, replaces checkpoint_during=False
|
"sync" means checkpoints are persisted sync after graph step executes, replaces checkpoint_during=False
|
||||||
"exit" means checkpoints are only persisted when the run exits, does not save intermediate steps
|
"exit" means checkpoints are only persisted when the run exits, does not save intermediate steps
|
||||||
|
stream_version: Stream format version. "v1" (default) returns raw SSE StreamPart
|
||||||
|
NamedTuples. "v2" returns typed dicts with `type`, `ns`, and `data` keys.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Iterator of stream results.
|
Iterator of stream results.
|
||||||
@@ -218,7 +296,7 @@ class SyncRunsClient:
|
|||||||
DeprecationWarning,
|
DeprecationWarning,
|
||||||
stacklevel=2,
|
stacklevel=2,
|
||||||
)
|
)
|
||||||
payload = {
|
payload: dict[str, Any] = {
|
||||||
"input": input,
|
"input": input,
|
||||||
"command": (
|
"command": (
|
||||||
{k: v for k, v in command.items() if v is not None} if command else None
|
{k: v for k, v in command.items() if v is not None} if command else None
|
||||||
@@ -255,7 +333,7 @@ class SyncRunsClient:
|
|||||||
if on_run_created and (metadata := _get_run_metadata_from_response(res)):
|
if on_run_created and (metadata := _get_run_metadata_from_response(res)):
|
||||||
on_run_created(metadata)
|
on_run_created(metadata)
|
||||||
|
|
||||||
return self.http.stream(
|
raw = self.http.stream(
|
||||||
endpoint,
|
endpoint,
|
||||||
"POST",
|
"POST",
|
||||||
json={k: v for k, v in payload.items() if v is not None},
|
json={k: v for k, v in payload.items() if v is not None},
|
||||||
@@ -263,6 +341,9 @@ class SyncRunsClient:
|
|||||||
headers=headers,
|
headers=headers,
|
||||||
on_response=on_response if on_run_created else None,
|
on_response=on_response if on_run_created else None,
|
||||||
)
|
)
|
||||||
|
if stream_version == "v2":
|
||||||
|
return _wrap_stream_v2_sync(raw)
|
||||||
|
return raw
|
||||||
|
|
||||||
@overload
|
@overload
|
||||||
def create(
|
def create(
|
||||||
|
|||||||
@@ -588,6 +588,206 @@ class StreamPart(NamedTuple):
|
|||||||
"""The ID of the event."""
|
"""The ID of the event."""
|
||||||
|
|
||||||
|
|
||||||
|
StreamVersion = Literal["v1", "v2"]
|
||||||
|
"""Stream format version.
|
||||||
|
|
||||||
|
- ``"v1"``: Traditional format — raw SSE ``StreamPart`` NamedTuples.
|
||||||
|
- ``"v2"``: Each event is a typed dict with ``type``, ``ns``, and ``data`` keys.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
# --- Typed payload dicts (JSON-deserialized from the server) ---
|
||||||
|
|
||||||
|
|
||||||
|
class TaskPayload(TypedDict):
|
||||||
|
"""Payload for a task start event."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
input: Any
|
||||||
|
triggers: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
class TaskResultPayload(TypedDict):
|
||||||
|
"""Payload for a task result event."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
error: str | None
|
||||||
|
interrupts: list[dict[str, Any]]
|
||||||
|
result: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointTaskPayload(TypedDict):
|
||||||
|
"""A task entry within a ``CheckpointPayload``.
|
||||||
|
|
||||||
|
The keys present depend on the task's state:
|
||||||
|
|
||||||
|
- **Error:** ``id``, ``name``, ``error``, ``state``
|
||||||
|
- **Has result:** ``id``, ``name``, ``result``, ``interrupts``, ``state``
|
||||||
|
- **Pending:** ``id``, ``name``, ``interrupts``, ``state``
|
||||||
|
"""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
name: str
|
||||||
|
error: NotRequired[str]
|
||||||
|
result: NotRequired[Any]
|
||||||
|
interrupts: NotRequired[list[dict[str, Any]]]
|
||||||
|
state: dict[str, Any] | None
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointPayload(TypedDict):
|
||||||
|
"""Payload for a checkpoint event."""
|
||||||
|
|
||||||
|
config: dict[str, Any] | None
|
||||||
|
metadata: dict[str, Any]
|
||||||
|
values: dict[str, Any]
|
||||||
|
next: list[str]
|
||||||
|
parent_config: dict[str, Any] | None
|
||||||
|
tasks: list[CheckpointTaskPayload]
|
||||||
|
|
||||||
|
|
||||||
|
class _DebugCheckpointPayload(TypedDict):
|
||||||
|
step: int
|
||||||
|
timestamp: str
|
||||||
|
type: Literal["checkpoint"]
|
||||||
|
payload: CheckpointPayload
|
||||||
|
|
||||||
|
|
||||||
|
class _DebugTaskPayload(TypedDict):
|
||||||
|
step: int
|
||||||
|
timestamp: str
|
||||||
|
type: Literal["task"]
|
||||||
|
payload: TaskPayload
|
||||||
|
|
||||||
|
|
||||||
|
class _DebugTaskResultPayload(TypedDict):
|
||||||
|
step: int
|
||||||
|
timestamp: str
|
||||||
|
type: Literal["task_result"]
|
||||||
|
payload: TaskResultPayload
|
||||||
|
|
||||||
|
|
||||||
|
DebugPayload = _DebugCheckpointPayload | _DebugTaskPayload | _DebugTaskResultPayload
|
||||||
|
"""Wrapper payload for debug events. Discriminate on ``type``."""
|
||||||
|
|
||||||
|
|
||||||
|
class RunMetadataPayload(TypedDict):
|
||||||
|
"""Payload for the ``metadata`` control event."""
|
||||||
|
|
||||||
|
run_id: str
|
||||||
|
|
||||||
|
|
||||||
|
# --- v2 stream part TypedDicts ---
|
||||||
|
|
||||||
|
|
||||||
|
class ValuesStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="values"``."""
|
||||||
|
|
||||||
|
type: Literal["values"]
|
||||||
|
ns: list[str]
|
||||||
|
data: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class UpdatesStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="updates"``."""
|
||||||
|
|
||||||
|
type: Literal["updates"]
|
||||||
|
ns: list[str]
|
||||||
|
data: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class MessagesPartialStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for partial message chunks (``messages/partial``)."""
|
||||||
|
|
||||||
|
type: Literal["messages/partial"]
|
||||||
|
ns: list[str]
|
||||||
|
data: list[dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
class MessagesCompleteStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for complete messages (``messages/complete``)."""
|
||||||
|
|
||||||
|
type: Literal["messages/complete"]
|
||||||
|
ns: list[str]
|
||||||
|
data: list[dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
class MessagesMetadataStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for message metadata (``messages/metadata``)."""
|
||||||
|
|
||||||
|
type: Literal["messages/metadata"]
|
||||||
|
ns: list[str]
|
||||||
|
data: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class MessagesTupleStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="messages"`` (raw message+metadata pair)."""
|
||||||
|
|
||||||
|
type: Literal["messages"]
|
||||||
|
ns: list[str]
|
||||||
|
data: list[dict[str, Any]]
|
||||||
|
|
||||||
|
|
||||||
|
class CustomStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="custom"``."""
|
||||||
|
|
||||||
|
type: Literal["custom"]
|
||||||
|
ns: list[str]
|
||||||
|
data: Any
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointsStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="checkpoints"``."""
|
||||||
|
|
||||||
|
type: Literal["checkpoints"]
|
||||||
|
ns: list[str]
|
||||||
|
data: CheckpointPayload
|
||||||
|
|
||||||
|
|
||||||
|
class TasksStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="tasks"``."""
|
||||||
|
|
||||||
|
type: Literal["tasks"]
|
||||||
|
ns: list[str]
|
||||||
|
data: TaskPayload | TaskResultPayload
|
||||||
|
|
||||||
|
|
||||||
|
class DebugStreamPart(TypedDict):
|
||||||
|
"""Stream part emitted for ``stream_mode="debug"``."""
|
||||||
|
|
||||||
|
type: Literal["debug"]
|
||||||
|
ns: list[str]
|
||||||
|
data: DebugPayload
|
||||||
|
|
||||||
|
|
||||||
|
class MetadataStreamPart(TypedDict):
|
||||||
|
"""Control event with ``run_id`` and other run metadata."""
|
||||||
|
|
||||||
|
type: Literal["metadata"]
|
||||||
|
ns: list[str]
|
||||||
|
data: RunMetadataPayload
|
||||||
|
|
||||||
|
|
||||||
|
StreamPartV2 = (
|
||||||
|
ValuesStreamPart
|
||||||
|
| UpdatesStreamPart
|
||||||
|
| MessagesPartialStreamPart
|
||||||
|
| MessagesCompleteStreamPart
|
||||||
|
| MessagesMetadataStreamPart
|
||||||
|
| MessagesTupleStreamPart
|
||||||
|
| CustomStreamPart
|
||||||
|
| CheckpointsStreamPart
|
||||||
|
| TasksStreamPart
|
||||||
|
| DebugStreamPart
|
||||||
|
| MetadataStreamPart
|
||||||
|
)
|
||||||
|
"""Discriminated union of all v2 stream part types.
|
||||||
|
|
||||||
|
Use ``part["type"]`` to narrow the type.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
class Send(TypedDict):
|
class Send(TypedDict):
|
||||||
"""Represents a message to be sent to a specific node in the graph.
|
"""Represents a message to be sent to a specific node in the graph.
|
||||||
|
|
||||||
|
|||||||
@@ -2,18 +2,39 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from collections.abc import Iterator, Sequence
|
from collections.abc import Iterator, Sequence
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
|
from typing_extensions import assert_type
|
||||||
|
|
||||||
|
from langgraph_sdk._shared.utilities import _sse_to_v2_dict
|
||||||
from langgraph_sdk.client import HttpClient, SyncHttpClient
|
from langgraph_sdk.client import HttpClient, SyncHttpClient
|
||||||
from langgraph_sdk.schema import StreamPart
|
from langgraph_sdk.schema import (
|
||||||
|
CheckpointPayload,
|
||||||
|
CheckpointsStreamPart,
|
||||||
|
CustomStreamPart,
|
||||||
|
DebugPayload,
|
||||||
|
DebugStreamPart,
|
||||||
|
MetadataStreamPart,
|
||||||
|
RunMetadataPayload,
|
||||||
|
StreamPart,
|
||||||
|
StreamPartV2,
|
||||||
|
TaskPayload,
|
||||||
|
TaskResultPayload,
|
||||||
|
TasksStreamPart,
|
||||||
|
UpdatesStreamPart,
|
||||||
|
ValuesStreamPart,
|
||||||
|
)
|
||||||
from langgraph_sdk.sse import BytesLike, BytesLineDecoder, SSEDecoder
|
from langgraph_sdk.sse import BytesLike, BytesLineDecoder, SSEDecoder
|
||||||
|
|
||||||
with open(Path(__file__).parent / "fixtures" / "response.txt", "rb") as f:
|
with open(Path(__file__).parent / "fixtures" / "response.txt", "rb") as f:
|
||||||
RESPONSE_PAYLOAD = f.read()
|
RESPONSE_PAYLOAD = f.read()
|
||||||
|
|
||||||
|
|
||||||
|
# --- test helpers ---
|
||||||
|
|
||||||
|
|
||||||
class AsyncListByteStream(httpx.AsyncByteStream):
|
class AsyncListByteStream(httpx.AsyncByteStream):
|
||||||
def __init__(self, chunks: Sequence[bytes], exc: Exception | None = None) -> None:
|
def __init__(self, chunks: Sequence[bytes], exc: Exception | None = None) -> None:
|
||||||
self._chunks = list(chunks)
|
self._chunks = list(chunks)
|
||||||
@@ -50,6 +71,24 @@ def iter_lines_raw(payload: list[bytes]) -> Iterator[BytesLike]:
|
|||||||
yield from decoder.flush()
|
yield from decoder.flush()
|
||||||
|
|
||||||
|
|
||||||
|
_V2_REQUIRED_KEYS = {"type", "ns", "data"}
|
||||||
|
|
||||||
|
|
||||||
|
def _assert_v2_shape(part: Any) -> None:
|
||||||
|
"""Assert a v2 stream part has the required keys and types."""
|
||||||
|
assert isinstance(part, dict), f"Expected dict, got {type(part)}"
|
||||||
|
assert part.keys() >= _V2_REQUIRED_KEYS, (
|
||||||
|
f"Missing keys: {_V2_REQUIRED_KEYS - part.keys()}"
|
||||||
|
)
|
||||||
|
assert isinstance(part["type"], str)
|
||||||
|
assert isinstance(part["ns"], list)
|
||||||
|
for elem in part["ns"]:
|
||||||
|
assert isinstance(elem, str)
|
||||||
|
|
||||||
|
|
||||||
|
# --- SSE parsing ---
|
||||||
|
|
||||||
|
|
||||||
def test_stream_sse():
|
def test_stream_sse():
|
||||||
for groups in (
|
for groups in (
|
||||||
[RESPONSE_PAYLOAD],
|
[RESPONSE_PAYLOAD],
|
||||||
@@ -69,6 +108,9 @@ def test_stream_sse():
|
|||||||
assert len(parts) == 79
|
assert len(parts) == 79
|
||||||
|
|
||||||
|
|
||||||
|
# --- HTTP client streaming ---
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_http_client_stream_flushes_trailing_event():
|
async def test_http_client_stream_flushes_trailing_event():
|
||||||
payload = b'event: foo\ndata: {"bar": 1}\n'
|
payload = b'event: foo\ndata: {"bar": 1}\n'
|
||||||
@@ -92,6 +134,26 @@ async def test_http_client_stream_flushes_trailing_event():
|
|||||||
assert parts == [StreamPart(event="foo", data={"bar": 1})]
|
assert parts == [StreamPart(event="foo", data={"bar": 1})]
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_http_client_stream_flushes_trailing_event():
|
||||||
|
payload = b'event: foo\ndata: {"bar": 1}\n'
|
||||||
|
|
||||||
|
def handler(request: httpx.Request) -> httpx.Response:
|
||||||
|
assert request.headers["accept"] == "text/event-stream"
|
||||||
|
assert request.headers["cache-control"] == "no-store"
|
||||||
|
return httpx.Response(
|
||||||
|
200,
|
||||||
|
headers={"Content-Type": "text/event-stream"},
|
||||||
|
content=payload,
|
||||||
|
)
|
||||||
|
|
||||||
|
transport = httpx.MockTransport(handler)
|
||||||
|
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||||
|
http_client = SyncHttpClient(client)
|
||||||
|
parts = list(http_client.stream("/stream", "GET"))
|
||||||
|
|
||||||
|
assert parts == [StreamPart(event="foo", data={"bar": 1})]
|
||||||
|
|
||||||
|
|
||||||
def test_sync_http_client_stream_recovers_after_disconnect():
|
def test_sync_http_client_stream_recovers_after_disconnect():
|
||||||
reconnect_path = "/reconnect"
|
reconnect_path = "/reconnect"
|
||||||
first_chunks = [
|
first_chunks = [
|
||||||
@@ -228,21 +290,138 @@ async def test_http_client_stream_recovers_after_disconnect():
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_sync_http_client_stream_flushes_trailing_event():
|
# --- _sse_to_v2_dict conversion ---
|
||||||
payload = b'event: foo\ndata: {"bar": 1}\n'
|
|
||||||
|
|
||||||
def handler(request: httpx.Request) -> httpx.Response:
|
|
||||||
assert request.headers["accept"] == "text/event-stream"
|
def test_sse_to_v2_dict_basic() -> None:
|
||||||
assert request.headers["cache-control"] == "no-store"
|
result = _sse_to_v2_dict("values", {"messages": [{"role": "user"}]})
|
||||||
return httpx.Response(
|
assert result is not None
|
||||||
200,
|
_assert_v2_shape(result)
|
||||||
headers={"Content-Type": "text/event-stream"},
|
assert result == {
|
||||||
content=payload,
|
"type": "values",
|
||||||
|
"ns": [],
|
||||||
|
"data": {"messages": [{"role": "user"}]},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_sse_to_v2_dict_with_namespace() -> None:
|
||||||
|
result = _sse_to_v2_dict("updates|sub:abc", {"key": "val"})
|
||||||
|
assert result is not None
|
||||||
|
_assert_v2_shape(result)
|
||||||
|
assert result == {
|
||||||
|
"type": "updates",
|
||||||
|
"ns": ["sub:abc"],
|
||||||
|
"data": {"key": "val"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_sse_to_v2_dict_with_multiple_ns() -> None:
|
||||||
|
result = _sse_to_v2_dict("custom|parent|child:123", "hello")
|
||||||
|
assert result is not None
|
||||||
|
_assert_v2_shape(result)
|
||||||
|
assert result == {
|
||||||
|
"type": "custom",
|
||||||
|
"ns": ["parent", "child:123"],
|
||||||
|
"data": "hello",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_sse_to_v2_dict_end_event() -> None:
|
||||||
|
assert _sse_to_v2_dict("end", None) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_sse_to_v2_dict_metadata_event() -> None:
|
||||||
|
result = _sse_to_v2_dict("metadata", {"run_id": "abc-123"})
|
||||||
|
assert result is not None
|
||||||
|
_assert_v2_shape(result)
|
||||||
|
assert result == {
|
||||||
|
"type": "metadata",
|
||||||
|
"ns": [],
|
||||||
|
"data": {"run_id": "abc-123"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_sse_to_v2_dict_messages_partial() -> None:
|
||||||
|
result = _sse_to_v2_dict("messages/partial", [{"type": "ai", "content": "hi"}])
|
||||||
|
assert result is not None
|
||||||
|
_assert_v2_shape(result)
|
||||||
|
assert result == {
|
||||||
|
"type": "messages/partial",
|
||||||
|
"ns": [],
|
||||||
|
"data": [{"type": "ai", "content": "hi"}],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# --- client-side v2 stream wrapping ---
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_async_stream_v2_client_side_conversion() -> None:
|
||||||
|
from langgraph_sdk._async.runs import _wrap_stream_v2
|
||||||
|
|
||||||
|
async def mock_stream() -> Any:
|
||||||
|
yield StreamPart(event="metadata", data={"run_id": "r1"})
|
||||||
|
yield StreamPart(
|
||||||
|
event="values", data={"messages": [{"role": "user", "content": "hi"}]}
|
||||||
)
|
)
|
||||||
|
yield StreamPart(event="updates|sub:abc", data={"node": {"out": 1}})
|
||||||
|
yield StreamPart(event="end", data=None) # type: ignore[arg-type]
|
||||||
|
|
||||||
transport = httpx.MockTransport(handler)
|
parts: list[StreamPartV2] = [part async for part in _wrap_stream_v2(mock_stream())]
|
||||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
assert len(parts) == 3
|
||||||
http_client = SyncHttpClient(client)
|
for part in parts:
|
||||||
parts = list(http_client.stream("/stream", "GET"))
|
_assert_v2_shape(part)
|
||||||
|
assert parts[0] == {"type": "metadata", "ns": [], "data": {"run_id": "r1"}}
|
||||||
|
assert parts[1] == {
|
||||||
|
"type": "values",
|
||||||
|
"ns": [],
|
||||||
|
"data": {"messages": [{"role": "user", "content": "hi"}]},
|
||||||
|
}
|
||||||
|
assert parts[2] == {
|
||||||
|
"type": "updates",
|
||||||
|
"ns": ["sub:abc"],
|
||||||
|
"data": {"node": {"out": 1}},
|
||||||
|
}
|
||||||
|
|
||||||
assert parts == [StreamPart(event="foo", data={"bar": 1})]
|
|
||||||
|
def test_sync_stream_v2_client_side_conversion() -> None:
|
||||||
|
from langgraph_sdk._sync.runs import _wrap_stream_v2_sync
|
||||||
|
|
||||||
|
def mock_stream() -> Any:
|
||||||
|
yield StreamPart(event="metadata", data={"run_id": "r1"})
|
||||||
|
yield StreamPart(event="values", data={"state": "full"})
|
||||||
|
yield StreamPart(event="end", data=None) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
parts: list[StreamPartV2] = list(_wrap_stream_v2_sync(mock_stream()))
|
||||||
|
assert len(parts) == 2
|
||||||
|
for part in parts:
|
||||||
|
_assert_v2_shape(part)
|
||||||
|
assert parts[0] == {"type": "metadata", "ns": [], "data": {"run_id": "r1"}}
|
||||||
|
assert parts[1] == {"type": "values", "ns": [], "data": {"state": "full"}}
|
||||||
|
|
||||||
|
|
||||||
|
# --- type narrowing compile-time checks ---
|
||||||
|
|
||||||
|
|
||||||
|
def _check_v2_type_narrowing(part: StreamPartV2) -> None:
|
||||||
|
"""Compile-time type narrowing checks — validates mypy narrows the union."""
|
||||||
|
if part["type"] == "values":
|
||||||
|
assert_type(part, ValuesStreamPart)
|
||||||
|
assert_type(part["data"], dict[str, Any])
|
||||||
|
elif part["type"] == "updates":
|
||||||
|
assert_type(part, UpdatesStreamPart)
|
||||||
|
assert_type(part["data"], dict[str, Any])
|
||||||
|
elif part["type"] == "custom":
|
||||||
|
assert_type(part, CustomStreamPart)
|
||||||
|
elif part["type"] == "checkpoints":
|
||||||
|
assert_type(part, CheckpointsStreamPart)
|
||||||
|
assert_type(part["data"], CheckpointPayload)
|
||||||
|
elif part["type"] == "tasks":
|
||||||
|
assert_type(part, TasksStreamPart)
|
||||||
|
assert_type(part["data"], TaskPayload | TaskResultPayload)
|
||||||
|
elif part["type"] == "debug":
|
||||||
|
assert_type(part, DebugStreamPart)
|
||||||
|
assert_type(part["data"], DebugPayload)
|
||||||
|
elif part["type"] == "metadata":
|
||||||
|
assert_type(part, MetadataStreamPart)
|
||||||
|
assert_type(part["data"], RunMetadataPayload)
|
||||||
|
|||||||
Reference in New Issue
Block a user