mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-09 19:27:54 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
844373e9b3 |
@@ -631,6 +631,26 @@ class RemoteGraph(PregelProtocol):
|
|||||||
)
|
)
|
||||||
return self._get_config(response["checkpoint"])
|
return self._get_config(response["checkpoint"])
|
||||||
|
|
||||||
|
def _prepare_run_input(
|
||||||
|
self,
|
||||||
|
input: dict[str, Any] | Any,
|
||||||
|
config: RunnableConfig | None,
|
||||||
|
) -> tuple[RunnableConfig, dict[str, Any] | Any, CommandSDK | None, str | None]:
|
||||||
|
"""Prepare input for run calls.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (sanitized_config, input, command, thread_id)
|
||||||
|
"""
|
||||||
|
merged_config = merge_configs(self.config, config)
|
||||||
|
sanitized_config = self._sanitize_config(merged_config)
|
||||||
|
if isinstance(input, Command):
|
||||||
|
command: CommandSDK | None = cast(CommandSDK, asdict(input))
|
||||||
|
input = None
|
||||||
|
else:
|
||||||
|
command = None
|
||||||
|
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
||||||
|
return sanitized_config, input, command, thread_id
|
||||||
|
|
||||||
def _get_stream_modes(
|
def _get_stream_modes(
|
||||||
self,
|
self,
|
||||||
stream_mode: StreamMode | list[StreamMode] | None,
|
stream_mode: StreamMode | list[StreamMode] | None,
|
||||||
@@ -715,17 +735,12 @@ class RemoteGraph(PregelProtocol):
|
|||||||
The output of the graph.
|
The output of the graph.
|
||||||
"""
|
"""
|
||||||
sync_client = self._validate_sync_client()
|
sync_client = self._validate_sync_client()
|
||||||
merged_config = merge_configs(self.config, config)
|
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||||
sanitized_config = self._sanitize_config(merged_config)
|
input, config
|
||||||
|
)
|
||||||
stream_modes, requested, req_single, stream = self._get_stream_modes(
|
stream_modes, requested, req_single, stream = self._get_stream_modes(
|
||||||
stream_mode, config
|
stream_mode, config
|
||||||
)
|
)
|
||||||
if isinstance(input, Command):
|
|
||||||
command: CommandSDK | None = cast(CommandSDK, asdict(input))
|
|
||||||
input = None
|
|
||||||
else:
|
|
||||||
command = None
|
|
||||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
|
||||||
|
|
||||||
for chunk in sync_client.runs.stream(
|
for chunk in sync_client.runs.stream(
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
@@ -825,17 +840,12 @@ class RemoteGraph(PregelProtocol):
|
|||||||
The output of the graph.
|
The output of the graph.
|
||||||
"""
|
"""
|
||||||
client = self._validate_client()
|
client = self._validate_client()
|
||||||
merged_config = merge_configs(self.config, config)
|
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||||
sanitized_config = self._sanitize_config(merged_config)
|
input, config
|
||||||
|
)
|
||||||
stream_modes, requested, req_single, stream = self._get_stream_modes(
|
stream_modes, requested, req_single, stream = self._get_stream_modes(
|
||||||
stream_mode, config
|
stream_mode, config
|
||||||
)
|
)
|
||||||
if isinstance(input, Command):
|
|
||||||
command: CommandSDK | None = cast(CommandSDK, asdict(input))
|
|
||||||
input = None
|
|
||||||
else:
|
|
||||||
command = None
|
|
||||||
thread_id = sanitized_config.get("configurable", {}).pop("thread_id", None)
|
|
||||||
|
|
||||||
async for chunk in client.runs.stream(
|
async for chunk in client.runs.stream(
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
@@ -937,27 +947,31 @@ 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.
|
||||||
**kwargs: Additional params to pass to RemoteGraph.stream.
|
**kwargs: Additional params to pass to client.runs.wait.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The output of the graph.
|
The output of the graph.
|
||||||
"""
|
"""
|
||||||
for chunk in self.stream(
|
sync_client = self._validate_sync_client()
|
||||||
input,
|
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||||
config=config,
|
input, config
|
||||||
|
)
|
||||||
|
|
||||||
|
return sync_client.runs.wait( # type: ignore
|
||||||
|
thread_id=thread_id,
|
||||||
|
assistant_id=self.assistant_id,
|
||||||
|
input=input,
|
||||||
|
command=command,
|
||||||
|
config=sanitized_config,
|
||||||
interrupt_before=interrupt_before,
|
interrupt_before=interrupt_before,
|
||||||
interrupt_after=interrupt_after,
|
interrupt_after=interrupt_after,
|
||||||
headers=headers,
|
if_not_exists="create",
|
||||||
stream_mode="values",
|
headers=(
|
||||||
|
_merge_tracing_headers(headers) if self.distributed_tracing else headers
|
||||||
|
),
|
||||||
params=params,
|
params=params,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
)
|
||||||
pass
|
|
||||||
try:
|
|
||||||
return chunk
|
|
||||||
except UnboundLocalError:
|
|
||||||
logger.warning("No events received from remote graph")
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def ainvoke(
|
async def ainvoke(
|
||||||
self,
|
self,
|
||||||
@@ -978,27 +992,31 @@ 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.
|
||||||
**kwargs: Additional params to pass to RemoteGraph.astream.
|
**kwargs: Additional params to pass to client.runs.wait.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The output of the graph.
|
The output of the graph.
|
||||||
"""
|
"""
|
||||||
async for chunk in self.astream(
|
client = self._validate_client()
|
||||||
input,
|
sanitized_config, input, command, thread_id = self._prepare_run_input(
|
||||||
config=config,
|
input, config
|
||||||
|
)
|
||||||
|
|
||||||
|
return await client.runs.wait(
|
||||||
|
thread_id=thread_id,
|
||||||
|
assistant_id=self.assistant_id,
|
||||||
|
input=input,
|
||||||
|
command=command,
|
||||||
|
config=sanitized_config,
|
||||||
interrupt_before=interrupt_before,
|
interrupt_before=interrupt_before,
|
||||||
interrupt_after=interrupt_after,
|
interrupt_after=interrupt_after,
|
||||||
headers=headers,
|
if_not_exists="create",
|
||||||
stream_mode="values",
|
headers=(
|
||||||
|
_merge_tracing_headers(headers) if self.distributed_tracing else headers
|
||||||
|
),
|
||||||
params=params,
|
params=params,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
)
|
||||||
pass
|
|
||||||
try:
|
|
||||||
return chunk
|
|
||||||
except UnboundLocalError:
|
|
||||||
logger.warning("No events received from remote graph")
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _merge_tracing_headers(headers: dict[str, str] | None) -> dict[str, str] | None:
|
def _merge_tracing_headers(headers: dict[str, str] | None) -> dict[str, str] | None:
|
||||||
|
|||||||
@@ -818,13 +818,9 @@ async def test_astream():
|
|||||||
def test_invoke():
|
def test_invoke():
|
||||||
# set up test
|
# set up test
|
||||||
mock_sync_client = MagicMock()
|
mock_sync_client = MagicMock()
|
||||||
mock_sync_client.runs.stream.return_value = [
|
mock_sync_client.runs.wait.return_value = {
|
||||||
StreamPart(event="values", data={"chunk": "data1"}),
|
"messages": [{"type": "human", "content": "world"}]
|
||||||
StreamPart(event="values", data={"chunk": "data2"}),
|
}
|
||||||
StreamPart(
|
|
||||||
event="values", data={"messages": [{"type": "human", "content": "world"}]}
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
# call method / assertions
|
# call method / assertions
|
||||||
remote_pregel = RemoteGraph(
|
remote_pregel = RemoteGraph(
|
||||||
@@ -838,13 +834,19 @@ def test_invoke():
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
||||||
|
# verify runs.wait was called with expected args
|
||||||
|
assert mock_sync_client.runs.wait.called
|
||||||
|
_, kwargs = mock_sync_client.runs.wait.call_args
|
||||||
|
assert kwargs.get("thread_id") == "thread_1"
|
||||||
|
assert kwargs.get("assistant_id") == "test_graph_id"
|
||||||
|
assert kwargs.get("if_not_exists") == "create"
|
||||||
|
|
||||||
|
|
||||||
def test_invoke_sanitizes_thread_id():
|
def test_invoke_sanitizes_thread_id():
|
||||||
# Ensure that invoking with thread_id passes thread_id as a top-level arg
|
# Ensure that invoking with thread_id passes thread_id as a top-level arg
|
||||||
# and removes it from the config body.
|
# and removes it from the config body.
|
||||||
mock_sync_client = MagicMock()
|
mock_sync_client = MagicMock()
|
||||||
mock_sync_client.runs.stream.return_value = []
|
mock_sync_client.runs.wait.return_value = {}
|
||||||
remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client)
|
remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client)
|
||||||
|
|
||||||
config = {"configurable": {"thread_id": "thread_1"}}
|
config = {"configurable": {"thread_id": "thread_1"}}
|
||||||
@@ -852,8 +854,8 @@ def test_invoke_sanitizes_thread_id():
|
|||||||
{"input": {"messages": [{"type": "human", "content": "hello"}]}}, config
|
{"input": {"messages": [{"type": "human", "content": "hello"}]}}, config
|
||||||
)
|
)
|
||||||
|
|
||||||
assert mock_sync_client.runs.stream.called
|
assert mock_sync_client.runs.wait.called
|
||||||
_, kwargs = mock_sync_client.runs.stream.call_args
|
_, kwargs = mock_sync_client.runs.wait.call_args
|
||||||
assert kwargs.get("thread_id") == "thread_1"
|
assert kwargs.get("thread_id") == "thread_1"
|
||||||
passed_config = kwargs.get("config") or {}
|
passed_config = kwargs.get("config") or {}
|
||||||
assert "configurable" in passed_config
|
assert "configurable" in passed_config
|
||||||
@@ -883,16 +885,10 @@ def test_stream_sanitizes_thread_id():
|
|||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_ainvoke():
|
async def test_ainvoke():
|
||||||
# set up test
|
# set up test
|
||||||
mock_async_client = MagicMock()
|
mock_async_client = AsyncMock()
|
||||||
async_iter = MagicMock()
|
mock_async_client.runs.wait.return_value = {
|
||||||
async_iter.__aiter__.return_value = [
|
"messages": [{"type": "human", "content": "world"}]
|
||||||
StreamPart(event="values", data={"chunk": "data1"}),
|
}
|
||||||
StreamPart(event="values", data={"chunk": "data2"}),
|
|
||||||
StreamPart(
|
|
||||||
event="values", data={"messages": [{"type": "human", "content": "world"}]}
|
|
||||||
),
|
|
||||||
]
|
|
||||||
mock_async_client.runs.stream.return_value = async_iter
|
|
||||||
|
|
||||||
# call method / assertions
|
# call method / assertions
|
||||||
remote_pregel = RemoteGraph(
|
remote_pregel = RemoteGraph(
|
||||||
@@ -906,6 +902,12 @@ async def test_ainvoke():
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
assert result == {"messages": [{"type": "human", "content": "world"}]}
|
||||||
|
# verify runs.wait was called with expected args
|
||||||
|
assert mock_async_client.runs.wait.called
|
||||||
|
_, kwargs = mock_async_client.runs.wait.call_args
|
||||||
|
assert kwargs.get("thread_id") == "thread_1"
|
||||||
|
assert kwargs.get("assistant_id") == "test_graph_id"
|
||||||
|
assert kwargs.get("if_not_exists") == "create"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skip(
|
@pytest.mark.skip(
|
||||||
@@ -1241,12 +1243,18 @@ async def test_include_headers(
|
|||||||
async_iter.__aiter__.return_value = return_value
|
async_iter.__aiter__.return_value = return_value
|
||||||
astream_mock = mock_async_client.runs.stream
|
astream_mock = mock_async_client.runs.stream
|
||||||
astream_mock.return_value = async_iter
|
astream_mock.return_value = async_iter
|
||||||
|
# Mock for ainvoke which uses runs.wait
|
||||||
|
await_mock = AsyncMock(return_value={"chunk": "data1"})
|
||||||
|
mock_async_client.runs.wait = await_mock
|
||||||
|
|
||||||
mock_sync_client = MagicMock()
|
mock_sync_client = MagicMock()
|
||||||
sync_iter = MagicMock()
|
sync_iter = MagicMock()
|
||||||
sync_iter.__iter__.return_value = return_value
|
sync_iter.__iter__.return_value = return_value
|
||||||
stream_mock = mock_sync_client.runs.stream
|
stream_mock = mock_sync_client.runs.stream
|
||||||
stream_mock.return_value = async_iter
|
stream_mock.return_value = async_iter
|
||||||
|
# Mock for invoke which uses runs.wait
|
||||||
|
wait_mock = MagicMock(return_value={"chunk": "data1"})
|
||||||
|
mock_sync_client.runs.wait = wait_mock
|
||||||
|
|
||||||
remote_pregel = RemoteGraph(
|
remote_pregel = RemoteGraph(
|
||||||
"test_graph_id",
|
"test_graph_id",
|
||||||
@@ -1279,8 +1287,12 @@ async def test_include_headers(
|
|||||||
expected["langsmith-trace"] = AnyStr()
|
expected["langsmith-trace"] = AnyStr()
|
||||||
expected["baggage"] = AnyStr("langsmith-metadata=")
|
expected["baggage"] = AnyStr("langsmith-metadata=")
|
||||||
|
|
||||||
assert astream_mock.call_args.kwargs["headers"] == expected
|
if stream:
|
||||||
|
assert astream_mock.call_args.kwargs["headers"] == expected
|
||||||
|
else:
|
||||||
|
assert await_mock.call_args.kwargs["headers"] == expected
|
||||||
stream_mock.assert_not_called()
|
stream_mock.assert_not_called()
|
||||||
|
wait_mock.assert_not_called()
|
||||||
|
|
||||||
with ls.tracing_context(enabled=True, client=MagicMock()):
|
with ls.tracing_context(enabled=True, client=MagicMock()):
|
||||||
with ls.trace("foo"):
|
with ls.trace("foo"):
|
||||||
@@ -1298,4 +1310,7 @@ async def test_include_headers(
|
|||||||
config,
|
config,
|
||||||
headers=headers,
|
headers=headers,
|
||||||
)
|
)
|
||||||
assert stream_mock.call_args.kwargs["headers"] == expected
|
if stream:
|
||||||
|
assert stream_mock.call_args.kwargs["headers"] == expected
|
||||||
|
else:
|
||||||
|
assert wait_mock.call_args.kwargs["headers"] == expected
|
||||||
|
|||||||
Reference in New Issue
Block a user