From 87d57b434a365ce07da6982a9efce35a0f91e9d4 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 6 Nov 2024 08:48:30 -0800 Subject: [PATCH] lib: Split out _match_writes util in PregelLoop --- libs/langgraph/langgraph/pregel/loop.py | 29 +++++++++++++------------ 1 file changed, 15 insertions(+), 14 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 2a9b29599..0dc349151 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -381,20 +381,7 @@ class PregelLoop(LoopProtocol): # if there are pending writes from a previous loop, apply them if self.skip_done_tasks and self.checkpoint_pending_writes: - for tid, k, v in self.checkpoint_pending_writes: - if k in (ERROR, INTERRUPT): - continue - if task := self.tasks.get(tid): - if k == SCHEDULED: - if v == max( - self.checkpoint["versions_seen"] - .get(INTERRUPT, {}) - .values(), - default=None, - ): - self.tasks[tid] = task._replace(scheduled=True) - else: - task.writes.append((k, v)) + self._match_writes(self.tasks) elif not self.skip_done_tasks: # "not skip_done_tasks" only applies to first tick after resuming self.skip_done_tasks = True @@ -429,6 +416,20 @@ class PregelLoop(LoopProtocol): # private + def _match_writes(self, tasks: Mapping[str, PregelExecutableTask]) -> None: + for tid, k, v in self.checkpoint_pending_writes: + if k in (ERROR, INTERRUPT): + continue + if task := tasks.get(tid): + if k == SCHEDULED: + if v == max( + self.checkpoint["versions_seen"].get(INTERRUPT, {}).values(), + default=None, + ): + self.tasks[tid] = task._replace(scheduled=True) + else: + task.writes.append((k, v)) + def _first(self, *, input_keys: Union[str, Sequence[str]]) -> None: # resuming from previous checkpoint requires # - finding a previous checkpoint