mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
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:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user