Merge pull request #2235 from langchain-ai/nc/30oct/remote-invoke

lib: Ensure RemoteGraph.invoke emits all subgraph events as expected
This commit is contained in:
Nuno Campos
2024-10-30 13:26:56 -07:00
committed by GitHub
2 changed files with 39 additions and 31 deletions
+20 -22
View File
@@ -695,19 +695,18 @@ class RemoteGraph(PregelProtocol):
Returns:
The output of the graph.
"""
sync_client = self._validate_sync_client()
merged_config = merge_configs(self.config, config)
sanitized_config = self._sanitize_config(merged_config)
return sync_client.runs.wait(
thread_id=sanitized_config["configurable"].get("thread_id"),
assistant_id=self.name,
input=input,
config=sanitized_config,
for chunk in self.stream(
input,
config=config,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
if_not_exists="create",
)
stream_mode="values",
):
pass
try:
return chunk
except UnboundLocalError:
return None
async def ainvoke(
self,
@@ -732,16 +731,15 @@ class RemoteGraph(PregelProtocol):
Returns:
The output of the graph.
"""
client = self._validate_client()
merged_config = merge_configs(self.config, config)
sanitized_config = self._sanitize_config(merged_config)
return await client.runs.wait(
thread_id=sanitized_config["configurable"].get("thread_id"),
assistant_id=self.name,
input=input,
config=sanitized_config,
async for chunk in self.astream(
input,
config=config,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
if_not_exists="create",
)
stream_mode="values",
):
pass
try:
return chunk
except UnboundLocalError:
return None
+19 -9
View File
@@ -650,9 +650,13 @@ async def test_astream():
def test_invoke():
# set up test
mock_sync_client = MagicMock()
mock_sync_client.runs.wait.return_value = {
"values": {"messages": [{"type": "human", "content": "world"}]}
}
mock_sync_client.runs.stream.return_value = [
StreamPart(event="values", data={"chunk": "data1"}),
StreamPart(event="values", data={"chunk": "data2"}),
StreamPart(
event="values", data={"messages": [{"type": "human", "content": "world"}]}
),
]
# call method / assertions
remote_pregel = RemoteGraph(
@@ -665,16 +669,22 @@ def test_invoke():
{"input": {"messages": [{"type": "human", "content": "hello"}]}}, config
)
assert result == {"values": {"messages": [{"type": "human", "content": "world"}]}}
assert result == {"messages": [{"type": "human", "content": "world"}]}
@pytest.mark.anyio
async def test_ainvoke():
# set up test
mock_async_client = AsyncMock()
mock_async_client.runs.wait.return_value = {
"values": {"messages": [{"type": "human", "content": "world"}]}
}
mock_async_client = MagicMock()
async_iter = MagicMock()
async_iter.__aiter__.return_value = [
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
remote_pregel = RemoteGraph(
@@ -687,7 +697,7 @@ async def test_ainvoke():
{"input": {"messages": [{"type": "human", "content": "hello"}]}}, config
)
assert result == {"values": {"messages": [{"type": "human", "content": "world"}]}}
assert result == {"messages": [{"type": "human", "content": "world"}]}
@pytest.mark.skip("Unskip this test to manually test the LangGraph Cloud integration")