diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 86b9ba8dc..772564a50 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -63,6 +63,7 @@ from langgraph.constants import ( RESUME, SCHEDULED, TAG_HIDDEN, + TASKS, ) from langgraph.errors import ( CheckpointNotLatest, @@ -295,6 +296,7 @@ class PregelLoop(LoopProtocol): """Put writes for a task, to be read by the next tick.""" if not writes: return + checkpoint_during = self.checkpoint_during or any(w[0] == TASKS for w in writes) # deduplicate writes to special channels, last write wins if all(w[0] in WRITES_IDX_MAP for w in writes): writes = list({w[0]: w for w in writes}.values()) @@ -304,7 +306,7 @@ class PregelLoop(LoopProtocol): ] # save writes self.checkpoint_pending_writes.extend((task_id, c, v) for c, v in writes) - if self.checkpoint_during and self.checkpointer_put_writes is not None: + if checkpoint_during and self.checkpointer_put_writes is not None: config = patch_configurable( self.checkpoint_config, { @@ -342,6 +344,16 @@ class PregelLoop(LoopProtocol): return if not self.checkpoint_pending_writes: return + # patch config + config = patch_configurable( + self.checkpoint_config, + { + CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get( + CONFIG_KEY_CHECKPOINT_NS, "" + ), + CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"], + }, + ) # group by task id by_task = defaultdict(list) for task_id, channel, value in self.checkpoint_pending_writes: @@ -354,7 +366,7 @@ class PregelLoop(LoopProtocol): task = self.tasks.get(task_id) self.submit( self.checkpointer_put_writes, - self.checkpoint_config, + config, writes, task_id, task_path_str(task.path) if task else "", @@ -362,7 +374,7 @@ class PregelLoop(LoopProtocol): else: self.submit( self.checkpointer_put_writes, - self.checkpoint_config, + config, writes, task_id, ) @@ -748,6 +760,7 @@ class PregelLoop(LoopProtocol): else self.stream_keys ), ) + self.checkpoint_id_prev = self.checkpoint["id"] if self.step > -1 else None # do checkpoint? do_checkpoint = self._checkpointer_put_after_previous is not None and ( exiting or self.checkpoint_during @@ -776,6 +789,7 @@ class PregelLoop(LoopProtocol): **self.checkpoint_config, CONF: { **self.checkpoint_config[CONF], + CONFIG_KEY_CHECKPOINT_ID: self.checkpoint_id_prev, CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get( CONFIG_KEY_CHECKPOINT_NS, "" ), diff --git a/libs/langgraph/tests/test_large_cases.py b/libs/langgraph/tests/test_large_cases.py index f15cd5c23..f5be91cee 100644 --- a/libs/langgraph/tests/test_large_cases.py +++ b/libs/langgraph/tests/test_large_cases.py @@ -7258,9 +7258,10 @@ def test_branch_then( ) -@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) +@pytest.mark.parametrize("checkpoint_during", [True, False]) +@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_SYNC) def test_send_dedupe_on_resume( - request: pytest.FixtureRequest, checkpointer_name: str + request: pytest.FixtureRequest, checkpointer_name: str, checkpoint_during: bool ) -> None: checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") @@ -7316,7 +7317,7 @@ def test_send_dedupe_on_resume( graph = builder.compile(checkpointer=checkpointer) thread1 = {"configurable": {"thread_id": "1"}} - assert graph.invoke(["0"], thread1, debug=1) == [ + assert graph.invoke(["0"], thread1, checkpoint_during=checkpoint_during) == [ "0", "1", "3.1", @@ -7333,12 +7334,11 @@ def test_send_dedupe_on_resume( pytest.xfail("TODO: shallow checkpointer reports wrong next set") assert state.next == ("flaky",) # check history - if "shallow" not in checkpointer_name: - history = [c for c in graph.get_state_history(thread1)] - assert len(history) == 4 + history = [c for c in graph.get_state_history(thread1)] + assert len(history) == (4 if checkpoint_during else 1) # resume execution - assert graph.invoke(None, thread1, debug=1) == [ + assert graph.invoke(None, thread1, checkpoint_during=checkpoint_during) == [ "0", "1", "3.1", @@ -7358,6 +7358,7 @@ def test_send_dedupe_on_resume( assert state.next == () # check history history = [c for c in graph.get_state_history(thread1)] + assert len(history) == (6 if checkpoint_during else 2) expected_history = [ StateSnapshot( values=[ @@ -7494,13 +7495,9 @@ def test_send_dedupe_on_resume( name="flaky", path=("__pregel_push", 1), error=None, - interrupts=( - Interrupt( - value="Bahh", resumable=False, ns=None, when="during" - ), - ), + interrupts=(Interrupt(value="Bahh", resumable=False, ns=None),), state=None, - result=["flaky|4"], + result=["flaky|4"] if checkpoint_during else None, ), PregelTask( id=AnyStr(), @@ -7637,10 +7634,11 @@ def test_send_dedupe_on_resume( ), ), ] - if "shallow" in checkpointer_name: - expected_history = expected_history[:1] - - assert history == expected_history + if checkpoint_during: + assert history == expected_history + else: + assert history[0] == expected_history[0] + assert history[1] == expected_history[2] @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 97405e41c..f1d45c192 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -1333,11 +1333,11 @@ def test_pending_writes_resume( "configurable": { "thread_id": "1", "checkpoint_ns": "", - "checkpoint_id": checkpoints[2].config["configurable"]["checkpoint_id"], + "checkpoint_id": checkpoints[2].config["configurable"]["checkpoint_id"] + if checkpoint_during + else AnyStr(), } - } - if checkpoint_during - else None, + }, pending_writes=UnsortedSequence( (AnyStr(), "value", 2), (AnyStr(), "__error__", 'ConnectionError("I\'m not good")'), @@ -1608,10 +1608,14 @@ def test_imp_task( assert mapper_calls == 2 +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_imp_nested( - request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion + request: pytest.FixtureRequest, checkpointer_name: str, checkpoint_during: bool ) -> None: + if not checkpoint_during and "shallow" in checkpointer_name: + pytest.skip("Checkpointing during execution not supported") + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") def mynode(input: list[str]) -> list[str]: @@ -1653,7 +1657,7 @@ def test_imp_nested( } thread1 = {"configurable": {"thread_id": "1"}} - assert [*graph.stream([0, 1], thread1)] == [ + assert [*graph.stream([0, 1], thread1, checkpoint_during=checkpoint_during)] == [ {"submapper": "0"}, {"mapper": "00"}, {"submapper": "1"}, @@ -1670,16 +1674,22 @@ def test_imp_nested( }, ] - assert graph.invoke(Command(resume="answer"), thread1) == [ + assert graph.invoke( + Command(resume="answer"), thread1, checkpoint_during=checkpoint_during + ) == [ "00answera", "11answera", ] +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_imp_stream_order( - request: pytest.FixtureRequest, checkpointer_name: str, snapshot: SnapshotAssertion + request: pytest.FixtureRequest, checkpointer_name: str, checkpoint_during: bool ) -> None: + if not checkpoint_during and "shallow" in checkpointer_name: + pytest.skip("Checkpointing during execution not supported") + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") @task() @@ -1702,7 +1712,10 @@ def test_imp_stream_order( return fut_baz.result() thread1 = {"configurable": {"thread_id": "1"}} - assert [c for c in graph.stream({"a": "0"}, thread1)] == [ + assert [ + c + for c in graph.stream({"a": "0"}, thread1, checkpoint_during=checkpoint_during) + ] == [ { "foo": ( "0foo", diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 467e97f53..4cc115f6f 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -2173,11 +2173,11 @@ async def test_pending_writes_resume( "checkpoint_ns": "", "checkpoint_id": checkpoints[2].config["configurable"][ "checkpoint_id" - ], + ] + if checkpoint_during + else AnyStr(), } - } - if checkpoint_during - else None, + }, pending_writes=UnsortedSequence( (AnyStr(), "value", 2), (AnyStr(), "__error__", 'ConnectionError("I\'m not good")'), @@ -2517,8 +2517,12 @@ async def test_imp_task(checkpointer_name: str, checkpoint_during: bool) -> None @NEEDS_CONTEXTVARS +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) -async def test_imp_nested(checkpointer_name: str) -> None: +async def test_imp_nested(checkpointer_name: str, checkpoint_during: bool) -> None: + if not checkpoint_during and "shallow" in checkpointer_name: + pytest.skip("Checkpointing during execution not supported") + async def mynode(input: list[str]) -> list[str]: return [it + "a" for it in input] @@ -2558,7 +2562,12 @@ async def test_imp_nested(checkpointer_name: str) -> None: } thread1 = {"configurable": {"thread_id": "1"}} - assert [c async for c in graph.astream([0, 1], thread1)] == [ + assert [ + c + async for c in graph.astream( + [0, 1], thread1, checkpoint_during=checkpoint_during + ) + ] == [ {"submapper": "0"}, {"mapper": "00"}, {"submapper": "1"}, @@ -2575,15 +2584,21 @@ async def test_imp_nested(checkpointer_name: str) -> None: }, ] - assert await graph.ainvoke(Command(resume="answer"), thread1) == [ + assert await graph.ainvoke( + Command(resume="answer"), thread1, checkpoint_during=checkpoint_during + ) == [ "00answera", "11answera", ] @NEEDS_CONTEXTVARS +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) -async def test_imp_task_cancel(checkpointer_name: str) -> None: +async def test_imp_task_cancel(checkpointer_name: str, checkpoint_during: bool) -> None: + if not checkpoint_during and "shallow" in checkpointer_name: + pytest.skip("Checkpointing during execution not supported") + async with awith_checkpointer(checkpointer_name) as checkpointer: mapper_calls = 0 mapper_cancels = 0 @@ -2609,7 +2624,12 @@ async def test_imp_task_cancel(checkpointer_name: str) -> None: return [m + answer for m in mapped] thread1 = {"configurable": {"thread_id": "1"}} - assert [c async for c in graph.astream([0, 1], thread1)] == [ + assert [ + c + async for c in graph.astream( + [0, 1], thread1, checkpoint_during=checkpoint_during + ) + ] == [ {"mapper": "00"}, { "__interrupt__": ( @@ -2625,7 +2645,9 @@ async def test_imp_task_cancel(checkpointer_name: str) -> None: assert mapper_calls == 2 assert mapper_cancels == 1 - assert await graph.ainvoke(Command(resume="answer"), thread1) == [ + assert await graph.ainvoke( + Command(resume="answer"), thread1, checkpoint_during=checkpoint_during + ) == [ "00answer", ] assert mapper_calls == 3 @@ -2633,8 +2655,14 @@ async def test_imp_task_cancel(checkpointer_name: str) -> None: @NEEDS_CONTEXTVARS +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) -async def test_imp_sync_from_async(checkpointer_name: str) -> None: +async def test_imp_sync_from_async( + checkpointer_name: str, checkpoint_during: bool +) -> None: + if not checkpoint_during and "shallow" in checkpointer_name: + pytest.skip("Checkpointing during execution not supported") + async with awith_checkpointer(checkpointer_name) as checkpointer: @task() @@ -2657,7 +2685,12 @@ async def test_imp_sync_from_async(checkpointer_name: str) -> None: return fut_baz.result() thread1 = {"configurable": {"thread_id": "1"}} - assert [c async for c in graph.astream({"a": "0"}, thread1)] == [ + assert [ + c + async for c in graph.astream( + {"a": "0"}, thread1, checkpoint_during=checkpoint_during + ) + ] == [ {"foo": {"a": "0foo", "b": "bar"}}, {"bar": {"a": "0foobar", "c": "bark"}}, {"baz": {"a": "0foobarbaz", "c": "something else"}}, @@ -2666,8 +2699,14 @@ async def test_imp_sync_from_async(checkpointer_name: str) -> None: @NEEDS_CONTEXTVARS +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) -async def test_imp_stream_order(checkpointer_name: str) -> None: +async def test_imp_stream_order( + checkpointer_name: str, checkpoint_during: bool +) -> None: + if not checkpoint_during and "shallow" in checkpointer_name: + pytest.skip("Checkpointing during execution not supported") + async with awith_checkpointer(checkpointer_name) as checkpointer: @task() @@ -2691,7 +2730,12 @@ async def test_imp_stream_order(checkpointer_name: str) -> None: return await fut_baz thread1 = {"configurable": {"thread_id": "1"}} - assert [c async for c in graph.astream({"a": "0"}, thread1)] == [ + assert [ + c + async for c in graph.astream( + {"a": "0"}, thread1, checkpoint_during=checkpoint_during + ) + ] == [ {"foo": {"a": "0foo", "b": "bar"}}, {"bar": {"a": "0foobar", "c": "bark"}}, {"baz": {"a": "0foobarbaz", "c": "something else"}}, @@ -2699,8 +2743,11 @@ async def test_imp_stream_order(checkpointer_name: str) -> None: ] +@pytest.mark.parametrize("checkpoint_during", [True, False]) @pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC) -async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: +async def test_send_dedupe_on_resume( + checkpointer_name: str, checkpoint_during: bool +) -> None: class InterruptOnce: ticks: int = 0 @@ -2751,7 +2798,9 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: async with awith_checkpointer(checkpointer_name) as checkpointer: graph = builder.compile(checkpointer=checkpointer) thread1 = {"configurable": {"thread_id": "1"}} - assert await graph.ainvoke(["0"], thread1, debug=1) == [ + assert await graph.ainvoke( + ["0"], thread1, checkpoint_during=checkpoint_during + ) == [ "0", "1", "3.1", @@ -2763,7 +2812,9 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: assert builder.nodes["2"].runnable.func.ticks == 3 assert builder.nodes["flaky"].runnable.func.ticks == 1 # resume execution - assert await graph.ainvoke(None, thread1, debug=1) == [ + assert await graph.ainvoke( + None, thread1, checkpoint_during=checkpoint_during + ) == [ "0", "1", "3.1", @@ -2780,7 +2831,8 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: assert builder.nodes["flaky"].runnable.func.ticks == 2 # check history history = [c async for c in graph.aget_state_history(thread1)] - assert history == [ + assert len(history) == (6 if checkpoint_during else 2) + expected_history = [ StateSnapshot( values=[ "0", @@ -2916,13 +2968,9 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: name="flaky", path=("__pregel_push", 1), error=None, - interrupts=( - Interrupt( - value="Bahh", resumable=False, ns=None, when="during" - ), - ), + interrupts=(Interrupt(value="Bahh", resumable=False, ns=None),), state=None, - result=["flaky|4"], + result=["flaky|4"] if checkpoint_during else None, ), PregelTask( id=AnyStr(), @@ -3059,6 +3107,11 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None: ), ), ] + if checkpoint_during: + assert history == expected_history + else: + assert history[0] == expected_history[0] + assert history[1] == expected_history[2] @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)