Improve how we match cached writes for async imperative tasks

- remove match_cached_writes from PregelRunner args (now called by PregelLoop internally)
- this will be helpful when implementing distributed runner classes
This commit is contained in:
Nuno Campos
2025-05-14 11:25:49 -07:00
parent c7691081d1
commit f51e5e2bd7
3 changed files with 153 additions and 124 deletions
+2 -4
View File
@@ -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():
+12
View File
@@ -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)
+139 -120
View File
@@ -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)