mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-30 11:49:38 +02:00
wip
This commit is contained in:
@@ -1020,6 +1020,99 @@ class RemoteGraph(PregelProtocol):
|
||||
else:
|
||||
yield chunk
|
||||
|
||||
def stream_v2(
|
||||
self,
|
||||
input: dict[str, Any] | Command | None,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
transformers: Sequence[Any] | None = None,
|
||||
stream_modes: Sequence[StreamMode] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""Start a sync v2 remote run driven by transformer projections."""
|
||||
from langgraph.pregel.main import (
|
||||
_build_stream_factories,
|
||||
_collect_stream_modes,
|
||||
_merge_v2_messages_flag,
|
||||
)
|
||||
from langgraph.stream._convert import convert_to_protocol_event
|
||||
from langgraph.stream._mux import StreamMux
|
||||
from langgraph.stream.run_stream import RemoteGraphRunStream
|
||||
|
||||
factories = _build_stream_factories((), transformers)
|
||||
mux = StreamMux(factories=factories, is_async=False)
|
||||
requested_stream_modes = set(_collect_stream_modes(mux))
|
||||
requested_stream_modes.update(stream_modes or ())
|
||||
remote_iter = (
|
||||
convert_to_protocol_event(part)
|
||||
for part in self.stream(
|
||||
input,
|
||||
_merge_v2_messages_flag(config),
|
||||
context=context,
|
||||
stream_mode=list(requested_stream_modes),
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
subgraphs=True,
|
||||
headers=headers,
|
||||
params=params,
|
||||
version="v2",
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
return RemoteGraphRunStream(iter(remote_iter), mux)
|
||||
|
||||
async def astream_v2(
|
||||
self,
|
||||
input: dict[str, Any] | Command | None,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
context: Context | None = None,
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
transformers: Sequence[Any] | None = None,
|
||||
stream_modes: Sequence[StreamMode] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""Async counterpart to ``stream_v2`` for remote graphs."""
|
||||
from langgraph.pregel.main import (
|
||||
_build_stream_factories,
|
||||
_collect_stream_modes,
|
||||
_merge_v2_messages_flag,
|
||||
)
|
||||
from langgraph.stream._convert import convert_to_protocol_event
|
||||
from langgraph.stream._mux import StreamMux
|
||||
from langgraph.stream.run_stream import AsyncRemoteGraphRunStream
|
||||
|
||||
factories = _build_stream_factories((), transformers)
|
||||
mux = StreamMux(factories=factories, is_async=True)
|
||||
requested_stream_modes = set(_collect_stream_modes(mux))
|
||||
requested_stream_modes.update(stream_modes or ())
|
||||
|
||||
async def remote_events() -> AsyncIterator[Any]:
|
||||
async for part in self.astream(
|
||||
input,
|
||||
_merge_v2_messages_flag(config),
|
||||
context=context,
|
||||
stream_mode=list(requested_stream_modes),
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
subgraphs=True,
|
||||
headers=headers,
|
||||
params=params,
|
||||
version="v2",
|
||||
**kwargs,
|
||||
):
|
||||
yield convert_to_protocol_event(part)
|
||||
|
||||
return AsyncRemoteGraphRunStream(remote_events().__aiter__(), mux)
|
||||
|
||||
async def astream_events(
|
||||
self,
|
||||
input: Any,
|
||||
|
||||
@@ -7,13 +7,20 @@ graph's raw events into ergonomic per-channel streams.
|
||||
|
||||
from langgraph.stream._event_log import EventLog
|
||||
from langgraph.stream._types import ProtocolEvent, StreamTransformer
|
||||
from langgraph.stream.run_stream import AsyncGraphRunStream, GraphRunStream
|
||||
from langgraph.stream.run_stream import (
|
||||
AsyncGraphRunStream,
|
||||
AsyncRemoteGraphRunStream,
|
||||
GraphRunStream,
|
||||
RemoteGraphRunStream,
|
||||
)
|
||||
from langgraph.stream.stream_channel import StreamChannel
|
||||
|
||||
__all__ = [
|
||||
"AsyncGraphRunStream",
|
||||
"AsyncRemoteGraphRunStream",
|
||||
"EventLog",
|
||||
"GraphRunStream",
|
||||
"RemoteGraphRunStream",
|
||||
"ProtocolEvent",
|
||||
"StreamChannel",
|
||||
"StreamTransformer",
|
||||
|
||||
@@ -410,3 +410,62 @@ class AsyncGraphRunStream(BaseRunStream):
|
||||
if (err := self._values_transformer.error) is not None:
|
||||
raise err
|
||||
return self._values_transformer._interrupts
|
||||
|
||||
|
||||
class RemoteGraphRunStream(GraphRunStream):
|
||||
"""Sync run stream fed by already-normalized remote protocol events."""
|
||||
|
||||
def __init__(self, events: Iterator[ProtocolEvent], mux: StreamMux) -> None:
|
||||
super().__init__(events, mux)
|
||||
|
||||
def _pump_next(self) -> bool:
|
||||
"""Pull one remote protocol event and push it through the mux."""
|
||||
if self._exhausted:
|
||||
return False
|
||||
try:
|
||||
event = next(self._graph_iter)
|
||||
except StopIteration:
|
||||
self._mux.close()
|
||||
self._exhausted = True
|
||||
return False
|
||||
except Exception as e:
|
||||
self._mux.fail(e)
|
||||
self._exhausted = True
|
||||
return False
|
||||
self._mux.push(event)
|
||||
return True
|
||||
|
||||
|
||||
class AsyncRemoteGraphRunStream(AsyncGraphRunStream):
|
||||
"""Async run stream fed by already-normalized remote protocol events."""
|
||||
|
||||
def __init__(self, events: AsyncIterator[ProtocolEvent], mux: StreamMux) -> None:
|
||||
super().__init__(events, mux)
|
||||
|
||||
async def _apump_next(self) -> bool:
|
||||
"""Pull one remote protocol event and push it through the mux."""
|
||||
async with self._pump_cond:
|
||||
if self._exhausted:
|
||||
return False
|
||||
if self._pumping:
|
||||
await self._pump_cond.wait()
|
||||
return not self._exhausted
|
||||
self._pumping = True
|
||||
|
||||
try:
|
||||
try:
|
||||
event = await self._graph_aiter.__anext__()
|
||||
except StopAsyncIteration:
|
||||
self._exhausted = True
|
||||
await self._mux.aclose()
|
||||
return False
|
||||
except Exception as e:
|
||||
self._exhausted = True
|
||||
await self._mux.afail(e)
|
||||
return False
|
||||
await self._mux.apush(event)
|
||||
return True
|
||||
finally:
|
||||
async with self._pump_cond:
|
||||
self._pumping = False
|
||||
self._pump_cond.notify_all()
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator, Iterator, Sequence
|
||||
from typing import Any, Literal
|
||||
|
||||
import pytest
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
from langgraph.pregel.remote import RemoteGraph
|
||||
from langgraph.types import All, StreamMode
|
||||
|
||||
|
||||
class _FakeRemoteGraph(RemoteGraph):
|
||||
def __init__(self) -> None:
|
||||
super().__init__("agent", url="http://unused")
|
||||
self.last_stream_modes: list[StreamMode] | None = None
|
||||
|
||||
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,
|
||||
version: Literal["v2"],
|
||||
**kwargs: Any,
|
||||
) -> Iterator[dict[str, Any]]:
|
||||
assert version == "v2"
|
||||
assert subgraphs is True
|
||||
self.last_stream_modes = (
|
||||
[stream_mode] if isinstance(stream_mode, str) else list(stream_mode or [])
|
||||
)
|
||||
yield {
|
||||
"type": "values",
|
||||
"ns": (),
|
||||
"data": {"value": input["value"] + "A"},
|
||||
"interrupts": (),
|
||||
}
|
||||
yield {
|
||||
"type": "values",
|
||||
"ns": (),
|
||||
"data": {"value": input["value"] + "AB"},
|
||||
"interrupts": (),
|
||||
}
|
||||
|
||||
async 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,
|
||||
version: Literal["v2"],
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
assert version == "v2"
|
||||
assert subgraphs is True
|
||||
self.last_stream_modes = (
|
||||
[stream_mode] if isinstance(stream_mode, str) else list(stream_mode or [])
|
||||
)
|
||||
yield {
|
||||
"type": "values",
|
||||
"ns": (),
|
||||
"data": {"value": input["value"] + "A"},
|
||||
"interrupts": (),
|
||||
}
|
||||
yield {
|
||||
"type": "values",
|
||||
"ns": (),
|
||||
"data": {"value": input["value"] + "AB"},
|
||||
"interrupts": (),
|
||||
}
|
||||
|
||||
|
||||
def test_remote_stream_v2_values_and_output() -> None:
|
||||
remote = _FakeRemoteGraph()
|
||||
run = remote.stream_v2({"value": "x"})
|
||||
|
||||
assert list(run.values) == [{"value": "xA"}, {"value": "xAB"}]
|
||||
assert "values" in (remote.last_stream_modes or [])
|
||||
|
||||
|
||||
def test_remote_stream_v2_output_drains_remote_events() -> None:
|
||||
remote = _FakeRemoteGraph()
|
||||
run = remote.stream_v2({"value": "x"})
|
||||
|
||||
assert run.output == {"value": "xAB"}
|
||||
|
||||
|
||||
def test_remote_stream_v2_raw_events() -> None:
|
||||
remote = _FakeRemoteGraph()
|
||||
run = remote.stream_v2({"value": "x"})
|
||||
|
||||
events = list(run)
|
||||
assert [event["method"] for event in events] == ["values", "values"]
|
||||
assert [event["seq"] for event in events] == [1, 2]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_remote_astream_v2_values_and_output() -> None:
|
||||
remote = _FakeRemoteGraph()
|
||||
run = await remote.astream_v2({"value": "x"})
|
||||
|
||||
assert [item async for item in run.values] == [
|
||||
{"value": "xA"},
|
||||
{"value": "xAB"},
|
||||
]
|
||||
assert "values" in (remote.last_stream_modes or [])
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_remote_astream_v2_output_drains_remote_events() -> None:
|
||||
remote = _FakeRemoteGraph()
|
||||
run = await remote.astream_v2({"value": "x"})
|
||||
|
||||
assert await run.output() == {"value": "xAB"}
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
"""Async thread-centric streaming primitives."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
from langgraph_sdk._async.http import HttpClient
|
||||
from langgraph_sdk.protocol import (
|
||||
Channel,
|
||||
Command,
|
||||
CommandMethod,
|
||||
CommandResponse,
|
||||
ErrorResponse,
|
||||
Event,
|
||||
SubscribeParams,
|
||||
)
|
||||
from langgraph_sdk.schema import QueryParamTypes, StreamPart
|
||||
|
||||
|
||||
def _stream_part_to_event(part: StreamPart) -> Event:
|
||||
"""Normalize an SSE ``StreamPart`` into a protocol event envelope."""
|
||||
if isinstance(part.data, dict) and part.data.get("type") == "event":
|
||||
event = cast(Event, part.data)
|
||||
else:
|
||||
event = {
|
||||
"type": "event",
|
||||
"method": part.event,
|
||||
"params": {
|
||||
"namespace": [],
|
||||
"timestamp": int(time.time() * 1000),
|
||||
"data": part.data,
|
||||
},
|
||||
}
|
||||
if part.id is not None and "event_id" not in event:
|
||||
event["event_id"] = part.id
|
||||
return event
|
||||
|
||||
|
||||
class EventSubscription:
|
||||
"""Async iterable handle for a filtered event subscription."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
subscription_id: str,
|
||||
params: SubscribeParams,
|
||||
events: AsyncIterator[Event],
|
||||
on_unsubscribe: Any,
|
||||
) -> None:
|
||||
self.subscription_id = subscription_id
|
||||
self.params = params
|
||||
self._events = events
|
||||
self._on_unsubscribe = on_unsubscribe
|
||||
|
||||
def __aiter__(self) -> AsyncIterator[Event]:
|
||||
return self._events
|
||||
|
||||
async def unsubscribe(self) -> None:
|
||||
await self._on_unsubscribe(self.subscription_id)
|
||||
|
||||
|
||||
class ProtocolSseTransport:
|
||||
"""SSE transport for the thread-centric protocol."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
http: HttpClient,
|
||||
thread_id: str,
|
||||
*,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
) -> None:
|
||||
self.http = http
|
||||
self.thread_id = thread_id
|
||||
self.headers = headers
|
||||
self.params = params
|
||||
self.commands_path = f"/v2/threads/{thread_id}/commands"
|
||||
self.events_path = f"/v2/threads/{thread_id}/events"
|
||||
|
||||
async def send(self, command: Command) -> CommandResponse | ErrorResponse | None:
|
||||
return await self.http.post(
|
||||
self.commands_path,
|
||||
json=cast(dict[str, Any], command),
|
||||
headers=self.headers,
|
||||
params=self.params,
|
||||
)
|
||||
|
||||
def open_event_stream(self, params: SubscribeParams) -> AsyncIterator[Event]:
|
||||
async def iterate() -> AsyncIterator[Event]:
|
||||
async for part in self.http.stream(
|
||||
self.events_path,
|
||||
"POST",
|
||||
json=cast(dict[str, Any], params),
|
||||
headers=self.headers,
|
||||
params=self.params,
|
||||
):
|
||||
yield _stream_part_to_event(part)
|
||||
|
||||
return iterate()
|
||||
|
||||
|
||||
class RunModule:
|
||||
"""Run commands exposed by ``ThreadStream.run``."""
|
||||
|
||||
def __init__(self, stream: ThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
async def input(self, params: Mapping[str, Any]) -> Any:
|
||||
return await self._stream.command("run.input", dict(params))
|
||||
|
||||
|
||||
class InputModule:
|
||||
"""Human-input commands exposed by ``ThreadStream.input``."""
|
||||
|
||||
def __init__(self, stream: ThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
async def respond(self, params: Mapping[str, Any]) -> Any:
|
||||
return await self._stream.command("input.respond", dict(params))
|
||||
|
||||
async def inject(self, params: Mapping[str, Any]) -> Any:
|
||||
return await self._stream.command("input.inject", dict(params))
|
||||
|
||||
|
||||
class StateModule:
|
||||
"""State commands exposed by ``ThreadStream.state``."""
|
||||
|
||||
def __init__(self, stream: ThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
async def get(self, params: Mapping[str, Any] | None = None) -> Any:
|
||||
return await self._stream.command("state.get", dict(params or {}))
|
||||
|
||||
async def list_checkpoints(self, params: Mapping[str, Any] | None = None) -> Any:
|
||||
return await self._stream.command("state.listCheckpoints", dict(params or {}))
|
||||
|
||||
async def fork(self, params: Mapping[str, Any]) -> Any:
|
||||
return await self._stream.command("state.fork", dict(params))
|
||||
|
||||
|
||||
class AgentModule:
|
||||
"""Agent commands exposed by ``ThreadStream.agent``."""
|
||||
|
||||
def __init__(self, stream: ThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
async def get_tree(self, params: Mapping[str, Any] | None = None) -> Any:
|
||||
return await self._stream.command("agent.getTree", dict(params or {}))
|
||||
|
||||
|
||||
class ThreadStream:
|
||||
"""High-level async wrapper around a thread protocol transport."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transport: ProtocolSseTransport,
|
||||
*,
|
||||
assistant_id: str,
|
||||
starting_command_id: int = 0,
|
||||
) -> None:
|
||||
if not assistant_id:
|
||||
raise ValueError("assistant_id is required")
|
||||
self.transport = transport
|
||||
self.assistant_id = assistant_id
|
||||
self._next_command_id = starting_command_id
|
||||
self._next_subscription_id = 0
|
||||
self.run = RunModule(self)
|
||||
self.input = InputModule(self)
|
||||
self.state = StateModule(self)
|
||||
self.agent = AgentModule(self)
|
||||
|
||||
@property
|
||||
def thread_id(self) -> str:
|
||||
return self.transport.thread_id
|
||||
|
||||
def _command_id(self) -> str:
|
||||
self._next_command_id += 1
|
||||
return str(self._next_command_id)
|
||||
|
||||
def _subscription_id(self) -> str:
|
||||
self._next_subscription_id += 1
|
||||
return f"sub-{self._next_subscription_id}"
|
||||
|
||||
async def command(self, method: CommandMethod, params: dict[str, Any]) -> Any:
|
||||
command: Command = {
|
||||
"id": self._command_id(),
|
||||
"method": method,
|
||||
"params": params,
|
||||
}
|
||||
if method == "run.input":
|
||||
command["params"] = {"assistant_id": self.assistant_id, **params}
|
||||
response = await self.transport.send(command)
|
||||
if response is None:
|
||||
return None
|
||||
if "error" in response:
|
||||
raise RuntimeError(response["error"])
|
||||
return response.get("result")
|
||||
|
||||
async def subscribe(
|
||||
self,
|
||||
channels: Sequence[Channel | str] | SubscribeParams,
|
||||
*,
|
||||
namespaces: Sequence[Sequence[str]] | None = None,
|
||||
depth: int | None = None,
|
||||
) -> EventSubscription:
|
||||
if isinstance(channels, dict):
|
||||
params = SubscribeParams(**channels)
|
||||
else:
|
||||
params = SubscribeParams(channels=list(channels))
|
||||
if namespaces is not None:
|
||||
params["namespaces"] = [list(ns) for ns in namespaces]
|
||||
if depth is not None:
|
||||
params["depth"] = depth
|
||||
|
||||
subscription_id = self._subscription_id()
|
||||
events = self.transport.open_event_stream(params)
|
||||
return EventSubscription(subscription_id, params, events, self._unsubscribe)
|
||||
|
||||
async def _unsubscribe(self, subscription_id: str) -> None:
|
||||
await self.command(
|
||||
"subscription.unsubscribe",
|
||||
{"subscription_id": subscription_id},
|
||||
)
|
||||
|
||||
async def close(self) -> None:
|
||||
return None
|
||||
|
||||
@@ -6,6 +6,7 @@ from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk._async.http import HttpClient
|
||||
from langgraph_sdk._async.stream import ProtocolSseTransport, ThreadStream
|
||||
from langgraph_sdk.schema import (
|
||||
Checkpoint,
|
||||
Json,
|
||||
@@ -42,6 +43,28 @@ class ThreadsClient:
|
||||
def __init__(self, http: HttpClient) -> None:
|
||||
self.http = http
|
||||
|
||||
def stream(
|
||||
self,
|
||||
thread_id: str,
|
||||
*,
|
||||
assistant_id: str,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
) -> ThreadStream:
|
||||
"""Create a thread-centric protocol stream.
|
||||
|
||||
This mirrors the JavaScript SDK's ``client.threads.stream(...)`` API
|
||||
and is intentionally separate from ``join_stream()``, which follows
|
||||
the older thread event stream endpoint.
|
||||
"""
|
||||
transport = ProtocolSseTransport(
|
||||
self.http,
|
||||
thread_id,
|
||||
headers=headers,
|
||||
params=params,
|
||||
)
|
||||
return ThreadStream(transport, assistant_id=assistant_id)
|
||||
|
||||
async def get(
|
||||
self,
|
||||
thread_id: str,
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
"""Synchronous thread-centric streaming primitives."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from typing import Any, cast
|
||||
|
||||
from langgraph_sdk._sync.http import SyncHttpClient
|
||||
from langgraph_sdk.protocol import (
|
||||
Channel,
|
||||
Command,
|
||||
CommandMethod,
|
||||
CommandResponse,
|
||||
ErrorResponse,
|
||||
Event,
|
||||
SubscribeParams,
|
||||
)
|
||||
from langgraph_sdk.schema import QueryParamTypes, StreamPart
|
||||
|
||||
|
||||
def _stream_part_to_event(part: StreamPart) -> Event:
|
||||
"""Normalize an SSE ``StreamPart`` into a protocol event envelope."""
|
||||
if isinstance(part.data, dict) and part.data.get("type") == "event":
|
||||
event = cast(Event, part.data)
|
||||
else:
|
||||
event = {
|
||||
"type": "event",
|
||||
"method": part.event,
|
||||
"params": {
|
||||
"namespace": [],
|
||||
"timestamp": int(time.time() * 1000),
|
||||
"data": part.data,
|
||||
},
|
||||
}
|
||||
if part.id is not None and "event_id" not in event:
|
||||
event["event_id"] = part.id
|
||||
return event
|
||||
|
||||
|
||||
class SyncEventSubscription:
|
||||
"""Iterator handle for a filtered event subscription."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
subscription_id: str,
|
||||
params: SubscribeParams,
|
||||
events: Iterator[Event],
|
||||
on_unsubscribe: Any,
|
||||
) -> None:
|
||||
self.subscription_id = subscription_id
|
||||
self.params = params
|
||||
self._events = events
|
||||
self._on_unsubscribe = on_unsubscribe
|
||||
|
||||
def __iter__(self) -> Iterator[Event]:
|
||||
return self._events
|
||||
|
||||
def unsubscribe(self) -> None:
|
||||
self._on_unsubscribe(self.subscription_id)
|
||||
|
||||
|
||||
class SyncProtocolSseTransport:
|
||||
"""SSE transport for the thread-centric protocol."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
http: SyncHttpClient,
|
||||
thread_id: str,
|
||||
*,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
) -> None:
|
||||
self.http = http
|
||||
self.thread_id = thread_id
|
||||
self.headers = headers
|
||||
self.params = params
|
||||
self.commands_path = f"/v2/threads/{thread_id}/commands"
|
||||
self.events_path = f"/v2/threads/{thread_id}/events"
|
||||
|
||||
def send(self, command: Command) -> CommandResponse | ErrorResponse | None:
|
||||
return self.http.post(
|
||||
self.commands_path,
|
||||
json=cast(dict[str, Any], command),
|
||||
headers=self.headers,
|
||||
params=self.params,
|
||||
)
|
||||
|
||||
def open_event_stream(self, params: SubscribeParams) -> Iterator[Event]:
|
||||
for part in self.http.stream(
|
||||
self.events_path,
|
||||
"POST",
|
||||
json=cast(dict[str, Any], params),
|
||||
headers=self.headers,
|
||||
params=self.params,
|
||||
):
|
||||
yield _stream_part_to_event(part)
|
||||
|
||||
|
||||
class SyncRunModule:
|
||||
"""Run commands exposed by ``SyncThreadStream.run``."""
|
||||
|
||||
def __init__(self, stream: SyncThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
def input(self, params: Mapping[str, Any]) -> Any:
|
||||
return self._stream.command("run.input", dict(params))
|
||||
|
||||
|
||||
class SyncInputModule:
|
||||
"""Human-input commands exposed by ``SyncThreadStream.input``."""
|
||||
|
||||
def __init__(self, stream: SyncThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
def respond(self, params: Mapping[str, Any]) -> Any:
|
||||
return self._stream.command("input.respond", dict(params))
|
||||
|
||||
def inject(self, params: Mapping[str, Any]) -> Any:
|
||||
return self._stream.command("input.inject", dict(params))
|
||||
|
||||
|
||||
class SyncStateModule:
|
||||
"""State commands exposed by ``SyncThreadStream.state``."""
|
||||
|
||||
def __init__(self, stream: SyncThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
def get(self, params: Mapping[str, Any] | None = None) -> Any:
|
||||
return self._stream.command("state.get", dict(params or {}))
|
||||
|
||||
def list_checkpoints(self, params: Mapping[str, Any] | None = None) -> Any:
|
||||
return self._stream.command("state.listCheckpoints", dict(params or {}))
|
||||
|
||||
def fork(self, params: Mapping[str, Any]) -> Any:
|
||||
return self._stream.command("state.fork", dict(params))
|
||||
|
||||
|
||||
class SyncAgentModule:
|
||||
"""Agent commands exposed by ``SyncThreadStream.agent``."""
|
||||
|
||||
def __init__(self, stream: SyncThreadStream) -> None:
|
||||
self._stream = stream
|
||||
|
||||
def get_tree(self, params: Mapping[str, Any] | None = None) -> Any:
|
||||
return self._stream.command("agent.getTree", dict(params or {}))
|
||||
|
||||
|
||||
class SyncThreadStream:
|
||||
"""High-level sync wrapper around a thread protocol transport."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transport: SyncProtocolSseTransport,
|
||||
*,
|
||||
assistant_id: str,
|
||||
starting_command_id: int = 0,
|
||||
) -> None:
|
||||
if not assistant_id:
|
||||
raise ValueError("assistant_id is required")
|
||||
self.transport = transport
|
||||
self.assistant_id = assistant_id
|
||||
self._next_command_id = starting_command_id
|
||||
self._next_subscription_id = 0
|
||||
self.run = SyncRunModule(self)
|
||||
self.input = SyncInputModule(self)
|
||||
self.state = SyncStateModule(self)
|
||||
self.agent = SyncAgentModule(self)
|
||||
|
||||
@property
|
||||
def thread_id(self) -> str:
|
||||
return self.transport.thread_id
|
||||
|
||||
def _command_id(self) -> str:
|
||||
self._next_command_id += 1
|
||||
return str(self._next_command_id)
|
||||
|
||||
def _subscription_id(self) -> str:
|
||||
self._next_subscription_id += 1
|
||||
return f"sub-{self._next_subscription_id}"
|
||||
|
||||
def command(self, method: CommandMethod, params: dict[str, Any]) -> Any:
|
||||
command: Command = {
|
||||
"id": self._command_id(),
|
||||
"method": method,
|
||||
"params": params,
|
||||
}
|
||||
if method == "run.input":
|
||||
command["params"] = {"assistant_id": self.assistant_id, **params}
|
||||
response = self.transport.send(command)
|
||||
if response is None:
|
||||
return None
|
||||
if "error" in response:
|
||||
raise RuntimeError(response["error"])
|
||||
return response.get("result")
|
||||
|
||||
def subscribe(
|
||||
self,
|
||||
channels: Sequence[Channel | str] | SubscribeParams,
|
||||
*,
|
||||
namespaces: Sequence[Sequence[str]] | None = None,
|
||||
depth: int | None = None,
|
||||
) -> SyncEventSubscription:
|
||||
if isinstance(channels, dict):
|
||||
params = SubscribeParams(**channels)
|
||||
else:
|
||||
params = SubscribeParams(channels=list(channels))
|
||||
if namespaces is not None:
|
||||
params["namespaces"] = [list(ns) for ns in namespaces]
|
||||
if depth is not None:
|
||||
params["depth"] = depth
|
||||
|
||||
subscription_id = self._subscription_id()
|
||||
events = self.transport.open_event_stream(params)
|
||||
return SyncEventSubscription(subscription_id, params, events, self._unsubscribe)
|
||||
|
||||
def _unsubscribe(self, subscription_id: str) -> None:
|
||||
self.command(
|
||||
"subscription.unsubscribe",
|
||||
{"subscription_id": subscription_id},
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
return None
|
||||
|
||||
@@ -6,6 +6,7 @@ from collections.abc import Iterator, Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk._sync.http import SyncHttpClient
|
||||
from langgraph_sdk._sync.stream import SyncProtocolSseTransport, SyncThreadStream
|
||||
from langgraph_sdk.schema import (
|
||||
Checkpoint,
|
||||
Json,
|
||||
@@ -41,6 +42,28 @@ class SyncThreadsClient:
|
||||
def __init__(self, http: SyncHttpClient) -> None:
|
||||
self.http = http
|
||||
|
||||
def stream(
|
||||
self,
|
||||
thread_id: str,
|
||||
*,
|
||||
assistant_id: str,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
params: QueryParamTypes | None = None,
|
||||
) -> SyncThreadStream:
|
||||
"""Create a thread-centric protocol stream.
|
||||
|
||||
This mirrors the JavaScript SDK's ``client.threads.stream(...)`` API
|
||||
and is intentionally separate from ``join_stream()``, which follows
|
||||
the older thread event stream endpoint.
|
||||
"""
|
||||
transport = SyncProtocolSseTransport(
|
||||
self.http,
|
||||
thread_id,
|
||||
headers=headers,
|
||||
params=params,
|
||||
)
|
||||
return SyncThreadStream(transport, assistant_id=assistant_id)
|
||||
|
||||
def get(
|
||||
self,
|
||||
thread_id: str,
|
||||
|
||||
@@ -17,6 +17,11 @@ from langgraph_sdk._async.client import LangGraphClient, get_client
|
||||
from langgraph_sdk._async.cron import CronClient
|
||||
from langgraph_sdk._async.http import HttpClient, _adecode_json, _aencode_json
|
||||
from langgraph_sdk._async.runs import RunsClient
|
||||
from langgraph_sdk._async.stream import (
|
||||
EventSubscription,
|
||||
ProtocolSseTransport,
|
||||
ThreadStream,
|
||||
)
|
||||
from langgraph_sdk._async.store import StoreClient
|
||||
from langgraph_sdk._async.threads import ThreadsClient
|
||||
from langgraph_sdk._shared.utilities import configure_loopback_transports
|
||||
@@ -27,6 +32,11 @@ from langgraph_sdk._sync.client import SyncLangGraphClient, get_sync_client
|
||||
from langgraph_sdk._sync.cron import SyncCronClient
|
||||
from langgraph_sdk._sync.http import SyncHttpClient, _decode_json, _encode_json
|
||||
from langgraph_sdk._sync.runs import SyncRunsClient
|
||||
from langgraph_sdk._sync.stream import (
|
||||
SyncEventSubscription,
|
||||
SyncProtocolSseTransport,
|
||||
SyncThreadStream,
|
||||
)
|
||||
from langgraph_sdk._sync.store import SyncStoreClient
|
||||
from langgraph_sdk._sync.threads import SyncThreadsClient
|
||||
|
||||
@@ -35,15 +45,21 @@ __all__ = [
|
||||
"CronClient",
|
||||
"HttpClient",
|
||||
"LangGraphClient",
|
||||
"EventSubscription",
|
||||
"ProtocolSseTransport",
|
||||
"RunsClient",
|
||||
"StoreClient",
|
||||
"SyncAssistantsClient",
|
||||
"SyncCronClient",
|
||||
"SyncHttpClient",
|
||||
"SyncLangGraphClient",
|
||||
"SyncEventSubscription",
|
||||
"SyncProtocolSseTransport",
|
||||
"SyncRunsClient",
|
||||
"SyncStoreClient",
|
||||
"SyncThreadStream",
|
||||
"SyncThreadsClient",
|
||||
"ThreadStream",
|
||||
"ThreadsClient",
|
||||
"_adecode_json",
|
||||
"_aencode_json",
|
||||
|
||||
@@ -0,0 +1,105 @@
|
||||
"""Typed protocol messages for thread-centric remote streaming.
|
||||
|
||||
These shapes mirror the JSON protocol used by the JavaScript SDK's
|
||||
``ThreadStream`` layer. The SDK keeps them lightweight and dependency-free so
|
||||
``langgraph`` can build higher-level projections on top without creating a
|
||||
package cycle.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Literal
|
||||
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
Channel = Literal[
|
||||
"values",
|
||||
"updates",
|
||||
"messages",
|
||||
"tools",
|
||||
"custom",
|
||||
"lifecycle",
|
||||
"input",
|
||||
"debug",
|
||||
"checkpoints",
|
||||
"tasks",
|
||||
]
|
||||
"""Built-in subscribable protocol channels."""
|
||||
|
||||
CommandMethod = Literal[
|
||||
"run.input",
|
||||
"subscription.subscribe",
|
||||
"subscription.unsubscribe",
|
||||
"agent.getTree",
|
||||
"input.respond",
|
||||
"input.inject",
|
||||
"state.get",
|
||||
"state.listCheckpoints",
|
||||
"state.fork",
|
||||
]
|
||||
"""Command methods understood by the thread stream protocol."""
|
||||
|
||||
|
||||
class SubscribeParams(TypedDict, total=False):
|
||||
"""Filter used when subscribing to thread protocol events."""
|
||||
|
||||
channels: list[Channel | str]
|
||||
namespaces: NotRequired[list[list[str]]]
|
||||
depth: NotRequired[int]
|
||||
|
||||
|
||||
class Command(TypedDict):
|
||||
"""Command sent to a thread protocol transport."""
|
||||
|
||||
id: str
|
||||
method: CommandMethod
|
||||
params: dict[str, Any]
|
||||
|
||||
|
||||
class CommandResponse(TypedDict):
|
||||
"""Successful command response."""
|
||||
|
||||
id: str
|
||||
result: Any
|
||||
|
||||
|
||||
class ErrorResponse(TypedDict):
|
||||
"""Error command response."""
|
||||
|
||||
id: str
|
||||
error: dict[str, Any]
|
||||
|
||||
|
||||
class EventParams(TypedDict):
|
||||
"""Protocol event parameters."""
|
||||
|
||||
namespace: list[str]
|
||||
timestamp: int
|
||||
data: Any
|
||||
node: NotRequired[str]
|
||||
run_id: NotRequired[str]
|
||||
interrupts: NotRequired[list[Any]]
|
||||
|
||||
|
||||
class Event(TypedDict):
|
||||
"""Protocol event envelope."""
|
||||
|
||||
type: Literal["event"]
|
||||
method: str
|
||||
params: EventParams
|
||||
event_id: NotRequired[str]
|
||||
seq: NotRequired[int]
|
||||
|
||||
|
||||
class RunInputResult(TypedDict, total=False):
|
||||
"""Result returned by ``run.input`` commands."""
|
||||
|
||||
run_id: str
|
||||
thread_id: str
|
||||
|
||||
|
||||
class SubscribeResult(TypedDict):
|
||||
"""Result returned by ``subscription.subscribe`` commands."""
|
||||
|
||||
subscription_id: str
|
||||
|
||||
@@ -41,6 +41,7 @@ def _normalize_return_annotation(ann: object) -> str:
|
||||
s = str(ann)
|
||||
s = re.sub(r"\s+", "", s)
|
||||
s = s.replace("typing.", "").replace("collections.abc.", "")
|
||||
s = s.replace("SyncThreadStream", "ThreadStream")
|
||||
s = re.sub(r"AsyncGenerator\[([^,\]]+)(?:,[^\]]*)?\]", r"Iterator[\1]", s)
|
||||
s = re.sub(r"Generator\[([^,\]]+)(?:,[^\]]*)?\]", r"Iterator[\1]", s)
|
||||
s = re.sub(r"AsyncIterator\[(.+)\]", r"Iterator[\1]", s)
|
||||
|
||||
@@ -10,17 +10,23 @@ from langgraph_sdk import get_sync_client as public_get_sync_client
|
||||
from langgraph_sdk.client import (
|
||||
AssistantsClient,
|
||||
CronClient,
|
||||
EventSubscription,
|
||||
HttpClient,
|
||||
LangGraphClient,
|
||||
RunsClient,
|
||||
StoreClient,
|
||||
ProtocolSseTransport,
|
||||
SyncAssistantsClient,
|
||||
SyncCronClient,
|
||||
SyncEventSubscription,
|
||||
SyncHttpClient,
|
||||
SyncLangGraphClient,
|
||||
SyncProtocolSseTransport,
|
||||
SyncRunsClient,
|
||||
SyncStoreClient,
|
||||
SyncThreadStream,
|
||||
SyncThreadsClient,
|
||||
ThreadStream,
|
||||
ThreadsClient,
|
||||
_adecode_json,
|
||||
_aencode_json,
|
||||
@@ -52,6 +58,9 @@ def test_client_exports():
|
||||
assert RunsClient is not None
|
||||
assert CronClient is not None
|
||||
assert StoreClient is not None
|
||||
assert EventSubscription is not None
|
||||
assert ProtocolSseTransport is not None
|
||||
assert ThreadStream is not None
|
||||
|
||||
# Resource client classes - Sync
|
||||
assert SyncAssistantsClient is not None
|
||||
@@ -59,6 +68,9 @@ def test_client_exports():
|
||||
assert SyncRunsClient is not None
|
||||
assert SyncCronClient is not None
|
||||
assert SyncStoreClient is not None
|
||||
assert SyncEventSubscription is not None
|
||||
assert SyncProtocolSseTransport is not None
|
||||
assert SyncThreadStream is not None
|
||||
|
||||
# Internal utilities (used by tests)
|
||||
assert callable(_aencode_json)
|
||||
|
||||
@@ -9,7 +9,14 @@ 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,
|
||||
SyncThreadStream,
|
||||
SyncThreadsClient,
|
||||
ThreadStream,
|
||||
ThreadsClient,
|
||||
)
|
||||
from langgraph_sdk.schema import (
|
||||
CheckpointPayload,
|
||||
CheckpointsStreamPart,
|
||||
@@ -154,6 +161,107 @@ def test_sync_http_client_stream_flushes_trailing_event():
|
||||
assert parts == [StreamPart(event="foo", data={"bar": 1})]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_threads_stream_sends_commands_and_subscribes() -> None:
|
||||
requests: list[tuple[str, str]] = []
|
||||
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = await request.aread()
|
||||
requests.append((request.method, request.url.path))
|
||||
if request.url.path == "/v2/threads/thread-1/commands":
|
||||
assert request.method == "POST"
|
||||
payload = httpx.Response(200, content=body).json()
|
||||
if payload["method"] == "run.input":
|
||||
assert payload["params"]["assistant_id"] == "agent"
|
||||
return httpx.Response(200, json={"id": payload["id"], "result": "ok"})
|
||||
if payload["method"] == "subscription.unsubscribe":
|
||||
return httpx.Response(200, json={"id": payload["id"], "result": {}})
|
||||
if request.url.path == "/v2/threads/thread-1/events":
|
||||
assert request.method == "POST"
|
||||
assert httpx.Response(200, content=body).json() == {
|
||||
"channels": ["messages"]
|
||||
}
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"Content-Type": "text/event-stream"},
|
||||
content=(
|
||||
b"id: evt-1\n"
|
||||
b"event: messages\n"
|
||||
b'data: {"type":"event","method":"messages","params":{"namespace":[],"timestamp":1,"data":{"event":"message-start","id":"m1"}}}\n\n'
|
||||
),
|
||||
)
|
||||
raise AssertionError(f"unexpected request: {request.method} {request.url.path}")
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
async with httpx.AsyncClient(
|
||||
transport=transport, base_url="https://example.com"
|
||||
) as client:
|
||||
threads = ThreadsClient(HttpClient(client))
|
||||
thread = threads.stream("thread-1", assistant_id="agent")
|
||||
assert isinstance(thread, ThreadStream)
|
||||
assert await thread.run.input({"input": {"messages": []}}) == "ok"
|
||||
subscription = await thread.subscribe(["messages"])
|
||||
events = [event async for event in subscription]
|
||||
await subscription.unsubscribe()
|
||||
|
||||
assert events[0]["method"] == "messages"
|
||||
assert events[0]["event_id"] == "evt-1"
|
||||
assert requests == [
|
||||
("POST", "/v2/threads/thread-1/commands"),
|
||||
("POST", "/v2/threads/thread-1/events"),
|
||||
("POST", "/v2/threads/thread-1/commands"),
|
||||
]
|
||||
|
||||
|
||||
def test_sync_threads_stream_sends_commands_and_subscribes() -> None:
|
||||
requests: list[tuple[str, str]] = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = request.read()
|
||||
requests.append((request.method, request.url.path))
|
||||
if request.url.path == "/v2/threads/thread-1/commands":
|
||||
assert request.method == "POST"
|
||||
payload = httpx.Response(200, content=body).json()
|
||||
if payload["method"] == "run.input":
|
||||
assert payload["params"]["assistant_id"] == "agent"
|
||||
return httpx.Response(200, json={"id": payload["id"], "result": "ok"})
|
||||
if payload["method"] == "subscription.unsubscribe":
|
||||
return httpx.Response(200, json={"id": payload["id"], "result": {}})
|
||||
if request.url.path == "/v2/threads/thread-1/events":
|
||||
assert request.method == "POST"
|
||||
assert httpx.Response(200, content=body).json() == {
|
||||
"channels": ["messages"]
|
||||
}
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"Content-Type": "text/event-stream"},
|
||||
content=(
|
||||
b"id: evt-1\n"
|
||||
b"event: messages\n"
|
||||
b'data: {"type":"event","method":"messages","params":{"namespace":[],"timestamp":1,"data":{"event":"message-start","id":"m1"}}}\n\n'
|
||||
),
|
||||
)
|
||||
raise AssertionError(f"unexpected request: {request.method} {request.url.path}")
|
||||
|
||||
transport = httpx.MockTransport(handler)
|
||||
with httpx.Client(transport=transport, base_url="https://example.com") as client:
|
||||
threads = SyncThreadsClient(SyncHttpClient(client))
|
||||
thread = threads.stream("thread-1", assistant_id="agent")
|
||||
assert isinstance(thread, SyncThreadStream)
|
||||
assert thread.run.input({"input": {"messages": []}}) == "ok"
|
||||
subscription = thread.subscribe(["messages"])
|
||||
events = list(subscription)
|
||||
subscription.unsubscribe()
|
||||
|
||||
assert events[0]["method"] == "messages"
|
||||
assert events[0]["event_id"] == "evt-1"
|
||||
assert requests == [
|
||||
("POST", "/v2/threads/thread-1/commands"),
|
||||
("POST", "/v2/threads/thread-1/events"),
|
||||
("POST", "/v2/threads/thread-1/commands"),
|
||||
]
|
||||
|
||||
|
||||
def test_sync_http_client_stream_recovers_after_disconnect():
|
||||
reconnect_path = "/reconnect"
|
||||
first_chunks = [
|
||||
|
||||
Reference in New Issue
Block a user