From 1edf5cee89d6a322722ef133031dc7e07d791e5e Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 2 May 2025 10:59:34 -0700 Subject: [PATCH] Accept default cache_policy for graph/entrypoint/pregel --- libs/langgraph/langgraph/func/__init__.py | 6 +++ libs/langgraph/langgraph/graph/graph.py | 7 ++++ libs/langgraph/langgraph/graph/state.py | 3 ++ libs/langgraph/langgraph/pregel/__init__.py | 26 ++++++++----- libs/langgraph/langgraph/pregel/algo.py | 41 +++++++++++++-------- libs/langgraph/langgraph/pregel/loop.py | 20 ++++++++++ libs/langgraph/langgraph/pregel/retry.py | 4 +- libs/langgraph/langgraph/types.py | 2 +- 8 files changed, 82 insertions(+), 27 deletions(-) diff --git a/libs/langgraph/langgraph/func/__init__.py b/libs/langgraph/langgraph/func/__init__.py index dd6f5f4a5..b2a4081fb 100644 --- a/libs/langgraph/langgraph/func/__init__.py +++ b/libs/langgraph/langgraph/func/__init__.py @@ -323,11 +323,15 @@ class entrypoint: store: Optional[BaseStore] = None, cache: Optional[BaseCache] = None, config_schema: Optional[type[Any]] = None, + cache_policy: Optional[CachePolicy] = None, + retry: Union[RetryPolicy, Sequence[RetryPolicy]] = (), ) -> None: """Initialize the entrypoint decorator.""" self.checkpointer = checkpointer self.store = store self.cache = cache + self.cache_policy = cache_policy + self.retry = retry self.config_schema = config_schema @dataclass(**_DC_KWARGS) @@ -458,5 +462,7 @@ class entrypoint: checkpointer=self.checkpointer, store=self.store, cache=self.cache, + cache_policy=self.cache_policy, + retry_policy=self.retry, config_type=self.config_schema, ) diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index e8250ba57..6c4c069b7 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -14,6 +14,7 @@ from typing import ( from langchain_core.runnables import Runnable from typing_extensions import Self +from langgraph.cache.base import BaseCache from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.constants import ( EMPTY_SEQ, @@ -28,6 +29,7 @@ from langgraph.graph.branch import Branch from langgraph.pregel import Channel, Pregel from langgraph.pregel.read import PregelNode from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry +from langgraph.store.base import BaseStore from langgraph.types import All, Checkpointer from langgraph.utils.runnable import RunnableLike, coerce_to_runnable @@ -316,6 +318,9 @@ class Graph: interrupt_after: Optional[Union[All, list[str]]] = None, debug: bool = False, name: Optional[str] = None, + *, + cache: Optional[BaseCache] = None, + store: Optional[BaseStore] = None, ) -> "CompiledGraph": """Compiles the graph into a `CompiledGraph` object. @@ -364,6 +369,8 @@ class Graph: auto_validate=False, debug=debug, name=name or "LangGraph", + cache=cache, + store=store, ) # attach nodes, edges, and branches diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index c86fc6b14..45138de0f 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -26,6 +26,7 @@ from pydantic import BaseModel from typing_extensions import Self from langgraph._api.deprecation import LangGraphDeprecationWarning +from langgraph.cache.base import BaseCache from langgraph.channels.base import BaseChannel from langgraph.channels.binop import BinaryOperatorAggregate from langgraph.channels.dynamic_barrier_value import ( @@ -571,6 +572,7 @@ class StateGraph(Graph): self, checkpointer: Checkpointer = None, *, + cache: Optional[BaseCache] = None, store: Optional[BaseStore] = None, interrupt_before: Optional[Union[All, list[str]]] = None, interrupt_after: Optional[Union[All, list[str]]] = None, @@ -655,6 +657,7 @@ class StateGraph(Graph): auto_validate=False, debug=debug, store=store, + cache=cache, name=name or "LangGraph", ) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 6d8e7adb1..6aabc9a61 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -104,6 +104,7 @@ from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore from langgraph.types import ( All, + CachePolicy, Checkpointer, Interrupt, LoopProtocol, @@ -500,8 +501,12 @@ class Pregel(PregelProtocol): cache: BaseCache | None = None """Cache to use for storing node results. Defaults to None.""" - retry_policy: Sequence[RetryPolicy] | None = None - """Retry policies to use when running tasks. Set to None to disable.""" + retry_policy: Sequence[RetryPolicy] = () + """Retry policies to use when running tasks. Empty set disables retries.""" + + cache_policy: CachePolicy | None = None + """Cache policy to use for all nodes. Can be overridden by individual nodes. + Defaults to None.""" config_type: type[Any] | None = None @@ -531,7 +536,8 @@ class Pregel(PregelProtocol): checkpointer: BaseCheckpointSaver | None = None, store: BaseStore | None = None, cache: BaseCache | None = None, - retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + retry_policy: RetryPolicy | Sequence[RetryPolicy] = (), + cache_policy: CachePolicy | None = None, config_type: type[Any] | None = None, input_model: type[BaseModel] | None = None, config: RunnableConfig | None = None, @@ -552,10 +558,10 @@ class Pregel(PregelProtocol): self.checkpointer = checkpointer self.store = store self.cache = cache - if isinstance(retry_policy, RetryPolicy): - self.retry_policy: Sequence[RetryPolicy] = (retry_policy,) - else: - self.retry_policy = retry_policy + self.retry_policy = ( + (retry_policy,) if isinstance(retry_policy, RetryPolicy) else retry_policy + ) + self.cache_policy = cache_policy self.config_type = config_type self.input_model = input_model self.config = config @@ -2465,6 +2471,8 @@ class Pregel(PregelProtocol): else config[CONF].get(CONFIG_KEY_CHECKPOINT_DURING, True), trigger_to_nodes=self.trigger_to_nodes, migrate_checkpoint=self._migrate_checkpoint, + retry_policy=self.retry_policy, + cache_policy=self.cache_policy, ) as loop: # create runner runner = PregelRunner( @@ -2514,7 +2522,6 @@ class Pregel(PregelProtocol): for _ in runner.tick( [t for t in loop.tasks.values() if not t.writes], timeout=self.step_timeout, - retry_policy=self.retry_policy, get_waiter=get_waiter, match_cached_writes=loop.match_cached_writes, ): @@ -2777,6 +2784,8 @@ class Pregel(PregelProtocol): else config[CONF].get(CONFIG_KEY_CHECKPOINT_DURING, True), trigger_to_nodes=self.trigger_to_nodes, migrate_checkpoint=self._migrate_checkpoint, + retry_policy=self.retry_policy, + cache_policy=self.cache_policy, ) as loop: # create runner runner = PregelRunner( @@ -2817,7 +2826,6 @@ class Pregel(PregelProtocol): async for _ in runner.atick( [t for t in loop.tasks.values() if not t.writes], timeout=self.step_timeout, - retry_policy=self.retry_policy, get_waiter=get_waiter, # TODO pass match_cached_writes ): diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index c0b15fd18..7449cd01d 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -424,6 +424,8 @@ def prepare_next_tasks( manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None, updated_channels: Optional[set[str]] = None, + retry_policy: Sequence[RetryPolicy] = (), + cache_policy: Optional[CachePolicy] = None, ) -> Union[dict[str, PregelTask], dict[str, PregelExecutableTask]]: """Prepare the set of tasks that will make up the next Pregel step. @@ -474,6 +476,8 @@ def prepare_next_tasks( checkpointer=checkpointer, manager=manager, input_cache=input_cache, + cache_policy=cache_policy, + retry_policy=retry_policy, ): tasks.append(task) @@ -517,6 +521,8 @@ def prepare_next_tasks( checkpointer=checkpointer, manager=manager, input_cache=input_cache, + cache_policy=cache_policy, + retry_policy=retry_policy, ): tasks.append(task) return {t.id: t for t in tasks} @@ -543,6 +549,8 @@ def prepare_single_task( checkpointer: Optional[BaseCheckpointSaver] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, input_cache: Optional[dict[INPUT_CACHE_KEY_TYPE, Any]] = None, + cache_policy: Optional[CachePolicy] = None, + retry_policy: Sequence[RetryPolicy] = (), ) -> Union[None, PregelTask, PregelExecutableTask]: """Prepares a single task for the next Pregel step, given a task path, which uniquely identifies a PUSH or PULL task within the graph.""" @@ -585,8 +593,9 @@ def prepare_single_task( assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}" if for_execution: writes: deque[tuple[str, Any]] = deque() - if call.cache_policy: - args_key = call.cache_policy.key(*call.input[0], **call.input[1]) + cache_policy = call.cache_policy or cache_policy + if cache_policy: + args_key = cache_policy.key(*call.input[0], **call.input[1]) cache_key: Optional[CacheKey] = CacheKey( xxh3_128_hexdigest( b"".join( @@ -602,8 +611,8 @@ def prepare_single_task( ) ) ), - call.cache_policy.ttl, - call.cache_policy.refresh, + cache_policy.ttl, + cache_policy.refresh, ) else: cache_key = None @@ -651,7 +660,7 @@ def prepare_single_task( }, ), triggers, - call.retry, + call.retry or retry_policy, cache_key, task_id, task_path, @@ -714,8 +723,9 @@ def prepare_single_task( if proc.metadata: metadata.update(proc.metadata) writes = deque() - if proc.cache_policy: - args_key = proc.cache_policy.key(packet.arg) + cache_policy = proc.cache_policy or cache_policy + if cache_policy: + args_key = cache_policy.key(packet.arg) cache_key = CacheKey( xxh3_128_hexdigest( b"".join( @@ -731,8 +741,8 @@ def prepare_single_task( ) ) ), - proc.cache_policy.ttl, - proc.cache_policy.refresh, + cache_policy.ttl, + cache_policy.refresh, ) else: cache_key = None @@ -784,7 +794,7 @@ def prepare_single_task( }, ), triggers, - proc.retry_policy, + proc.retry_policy or retry_policy, cache_key, task_id, task_path, @@ -852,8 +862,9 @@ def prepare_single_task( if proc.metadata: metadata.update(proc.metadata) writes = deque() - if proc.cache_policy: - args_key = proc.cache_policy.key(val) + cache_policy = proc.cache_policy or cache_policy + if cache_policy: + args_key = cache_policy.key(val) cache_key = CacheKey( xxh3_128_hexdigest( b"".join( @@ -869,8 +880,8 @@ def prepare_single_task( ) ) ), - proc.cache_policy.ttl, - proc.cache_policy.refresh, + cache_policy.ttl, + cache_policy.refresh, ) else: cache_key = None @@ -934,7 +945,7 @@ def prepare_single_task( }, ), triggers, - proc.retry_policy, + proc.retry_policy or retry_policy, cache_key, task_id, task_path[:3], diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 95fd863b7..00feb2650 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -118,10 +118,12 @@ from langgraph.pregel.utils import get_new_channel_versions, is_xxh3_128_hexdige from langgraph.store.base import BaseStore from langgraph.types import ( All, + CachePolicy, Command, LoopProtocol, PregelExecutableTask, PregelScratchpad, + RetryPolicy, StreamChunk, StreamProtocol, ) @@ -162,6 +164,8 @@ class PregelLoop(LoopProtocol): interrupt_before: Union[All, Sequence[str]] checkpoint_during: bool debug: bool + retry_policy: Sequence[RetryPolicy] + cache_policy: Optional[CachePolicy] checkpointer_get_next_version: GetNextVersion checkpointer_put_writes: Optional[Callable[[RunnableConfig, WritesT, str], Any]] @@ -220,6 +224,8 @@ class PregelLoop(LoopProtocol): input_model: Optional[type[BaseModel]] = None, debug: bool = False, migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None, + retry_policy: Sequence[RetryPolicy] = (), + cache_policy: Optional[CachePolicy] = None, checkpoint_during: bool = True, ) -> None: super().__init__( @@ -247,6 +253,8 @@ class PregelLoop(LoopProtocol): ) self._migrate_checkpoint = migrate_checkpoint self.trigger_to_nodes = trigger_to_nodes + self.retry_policy = retry_policy + self.cache_policy = cache_policy self.checkpoint_during = checkpoint_during self.debug = debug if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]: @@ -421,6 +429,8 @@ class PregelLoop(LoopProtocol): store=self.store, checkpointer=self.checkpointer, manager=self.manager, + retry_policy=self.retry_policy, + cache_policy=self.cache_policy, ), ): # don't start if we should interrupt *before* the new task @@ -550,6 +560,8 @@ class PregelLoop(LoopProtocol): checkpointer=self.checkpointer, trigger_to_nodes=self.trigger_to_nodes, updated_channels=updated_channels, + retry_policy=self.retry_policy, + cache_policy=self.cache_policy, ) self.to_interrupt = [] @@ -989,6 +1001,8 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): input_model: Optional[type[BaseModel]] = None, debug: bool = False, migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None, + retry_policy: Sequence[RetryPolicy] = (), + cache_policy: Optional[CachePolicy] = None, checkpoint_during: bool = True, ) -> None: super().__init__( @@ -1009,6 +1023,8 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): debug=debug, migrate_checkpoint=migrate_checkpoint, trigger_to_nodes=trigger_to_nodes, + retry_policy=retry_policy, + cache_policy=cache_policy, checkpoint_during=checkpoint_during, ) self.stack = ExitStack() @@ -1169,6 +1185,8 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): input_model: Optional[type[BaseModel]] = None, debug: bool = False, migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None, + retry_policy: Sequence[RetryPolicy] = (), + cache_policy: Optional[CachePolicy] = None, checkpoint_during: bool = True, ) -> None: super().__init__( @@ -1189,6 +1207,8 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): debug=debug, migrate_checkpoint=migrate_checkpoint, trigger_to_nodes=trigger_to_nodes, + retry_policy=retry_policy, + cache_policy=cache_policy, checkpoint_during=checkpoint_during, ) self.stack = AsyncExitStack() diff --git a/libs/langgraph/langgraph/pregel/retry.py b/libs/langgraph/langgraph/pregel/retry.py index 78c2f8b0c..be5c0a91c 100644 --- a/libs/langgraph/langgraph/pregel/retry.py +++ b/libs/langgraph/langgraph/pregel/retry.py @@ -61,7 +61,7 @@ def run_with_retry( except Exception as exc: if SUPPORTS_EXC_NOTES: exc.add_note(f"During task with name '{task.name}' and id '{task.id}'") - if retry_policy is None: + if not retry_policy: raise # Check which retry policy applies to this exception @@ -149,7 +149,7 @@ async def arun_with_retry( except Exception as exc: if SUPPORTS_EXC_NOTES: exc.add_note(f"During task with name '{task.name}' and id '{task.id}'") - if retry_policies is None: + if not retry_policies: raise # Check which retry policy applies to this exception diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 9fbb92217..79da90b67 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -206,7 +206,7 @@ class PregelExecutableTask: writes: deque[tuple[str, Any]] config: RunnableConfig triggers: Sequence[str] - retry_policy: Optional[Sequence[RetryPolicy]] + retry_policy: Sequence[RetryPolicy] cache_key: Optional[CacheKey] id: str path: tuple[Union[str, int, tuple], ...]