diff --git a/libs/checkpoint/langgraph/cache/base/__init__.py b/libs/checkpoint/langgraph/cache/base/__init__.py index abf14b1b6..8f8660cca 100644 --- a/libs/checkpoint/langgraph/cache/base/__init__.py +++ b/libs/checkpoint/langgraph/cache/base/__init__.py @@ -8,6 +8,8 @@ from langgraph.checkpoint.serde.base import SerializerProtocol from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer ValueT = TypeVar("ValueT") +Namespace = tuple[str, ...] +FullKey = tuple[Namespace, str] class BaseCache(ABC, Generic[ValueT]): @@ -20,25 +22,25 @@ class BaseCache(ABC, Generic[ValueT]): self.serde = serde or self.serde @abstractmethod - def get(self, keys: Sequence[str]) -> dict[str, ValueT]: + def get(self, keys: Sequence[FullKey]) -> dict[FullKey, ValueT]: """Get the cached values for the given keys.""" @abstractmethod - async def aget(self, keys: Sequence[str]) -> dict[str, ValueT]: + async def aget(self, keys: Sequence[FullKey]) -> dict[FullKey, ValueT]: """Asynchronously get the cached values for the given keys.""" @abstractmethod - def set(self, mapping: Mapping[str, tuple[ValueT, int | None]]) -> None: + def set(self, pairs: Mapping[FullKey, tuple[ValueT, int | None]]) -> None: """Set the cached values for the given keys and TTLs.""" @abstractmethod - async def aset(self, mapping: Mapping[str, tuple[ValueT, int | None]]) -> None: + async def aset(self, pairs: Mapping[FullKey, tuple[ValueT, int | None]]) -> None: """Asynchronously set the cached values for the given keys and TTLs.""" @abstractmethod - def delete(self, keys: Sequence[str]) -> None: + def delete(self, keys: Sequence[Namespace]) -> None: """Delete the cached values for the given keys.""" @abstractmethod - async def adelete(self, keys: Sequence[str]) -> None: + async def adelete(self, keys: Sequence[Namespace]) -> None: """Asynchronously delete the cached values for the given keys.""" diff --git a/libs/checkpoint/langgraph/cache/file/__init__.py b/libs/checkpoint/langgraph/cache/file/__init__.py index 080d24c66..29231b206 100644 --- a/libs/checkpoint/langgraph/cache/file/__init__.py +++ b/libs/checkpoint/langgraph/cache/file/__init__.py @@ -6,7 +6,7 @@ import sqlite3 import threading from collections.abc import Mapping, Sequence -from langgraph.cache.base import BaseCache +from langgraph.cache.base import BaseCache, FullKey, Namespace from langgraph.checkpoint.serde.base import SerializerProtocol @@ -33,40 +33,49 @@ class FileCache(BaseCache): # Schema: key -> (expiry, encoding, value) self._conn.execute( """CREATE TABLE IF NOT EXISTS cache ( - key TEXT PRIMARY KEY, + ns TEXT, + key TEXT, expiry REAL, encoding TEXT NOT NULL, - val BLOB NOT NULL + val BLOB NOT NULL, + PRIMARY KEY (ns, key) )""" ) self._conn.commit() - def get(self, keys: Sequence[str]) -> dict[str, bytes]: + def get(self, keys: Sequence[FullKey]) -> dict[FullKey, bytes]: """Get the cached values for the given keys.""" with self._lock, self._conn: now = datetime.datetime.now(datetime.timezone.utc).timestamp() if not keys: return {} - placeholders = ",".join("?" for _ in keys) + placeholders = ",".join("(?, ?)" for _ in keys) + params: list[str] = [] + for ns_tuple, key in keys: + params.extend((",".join(ns_tuple), key)) cursor = self._conn.execute( - f"SELECT key, expiry, encoding, val FROM cache WHERE key IN ({placeholders})", - tuple(keys), + f"SELECT ns, key, expiry, encoding, val FROM cache WHERE (ns, key) IN ({placeholders})", + tuple(params), ) - values: dict[str, bytes] = {} + values: dict[FullKey, bytes] = {} rows = cursor.fetchall() - for key, expiry, encoding, raw in rows: + for ns, key, expiry, encoding, raw in rows: if expiry is not None and now > expiry: # purge expired entry - self._conn.execute("DELETE FROM cache WHERE key = ?", (key,)) + self._conn.execute( + "DELETE FROM cache WHERE (ns, key) = (?, ?)", (ns, key) + ) continue - values[key] = self.serde.loads_typed((encoding, raw)) + values[(tuple(ns.split(",")), key)] = self.serde.loads_typed( + (encoding, raw) + ) return values - async def aget(self, keys: Sequence[str]) -> dict[str, bytes]: + async def aget(self, keys: Sequence[FullKey]) -> dict[FullKey, bytes]: """Asynchronously get the cached values for the given keys.""" return await asyncio.to_thread(self.get, keys) - def set(self, mapping: Mapping[str, tuple[bytes, int | None]]) -> None: + def set(self, mapping: Mapping[FullKey, tuple[bytes, int | None]]) -> None: """Set the cached values for the given keys and TTLs.""" with self._lock, self._conn: now = datetime.datetime.now(datetime.timezone.utc) @@ -78,26 +87,27 @@ class FileCache(BaseCache): expiry = None encoding, raw = self.serde.dumps_typed(value) self._conn.execute( - "INSERT OR REPLACE INTO cache (key, expiry, encoding, val) VALUES (?, ?, ?, ?)", - (key, expiry, encoding, raw), + "INSERT OR REPLACE INTO cache (ns, key, expiry, encoding, val) VALUES (?, ?, ?, ?, ?)", + (",".join(key[0]), key[1], expiry, encoding, raw), ) - async def aset(self, mapping: Mapping[str, tuple[bytes, int | None]]) -> None: + async def aset(self, mapping: Mapping[FullKey, tuple[bytes, int | None]]) -> None: """Asynchronously set the cached values for the given keys and TTLs.""" await asyncio.to_thread(self.set, mapping) - def delete(self, keys: Sequence[str]) -> None: - """Delete the cached values for the given keys.""" + def delete(self, keys: Sequence[Namespace]) -> None: + """Delete the cached values for the given namespaces.""" if not keys: return with self._lock, self._conn: placeholders = ",".join("?" for _ in keys) self._conn.execute( - f"DELETE FROM cache WHERE key IN ({placeholders})", tuple(keys) + f"DELETE FROM cache WHERE (ns) IN ({placeholders})", + tuple(",".join(key) for key in keys), ) - async def adelete(self, keys: Sequence[str]) -> None: - """Asynchronously delete the cached values for the given keys.""" + async def adelete(self, keys: Sequence[Namespace]) -> None: + """Asynchronously delete the cached values for the given namespaces.""" await asyncio.to_thread(self.delete, keys) def __del__(self) -> None: diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index c3ca2821b..e361c8f1a 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -46,6 +46,10 @@ TASKS = sys.intern("__pregel_tasks") RETURN = sys.intern("__return__") # for writes of a task where we simply record the return value +# --- Reserved cache namespaces --- +CACHE_NS_WRITES = sys.intern("__pregel_ns_writes") +# cache namespace for node writes + # --- Reserved config.configurable keys --- CONFIG_KEY_SEND = sys.intern("__pregel_send") # holds the `write` function that accepts writes to state/edges/reserved keys diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 4d7d3075c..a6b9ca52e 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -33,6 +33,7 @@ from langgraph.checkpoint.base import ( V, ) from langgraph.constants import ( + CACHE_NS_WRITES, CONF, CONFIG_KEY_CHECKPOINT_ID, CONFIG_KEY_CHECKPOINT_MAP, @@ -79,7 +80,7 @@ from langgraph.types import ( PregelTask, RetryPolicy, ) -from langgraph.utils.config import merge_configs, patch_config, recast_checkpoint_ns +from langgraph.utils.config import merge_configs, patch_config GetNextVersion = Callable[[Optional[V], BaseChannel], V] SUPPORTS_EXC_NOTES = sys.version_info >= (3, 11) @@ -601,19 +602,12 @@ def prepare_single_task( if cache_policy: args_key = cache_policy.key_func(*call.input[0], **call.input[1]) cache_key: Optional[CacheKey] = CacheKey( + ( + CACHE_NS_WRITES, + (identifier(call.func) or "__dynamic__"), + ), xxh3_128_hexdigest( - b"".join( - ( - b"__pregel_cache", - recast_checkpoint_ns(parent_ns).encode() - if parent_ns - else b"", - (identifier(call.func) or "__dynamic__").encode(), - args_key.encode() - if isinstance(args_key, str) - else args_key, - ) - ) + args_key.encode() if isinstance(args_key, str) else args_key, ), cache_policy.ttl, cache_policy.refresh, @@ -731,19 +725,13 @@ def prepare_single_task( if cache_policy: args_key = cache_policy.key_func(packet.arg) cache_key = CacheKey( + ( + CACHE_NS_WRITES, + (identifier(proc) or "__dynamic__"), + packet.node, + ), xxh3_128_hexdigest( - b"".join( - ( - b"__pregel_cache", - recast_checkpoint_ns(parent_ns).encode() - if parent_ns - else b"", - packet.node.encode(), - args_key.encode() - if isinstance(args_key, str) - else args_key, - ) - ) + args_key.encode() if isinstance(args_key, str) else args_key, ), cache_policy.ttl, cache_policy.refresh, @@ -870,19 +858,15 @@ def prepare_single_task( if cache_policy: args_key = cache_policy.key_func(val) cache_key = CacheKey( + ( + CACHE_NS_WRITES, + (identifier(proc) or "__dynamic__"), + name, + ), xxh3_128_hexdigest( - b"".join( - ( - b"__pregel_cache", - recast_checkpoint_ns(parent_ns).encode() - if parent_ns - else b"", - name.encode(), - args_key.encode() - if isinstance(args_key, str) - else args_key, - ) - ) + args_key.encode() + if isinstance(args_key, str) + else args_key, ), cache_policy.ttl, cache_policy.refresh, diff --git a/libs/langgraph/langgraph/pregel/call.py b/libs/langgraph/langgraph/pregel/call.py index d9abc6cf3..d18d09b4b 100644 --- a/libs/langgraph/langgraph/pregel/call.py +++ b/libs/langgraph/langgraph/pregel/call.py @@ -75,6 +75,16 @@ def _whichmodule(obj: Any, name: str) -> Optional[str]: def identifier(obj: Any, name: Optional[str] = None) -> Optional[str]: + """Return the module and name of an object.""" + from langgraph.pregel.read import PregelNode + from langgraph.utils.runnable import RunnableCallable, RunnableSeq + + if isinstance(obj, PregelNode): + obj = obj.bound + if isinstance(obj, RunnableSeq): + obj = obj.steps[0] + if isinstance(obj, RunnableCallable): + obj = obj.func if name is None: name = getattr(obj, "__qualname__", None) if name is None: # pragma: no cover diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 00feb2650..2c3e26259 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -1069,7 +1069,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): return () matched: list[PregelExecutableTask] = [] if cached := { - t.cache_key.key: t + (t.cache_key.ns, t.cache_key.key): t for t in self.tasks.values() if t.cache_key and not t.cache_key.refresh and not t.writes }: @@ -1089,7 +1089,12 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): return self.submit( self.cache.set, - {task.cache_key.key: (task.writes, task.cache_key.ttl)}, + { + (task.cache_key.ns, task.cache_key.key): ( + task.writes, + task.cache_key.ttl, + ) + }, ) # context manager @@ -1253,7 +1258,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): return [] matched: list[PregelExecutableTask] = [] if cached := { - t.cache_key.key: t + (t.cache_key.ns, t.cache_key.key): t for t in self.tasks.values() if t.cache_key and not t.cache_key.refresh and not t.writes }: @@ -1276,7 +1281,12 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): return self.submit( self.cache.aset, - {task.cache_key.key: (task.writes, task.cache_key.ttl)}, + { + (task.cache_key.ns, task.cache_key.key): ( + task.writes, + task.cache_key.ttl, + ) + }, ) # context manager diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index eba426deb..c77a2bc3b 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -187,6 +187,8 @@ else: class CacheKey(NamedTuple): """Cache key for a task.""" + ns: tuple[str, ...] + """Namespace for the cache entry.""" key: str """Key for the cache entry.""" ttl: Optional[int]