Accept default cache_policy for graph/entrypoint/pregel

This commit is contained in:
Nuno Campos
2025-05-08 16:49:38 -07:00
parent 0e81699fec
commit 1edf5cee89
8 changed files with 82 additions and 27 deletions
@@ -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,
)
+7
View File
@@ -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
+3
View File
@@ -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",
)
+17 -9
View File
@@ -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
):
+26 -15
View File
@@ -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],
+20
View File
@@ -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()
+2 -2
View File
@@ -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
+1 -1
View File
@@ -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], ...]