diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index e6af24163..7c70e6bd1 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -2486,7 +2486,6 @@ class Pregel(PregelProtocol): CONFIG_KEY_RUNNER_SUBMIT, weakref.WeakMethod(loop.submit) ), put_writes=weakref.WeakMethod(loop.put_writes), - schedule_task=weakref.WeakMethod(loop.accept_push), node_finished=config[CONF].get(CONFIG_KEY_NODE_FINISHED), ) # enable subgraph streaming @@ -2529,7 +2528,7 @@ class Pregel(PregelProtocol): [t for t in loop.tasks.values() if not t.writes], timeout=self.step_timeout, get_waiter=get_waiter, - match_cached_writes=loop.match_cached_writes, + schedule_task=loop.accept_push, ): # emit output yield from output() @@ -2799,7 +2798,6 @@ class Pregel(PregelProtocol): CONFIG_KEY_RUNNER_SUBMIT, weakref.WeakMethod(loop.submit) ), put_writes=weakref.WeakMethod(loop.put_writes), - schedule_task=weakref.WeakMethod(loop.accept_push), use_astream=do_stream, node_finished=config[CONF].get(CONFIG_KEY_NODE_FINISHED), ) @@ -2833,7 +2831,7 @@ class Pregel(PregelProtocol): [t for t in loop.tasks.values() if not t.writes], timeout=self.step_timeout, get_waiter=get_waiter, - match_cached_writes=loop.amatch_cached_writes, + schedule_task=loop.aaccept_push, ): # emit output for o in output(): diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 204dbdd3e..a10c06d72 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -1079,6 +1079,11 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): matched.append(task) return matched + def accept_push(self, task, write_idx, call=None): + if pushed := super().accept_push(task, write_idx, call): + self.match_cached_writes() + return pushed + def put_writes(self, task_id: str, writes: WritesT) -> None: """Put writes for a task, to be read by the next tick.""" super().put_writes(task_id, writes) @@ -1268,6 +1273,13 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): matched.append(task) return matched + async def aaccept_push( + self, task: PregelExecutableTask, write_idx: int, call: Optional[Call] = None + ) -> Optional[PregelExecutableTask]: + if pushed := super().accept_push(task, write_idx, call): + await self.amatch_cached_writes() + return pushed + def put_writes(self, task_id: str, writes: WritesT) -> None: """Put writes for a task, to be read by the next tick.""" super().put_writes(task_id, writes) diff --git a/libs/langgraph/langgraph/pregel/runner.py b/libs/langgraph/langgraph/pregel/runner.py index b0b2f31a9..a956b72b2 100644 --- a/libs/langgraph/langgraph/pregel/runner.py +++ b/libs/langgraph/langgraph/pregel/runner.py @@ -39,7 +39,7 @@ from langgraph.types import ( PregelScratchpad, RetryPolicy, ) -from langgraph.utils.future import chain_future +from langgraph.utils.future import chain_future, run_coroutine_threadsafe F = TypeVar("F", concurrent.futures.Future, asyncio.Future) E = TypeVar("E", threading.Event, asyncio.Event) @@ -119,12 +119,6 @@ class PregelRunner: *, submit: weakref.ref[Submit], put_writes: weakref.ref[Callable[[str, Sequence[tuple[str, Any]]], None]], - schedule_task: weakref.ref[ - Callable[ - [PregelExecutableTask, int, Optional[Call]], - Optional[PregelExecutableTask], - ] - ], use_astream: bool = False, node_finished: Optional[Callable[[str], None]] = None, ) -> None: @@ -132,7 +126,6 @@ class PregelRunner: self.put_writes = put_writes self.use_astream = use_astream self.node_finished = node_finished - self.schedule_task = schedule_task def tick( self, @@ -142,9 +135,10 @@ class PregelRunner: timeout: Optional[float] = None, retry_policy: Optional[Sequence[RetryPolicy]] = None, get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None, - match_cached_writes: Optional[ - Callable[[], Sequence[PregelExecutableTask]] - ] = None, + schedule_task: Callable[ + [PregelExecutableTask, int, Optional[Call]], + Optional[PregelExecutableTask], + ], ) -> Iterator[None]: tasks = tuple(tasks) futures = FuturesDict( @@ -169,8 +163,7 @@ class PregelRunner: weakref.ref(t), retry=retry_policy, futures=weakref.ref(futures), - schedule_task=self.schedule_task, - match_cached_writes=match_cached_writes, + schedule_task=schedule_task, submit=self.submit, reraise=reraise, ), @@ -212,8 +205,7 @@ class PregelRunner: weakref.ref(t), retry=retry_policy, futures=weakref.ref(futures), - schedule_task=self.schedule_task, - match_cached_writes=match_cached_writes, + schedule_task=schedule_task, submit=self.submit, reraise=reraise, ), @@ -277,9 +269,10 @@ class PregelRunner: timeout: Optional[float] = None, retry_policy: Optional[Sequence[RetryPolicy]] = None, get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None, - match_cached_writes: Optional[ - Callable[[], Awaitable[Sequence[PregelExecutableTask]]] - ] = None, + schedule_task: Callable[ + [PregelExecutableTask, int, Optional[Call]], + Awaitable[Optional[PregelExecutableTask]], + ], ) -> AsyncIterator[None]: loop = asyncio.get_event_loop() tasks = tuple(tasks) @@ -307,8 +300,7 @@ class PregelRunner: stream=self.use_astream, retry=retry_policy, futures=weakref.ref(futures), - schedule_task=self.schedule_task, - match_cached_writes=match_cached_writes, + schedule_task=schedule_task, submit=self.submit, reraise=reraise, loop=loop, @@ -355,8 +347,7 @@ class PregelRunner: retry=retry_policy, stream=self.use_astream, futures=weakref.ref(futures), - schedule_task=self.schedule_task, - match_cached_writes=match_cached_writes, + schedule_task=schedule_task, submit=self.submit, reraise=reraise, loop=loop, @@ -535,12 +526,9 @@ def _call( cache_policy: Optional[CachePolicy] = None, callbacks: Callbacks = None, futures: weakref.ref[FuturesDict], - schedule_task: weakref.ref[ - Callable[ - [PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask] - ] + schedule_task: Callable[ + [PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask] ], - match_cached_writes: Optional[Callable[[], Sequence[PregelExecutableTask]]], submit: weakref.ref[Submit], reraise: bool, ) -> concurrent.futures.Future[Any]: @@ -551,13 +539,11 @@ def _call( # schedule PUSH tasks, collect futures scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr] # schedule the next task, if the callback returns one - if next_task := schedule_task()( # type: ignore[misc] + if next_task := schedule_task( # type: ignore[misc] task(), # type: ignore[arg-type] scratchpad.call_counter(), Call(func, input, retry=retry, cache_policy=cache_policy, callbacks=callbacks), ): - if match_cached_writes: - match_cached_writes() if fut := next( ( f @@ -595,7 +581,6 @@ def _call( retry=retry, callbacks=callbacks, schedule_task=schedule_task, - match_cached_writes=match_cached_writes, submit=submit, reraise=reraise, ), @@ -622,106 +607,140 @@ def _acall( callbacks: Callbacks = None, # injected dependencies futures: weakref.ref[FuturesDict], - schedule_task: weakref.ref[ - Callable[ - [PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask] - ] + schedule_task: Callable[ + [PregelExecutableTask, int, Optional[Call]], + Awaitable[Optional[PregelExecutableTask]], ], - match_cached_writes: Optional[ - Callable[[], Awaitable[Sequence[PregelExecutableTask]]] - ] = None, submit: weakref.ref[Submit], loop: asyncio.AbstractEventLoop, reraise: bool = False, stream: bool = False, ) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]: - fut: Optional[asyncio.Future] = None - # schedule PUSH tasks, collect futures - scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr] - # schedule the next task, if the callback returns one - if next_task := schedule_task()( # type: ignore[misc] - task(), # type: ignore[arg-type] - scratchpad.call_counter(), - Call(func, input, retry=retry, cache_policy=cache_policy, callbacks=callbacks), - ): - if fut := next( - ( - f - for f, t in futures().items() # type: ignore[union-attr] - if t is not None and t == next_task.id - ), - None, - ): - # if the parent task was retried, - # the next task might already be running - - pass - elif next_task.writes: - # if it already ran, return the result - fut = asyncio.Future(loop=loop) - ret = next((v for c, v in next_task.writes if c == RETURN), MISSING) - if ret is not MISSING: - fut.set_result(ret) - elif exc := next((v for c, v in next_task.writes if c == ERROR), None): - fut.set_exception( - exc if isinstance(exc, BaseException) else Exception(exc) - ) - else: - fut.set_result(None) - futures()[fut] = next_task # type: ignore[index] - else: - # schedule the next task - fut = cast( - asyncio.Future, - submit()( # type: ignore[misc] - arun_with_retry, - next_task, - retry, - stream=stream, - match_cached_writes=match_cached_writes, - configurable={ - CONFIG_KEY_CALL: partial( - _acall, - weakref.ref(next_task), - stream=stream, - futures=futures, - schedule_task=schedule_task, - match_cached_writes=match_cached_writes, - submit=submit, - loop=loop, - reraise=reraise, - ), - }, - __name__=task().name, # type: ignore[union-attr] - __cancel_on_exit__=True, - __reraise_on_exit__=reraise, - # starting a new task in the next tick ensures - # updates from this tick are committed/streamed first - __next_tick__=True, - ), - ) - futures()[fut] = next_task # type: ignore[index] - - fut = cast(Union[asyncio.Future, concurrent.futures.Future], fut) # return a chained future to ensure commit() callback is called # before the returned future is resolved, to ensure stream order etc try: in_async = asyncio.current_task() is not None except RuntimeError: in_async = False - # if in async context return an async future - # otherwise return a chained sync future + # if in async context return an async future, otherwise return a sync future if in_async: - if isinstance(fut, asyncio.Task): - sfut: Union[asyncio.Future[Any], concurrent.futures.Future[Any]] = ( - asyncio.Future(loop=loop) - ) - loop.call_soon_threadsafe(chain_future, fut, sfut) - return sfut - else: - # already wrapped in a future - return fut + fut: Union[asyncio.Future[Any], concurrent.futures.Future[Any]] = ( + asyncio.Future(loop=loop) + ) else: - sfut = concurrent.futures.Future() - loop.call_soon_threadsafe(chain_future, fut, sfut) - return sfut + fut = concurrent.futures.Future() + # schedule the next task + run_coroutine_threadsafe( + _acall_impl( + fut, + task, + func, + input, + retry=retry, + cache_policy=cache_policy, + callbacks=callbacks, + futures=futures, + schedule_task=schedule_task, + submit=submit, + loop=loop, + reraise=reraise, + stream=stream, + ), + loop, + lazy=False, + ) + return fut + + +async def _acall_impl( + destination: Union[asyncio.Future[Any], concurrent.futures.Future[Any]], + task: weakref.ref[PregelExecutableTask], + func: Callable[[Any], Union[Awaitable[Any], Any]], + input: Any, + *, + retry: Optional[Sequence[RetryPolicy]] = None, + cache_policy: Optional[CachePolicy] = None, + callbacks: Callbacks = None, + # injected dependencies + futures: weakref.ref[FuturesDict[asyncio.Future, asyncio.Event]], + schedule_task: Callable[ + [PregelExecutableTask, int, Optional[Call]], + Awaitable[Optional[PregelExecutableTask]], + ], + submit: weakref.ref[Submit], + loop: asyncio.AbstractEventLoop, + reraise: bool = False, + stream: bool = False, +) -> None: + try: + fut: Optional[asyncio.Future] = None + # schedule PUSH tasks, collect futures + scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr] + # schedule the next task, if the callback returns one + if next_task := await schedule_task( # type: ignore[misc] + task(), # type: ignore[arg-type] + scratchpad.call_counter(), + Call( + func, input, retry=retry, cache_policy=cache_policy, callbacks=callbacks + ), + ): + if fut := next( + ( + f + for f, t in futures().items() # type: ignore[union-attr] + if t is not None and t == next_task.id + ), + None, + ): + # if the parent task was retried, + # the next task might already be running + pass + elif next_task.writes: + # if it already ran, return the result + fut = asyncio.Future(loop=loop) + ret = next((v for c, v in next_task.writes if c == RETURN), MISSING) + if ret is not MISSING: + fut.set_result(ret) + elif exc := next((v for c, v in next_task.writes if c == ERROR), None): + fut.set_exception( + exc if isinstance(exc, BaseException) else Exception(exc) + ) + else: + fut.set_result(None) + futures()[fut] = next_task # type: ignore[index] + else: + # schedule the next task + fut = cast( + asyncio.Future, + submit()( # type: ignore[misc] + arun_with_retry, + next_task, + retry, + stream=stream, + configurable={ + CONFIG_KEY_CALL: partial( + _acall, + weakref.ref(next_task), + stream=stream, + futures=futures, + schedule_task=schedule_task, + submit=submit, + loop=loop, + reraise=reraise, + ), + }, + __name__=task().name, # type: ignore[union-attr] + __cancel_on_exit__=True, + __reraise_on_exit__=reraise, + # starting a new task in the next tick ensures + # updates from this tick are committed/streamed first + __next_tick__=True, + ), + ) + futures()[fut] = next_task # type: ignore[index] + if fut is not None: + chain_future(fut, destination) + else: + destination.set_exception(RuntimeError("Task not scheduled")) + except Exception as exc: + destination.set_exception(exc)