Interop of RemoteGraph w core lib

This commit is contained in:
Nuno Campos
2024-10-23 15:37:17 -07:00
parent 62a5ec509d
commit dc8260bb72
5 changed files with 87 additions and 35 deletions
+48 -18
View File
@@ -17,7 +17,7 @@ from langgraph.pregel.types import StateSnapshot
def test_with_config():
# set up test
remote_pregel = RemoteGraph(
graph_id="test_graph_id",
"test_graph_id",
config={
"configurable": {
"foo": "bar",
@@ -64,7 +64,7 @@ def test_get_graph():
],
}
remote_pregel = RemoteGraph(sync_client=mock_sync_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph("test_graph_id", sync_client=mock_sync_client)
# call method / assertions
drawable_graph = remote_pregel.get_graph()
@@ -111,7 +111,7 @@ async def test_aget_graph():
],
}
remote_pregel = RemoteGraph(client=mock_async_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph("test_graph_id", client=mock_async_client)
# call method / assertions
drawable_graph = await remote_pregel.aget_graph()
@@ -155,9 +155,7 @@ def test_get_subgraphs():
},
}
remote_pregel = RemoteGraph(
sync_client=mock_sync_client, graph_id="test_graph_id_1"
)
remote_pregel = RemoteGraph("test_graph_id_1", sync_client=mock_sync_client)
# call method / assertions
subgraphs = list(remote_pregel.get_subgraphs())
@@ -198,8 +196,8 @@ async def test_aget_subgraphs():
}
remote_pregel = RemoteGraph(
"test_graph_id_1",
client=mock_async_client,
graph_id="test_graph_id_1",
)
# call method / assertions
@@ -240,7 +238,10 @@ def test_get_state():
}
# call method / assertions
remote_pregel = RemoteGraph(sync_client=mock_sync_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
sync_client=mock_sync_client,
)
config = {"configurable": {"thread_id": "thread1"}}
state_snapshot = remote_pregel.get_state(config)
@@ -287,7 +288,10 @@ async def test_aget_state():
}
# call method / assertions
remote_pregel = RemoteGraph(client=mock_async_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
client=mock_async_client,
)
config = {"configurable": {"thread_id": "thread1"}}
state_snapshot = await remote_pregel.aget_state(config)
@@ -338,7 +342,10 @@ def test_get_state_history():
]
# call method / assertions
remote_pregel = RemoteGraph(sync_client=mock_sync_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
sync_client=mock_sync_client,
)
config = {"configurable": {"thread_id": "thread1"}}
state_history_snapshot = list(
@@ -386,7 +393,10 @@ async def test_aget_state_history():
]
# call method / assertions
remote_pregel = RemoteGraph(client=mock_async_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
client=mock_async_client,
)
config = {"configurable": {"thread_id": "thread1"}}
state_history_snapshot = []
@@ -427,7 +437,10 @@ def test_update_state():
}
# call method / assertions
remote_pregel = RemoteGraph(sync_client=mock_sync_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
sync_client=mock_sync_client,
)
config = {"configurable": {"thread_id": "thread1"}}
response = remote_pregel.update_state(config, {"key": "value"})
@@ -456,7 +469,10 @@ async def test_aupdate_state():
}
# call method / assertions
remote_pregel = RemoteGraph(client=mock_async_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
client=mock_async_client,
)
config = {"configurable": {"thread_id": "thread1"}}
response = await remote_pregel.aupdate_state(config, {"key": "value"})
@@ -483,7 +499,10 @@ def test_stream():
]
# call method / assertions
remote_pregel = RemoteGraph(sync_client=mock_sync_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
sync_client=mock_sync_client,
)
# stream modes doesn't include 'updates'
stream_parts = []
@@ -583,7 +602,10 @@ async def test_astream():
mock_async_client.runs.stream.return_value = async_iter
# call method / assertions
remote_pregel = RemoteGraph(client=mock_async_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
client=mock_async_client,
)
# stream modes doesn't include 'updates'
stream_parts = []
@@ -717,7 +739,10 @@ def test_invoke():
}
# call method / assertions
remote_pregel = RemoteGraph(sync_client=mock_sync_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
sync_client=mock_sync_client,
)
config = {"configurable": {"thread_id": "thread_1"}}
result = remote_pregel.invoke(
@@ -736,7 +761,10 @@ async def test_ainvoke():
}
# call method / assertions
remote_pregel = RemoteGraph(client=mock_async_client, graph_id="test_graph_id")
remote_pregel = RemoteGraph(
"test_graph_id",
client=mock_async_client,
)
config = {"configurable": {"thread_id": "thread_1"}}
result = await remote_pregel.ainvoke(
@@ -758,7 +786,9 @@ async def test_langgraph_cloud_integration():
client = get_client()
sync_client = get_sync_client()
remote_pregel = RemoteGraph(
client=client, sync_client=sync_client, graph_id="agent"
"agent",
client=client,
sync_client=sync_client,
)
# define graph