Merge pull request #1845 from langchain-ai/nc/25sep/fix-stream-time

Fix stream_mode=messages/custom buffering for sync calls with single node
This commit is contained in:
Nuno Campos
2024-09-25 11:12:05 -07:00
committed by GitHub
3 changed files with 66 additions and 2 deletions
+2 -2
View File
@@ -49,8 +49,8 @@ class PregelRunner:
tasks = tuple(tasks)
# give control back to the caller
yield
# fast path if single task with no timeout
if len(tasks) == 1 and timeout is None:
# fast path if single task with no timeout and no waiter
if len(tasks) == 1 and timeout is None and get_waiter is None:
t = tasks[0]
try:
run_with_retry(t, retry_policy)
+33
View File
@@ -8606,6 +8606,39 @@ def test_stream_subgraphs_during_execution(
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_stream_buffering_single_node(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue("checkpointer_" + checkpointer_name)
class State(TypedDict):
my_key: Annotated[str, operator.add]
def node(state: State, writer: StreamWriter):
writer("Before sleep")
time.sleep(0.2)
writer("After sleep")
return {"my_key": "got here"}
builder = StateGraph(State)
builder.add_node("node", node)
builder.add_edge(START, "node")
builder.add_edge("node", END)
graph = builder.compile(checkpointer=checkpointer)
start = time.perf_counter()
chunks: list[tuple[float, Any]] = []
config = {"configurable": {"thread_id": "2"}}
for c in graph.stream({"my_key": ""}, config, stream_mode="custom"):
chunks.append((round(time.perf_counter() - start, 1), c))
assert chunks == [
(FloatBetween(0.0, 0.1), "Before sleep"),
(FloatBetween(0.2, 0.3), "After sleep"),
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_nested_graph_interrupts_parallel(
request: pytest.FixtureRequest, checkpointer_name: str
+31
View File
@@ -7208,6 +7208,37 @@ async def test_stream_subgraphs_during_execution(checkpointer_name: str) -> None
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_stream_buffering_single_node(checkpointer_name: str) -> None:
class State(TypedDict):
my_key: Annotated[str, operator.add]
async def node(state: State, writer: StreamWriter):
writer("Before sleep")
await asyncio.sleep(0.2)
writer("After sleep")
return {"my_key": "got here"}
builder = StateGraph(State)
builder.add_node("node", node)
builder.add_edge(START, "node")
builder.add_edge("node", END)
async with awith_checkpointer(checkpointer_name) as checkpointer:
graph = builder.compile(checkpointer=checkpointer)
start = perf_counter()
chunks: list[tuple[float, Any]] = []
config = {"configurable": {"thread_id": "2"}}
async for c in graph.astream({"my_key": ""}, config, stream_mode="custom"):
chunks.append((round(perf_counter() - start, 1), c))
assert chunks == [
(FloatBetween(0.0, 0.1), "Before sleep"),
(FloatBetween(0.2, 0.3), "After sleep"),
]
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
async def test_nested_graph_interrupts_parallel(checkpointer_name: str) -> None:
class InnerState(TypedDict):