mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-31 04:09:49 +02:00
Accept default cache_policy for graph/entrypoint/pregel
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
):
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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], ...]
|
||||
|
||||
Reference in New Issue
Block a user