mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-08-05 20:38:41 +00:00
feat(checkpoint): checkpoint history cache (#4638)
* feat(checkpoint-cache): delta-mode checkpoint history cache with recursive compose
Read-only, invalidation-free cache for LangGraph delta-channel history
({writes, seed}) at the get_delta_channel_history choke point:
- database.checkpoint_cache config (memory|redis; max_entries 0=disabled;
redis bounded by TTL, Gateway/async only)
- memory LRU backend (copy-on-read, zero-serde hit path) and redis backend
(lazy import, degrades to all-miss on outage)
- CachedHistorySaver: recursive composition from the nearest warm ancestor
(depth budget 8), caching each level; depth-0 cold chains delegate one
inner fast-path walk. Entries keyed by immutable
(db, thread, ns, checkpoint_id, channel) — no invalidation, coherent
across workers
- provider wiring: wraps in delta mode only (async + sync), full mode
untouched; sync path is memory-only
- bench opt-in: DEERFLOW_CHECKPOINT_BENCH_HISTORY_CACHE=1
sqlite bench (500 updates, payload 2KB): write phase 2.28x at f=250,
1.32x at f=10; one delegated walk per thread cold start.
* chore(config): bump config_version to 32 for database.checkpoint_cache
The checkpoint history cache feature added the database.checkpoint_cache
section to config.example.yaml; bump the schema version so existing
deployments get the outdated-config warning and can run make config-upgrade.
* chore(helm): bump config_version to 32 in chart values and README
* fix(checkpoint-cache): purge thread history entries on delete paths
Addresses review on #4638: delete_thread/prune removed source-of-truth
checkpoints but left the thread's materialized history payloads in the
cache (memory: until LRU eviction; redis: until TTL, default 1 day) — a
data-lifecycle gap for tenant offboarding / GDPR-style erasure.
- Cache contract gains thread-scoped adelete_thread/delete_thread
(lifecycle purge, not invalidation; entries remain immutable)
- Memory backend: stem scan over the LRU map; redis: SCAN MATCH + UNLINK,
outage degrades to TTL-bounded retention without raising
- CachedHistorySaver purges on delete_thread/adelete_thread and
prune/aprune (prune rewrites chains, so pre-prune histories must go);
delete_for_runs stays delegation-only (run->thread mapping unavailable,
no in-tree callers), documented in code
- ttl_seconds description documents the residual-retention window
- Tests: thread-scoped purge on both backends, saver-level delete/prune
purge, prefix-safety (t1 vs t10), redis outage degradation, and the
pinned no-purge behavior of delete_for_runs
* fix(checkpoint-cache): stable db identity, prefix-aware sync singleton, explicit zero TTL
Addresses Copilot review on #4638:
- checkpoint_cache_db_hash now hashes the credential-free postgres
identity (host:port/database + schema): credential rotation no longer
changes the cache namespace (cold cache + orphaned keys until TTL).
Unparseable URLs fall back to the raw string.
- The sync-path memory cache singleton is also keyed by its key_prefix:
a namespace change (db identity change or operator override) recreates
the cache instead of leaving stale-prefix entries unreachable and
unpurgeable.
- ttl_seconds=0 is now an explicit, documented opt-out of redis expiry
(SET without EX; redis maxmemory policy only) instead of a silent
'ttl_seconds or None' coercion.
Tests: credential-rotation hash stability, unparseable-URL fallback,
prefix-change singleton recreation, same-prefix singleton reuse, and
zero-TTL wire behavior (ex=None).
---------
Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
This commit is contained in:
parent
7025ccee40
commit
c8cf1bf2fb
@ -112,6 +112,8 @@ This section accumulates work toward the **2.1.0** milestone
|
|||||||
- **runtime:** Dual-mode checkpoint storage with LangGraph `DeltaChannel` cuts
|
- **runtime:** Dual-mode checkpoint storage with LangGraph `DeltaChannel` cuts
|
||||||
thread storage from O(N²) to near-linear for long research/coding runs.
|
thread storage from O(N²) to near-linear for long research/coding runs.
|
||||||
([#4292])
|
([#4292])
|
||||||
|
- **runtime:** Delta-mode checkpoint history cache (memory/redis) with O(1)
|
||||||
|
incremental composition, configured via `database.checkpoint_cache`.
|
||||||
- **agent:** Config-declared lead-agent middlewares let deployments add custom
|
- **agent:** Config-declared lead-agent middlewares let deployments add custom
|
||||||
`AgentMiddleware` classes without patching the runtime chain. ([#3964])
|
`AgentMiddleware` classes without patching the runtime chain. ([#3964])
|
||||||
- **agents:** Per-agent model and generation settings (`temperature`,
|
- **agents:** Per-agent model and generation settings (`temperature`,
|
||||||
|
|||||||
@ -264,6 +264,14 @@ The checkpoint storage settings `database.checkpoint_channel_mode` and
|
|||||||
both are frozen when the process first builds an agent (including through
|
both are frozen when the process first builds an agent (including through
|
||||||
`DeerFlowClient`) and require a process restart to change safely.
|
`DeerFlowClient`) and require a process restart to change safely.
|
||||||
|
|
||||||
|
The optional `database.checkpoint_cache` section (delta channel mode only)
|
||||||
|
caches materialized checkpoint histories: `type` is `memory` (default) or
|
||||||
|
`redis`, and `max_entries: 0` disables the cache. The `redis` backend is
|
||||||
|
Gateway/async-only; the sync TUI/embedded path supports `memory` only. The
|
||||||
|
cache is performance-only — results are identical with it disabled — so it is
|
||||||
|
never frozen and workers sharing one checkpoint database may safely run
|
||||||
|
different cache settings.
|
||||||
|
|
||||||
> [!TIP]
|
> [!TIP]
|
||||||
> On Linux, if Docker-based commands fail with `permission denied while trying to connect to the Docker daemon socket at unix:///var/run/docker.sock`, add your user to the `docker` group and re-login before retrying. See [CONTRIBUTING.md](CONTRIBUTING.md#linux-docker-daemon-permission-denied) for the full fix.
|
> On Linux, if Docker-based commands fail with `permission denied while trying to connect to the Docker daemon socket at unix:///var/run/docker.sock`, add your user to the `docker` group and re-login before retrying. See [CONTRIBUTING.md](CONTRIBUTING.md#linux-docker-daemon-permission-denied) for the full fix.
|
||||||
|
|
||||||
|
|||||||
@ -1089,7 +1089,8 @@ Checkpointer storage runs in one of two channel modes, selected by `checkpoint_c
|
|||||||
- `checkpoint_patches.py` (package root) — checkpoint-machinery patches: delta-history folding for `InMemorySaver` (delegating to the base walk), stable message IDs across materialization, upstream first-write drop fix, and `BinaryOperatorAggregate` unwrapping an `Overwrite` first write into an empty (MISSING) channel — Union-typed reducer channels (`sandbox`/`goal`/`todos`/`promoted`) have no constructible default, so a replace-style write into a fresh branch thread or a never-written channel stored the wrapper literally and crashed the next consumer (#4380; probe-guarded, stands down if upstream fixes it)
|
- `checkpoint_patches.py` (package root) — checkpoint-machinery patches: delta-history folding for `InMemorySaver` (delegating to the base walk), stable message IDs across materialization, upstream first-write drop fix, and `BinaryOperatorAggregate` unwrapping an `Overwrite` first write into an empty (MISSING) channel — Union-typed reducer channels (`sandbox`/`goal`/`todos`/`promoted`) have no constructible default, so a replace-style write into a fresh branch thread or a never-written channel stored the wrapper literally and crashed the next consumer (#4380; probe-guarded, stands down if upstream fixes it)
|
||||||
- `agents/thread_state.py` — `ThreadState`/`DeltaThreadState`, `delta_messages_field` / `DELTA_MESSAGES_FIELD` (`DeltaChannel` at the configured `snapshot_frequency`, default 10), schema adaptation helpers
|
- `agents/thread_state.py` — `ThreadState`/`DeltaThreadState`, `delta_messages_field` / `DELTA_MESSAGES_FIELD` (`DeltaChannel` at the configured `snapshot_frequency`, default 10), schema adaptation helpers
|
||||||
- `runtime/context_compaction.py` — compaction via accessor + mutation graph (reference consumer)
|
- `runtime/context_compaction.py` — compaction via accessor + mutation graph (reference consumer)
|
||||||
- Tests: `tests/test_checkpoint_mode.py` (freeze/detect/gate), `tests/test_checkpoint_state.py` (accessor/mutation graph), `tests/test_delta_channel_checkpointers.py` (saver parity), `tests/test_threads_checkpoint_mode.py`, `tests/test_gateway_checkpoint_mode.py` (dual-mode e2e parity), `tests/test_context_compaction.py` (mutation-graph write, no scheduling), `tests/test_run_worker_rollback.py`
|
- `runtime/checkpoint_cache/` + `runtime/checkpointer/cached_saver.py` — delta-mode checkpoint history cache; checkpoint state reads MUST go through `CheckpointStateAccessor`, and the checkpointer may be a `CachedHistorySaver` wrapper — never rely on concrete saver types
|
||||||
|
- Tests: `tests/test_checkpoint_mode.py` (freeze/detect/gate), `tests/test_checkpoint_state.py` (accessor/mutation graph), `tests/test_delta_channel_checkpointers.py` (saver parity), `tests/test_threads_checkpoint_mode.py`, `tests/test_gateway_checkpoint_mode.py` (dual-mode e2e parity), `tests/test_context_compaction.py` (mutation-graph write, no scheduling), `tests/test_run_worker_rollback.py`, `tests/test_cached_history_saver.py` + `tests/test_cached_history_saver_integration.py` (history cache)
|
||||||
|
|
||||||
**Checkpoint channel benchmark**: `scripts/benchmark/checkpoint/bench_channels.py`
|
**Checkpoint channel benchmark**: `scripts/benchmark/checkpoint/bench_channels.py`
|
||||||
runs paired `full`/`delta` message-only StateGraphs in a fresh child process per
|
runs paired `full`/`delta` message-only StateGraphs in a fresh child process per
|
||||||
|
|||||||
@ -100,6 +100,45 @@ class CheckpointGraphCacheConfig(BaseModel):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointCacheConfig(BaseModel):
|
||||||
|
"""Delta-history cache policy. Performance-only: never frozen, never
|
||||||
|
required to match across processes sharing one checkpoint database.
|
||||||
|
|
||||||
|
Applies only when ``checkpoint_channel_mode`` is ``delta``. ``max_entries``
|
||||||
|
bounds the process-local memory backend; ``0`` disables the cache
|
||||||
|
entirely. The redis backend is bounded by ``ttl_seconds`` and the server's
|
||||||
|
own maxmemory policy.
|
||||||
|
"""
|
||||||
|
|
||||||
|
type: Literal["memory", "redis"] = Field(
|
||||||
|
default="memory",
|
||||||
|
description=("Checkpoint history cache backend. 'memory' = process-local LRU; 'redis' = shared cache for multi-worker deployments (async/Gateway path only; the sync embedded path rejects it)."),
|
||||||
|
)
|
||||||
|
max_entries: int = Field(
|
||||||
|
default=128,
|
||||||
|
ge=0,
|
||||||
|
description="LRU capacity of the memory backend. 0 disables the cache.",
|
||||||
|
)
|
||||||
|
redis_url: str | None = Field(
|
||||||
|
default=None,
|
||||||
|
description=("Redis URL for type=redis. If omitted, DEER_FLOW_CHECKPOINT_CACHE_REDIS_URL, REDIS_URL, or redis://localhost:6379/0 is used."),
|
||||||
|
)
|
||||||
|
ttl_seconds: int = Field(
|
||||||
|
default=86400,
|
||||||
|
ge=0,
|
||||||
|
description=(
|
||||||
|
"Redis entry TTL; a leak safety net, not a correctness mechanism (entries are immutable). "
|
||||||
|
"Thread deletion purges that thread's entries immediately; if the purge fails (redis outage), "
|
||||||
|
"residual copies of the thread's history persist until this TTL expires. "
|
||||||
|
"0 explicitly disables expiry — orphaned keys then rely on the redis maxmemory policy alone."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
key_prefix: str = Field(
|
||||||
|
default="",
|
||||||
|
description="Optional override for the redis key prefix; defaults to a hash of the database identity.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class DatabaseConfig(BaseModel):
|
class DatabaseConfig(BaseModel):
|
||||||
backend: Literal["memory", "sqlite", "postgres"] = Field(
|
backend: Literal["memory", "sqlite", "postgres"] = Field(
|
||||||
default="memory",
|
default="memory",
|
||||||
@ -122,6 +161,10 @@ class DatabaseConfig(BaseModel):
|
|||||||
default_factory=CheckpointGraphCacheConfig,
|
default_factory=CheckpointGraphCacheConfig,
|
||||||
description="Size caps for the compiled checkpoint graph caches. Hot-reloadable; not restart-required.",
|
description="Size caps for the compiled checkpoint graph caches. Hot-reloadable; not restart-required.",
|
||||||
)
|
)
|
||||||
|
checkpoint_cache: CheckpointCacheConfig = Field(
|
||||||
|
default_factory=CheckpointCacheConfig,
|
||||||
|
description="Delta-mode checkpoint history cache. Performance-only; safe to differ across workers.",
|
||||||
|
)
|
||||||
sqlite_dir: str = Field(
|
sqlite_dir: str = Field(
|
||||||
default=".deer-flow/data",
|
default=".deer-flow/data",
|
||||||
description=("Directory for the SQLite database file. Both checkpointer and application data share {sqlite_dir}/deerflow.db."),
|
description=("Directory for the SQLite database file. Both checkpointer and application data share {sqlite_dir}/deerflow.db."),
|
||||||
|
|||||||
@ -0,0 +1,19 @@
|
|||||||
|
"""Checkpoint delta-history cache backends (delta mode only)."""
|
||||||
|
|
||||||
|
from deerflow.runtime.checkpoint_cache.base import (
|
||||||
|
CACHE_FORMAT_VERSION,
|
||||||
|
CheckpointCacheStats,
|
||||||
|
CheckpointHistoryCache,
|
||||||
|
SyncCheckpointHistoryCache,
|
||||||
|
make_history_key,
|
||||||
|
)
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"CACHE_FORMAT_VERSION",
|
||||||
|
"CheckpointCacheStats",
|
||||||
|
"CheckpointHistoryCache",
|
||||||
|
"MemoryCheckpointHistoryCache",
|
||||||
|
"SyncCheckpointHistoryCache",
|
||||||
|
"make_history_key",
|
||||||
|
]
|
||||||
@ -0,0 +1,80 @@
|
|||||||
|
"""Cache backend contract for checkpoint delta-history entries.
|
||||||
|
|
||||||
|
Entries are ``DeltaChannelHistory``-shaped dicts (``{"writes": [...], "seed"?}``)
|
||||||
|
keyed by immutable (database, thread, namespace, checkpoint_id, channel)
|
||||||
|
tuples. Checkpoint lineage is append-only and a checkpoint's history excludes
|
||||||
|
its own pending writes, so entries never change once written: correctness
|
||||||
|
never requires invalidation, and a shared backend is coherent across
|
||||||
|
processes without any coordination.
|
||||||
|
|
||||||
|
The only delete API is thread-scoped (``adelete_thread``/``delete_thread``),
|
||||||
|
and it exists purely for data lifecycle, not correctness: when the source
|
||||||
|
checkpoints are erased (thread deletion, tenant offboarding, GDPR-style
|
||||||
|
erasure), the cached history payloads for that thread must go too instead of
|
||||||
|
lingering until LRU eviction or TTL expiry.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Protocol
|
||||||
|
|
||||||
|
CACHE_FORMAT_VERSION = 1
|
||||||
|
|
||||||
|
|
||||||
|
def make_history_key(
|
||||||
|
key_prefix: str,
|
||||||
|
thread_id: str,
|
||||||
|
checkpoint_ns: str,
|
||||||
|
checkpoint_id: str,
|
||||||
|
channel: str,
|
||||||
|
) -> str:
|
||||||
|
"""Build a collision-safe cache key.
|
||||||
|
|
||||||
|
``thread_id`` stays readable for ops debugging; the remaining components
|
||||||
|
are hashed with NUL separators so namespaces containing ':' cannot
|
||||||
|
produce ambiguous keys.
|
||||||
|
"""
|
||||||
|
digest = hashlib.sha256(f"{checkpoint_ns}\x00{checkpoint_id}\x00{channel}".encode()).hexdigest()[:24]
|
||||||
|
return f"{key_prefix}:{thread_id}:{digest}"
|
||||||
|
|
||||||
|
|
||||||
|
def thread_key_stem(key_prefix: str, thread_id: str) -> str:
|
||||||
|
"""Prefix matching every history key of one thread (see make_history_key)."""
|
||||||
|
return f"{key_prefix}:{thread_id}:"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class CheckpointCacheStats:
|
||||||
|
hits: int = 0
|
||||||
|
misses: int = 0
|
||||||
|
evictions: int = 0
|
||||||
|
entries: int = 0
|
||||||
|
|
||||||
|
def as_dict(self) -> dict[str, int]:
|
||||||
|
return {
|
||||||
|
"hits": self.hits,
|
||||||
|
"misses": self.misses,
|
||||||
|
"evictions": self.evictions,
|
||||||
|
"entries": self.entries,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class CheckpointHistoryCache(Protocol):
|
||||||
|
"""Async backend contract. Deletes are thread-scoped lifecycle purges only."""
|
||||||
|
|
||||||
|
async def aget_many(self, keys: list[str]) -> dict[str, dict[str, Any]]: ...
|
||||||
|
async def aset_many(self, entries: dict[str, dict[str, Any]]) -> None: ...
|
||||||
|
async def adelete_thread(self, key_prefix: str, thread_id: str) -> None: ...
|
||||||
|
def stats(self) -> CheckpointCacheStats: ...
|
||||||
|
async def aclose(self) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class SyncCheckpointHistoryCache(Protocol):
|
||||||
|
"""Sync backend contract (embedded/TUI path). Memory backend only."""
|
||||||
|
|
||||||
|
def get_many(self, keys: list[str]) -> dict[str, dict[str, Any]]: ...
|
||||||
|
def set_many(self, entries: dict[str, dict[str, Any]]) -> None: ...
|
||||||
|
def delete_thread(self, key_prefix: str, thread_id: str) -> None: ...
|
||||||
|
def stats(self) -> CheckpointCacheStats: ...
|
||||||
@ -0,0 +1,79 @@
|
|||||||
|
"""Process-local LRU backend. Zero serialization on the hit path."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections import OrderedDict
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from deerflow.runtime.checkpoint_cache.base import CheckpointCacheStats, thread_key_stem
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_entry(entry: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
"""Copy-on-read/write: fresh writes list; seed shared (never mutated in place)."""
|
||||||
|
copied: dict[str, Any] = {"writes": list(entry["writes"])}
|
||||||
|
if "seed" in entry:
|
||||||
|
copied["seed"] = entry["seed"]
|
||||||
|
return copied
|
||||||
|
|
||||||
|
|
||||||
|
class MemoryCheckpointHistoryCache:
|
||||||
|
def __init__(self, max_entries: int = 128) -> None:
|
||||||
|
if max_entries < 0:
|
||||||
|
raise ValueError("max_entries must be >= 0")
|
||||||
|
self._max_entries = max_entries
|
||||||
|
self._data: OrderedDict[str, dict[str, Any]] = OrderedDict()
|
||||||
|
self._hits = 0
|
||||||
|
self._misses = 0
|
||||||
|
self._evictions = 0
|
||||||
|
|
||||||
|
@property
|
||||||
|
def enabled(self) -> bool:
|
||||||
|
return self._max_entries > 0
|
||||||
|
|
||||||
|
def get_many(self, keys: list[str]) -> dict[str, dict[str, Any]]:
|
||||||
|
found: dict[str, dict[str, Any]] = {}
|
||||||
|
for key in keys:
|
||||||
|
entry = self._data.get(key)
|
||||||
|
if entry is None:
|
||||||
|
self._misses += 1
|
||||||
|
continue
|
||||||
|
self._data.move_to_end(key)
|
||||||
|
self._hits += 1
|
||||||
|
found[key] = _copy_entry(entry)
|
||||||
|
return found
|
||||||
|
|
||||||
|
def set_many(self, entries: dict[str, dict[str, Any]]) -> None:
|
||||||
|
if not self.enabled:
|
||||||
|
return
|
||||||
|
for key, entry in entries.items():
|
||||||
|
self._data[key] = _copy_entry(entry)
|
||||||
|
self._data.move_to_end(key)
|
||||||
|
while len(self._data) > self._max_entries:
|
||||||
|
self._data.popitem(last=False)
|
||||||
|
self._evictions += 1
|
||||||
|
|
||||||
|
async def aget_many(self, keys: list[str]) -> dict[str, dict[str, Any]]:
|
||||||
|
return self.get_many(keys)
|
||||||
|
|
||||||
|
async def aset_many(self, entries: dict[str, dict[str, Any]]) -> None:
|
||||||
|
self.set_many(entries)
|
||||||
|
|
||||||
|
def delete_thread(self, key_prefix: str, thread_id: str) -> None:
|
||||||
|
"""Purge every entry of one thread (lifecycle, not invalidation)."""
|
||||||
|
stem = thread_key_stem(key_prefix, thread_id)
|
||||||
|
for key in [k for k in self._data if k.startswith(stem)]:
|
||||||
|
del self._data[key]
|
||||||
|
|
||||||
|
async def adelete_thread(self, key_prefix: str, thread_id: str) -> None:
|
||||||
|
self.delete_thread(key_prefix, thread_id)
|
||||||
|
|
||||||
|
def stats(self) -> CheckpointCacheStats:
|
||||||
|
return CheckpointCacheStats(
|
||||||
|
hits=self._hits,
|
||||||
|
misses=self._misses,
|
||||||
|
evictions=self._evictions,
|
||||||
|
entries=len(self._data),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
self._data.clear()
|
||||||
@ -0,0 +1,101 @@
|
|||||||
|
"""Cache factory. Mirrors make_stream_bridge: config -> env fallback -> memory."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import hashlib
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from deerflow.config.app_config import AppConfig
|
||||||
|
from deerflow.runtime.checkpoint_cache.base import CACHE_FORMAT_VERSION, CheckpointHistoryCache
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_ENV_REDIS_URL = "DEER_FLOW_CHECKPOINT_CACHE_REDIS_URL"
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_redis_url(config: Any) -> str:
|
||||||
|
return config.redis_url or os.getenv(_ENV_REDIS_URL) or os.getenv("REDIS_URL") or "redis://localhost:6379/0"
|
||||||
|
|
||||||
|
|
||||||
|
def _stable_postgres_identity(postgres_url: str) -> str:
|
||||||
|
"""Credential-free database identity: host/port/database.
|
||||||
|
|
||||||
|
Hashing the raw URL would change the cache namespace on every credential
|
||||||
|
rotation (cold cache + orphaned keys until TTL) even though the database
|
||||||
|
— and thus every cached checkpoint history — is unchanged. Unparseable
|
||||||
|
URLs fall back to the raw string (still stable per deployment).
|
||||||
|
"""
|
||||||
|
if not postgres_url:
|
||||||
|
return ""
|
||||||
|
try:
|
||||||
|
from sqlalchemy.engine.url import make_url
|
||||||
|
|
||||||
|
parsed = make_url(postgres_url)
|
||||||
|
except Exception: # noqa: BLE001 - identity must never fail config load
|
||||||
|
return postgres_url
|
||||||
|
return f"{parsed.host or 'localhost'}:{parsed.port or 5432}/{parsed.database or ''}"
|
||||||
|
|
||||||
|
|
||||||
|
def checkpoint_cache_db_hash(db_config: Any) -> str:
|
||||||
|
"""Deployment-identity hash so two deployments sharing one Redis never collide."""
|
||||||
|
backend = getattr(db_config, "backend", "memory")
|
||||||
|
if backend == "postgres":
|
||||||
|
identity = f"postgres:{_stable_postgres_identity(getattr(db_config, 'postgres_url', ''))}:{getattr(db_config, 'postgres_schema', '')}"
|
||||||
|
elif backend == "sqlite":
|
||||||
|
identity = f"sqlite:{getattr(db_config, 'checkpointer_sqlite_path', '')}"
|
||||||
|
else:
|
||||||
|
identity = "memory"
|
||||||
|
return hashlib.sha256(identity.encode()).hexdigest()[:12]
|
||||||
|
|
||||||
|
|
||||||
|
def checkpoint_cache_key_prefix(app_config: AppConfig) -> str:
|
||||||
|
cache_config = app_config.database.checkpoint_cache
|
||||||
|
if cache_config.key_prefix:
|
||||||
|
return cache_config.key_prefix
|
||||||
|
return f"ckpt-hist:v{CACHE_FORMAT_VERSION}:{checkpoint_cache_db_hash(app_config.database)}"
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.asynccontextmanager
|
||||||
|
async def make_checkpoint_cache(
|
||||||
|
app_config: AppConfig | None = None,
|
||||||
|
*,
|
||||||
|
serde: Any,
|
||||||
|
) -> AsyncIterator[CheckpointHistoryCache]:
|
||||||
|
"""Yield a history cache for the caller's lifetime.
|
||||||
|
|
||||||
|
``max_entries == 0`` disables the cache uniformly (both types) via a
|
||||||
|
disabled memory backend, so the wrapper never needs a None check.
|
||||||
|
"""
|
||||||
|
config = app_config.database.checkpoint_cache if app_config is not None else None
|
||||||
|
|
||||||
|
if config is None or config.type == "memory" or config.max_entries == 0:
|
||||||
|
max_entries = config.max_entries if config is not None else 128
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=max_entries)
|
||||||
|
logger.info("Checkpoint history cache initialised: memory (max_entries=%d)", max_entries)
|
||||||
|
try:
|
||||||
|
yield cache
|
||||||
|
finally:
|
||||||
|
await cache.aclose()
|
||||||
|
return
|
||||||
|
|
||||||
|
if config.type == "redis":
|
||||||
|
from deerflow.runtime.checkpoint_cache.redis import RedisCheckpointHistoryCache
|
||||||
|
|
||||||
|
cache = RedisCheckpointHistoryCache(
|
||||||
|
_resolve_redis_url(config),
|
||||||
|
serde=serde,
|
||||||
|
ttl_seconds=config.ttl_seconds,
|
||||||
|
)
|
||||||
|
logger.info("Checkpoint history cache initialised: redis (ttl_seconds=%d)", config.ttl_seconds)
|
||||||
|
try:
|
||||||
|
yield cache
|
||||||
|
finally:
|
||||||
|
await cache.aclose()
|
||||||
|
return
|
||||||
|
|
||||||
|
raise ValueError(f"Unknown checkpoint cache type: {config.type!r}")
|
||||||
@ -0,0 +1,116 @@
|
|||||||
|
"""Shared Redis backend. Entries are immutable, so a multi-worker shared
|
||||||
|
cache needs no invalidation; the TTL is a leak safety net only.
|
||||||
|
|
||||||
|
Thread-scoped purge (``adelete_thread``) exists for data lifecycle: when a
|
||||||
|
thread's checkpoints are deleted, its cached history payloads are removed
|
||||||
|
immediately instead of lingering until TTL expiry.
|
||||||
|
|
||||||
|
The redis import is lazy (module is importable without the optional
|
||||||
|
``redis`` extra), mirroring runtime/stream_bridge/redis.py.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from deerflow.runtime.checkpoint_cache.base import CheckpointCacheStats, thread_key_stem
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
REDIS_INSTALL = "redis is required for the redis checkpoint cache backend. Install it with: uv sync --extra redis"
|
||||||
|
|
||||||
|
_TAG_SEPARATOR = b"\x00"
|
||||||
|
|
||||||
|
|
||||||
|
def _create_client(redis_url: str, *, max_connections: int | None) -> Any:
|
||||||
|
try:
|
||||||
|
import redis.asyncio as redis_async
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError(REDIS_INSTALL) from exc
|
||||||
|
kwargs: dict[str, Any] = {"decode_responses": False}
|
||||||
|
if max_connections is not None:
|
||||||
|
kwargs["max_connections"] = max_connections
|
||||||
|
return redis_async.from_url(redis_url, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def _redis_error() -> type[Exception]:
|
||||||
|
"""Lazy RedisError import, mirroring the lazy client creation above."""
|
||||||
|
try:
|
||||||
|
from redis.exceptions import RedisError
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError(REDIS_INSTALL) from exc
|
||||||
|
return RedisError
|
||||||
|
|
||||||
|
|
||||||
|
class RedisCheckpointHistoryCache:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
redis_url: str,
|
||||||
|
*,
|
||||||
|
serde: Any,
|
||||||
|
ttl_seconds: int,
|
||||||
|
max_connections: int | None = None,
|
||||||
|
) -> None:
|
||||||
|
self._client = _create_client(redis_url, max_connections=max_connections)
|
||||||
|
self._serde = serde
|
||||||
|
# ttl_seconds=0 is an explicit opt-out of expiry (no SETEX) — not the
|
||||||
|
# default, and leaked/orphaned keys then rely on redis maxmemory only.
|
||||||
|
self._ttl = ttl_seconds if ttl_seconds > 0 else None
|
||||||
|
self._hits = 0
|
||||||
|
self._misses = 0
|
||||||
|
|
||||||
|
async def aget_many(self, keys: list[str]) -> dict[str, dict[str, Any]]:
|
||||||
|
if not keys:
|
||||||
|
return {}
|
||||||
|
try:
|
||||||
|
raws = await self._client.mget(keys)
|
||||||
|
except _redis_error() as exc:
|
||||||
|
# Performance-only bypass: a redis outage costs hits, never availability.
|
||||||
|
logger.warning("checkpoint history cache mget failed; treating as all-miss: %s", exc)
|
||||||
|
self._misses += len(keys)
|
||||||
|
return {}
|
||||||
|
found: dict[str, dict[str, Any]] = {}
|
||||||
|
for key, raw in zip(keys, raws, strict=True):
|
||||||
|
if raw is None:
|
||||||
|
self._misses += 1
|
||||||
|
continue
|
||||||
|
self._hits += 1
|
||||||
|
tag, payload = raw.split(_TAG_SEPARATOR, 1)
|
||||||
|
found[key] = self._serde.loads_typed((tag.decode(), payload))
|
||||||
|
return found
|
||||||
|
|
||||||
|
async def aset_many(self, entries: dict[str, dict[str, Any]]) -> None:
|
||||||
|
if not entries:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
pipe = self._client.pipeline(transaction=False)
|
||||||
|
for key, entry in entries.items():
|
||||||
|
tag, data = self._serde.dumps_typed(entry)
|
||||||
|
pipe.set(key, tag.encode() + _TAG_SEPARATOR + data, ex=self._ttl)
|
||||||
|
await pipe.execute()
|
||||||
|
except _redis_error() as exc:
|
||||||
|
# Writes are optional; the next read simply recomputes the history.
|
||||||
|
logger.warning("checkpoint history cache write failed; skipping: %s", exc)
|
||||||
|
|
||||||
|
async def adelete_thread(self, key_prefix: str, thread_id: str) -> None:
|
||||||
|
"""SCAN+UNLINK every entry of one thread. Failure degrades to
|
||||||
|
TTL-bounded residual retention; the source-of-truth delete already
|
||||||
|
happened, so this never raises."""
|
||||||
|
stem = thread_key_stem(key_prefix, thread_id)
|
||||||
|
try:
|
||||||
|
cursor = 0
|
||||||
|
while True:
|
||||||
|
cursor, keys = await self._client.scan(cursor=cursor, match=stem + "*", count=500)
|
||||||
|
if keys:
|
||||||
|
await self._client.unlink(*keys)
|
||||||
|
if cursor == 0:
|
||||||
|
break
|
||||||
|
except _redis_error() as exc:
|
||||||
|
logger.warning("checkpoint history cache thread purge failed; residual entries expire via TTL: %s", exc)
|
||||||
|
|
||||||
|
def stats(self) -> CheckpointCacheStats:
|
||||||
|
return CheckpointCacheStats(hits=self._hits, misses=self._misses)
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
await self._client.aclose()
|
||||||
@ -185,24 +185,14 @@ async def _async_checkpointer_from_database(db_config) -> AsyncIterator[Checkpoi
|
|||||||
|
|
||||||
|
|
||||||
@contextlib.asynccontextmanager
|
@contextlib.asynccontextmanager
|
||||||
async def make_checkpointer(app_config: AppConfig | None = None) -> AsyncIterator[Checkpointer]:
|
async def _select_inner_checkpointer(app_config: AppConfig) -> AsyncIterator[Checkpointer]:
|
||||||
"""Async context manager that yields a checkpointer for the caller's lifetime.
|
"""Yield the raw checkpointer selected by *app_config* (no delta-cache wrapping).
|
||||||
Resources are opened on enter and closed on exit -- no global state::
|
|
||||||
|
|
||||||
async with make_checkpointer(app_config) as checkpointer:
|
|
||||||
app.state.checkpointer = checkpointer
|
|
||||||
|
|
||||||
Yields an ``InMemorySaver`` when no checkpointer is configured in *config.yaml*.
|
|
||||||
|
|
||||||
Priority:
|
Priority:
|
||||||
1. Legacy ``checkpointer:`` config section (backward compatible)
|
1. Legacy ``checkpointer:`` config section (backward compatible)
|
||||||
2. Unified ``database:`` config section
|
2. Unified ``database:`` config section
|
||||||
3. Default InMemorySaver
|
3. Default InMemorySaver
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if app_config is None:
|
|
||||||
app_config = get_app_config()
|
|
||||||
|
|
||||||
# Legacy: standalone checkpointer config takes precedence
|
# Legacy: standalone checkpointer config takes precedence
|
||||||
if app_config.checkpointer is not None:
|
if app_config.checkpointer is not None:
|
||||||
async with _async_checkpointer(app_config.checkpointer) as saver:
|
async with _async_checkpointer(app_config.checkpointer) as saver:
|
||||||
@ -220,3 +210,44 @@ async def make_checkpointer(app_config: AppConfig | None = None) -> AsyncIterato
|
|||||||
from langgraph.checkpoint.memory import InMemorySaver
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
|
|
||||||
yield InMemorySaver()
|
yield InMemorySaver()
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.asynccontextmanager
|
||||||
|
async def make_checkpointer(app_config: AppConfig | None = None) -> AsyncIterator[Checkpointer]:
|
||||||
|
"""Async context manager that yields a checkpointer for the caller's lifetime.
|
||||||
|
Resources are opened on enter and closed on exit -- no global state::
|
||||||
|
|
||||||
|
async with make_checkpointer(app_config) as checkpointer:
|
||||||
|
app.state.checkpointer = checkpointer
|
||||||
|
|
||||||
|
Yields an ``InMemorySaver`` when no checkpointer is configured in *config.yaml*.
|
||||||
|
|
||||||
|
Backend selection priority:
|
||||||
|
1. Legacy ``checkpointer:`` config section (backward compatible)
|
||||||
|
2. Unified ``database:`` config section
|
||||||
|
3. Default InMemorySaver
|
||||||
|
|
||||||
|
When the effective checkpoint channel mode is ``delta`` (the process-frozen
|
||||||
|
mode wins, falling back to ``database.checkpoint_channel_mode``), the raw
|
||||||
|
saver is wrapped in a :class:`CachedHistorySaver` backed by a history cache
|
||||||
|
whose lifetime equals this context manager's.
|
||||||
|
"""
|
||||||
|
from deerflow.runtime.checkpoint_mode import frozen_checkpoint_channel_mode
|
||||||
|
|
||||||
|
if app_config is None:
|
||||||
|
app_config = get_app_config()
|
||||||
|
|
||||||
|
async with _select_inner_checkpointer(app_config) as saver:
|
||||||
|
db_config = getattr(app_config, "database", None)
|
||||||
|
mode = frozen_checkpoint_channel_mode() or (db_config.checkpoint_channel_mode if db_config is not None else "full")
|
||||||
|
if mode == "delta":
|
||||||
|
from deerflow.runtime.checkpoint_cache.provider import (
|
||||||
|
checkpoint_cache_key_prefix,
|
||||||
|
make_checkpoint_cache,
|
||||||
|
)
|
||||||
|
from deerflow.runtime.checkpointer.cached_saver import CachedHistorySaver
|
||||||
|
|
||||||
|
async with make_checkpoint_cache(app_config, serde=saver.serde) as cache:
|
||||||
|
yield CachedHistorySaver(saver, cache, key_prefix=checkpoint_cache_key_prefix(app_config))
|
||||||
|
else:
|
||||||
|
yield saver
|
||||||
|
|||||||
@ -0,0 +1,328 @@
|
|||||||
|
"""Read-through delta-history cache wrapper for any BaseCheckpointSaver.
|
||||||
|
|
||||||
|
Correctness argument (spec §3): a checkpoint's delta history is a pure
|
||||||
|
function of its sealed ancestor chain — the LangGraph contract excludes the
|
||||||
|
target's own pending writes, parent links are fixed at creation, and an
|
||||||
|
ancestor's writes are sealed once its child exists. Entries keyed by
|
||||||
|
(thread, ns, checkpoint_id, channel) are therefore immutable: no
|
||||||
|
invalidation, and shared backends are coherent across processes.
|
||||||
|
|
||||||
|
The wrapper never caches the "latest checkpoint" resolution; only histories
|
||||||
|
keyed by resolved immutable checkpoint_ids.
|
||||||
|
|
||||||
|
Data lifecycle: thread deletion and prune purge the thread's cached entries
|
||||||
|
(source-of-truth removal must not leave residual history payloads in the
|
||||||
|
cache); run-scoped deletes cannot be mapped to threads cheaply and rely on
|
||||||
|
LRU/TTL bounds.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from collections.abc import Iterator, Sequence
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain_core.runnables import RunnableConfig
|
||||||
|
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointTuple, PendingWrite
|
||||||
|
|
||||||
|
from deerflow.runtime.checkpoint_cache.base import make_history_key
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Depth budget for recursive compose before falling back to a chain-warming
|
||||||
|
# walk. Steady-state runs need ~2 (one intermediate checkpoint per step);
|
||||||
|
# deeper cold chains are handled faster by one warming walk than by many
|
||||||
|
# recursive single-tuple fetches.
|
||||||
|
_COMPOSE_MAX_DEPTH = 8
|
||||||
|
|
||||||
|
|
||||||
|
def _checkpoint_ref(tup: CheckpointTuple) -> tuple[str, str, str]:
|
||||||
|
configurable = tup.config["configurable"]
|
||||||
|
return (
|
||||||
|
str(configurable["thread_id"]),
|
||||||
|
str(configurable.get("checkpoint_ns", "")),
|
||||||
|
str(configurable["checkpoint_id"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _channel_writes(tup: CheckpointTuple, channel: str) -> list[PendingWrite]:
|
||||||
|
"""Writes for one channel, oldest→newest (tuple storage order)."""
|
||||||
|
return [w for w in (tup.pending_writes or []) if w[1] == channel]
|
||||||
|
|
||||||
|
|
||||||
|
class CachedHistorySaver(BaseCheckpointSaver):
|
||||||
|
def __init__(self, inner: BaseCheckpointSaver, cache: Any, *, key_prefix: str) -> None:
|
||||||
|
# Instance attr shadows the base class JsonPlusSerializer default.
|
||||||
|
self.serde = inner.serde
|
||||||
|
self._inner = inner
|
||||||
|
self._cache = cache
|
||||||
|
self._key_prefix = key_prefix
|
||||||
|
self._compose_hits = 0
|
||||||
|
self._full_walks = 0
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
# Safety net for saver-specific extras (e.g. AsyncSqliteSaver.setup).
|
||||||
|
# Base-class methods are explicitly delegated below, so this only
|
||||||
|
# fires for attributes BaseCheckpointSaver does not define.
|
||||||
|
inner = self.__dict__.get("_inner")
|
||||||
|
if inner is None:
|
||||||
|
raise AttributeError(name)
|
||||||
|
return getattr(inner, name)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Key building
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _key(self, tup: CheckpointTuple, channel: str) -> str:
|
||||||
|
thread_id, ns, checkpoint_id = _checkpoint_ref(tup)
|
||||||
|
return make_history_key(self._key_prefix, thread_id, ns, checkpoint_id, channel)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Stats
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def stats(self) -> dict[str, int]:
|
||||||
|
backend = self._cache.stats().as_dict()
|
||||||
|
return {**backend, "compose_hits": self._compose_hits, "full_walks": self._full_walks}
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Delta history: the only overridden behavior
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def aget_delta_channel_history(self, *, config: RunnableConfig, channels: Sequence[str]) -> dict[str, Any]:
|
||||||
|
if not channels:
|
||||||
|
return {}
|
||||||
|
if not getattr(self._cache, "enabled", True):
|
||||||
|
# Disabled cache: pass straight through; composing over all-miss
|
||||||
|
# entries is strictly more work than the raw saver's walk.
|
||||||
|
return await self._walk_inner(config, channels)
|
||||||
|
target = await self._inner.aget_tuple(config)
|
||||||
|
if target is None:
|
||||||
|
return await self._walk_inner(config, channels)
|
||||||
|
|
||||||
|
keys = {ch: self._key(target, ch) for ch in channels}
|
||||||
|
hits = await self._cache.aget_many(list(keys.values()))
|
||||||
|
found: dict[str, dict[str, Any]] = {}
|
||||||
|
missing: list[str] = []
|
||||||
|
for ch in channels:
|
||||||
|
entry = hits.get(keys[ch])
|
||||||
|
if entry is None:
|
||||||
|
missing.append(ch)
|
||||||
|
else:
|
||||||
|
found[ch] = entry
|
||||||
|
|
||||||
|
computed: dict[str, dict[str, Any]] = {}
|
||||||
|
if missing:
|
||||||
|
computed = await self._compose_or_walk(config, target, missing)
|
||||||
|
|
||||||
|
new_entries = {keys[ch]: computed[ch] for ch in missing if ch in computed}
|
||||||
|
if new_entries:
|
||||||
|
await self._cache.aset_many(new_entries)
|
||||||
|
|
||||||
|
return {ch: found.get(ch) or computed.get(ch) or {"writes": []} for ch in channels}
|
||||||
|
|
||||||
|
async def _compose_or_walk(self, config: RunnableConfig, target: CheckpointTuple, missing: list[str]) -> dict[str, dict[str, Any]]:
|
||||||
|
return {ch: await self._aresolve(target, ch, _COMPOSE_MAX_DEPTH) for ch in missing}
|
||||||
|
|
||||||
|
async def _aresolve(self, tup: CheckpointTuple, channel: str, depth: int) -> dict[str, Any]:
|
||||||
|
"""Recursively compose history(tup) from the nearest warm ancestor.
|
||||||
|
|
||||||
|
Real runs create several checkpoints per super-step and only some are
|
||||||
|
ever materialized as targets, so the parent is usually an unwarmed
|
||||||
|
intermediate checkpoint (measured: 0 cache hits with single-level
|
||||||
|
compose on a 500-step sqlite run). Recursing one level per
|
||||||
|
intermediate lands on a warmed ancestor within ~2 levels in steady
|
||||||
|
state; each composed level is cached, so the warm frontier follows
|
||||||
|
the run. At depth 0 on a cold chain it delegates one inner fast-path
|
||||||
|
walk (2 SQL) rather than crawling ancestors tuple-by-tuple.
|
||||||
|
"""
|
||||||
|
parent_config = tup.parent_config
|
||||||
|
if parent_config is None:
|
||||||
|
return {"writes": []}
|
||||||
|
parent = await self._inner.aget_tuple(parent_config)
|
||||||
|
if parent is None:
|
||||||
|
return {"writes": []}
|
||||||
|
|
||||||
|
channel_values = parent.checkpoint.get("channel_values") or {}
|
||||||
|
writes = _channel_writes(parent, channel)
|
||||||
|
if channel in channel_values:
|
||||||
|
self._compose_hits += 1
|
||||||
|
return {"writes": writes, "seed": channel_values[channel]}
|
||||||
|
|
||||||
|
key = self._key(parent, channel)
|
||||||
|
hits = await self._cache.aget_many([key])
|
||||||
|
parent_history = hits.get(key)
|
||||||
|
if parent_history is None:
|
||||||
|
if depth > 0:
|
||||||
|
parent_history = await self._aresolve(parent, channel, depth - 1)
|
||||||
|
else:
|
||||||
|
# Depth budget exhausted on a cold chain: delegate ONE inner
|
||||||
|
# fast-path walk (2 SQL total) for this level instead of
|
||||||
|
# fetching every ancestor tuple individually. Ancestors below
|
||||||
|
# stay cold; resolving them later recurses up to the nearest
|
||||||
|
# warm level, so the frontier still follows the run.
|
||||||
|
self._full_walks += 1
|
||||||
|
walked = await self._inner.aget_delta_channel_history(config=parent.config, channels=[channel])
|
||||||
|
parent_history = walked.get(channel) or {"writes": []}
|
||||||
|
if parent_history is not None:
|
||||||
|
await self._cache.aset_many({key: parent_history})
|
||||||
|
|
||||||
|
self._compose_hits += 1
|
||||||
|
entry: dict[str, Any] = {"writes": list(parent_history["writes"]) + writes}
|
||||||
|
if "seed" in parent_history:
|
||||||
|
entry["seed"] = parent_history["seed"]
|
||||||
|
return entry
|
||||||
|
|
||||||
|
async def _walk_inner(self, config: RunnableConfig, channels: Sequence[str]) -> dict[str, Any]:
|
||||||
|
self._full_walks += 1
|
||||||
|
return dict(await self._inner.aget_delta_channel_history(config=config, channels=channels))
|
||||||
|
|
||||||
|
def get_delta_channel_history(self, *, config: RunnableConfig, channels: Sequence[str]) -> dict[str, Any]:
|
||||||
|
if not channels:
|
||||||
|
return {}
|
||||||
|
if not getattr(self._cache, "enabled", True):
|
||||||
|
return self._walk_inner_sync(config, channels)
|
||||||
|
get_many = getattr(self._cache, "get_many", None)
|
||||||
|
set_many = getattr(self._cache, "set_many", None)
|
||||||
|
if get_many is None or set_many is None:
|
||||||
|
raise TypeError("sync get_delta_channel_history requires a SyncCheckpointHistoryCache (memory backend)")
|
||||||
|
target = self._inner.get_tuple(config)
|
||||||
|
if target is None:
|
||||||
|
return self._walk_inner_sync(config, channels)
|
||||||
|
|
||||||
|
keys = {ch: self._key(target, ch) for ch in channels}
|
||||||
|
hits = get_many(list(keys.values()))
|
||||||
|
found: dict[str, dict[str, Any]] = {}
|
||||||
|
missing: list[str] = []
|
||||||
|
for ch in channels:
|
||||||
|
entry = hits.get(keys[ch])
|
||||||
|
if entry is None:
|
||||||
|
missing.append(ch)
|
||||||
|
else:
|
||||||
|
found[ch] = entry
|
||||||
|
|
||||||
|
computed: dict[str, dict[str, Any]] = {}
|
||||||
|
if missing:
|
||||||
|
computed = {ch: self._resolve_sync(target, ch, _COMPOSE_MAX_DEPTH) for ch in missing}
|
||||||
|
|
||||||
|
new_entries = {keys[ch]: computed[ch] for ch in missing if ch in computed}
|
||||||
|
if new_entries:
|
||||||
|
set_many(new_entries)
|
||||||
|
return {ch: found.get(ch) or computed.get(ch) or {"writes": []} for ch in channels}
|
||||||
|
|
||||||
|
def _resolve_sync(self, tup: CheckpointTuple, channel: str, depth: int) -> dict[str, Any]:
|
||||||
|
"""Sync twin of _aresolve (recursive compose, see its docstring)."""
|
||||||
|
parent_config = tup.parent_config
|
||||||
|
if parent_config is None:
|
||||||
|
return {"writes": []}
|
||||||
|
parent = self._inner.get_tuple(parent_config)
|
||||||
|
if parent is None:
|
||||||
|
return {"writes": []}
|
||||||
|
|
||||||
|
channel_values = parent.checkpoint.get("channel_values") or {}
|
||||||
|
writes = _channel_writes(parent, channel)
|
||||||
|
if channel in channel_values:
|
||||||
|
self._compose_hits += 1
|
||||||
|
return {"writes": writes, "seed": channel_values[channel]}
|
||||||
|
|
||||||
|
key = self._key(parent, channel)
|
||||||
|
parent_history = self._cache.get_many([key]).get(key)
|
||||||
|
if parent_history is None:
|
||||||
|
if depth > 0:
|
||||||
|
parent_history = self._resolve_sync(parent, channel, depth - 1)
|
||||||
|
else:
|
||||||
|
# See _aresolve: one inner fast-path walk, no per-tuple crawl.
|
||||||
|
self._full_walks += 1
|
||||||
|
parent_history = self._inner.get_delta_channel_history(config=parent.config, channels=[channel]).get(channel) or {"writes": []}
|
||||||
|
if parent_history is not None:
|
||||||
|
self._cache.set_many({key: parent_history})
|
||||||
|
|
||||||
|
self._compose_hits += 1
|
||||||
|
entry: dict[str, Any] = {"writes": list(parent_history["writes"]) + writes}
|
||||||
|
if "seed" in parent_history:
|
||||||
|
entry["seed"] = parent_history["seed"]
|
||||||
|
return entry
|
||||||
|
|
||||||
|
def _walk_inner_sync(self, config: RunnableConfig, channels: Sequence[str]) -> dict[str, Any]:
|
||||||
|
self._full_walks += 1
|
||||||
|
return dict(self._inner.get_delta_channel_history(config=config, channels=channels))
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Explicit delegation (BaseCheckpointSaver defines these concretely,
|
||||||
|
# so __getattr__ never fires for them)
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||||
|
return self._inner.get_tuple(config)
|
||||||
|
|
||||||
|
def list(self, config: RunnableConfig | None, *, filter: dict[str, Any] | None = None, before: RunnableConfig | None = None, limit: int | None = None) -> Iterator[CheckpointTuple]:
|
||||||
|
return self._inner.list(config, filter=filter, before=before, limit=limit)
|
||||||
|
|
||||||
|
def put(self, config: RunnableConfig, checkpoint: dict[str, Any], metadata: dict[str, Any], new_versions: dict[str, Any]) -> RunnableConfig:
|
||||||
|
return self._inner.put(config, checkpoint, metadata, new_versions)
|
||||||
|
|
||||||
|
def put_writes(self, config: RunnableConfig, writes: Sequence[tuple[str, str, Any]], task_id: str, task_path: str = "") -> None:
|
||||||
|
self._inner.put_writes(config, writes, task_id, task_path)
|
||||||
|
|
||||||
|
def delete_thread(self, thread_id: str) -> None:
|
||||||
|
self._inner.delete_thread(thread_id)
|
||||||
|
self._purge_thread_sync(thread_id)
|
||||||
|
|
||||||
|
def delete_for_runs(self, run_ids: Sequence[str]) -> None:
|
||||||
|
# Run-scoped deletes cannot be mapped back to threads without an extra
|
||||||
|
# query, so cached entries are left in place: they stay *correct* (the
|
||||||
|
# sealed-chain argument is unaffected by other chains) and residual
|
||||||
|
# retention is bounded by LRU/TTL. No in-tree callers today.
|
||||||
|
self._inner.delete_for_runs(run_ids)
|
||||||
|
|
||||||
|
def _purge_thread_sync(self, thread_id: str) -> None:
|
||||||
|
delete = getattr(self._cache, "delete_thread", None)
|
||||||
|
if delete is not None:
|
||||||
|
delete(self._key_prefix, thread_id)
|
||||||
|
|
||||||
|
def copy_thread(self, source_thread_id: str, target_thread_id: str) -> None:
|
||||||
|
self._inner.copy_thread(source_thread_id, target_thread_id)
|
||||||
|
|
||||||
|
def prune(self, thread_ids: Sequence[str], *, strategy: str = "keep_latest") -> None:
|
||||||
|
self._inner.prune(thread_ids, strategy=strategy)
|
||||||
|
# Pruning rewrites these threads' chains: purge so no cached history
|
||||||
|
# references a deleted ancestor (retention) or its pre-prune chain.
|
||||||
|
for thread_id in thread_ids:
|
||||||
|
self._purge_thread_sync(thread_id)
|
||||||
|
|
||||||
|
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||||
|
return await self._inner.aget_tuple(config)
|
||||||
|
|
||||||
|
def alist(self, config: RunnableConfig | None, *, filter: dict[str, Any] | None = None, before: RunnableConfig | None = None, limit: int | None = None) -> Any:
|
||||||
|
return self._inner.alist(config, filter=filter, before=before, limit=limit)
|
||||||
|
|
||||||
|
async def aput(self, config: RunnableConfig, checkpoint: dict[str, Any], metadata: dict[str, Any], new_versions: dict[str, Any]) -> RunnableConfig:
|
||||||
|
return await self._inner.aput(config, checkpoint, metadata, new_versions)
|
||||||
|
|
||||||
|
async def aput_writes(self, config: RunnableConfig, writes: Sequence[tuple[str, str, Any]], task_id: str, task_path: str = "") -> None:
|
||||||
|
await self._inner.aput_writes(config, writes, task_id, task_path)
|
||||||
|
|
||||||
|
async def adelete_thread(self, thread_id: str) -> None:
|
||||||
|
await self._inner.adelete_thread(thread_id)
|
||||||
|
await self._apurge_thread(thread_id)
|
||||||
|
|
||||||
|
async def adelete_for_runs(self, run_ids: Sequence[str]) -> None:
|
||||||
|
# See delete_for_runs: unscoped residual retention bounded by LRU/TTL.
|
||||||
|
await self._inner.adelete_for_runs(run_ids)
|
||||||
|
|
||||||
|
async def _apurge_thread(self, thread_id: str) -> None:
|
||||||
|
delete = getattr(self._cache, "adelete_thread", None)
|
||||||
|
if delete is not None:
|
||||||
|
await delete(self._key_prefix, thread_id)
|
||||||
|
|
||||||
|
async def acopy_thread(self, source_thread_id: str, target_thread_id: str) -> None:
|
||||||
|
await self._inner.acopy_thread(source_thread_id, target_thread_id)
|
||||||
|
|
||||||
|
async def aprune(self, thread_ids: Sequence[str], *, strategy: str = "keep_latest") -> None:
|
||||||
|
await self._inner.aprune(thread_ids, strategy=strategy)
|
||||||
|
# See prune: rewritten chains must not keep pre-prune cached histories.
|
||||||
|
for thread_id in thread_ids:
|
||||||
|
await self._apurge_thread(thread_id)
|
||||||
|
|
||||||
|
def get_next_version(self, current: Any, channel: Any) -> Any:
|
||||||
|
return self._inner.get_next_version(current, channel)
|
||||||
@ -29,6 +29,7 @@ from langgraph.types import Checkpointer
|
|||||||
from deerflow.config.app_config import AppConfig, get_app_config
|
from deerflow.config.app_config import AppConfig, get_app_config
|
||||||
from deerflow.config.checkpointer_config import CheckpointerConfig, ensure_config_loaded, get_checkpointer_config
|
from deerflow.config.checkpointer_config import CheckpointerConfig, ensure_config_loaded, get_checkpointer_config
|
||||||
from deerflow.persistence.postgres_schema import dsn_with_search_path, ensure_postgres_schema
|
from deerflow.persistence.postgres_schema import dsn_with_search_path, ensure_postgres_schema
|
||||||
|
from deerflow.runtime.checkpoint_mode import frozen_checkpoint_channel_mode
|
||||||
from deerflow.runtime.store._sqlite_utils import ensure_sqlite_parent_dir, resolve_sqlite_conn_str
|
from deerflow.runtime.store._sqlite_utils import ensure_sqlite_parent_dir, resolve_sqlite_conn_str
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@ -155,6 +156,41 @@ def _sync_checkpointer_cm(config: CheckpointerConfig) -> Iterator[Checkpointer]:
|
|||||||
_checkpointer: Checkpointer | None = None
|
_checkpointer: Checkpointer | None = None
|
||||||
_checkpointer_ctx = None # open context manager keeping the connection alive
|
_checkpointer_ctx = None # open context manager keeping the connection alive
|
||||||
_checkpointer_lock = threading.Lock()
|
_checkpointer_lock = threading.Lock()
|
||||||
|
_checkpointer_cache = None # MemoryCheckpointHistoryCache singleton shared by wrapped sync savers
|
||||||
|
_checkpointer_cache_prefix: str | None = None # key prefix the singleton was built for
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_sync_if_delta(saver: Checkpointer, app_config: AppConfig) -> Checkpointer:
|
||||||
|
"""Wrap *saver* in a delta-history cache when the effective mode is ``delta``.
|
||||||
|
|
||||||
|
The process-frozen mode wins; ``database.checkpoint_channel_mode`` is the
|
||||||
|
fallback when nothing is frozen yet. Only the memory cache backend is
|
||||||
|
supported on the sync path (TUI/embedded) — it is process-local anyway.
|
||||||
|
"""
|
||||||
|
global _checkpointer_cache, _checkpointer_cache_prefix
|
||||||
|
# The ``_checkpointer_cache`` singleton is reassigned here without holding
|
||||||
|
# ``_checkpointer_lock`` on the ``checkpointer_context()`` path (and under
|
||||||
|
# the lock on the ``get_checkpointer()`` path). The race is intentional
|
||||||
|
# and benign: worst case two wrappers get their own fresh memory cache —
|
||||||
|
# last writer wins, and the cache is performance-only.
|
||||||
|
db_config = getattr(app_config, "database", None)
|
||||||
|
mode = frozen_checkpoint_channel_mode() or (db_config.checkpoint_channel_mode if db_config is not None else "full")
|
||||||
|
if mode != "delta":
|
||||||
|
return saver
|
||||||
|
cache_config = app_config.database.checkpoint_cache
|
||||||
|
if cache_config.type == "redis":
|
||||||
|
raise ValueError("database.checkpoint_cache.type 'redis' is not supported on the sync checkpointer path (TUI/embedded); use 'memory'.")
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
from deerflow.runtime.checkpoint_cache.provider import checkpoint_cache_key_prefix
|
||||||
|
from deerflow.runtime.checkpointer.cached_saver import CachedHistorySaver
|
||||||
|
|
||||||
|
key_prefix = checkpoint_cache_key_prefix(app_config)
|
||||||
|
# Recreate on capacity OR namespace change: entries under a stale prefix
|
||||||
|
# would be unreachable and no longer covered by thread purges.
|
||||||
|
if _checkpointer_cache is None or _checkpointer_cache._max_entries != cache_config.max_entries or _checkpointer_cache_prefix != key_prefix:
|
||||||
|
_checkpointer_cache = MemoryCheckpointHistoryCache(max_entries=cache_config.max_entries)
|
||||||
|
_checkpointer_cache_prefix = key_prefix
|
||||||
|
return CachedHistorySaver(saver, _checkpointer_cache, key_prefix=key_prefix)
|
||||||
|
|
||||||
|
|
||||||
def get_checkpointer() -> Checkpointer:
|
def get_checkpointer() -> Checkpointer:
|
||||||
@ -177,12 +213,27 @@ def get_checkpointer() -> Checkpointer:
|
|||||||
# config outside this provider lock to avoid cross-provider lock-order inversion.
|
# config outside this provider lock to avoid cross-provider lock-order inversion.
|
||||||
config = _get_checkpointer_config()
|
config = _get_checkpointer_config()
|
||||||
|
|
||||||
|
# ``get_app_config()`` can trigger a config reload whose
|
||||||
|
# ``_apply_singleton_configs`` calls ``reset_checkpointer()`` — which takes
|
||||||
|
# ``_checkpointer_lock``. Resolve it (non-reentrant lock) BEFORE acquiring
|
||||||
|
# the lock below, exactly like ``_get_checkpointer_config()`` above.
|
||||||
|
try:
|
||||||
|
app_config = get_app_config()
|
||||||
|
except FileNotFoundError:
|
||||||
|
app_config = None
|
||||||
|
|
||||||
with _checkpointer_lock:
|
with _checkpointer_lock:
|
||||||
if _checkpointer is not None:
|
if _checkpointer is not None:
|
||||||
return _checkpointer
|
return _checkpointer
|
||||||
|
|
||||||
checkpointer_ctx = _sync_checkpointer_cm(config)
|
checkpointer_ctx = _sync_checkpointer_cm(config)
|
||||||
checkpointer = checkpointer_ctx.__enter__()
|
checkpointer = checkpointer_ctx.__enter__()
|
||||||
|
try:
|
||||||
|
if app_config is not None:
|
||||||
|
checkpointer = _wrap_sync_if_delta(checkpointer, app_config)
|
||||||
|
except Exception:
|
||||||
|
checkpointer_ctx.__exit__(None, None, None)
|
||||||
|
raise
|
||||||
_checkpointer_ctx = checkpointer_ctx
|
_checkpointer_ctx = checkpointer_ctx
|
||||||
_checkpointer = checkpointer
|
_checkpointer = checkpointer
|
||||||
|
|
||||||
@ -195,7 +246,7 @@ def reset_checkpointer() -> None:
|
|||||||
Closes any open backend connections and clears the cached instance.
|
Closes any open backend connections and clears the cached instance.
|
||||||
Useful in tests or after a configuration change.
|
Useful in tests or after a configuration change.
|
||||||
"""
|
"""
|
||||||
global _checkpointer, _checkpointer_ctx
|
global _checkpointer, _checkpointer_ctx, _checkpointer_cache, _checkpointer_cache_prefix
|
||||||
with _checkpointer_lock:
|
with _checkpointer_lock:
|
||||||
if _checkpointer_ctx is not None:
|
if _checkpointer_ctx is not None:
|
||||||
try:
|
try:
|
||||||
@ -204,6 +255,8 @@ def reset_checkpointer() -> None:
|
|||||||
logger.warning("Error during checkpointer cleanup", exc_info=True)
|
logger.warning("Error during checkpointer cleanup", exc_info=True)
|
||||||
_checkpointer_ctx = None
|
_checkpointer_ctx = None
|
||||||
_checkpointer = None
|
_checkpointer = None
|
||||||
|
_checkpointer_cache = None
|
||||||
|
_checkpointer_cache_prefix = None
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@ -227,6 +280,7 @@ def checkpointer_context() -> Iterator[Checkpointer]:
|
|||||||
``InMemorySaver`` when neither selects a persistent backend.
|
``InMemorySaver`` when neither selects a persistent backend.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
config = _resolve_checkpointer_config(get_app_config())
|
app_config = get_app_config()
|
||||||
|
config = _resolve_checkpointer_config(app_config)
|
||||||
with _sync_checkpointer_cm(config) as saver:
|
with _sync_checkpointer_cm(config) as saver:
|
||||||
yield saver
|
yield _wrap_sync_if_delta(saver, app_config)
|
||||||
|
|||||||
@ -350,13 +350,41 @@ def _validate_materialized(case: BenchmarkCase, expected: list[BaseMessage], war
|
|||||||
return len(cold), cold_digest
|
return len(cold), cold_digest
|
||||||
|
|
||||||
|
|
||||||
|
_HISTORY_CACHE_ENV = "DEERFLOW_CHECKPOINT_BENCH_HISTORY_CACHE"
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_history_cache(saver: Any) -> Any:
|
||||||
|
"""Wrap *saver* in a CachedHistorySaver with a fresh, unbounded memory cache.
|
||||||
|
|
||||||
|
Opt-in via DEERFLOW_CHECKPOINT_BENCH_HISTORY_CACHE=1 so default rows are
|
||||||
|
byte-identical to the pre-cache benchmark. A fresh wrapper per phase keeps
|
||||||
|
the cold read genuinely cold: the write-phase cache is discarded, mirroring
|
||||||
|
a process restart (cache lifetime == checkpointer CM lifetime).
|
||||||
|
"""
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
from deerflow.runtime.checkpointer.cached_saver import CachedHistorySaver
|
||||||
|
|
||||||
|
return CachedHistorySaver(
|
||||||
|
saver,
|
||||||
|
MemoryCheckpointHistoryCache(max_entries=1_000_000),
|
||||||
|
key_prefix="bench:v1:checkpoint-bench",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _history_cache_stats(wrapper: Any, prefix: str) -> dict[str, Any]:
|
||||||
|
return {f"{prefix}{key}": value for key, value in wrapper.stats().items()}
|
||||||
|
|
||||||
|
|
||||||
def _run_memory_case(case: BenchmarkCase, messages: list[BaseMessage]) -> dict[str, Any]:
|
def _run_memory_case(case: BenchmarkCase, messages: list[BaseMessage]) -> dict[str, Any]:
|
||||||
|
cache_opt_in = os.environ.get(_HISTORY_CACHE_ENV) == "1"
|
||||||
saver = InMemorySaver()
|
saver = InMemorySaver()
|
||||||
metrics, warm = _write_and_read(case, saver, messages)
|
write_saver = _wrap_history_cache(saver) if cache_opt_in else saver
|
||||||
|
metrics, warm = _write_and_read(case, write_saver, messages)
|
||||||
stats = _collect_storage_stats(lambda: _memory_storage_stats(saver, _config(case)["configurable"]["thread_id"]))
|
stats = _collect_storage_stats(lambda: _memory_storage_stats(saver, _config(case)["configurable"]["thread_id"]))
|
||||||
cold_read_ms, cold = _cold_read(case, saver)
|
cold_saver = _wrap_history_cache(saver) if cache_opt_in else saver
|
||||||
|
cold_read_ms, cold = _cold_read(case, cold_saver)
|
||||||
actual_count, digest = _validate_materialized(case, messages, warm, cold)
|
actual_count, digest = _validate_materialized(case, messages, warm, cold)
|
||||||
return {
|
result = {
|
||||||
**metrics,
|
**metrics,
|
||||||
**stats,
|
**stats,
|
||||||
"cold_read_ms": cold_read_ms,
|
"cold_read_ms": cold_read_ms,
|
||||||
@ -371,12 +399,19 @@ def _run_memory_case(case: BenchmarkCase, messages: list[BaseMessage]) -> dict[s
|
|||||||
"actual_message_count": actual_count,
|
"actual_message_count": actual_count,
|
||||||
"content_sha256": digest,
|
"content_sha256": digest,
|
||||||
}
|
}
|
||||||
|
if cache_opt_in:
|
||||||
|
result["history_cache_enabled"] = True
|
||||||
|
result.update(_history_cache_stats(write_saver, "history_cache_write_"))
|
||||||
|
result.update(_history_cache_stats(cold_saver, "history_cache_cold_"))
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _run_sqlite_case(case: BenchmarkCase, messages: list[BaseMessage], db_path: Path) -> dict[str, Any]:
|
def _run_sqlite_case(case: BenchmarkCase, messages: list[BaseMessage], db_path: Path) -> dict[str, Any]:
|
||||||
|
cache_opt_in = os.environ.get(_HISTORY_CACHE_ENV) == "1"
|
||||||
with SqliteSaver.from_conn_string(str(db_path)) as saver:
|
with SqliteSaver.from_conn_string(str(db_path)) as saver:
|
||||||
saver.setup()
|
saver.setup()
|
||||||
metrics, warm = _write_and_read(case, saver, messages)
|
write_saver = _wrap_history_cache(saver) if cache_opt_in else saver
|
||||||
|
metrics, warm = _write_and_read(case, write_saver, messages)
|
||||||
stats = _collect_storage_stats(lambda: _sqlite_storage_stats(saver, _config(case)["configurable"]["thread_id"]))
|
stats = _collect_storage_stats(lambda: _sqlite_storage_stats(saver, _config(case)["configurable"]["thread_id"]))
|
||||||
db_bytes = _file_size(db_path)
|
db_bytes = _file_size(db_path)
|
||||||
wal_bytes = _file_size(Path(f"{db_path}-wal"))
|
wal_bytes = _file_size(Path(f"{db_path}-wal"))
|
||||||
@ -387,10 +422,11 @@ def _run_sqlite_case(case: BenchmarkCase, messages: list[BaseMessage], db_path:
|
|||||||
with SqliteSaver.from_conn_string(str(db_path)) as reopened:
|
with SqliteSaver.from_conn_string(str(db_path)) as reopened:
|
||||||
reopened.setup()
|
reopened.setup()
|
||||||
saver_reopen_ms = (time.perf_counter() - reopen_start) * 1000
|
saver_reopen_ms = (time.perf_counter() - reopen_start) * 1000
|
||||||
cold_read_ms, cold = _cold_read(case, reopened)
|
cold_saver = _wrap_history_cache(reopened) if cache_opt_in else reopened
|
||||||
|
cold_read_ms, cold = _cold_read(case, cold_saver)
|
||||||
|
|
||||||
actual_count, digest = _validate_materialized(case, messages, warm, cold)
|
actual_count, digest = _validate_materialized(case, messages, warm, cold)
|
||||||
return {
|
result = {
|
||||||
**metrics,
|
**metrics,
|
||||||
**stats,
|
**stats,
|
||||||
"cold_read_ms": cold_read_ms,
|
"cold_read_ms": cold_read_ms,
|
||||||
@ -403,6 +439,11 @@ def _run_sqlite_case(case: BenchmarkCase, messages: list[BaseMessage], db_path:
|
|||||||
"actual_message_count": actual_count,
|
"actual_message_count": actual_count,
|
||||||
"content_sha256": digest,
|
"content_sha256": digest,
|
||||||
}
|
}
|
||||||
|
if cache_opt_in:
|
||||||
|
result["history_cache_enabled"] = True
|
||||||
|
result.update(_history_cache_stats(write_saver, "history_cache_write_"))
|
||||||
|
result.update(_history_cache_stats(cold_saver, "history_cache_cold_"))
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _run_case(case: BenchmarkCase, *, work_dir: Path) -> dict[str, Any]:
|
def _run_case(case: BenchmarkCase, *, work_dir: Path) -> dict[str, Any]:
|
||||||
|
|||||||
398
backend/tests/test_cached_history_saver.py
Normal file
398
backend/tests/test_cached_history_saver.py
Normal file
@ -0,0 +1,398 @@
|
|||||||
|
"""CachedHistorySaver composition vs. the saver's own full walk."""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointTuple
|
||||||
|
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
from deerflow.runtime.checkpointer.cached_saver import CachedHistorySaver
|
||||||
|
|
||||||
|
PREFIX = "ckpt-hist:v1:testdb"
|
||||||
|
|
||||||
|
|
||||||
|
class _DictSaver(BaseCheckpointSaver):
|
||||||
|
"""Minimal in-memory saver. Deliberately does NOT override the delta
|
||||||
|
history methods, so the base class full parent-chain walk is the
|
||||||
|
differential oracle."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.checkpoints: dict[str, tuple[CheckpointTuple, ...]] = {}
|
||||||
|
self.tuple_reads = 0
|
||||||
|
self.history_walks = 0
|
||||||
|
|
||||||
|
def put_tuple(self, tup: CheckpointTuple) -> None:
|
||||||
|
cid = tup.config["configurable"]["checkpoint_id"]
|
||||||
|
self.checkpoints[cid] = (tup,)
|
||||||
|
|
||||||
|
def get_tuple(self, config):
|
||||||
|
self.tuple_reads += 1
|
||||||
|
configurable = config["configurable"]
|
||||||
|
cid = configurable.get("checkpoint_id")
|
||||||
|
if cid is None:
|
||||||
|
cid = next(reversed(self.checkpoints), None)
|
||||||
|
if cid is None:
|
||||||
|
return None
|
||||||
|
stored = self.checkpoints.get(cid)
|
||||||
|
return stored[0] if stored else None
|
||||||
|
|
||||||
|
async def aget_tuple(self, config):
|
||||||
|
return self.get_tuple(config)
|
||||||
|
|
||||||
|
def get_delta_channel_history(self, *, config, channels):
|
||||||
|
self.history_walks += 1
|
||||||
|
return super().get_delta_channel_history(config=config, channels=channels)
|
||||||
|
|
||||||
|
async def aget_delta_channel_history(self, *, config, channels):
|
||||||
|
self.history_walks += 1
|
||||||
|
return await super().aget_delta_channel_history(config=config, channels=channels)
|
||||||
|
|
||||||
|
# Unused abstract surface.
|
||||||
|
def list(self, config, *, filter=None, before=None, limit=None):
|
||||||
|
yield from ()
|
||||||
|
|
||||||
|
def put(self, config, checkpoint, metadata, new_versions):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def put_writes(self, config, writes, task_id, task_path=""):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def delete_thread(self, thread_id):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
def _cfg(thread: str, cid: str | None) -> dict:
|
||||||
|
configurable: dict[str, Any] = {"thread_id": thread, "checkpoint_ns": ""}
|
||||||
|
if cid is not None:
|
||||||
|
configurable["checkpoint_id"] = cid
|
||||||
|
return {"configurable": configurable}
|
||||||
|
|
||||||
|
|
||||||
|
def _tup(thread: str, cid: str, parent: str | None, *, channel_values: dict, writes: list) -> CheckpointTuple:
|
||||||
|
config = _cfg(thread, cid)
|
||||||
|
parent_config = _cfg(thread, parent) if parent else None
|
||||||
|
checkpoint = {"v": 1, "id": cid, "channel_values": channel_values, "channel_versions": {}, "versions_seen": {}, "updated_at": None}
|
||||||
|
return CheckpointTuple(config=config, checkpoint=checkpoint, metadata={}, parent_config=parent_config, pending_writes=list(writes))
|
||||||
|
|
||||||
|
|
||||||
|
def _chain(saver: _DictSaver, thread: str = "t1") -> list[str]:
|
||||||
|
"""c0(seed snapshot) -> c1(writes w1) -> c2(writes w2)."""
|
||||||
|
saver.put_tuple(_tup(thread, "c0", None, channel_values={"messages": ["seed-msg"]}, writes=[]))
|
||||||
|
saver.put_tuple(_tup(thread, "c1", "c0", channel_values={}, writes=[("task1", "messages", "w1")]))
|
||||||
|
saver.put_tuple(_tup(thread, "c2", "c1", channel_values={}, writes=[("task2", "messages", "w2")]))
|
||||||
|
return ["c0", "c1", "c2"]
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap(saver: _DictSaver, cache) -> CachedHistorySaver:
|
||||||
|
return CachedHistorySaver(saver, cache, key_prefix=PREFIX)
|
||||||
|
|
||||||
|
|
||||||
|
class _DeletableSaver(_DictSaver):
|
||||||
|
"""Functional delete/prune so purge behavior is observable."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.deleted_threads: list[str] = []
|
||||||
|
self.deleted_run_ids: list[list[str]] = []
|
||||||
|
self.pruned_threads: list[list[str]] = []
|
||||||
|
|
||||||
|
def delete_for_runs(self, run_ids):
|
||||||
|
self.deleted_run_ids.append(list(run_ids))
|
||||||
|
|
||||||
|
async def adelete_for_runs(self, run_ids):
|
||||||
|
self.deleted_run_ids.append(list(run_ids))
|
||||||
|
|
||||||
|
def _drop(self, thread_id: str) -> None:
|
||||||
|
self.deleted_threads.append(thread_id)
|
||||||
|
self.checkpoints = {cid: stored for cid, stored in self.checkpoints.items() if stored[0].config["configurable"]["thread_id"] != thread_id}
|
||||||
|
|
||||||
|
def delete_thread(self, thread_id):
|
||||||
|
self._drop(thread_id)
|
||||||
|
|
||||||
|
async def adelete_thread(self, thread_id):
|
||||||
|
self._drop(thread_id)
|
||||||
|
|
||||||
|
def prune(self, thread_ids, *, strategy="keep_latest"):
|
||||||
|
self.pruned_threads.append(list(thread_ids))
|
||||||
|
|
||||||
|
async def aprune(self, thread_ids, *, strategy="keep_latest"):
|
||||||
|
self.pruned_threads.append(list(thread_ids))
|
||||||
|
|
||||||
|
|
||||||
|
def _thread_entries(cache: MemoryCheckpointHistoryCache, thread_id: str) -> int:
|
||||||
|
stem = f"{PREFIX}:{thread_id}:"
|
||||||
|
return sum(1 for key in cache._data if key.startswith(stem))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_adelete_thread_purges_only_that_threads_cache_entries():
|
||||||
|
inner = _DeletableSaver()
|
||||||
|
_chain(inner, "t1")
|
||||||
|
# _DictSaver keys tuples by checkpoint_id alone: t2 needs distinct cids.
|
||||||
|
inner.put_tuple(_tup("t2", "u0", None, channel_values={"messages": ["seed2"]}, writes=[]))
|
||||||
|
inner.put_tuple(_tup("t2", "u1", "u0", channel_values={}, writes=[("task1", "messages", "x1")]))
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=32)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t2", "u1"), channels=["messages"])
|
||||||
|
assert _thread_entries(cache, "t1") > 0
|
||||||
|
assert _thread_entries(cache, "t2") > 0
|
||||||
|
|
||||||
|
await saver.adelete_thread("t1")
|
||||||
|
|
||||||
|
assert inner.deleted_threads == ["t1"]
|
||||||
|
assert inner.get_tuple(_cfg("t1", "c2")) is None # source of truth gone
|
||||||
|
assert _thread_entries(cache, "t1") == 0 # residual history payloads purged
|
||||||
|
assert _thread_entries(cache, "t2") > 0 # other threads untouched
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_delete_thread_purges_cache_entries():
|
||||||
|
inner = _DeletableSaver()
|
||||||
|
_chain(inner, "t1")
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=32)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
saver.get_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
assert _thread_entries(cache, "t1") > 0
|
||||||
|
|
||||||
|
saver.delete_thread("t1")
|
||||||
|
|
||||||
|
assert inner.deleted_threads == ["t1"]
|
||||||
|
assert _thread_entries(cache, "t1") == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_prune_purges_rewritten_threads_cache_entries():
|
||||||
|
inner = _DeletableSaver()
|
||||||
|
_chain(inner, "t1")
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=32)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
assert _thread_entries(cache, "t1") > 0
|
||||||
|
|
||||||
|
await saver.aprune(["t1"], strategy="keep_latest")
|
||||||
|
|
||||||
|
assert inner.pruned_threads == [["t1"]]
|
||||||
|
# The chain was rewritten: pre-prune histories must not linger.
|
||||||
|
assert _thread_entries(cache, "t1") == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_delete_for_runs_delegates_without_cache_purge():
|
||||||
|
"""Run-scoped deletes cannot be mapped to threads cheaply; documented
|
||||||
|
behavior is delegation with LRU/TTL-bounded residual retention."""
|
||||||
|
inner = _DeletableSaver()
|
||||||
|
_chain(inner, "t1")
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=32)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
entries_before = cache.stats().entries
|
||||||
|
|
||||||
|
await saver.adelete_for_runs(["run-1"]) # base-class no-op on _DictSaver lineage
|
||||||
|
|
||||||
|
assert inner.deleted_run_ids == [["run-1"]]
|
||||||
|
assert cache.stats().entries == entries_before
|
||||||
|
|
||||||
|
|
||||||
|
class _RecordingSaver(_DictSaver):
|
||||||
|
"""Captures the config passed into each fallback walk."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.walk_configs: list[dict] = []
|
||||||
|
|
||||||
|
def get_delta_channel_history(self, *, config, channels):
|
||||||
|
self.walk_configs.append(config)
|
||||||
|
return super().get_delta_channel_history(config=config, channels=channels)
|
||||||
|
|
||||||
|
async def aget_delta_channel_history(self, *, config, channels):
|
||||||
|
self.walk_configs.append(config)
|
||||||
|
return await super().aget_delta_channel_history(config=config, channels=channels)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_composition_matches_full_walk_and_avoids_it():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=16)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
|
||||||
|
# Cold: c1 composes from snapshot parent c0 (channel_values hit) — no walk.
|
||||||
|
h1 = await saver.aget_delta_channel_history(config=_cfg("t1", "c1"), channels=["messages"])
|
||||||
|
assert h1["messages"]["writes"] == []
|
||||||
|
assert h1["messages"]["seed"] == ["seed-msg"]
|
||||||
|
assert inner.history_walks == 0
|
||||||
|
|
||||||
|
# c2 composes from cached history(c1) + c1's pending writes: no inner walk.
|
||||||
|
# NOTE: history(c2) excludes c2's OWN pending writes (they belong to the
|
||||||
|
# next super-step per the LangGraph contract) — on-path writes are c1's.
|
||||||
|
h2 = await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
assert h2["messages"]["writes"] == [("task1", "messages", "w1")]
|
||||||
|
assert h2["messages"]["seed"] == ["seed-msg"]
|
||||||
|
assert inner.history_walks == 0 # unchanged: composition, not a walk
|
||||||
|
|
||||||
|
# Differential oracle: identical to the inner saver's own full walk.
|
||||||
|
oracle = _DictSaver()
|
||||||
|
_chain(oracle)
|
||||||
|
for cid in ("c0", "c1", "c2"):
|
||||||
|
expected = await oracle.aget_delta_channel_history(config=_cfg("t1", cid), channels=["messages"])
|
||||||
|
actual = await saver.aget_delta_channel_history(config=_cfg("t1", cid), channels=["messages"])
|
||||||
|
assert actual == expected, cid
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_snapshot_parent_composes_without_parent_history_lookup():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
# c3 has c0-style snapshot directly at parent c1: rewrite c1 with channel_values.
|
||||||
|
inner.put_tuple(_tup("t1", "c1b", "c0", channel_values={"messages": ["snap"]}, writes=[("t", "messages", "wx")]))
|
||||||
|
inner.put_tuple(_tup("t1", "c2b", "c1b", channel_values={}, writes=[("t2", "messages", "wy")]))
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
h = await saver.aget_delta_channel_history(config=_cfg("t1", "c2b"), channels=["messages"])
|
||||||
|
assert h["messages"] == {"writes": [("t", "messages", "wx")], "seed": ["snap"]}
|
||||||
|
assert inner.history_walks == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_root_checkpoint_history_is_empty_writes():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
h = await saver.aget_delta_channel_history(config=_cfg("t1", "c0"), channels=["messages"])
|
||||||
|
assert h["messages"] == {"writes": []}
|
||||||
|
assert "seed" not in h["messages"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_latest_config_caches_under_resolved_checkpoint_id():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=16)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", None), channels=["messages"])
|
||||||
|
reads_after_first = inner.tuple_reads
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
# Second call resolves by id from cache: target tuple only, no parent refetch.
|
||||||
|
assert inner.tuple_reads == reads_after_first + 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_eviction_falls_back_to_walk_but_stays_correct():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=1)
|
||||||
|
saver = _wrap(inner, cache)
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"]) # cold: fallback walk
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c1"), channels=["messages"]) # compose; evicts c2 entry
|
||||||
|
h = await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"]) # recomposes via cached history(c1)
|
||||||
|
assert h["messages"]["writes"] == [("task1", "messages", "w1")]
|
||||||
|
assert h["messages"]["seed"] == ["seed-msg"]
|
||||||
|
assert saver.stats()["full_walks"] == 0 # cold read resolved by recursive compose
|
||||||
|
assert inner.history_walks == 0 # resolution never delegates to the inner walk
|
||||||
|
assert cache.stats().evictions >= 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_path_matches_async():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
h = saver.get_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"])
|
||||||
|
assert h["messages"]["writes"] == [("task1", "messages", "w1")]
|
||||||
|
assert h["messages"]["seed"] == ["seed-msg"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_cold_resolve_never_delegates_to_inner_history():
|
||||||
|
"""Cold reads compose recursively (or walk themselves via aget_tuple):
|
||||||
|
the inner saver's own history method is never called, so a 'latest'
|
||||||
|
config cannot be re-resolved mid-resolution (the old pinned race is gone
|
||||||
|
by construction)."""
|
||||||
|
inner = _RecordingSaver()
|
||||||
|
_chain(inner)
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
h = await saver.aget_delta_channel_history(config=_cfg("t1", None), channels=["messages"])
|
||||||
|
assert h["messages"]["writes"] == [("task1", "messages", "w1")]
|
||||||
|
assert h["messages"]["seed"] == ["seed-msg"]
|
||||||
|
assert inner.walk_configs == [], "resolution must not delegate to the inner history walk"
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_cold_resolve_never_delegates_to_inner_history():
|
||||||
|
inner = _RecordingSaver()
|
||||||
|
_chain(inner)
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
h = saver.get_delta_channel_history(config=_cfg("t1", None), channels=["messages"])
|
||||||
|
assert h["messages"]["writes"] == [("task1", "messages", "w1")]
|
||||||
|
assert inner.walk_configs == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_recursive_resolve_caches_intermediate_levels():
|
||||||
|
"""Cold short chains resolve by recursive compose: every intermediate
|
||||||
|
level is computed once and cached, later reads hit directly."""
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
inner.put_tuple(_tup("t1", "c3", "c2", channel_values={}, writes=[("task3", "messages", "w3")]))
|
||||||
|
oracle = _DictSaver()
|
||||||
|
_chain(oracle)
|
||||||
|
oracle.put_tuple(_tup("t1", "c3", "c2", channel_values={}, writes=[("task3", "messages", "w3")]))
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c3"), channels=["messages"])
|
||||||
|
reads_after_cold = inner.tuple_reads
|
||||||
|
|
||||||
|
for cid in ("c0", "c1", "c2", "c3"):
|
||||||
|
expected = await oracle.aget_delta_channel_history(config=_cfg("t1", cid), channels=["messages"])
|
||||||
|
actual = await saver.aget_delta_channel_history(config=_cfg("t1", cid), channels=["messages"])
|
||||||
|
assert actual == expected, cid
|
||||||
|
# Every level is now warm: each read costs exactly its target tuple fetch.
|
||||||
|
assert inner.tuple_reads - reads_after_cold == 4
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_deep_cold_chain_delegates_one_inner_walk_at_depth_limit():
|
||||||
|
"""A cold chain deeper than the compose budget resolves the deepest
|
||||||
|
reached level with ONE inner fast-path walk (2 SQL), caches every level
|
||||||
|
above it, and leaves deeper ancestors cold until asked."""
|
||||||
|
inner = _DictSaver()
|
||||||
|
# 12-deep chain: c0(seed) <- c1 <- ... <- c11, deeper than the budget.
|
||||||
|
inner.put_tuple(_tup("t1", "c0", None, channel_values={"messages": ["seed-msg"]}, writes=[]))
|
||||||
|
for i in range(1, 12):
|
||||||
|
inner.put_tuple(_tup("t1", f"c{i}", f"c{i - 1}", channel_values={}, writes=[("task", "messages", f"w{i}")]))
|
||||||
|
oracle = _DictSaver()
|
||||||
|
oracle.put_tuple(_tup("t1", "c0", None, channel_values={"messages": ["seed-msg"]}, writes=[]))
|
||||||
|
for i in range(1, 12):
|
||||||
|
oracle.put_tuple(_tup("t1", f"c{i}", f"c{i - 1}", channel_values={}, writes=[("task", "messages", f"w{i}")]))
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=64))
|
||||||
|
|
||||||
|
h = await saver.aget_delta_channel_history(config=_cfg("t1", "c11"), channels=["messages"])
|
||||||
|
assert saver.stats()["full_walks"] == 1
|
||||||
|
assert inner.history_walks == 1 # exactly one delegated fast-path walk
|
||||||
|
assert h["messages"]["writes"] == [("task", "messages", f"w{i}") for i in range(1, 11)]
|
||||||
|
assert h["messages"]["seed"] == ["seed-msg"]
|
||||||
|
|
||||||
|
reads_after_cold = inner.tuple_reads
|
||||||
|
for cid in [f"c{i}" for i in range(12)]:
|
||||||
|
expected = await oracle.aget_delta_channel_history(config=_cfg("t1", cid), channels=["messages"])
|
||||||
|
actual = await saver.aget_delta_channel_history(config=_cfg("t1", cid), channels=["messages"])
|
||||||
|
assert actual == expected, cid
|
||||||
|
# Warm after the cold read: c2..c11. Reads: c0 = 1 (root, no parent);
|
||||||
|
# c1 = 2 (target + snapshot parent c0); c2..c11 = 1 each.
|
||||||
|
assert inner.tuple_reads - reads_after_cold == 13
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_stats_expose_composition_counters():
|
||||||
|
inner = _DictSaver()
|
||||||
|
_chain(inner)
|
||||||
|
saver = _wrap(inner, MemoryCheckpointHistoryCache(max_entries=16))
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c2"), channels=["messages"]) # cold: recursive compose ×2, warms c1+c2
|
||||||
|
await saver.aget_delta_channel_history(config=_cfg("t1", "c1"), channels=["messages"]) # direct hit (warmed)
|
||||||
|
# Compose: add c3 (parent c2, not a snapshot) AFTER the cold read — c3 is
|
||||||
|
# not cached but history(c2) is, so c3 composes without a walk.
|
||||||
|
inner.put_tuple(_tup("t1", "c3", "c2", channel_values={}, writes=[("task3", "messages", "w3")]))
|
||||||
|
h3 = await saver.aget_delta_channel_history(config=_cfg("t1", "c3"), channels=["messages"])
|
||||||
|
assert h3["messages"]["writes"] == [("task1", "messages", "w1"), ("task2", "messages", "w2")]
|
||||||
|
stats = saver.stats()
|
||||||
|
assert stats["full_walks"] == 0
|
||||||
|
assert stats["compose_hits"] == 3 # c1-level + c2-level (cold) + c3
|
||||||
|
assert stats["hits"] >= 1
|
||||||
364
backend/tests/test_cached_history_saver_integration.py
Normal file
364
backend/tests/test_cached_history_saver_integration.py
Normal file
@ -0,0 +1,364 @@
|
|||||||
|
"""Behavioral integration tests for CachedHistorySaver on REAL LangGraph execution.
|
||||||
|
|
||||||
|
Unlike tests/test_cached_history_saver.py (fake saver, hand-built chains), these
|
||||||
|
tests drive compiled StateGraphs through pregel in delta mode
|
||||||
|
(``DeltaChannel(merge_message_writes, snapshot_frequency=2)``) and verify the
|
||||||
|
cache against a differential oracle: the identical scenario executed on a raw
|
||||||
|
``InMemorySaver`` in a fresh thread. Digests are (type, content, id) triples of
|
||||||
|
the materialized ``messages`` channel, so any history corruption shows up as a
|
||||||
|
digest mismatch.
|
||||||
|
|
||||||
|
Observed pregel call pattern on langgraph 1.2.9 (5-step linear graph, 7
|
||||||
|
checkpoints, snapshot cadence 2): one ``aget_delta_channel_history`` per run
|
||||||
|
start (empty-thread load), none for snapshot checkpoints, one per materialized
|
||||||
|
non-snapshot checkpoint on the raw saver. The cached saver pays the run-start
|
||||||
|
walk plus one cold fallback walk; every other materialization composes from a
|
||||||
|
parent snapshot seed or a cached parent history, and a second identical read
|
||||||
|
pass costs zero inner walks.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Annotated, Any, TypedDict
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langchain_core.messages import AIMessage, AnyMessage, HumanMessage
|
||||||
|
from langgraph.channels import DeltaChannel
|
||||||
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
|
from langgraph.graph import StateGraph
|
||||||
|
from langgraph.types import Command, interrupt
|
||||||
|
|
||||||
|
from deerflow.agents.thread_state import merge_message_writes
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
from deerflow.runtime.checkpoint_state import CheckpointStateAccessor
|
||||||
|
from deerflow.runtime.checkpointer.cached_saver import CachedHistorySaver
|
||||||
|
|
||||||
|
STEPS = 5
|
||||||
|
SNAPSHOT_FREQUENCY = 2
|
||||||
|
|
||||||
|
|
||||||
|
class _CountingInMemorySaver(InMemorySaver):
|
||||||
|
"""InMemorySaver that counts full delta-history walks.
|
||||||
|
|
||||||
|
Placed under ``CachedHistorySaver`` it records exactly the walks the cache
|
||||||
|
could not serve; used bare it is the uncached oracle's walk counter.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.history_walks = 0
|
||||||
|
|
||||||
|
def get_delta_channel_history(self, *, config: Any, channels: Any) -> Any:
|
||||||
|
self.history_walks += 1
|
||||||
|
return super().get_delta_channel_history(config=config, channels=channels)
|
||||||
|
|
||||||
|
async def aget_delta_channel_history(self, *, config: Any, channels: Any) -> Any:
|
||||||
|
self.history_walks += 1
|
||||||
|
return await super().aget_delta_channel_history(config=config, channels=channels)
|
||||||
|
|
||||||
|
|
||||||
|
def _state_schema() -> type:
|
||||||
|
class State(TypedDict):
|
||||||
|
messages: Annotated[
|
||||||
|
list[AnyMessage],
|
||||||
|
DeltaChannel(merge_message_writes, snapshot_frequency=SNAPSHOT_FREQUENCY),
|
||||||
|
]
|
||||||
|
|
||||||
|
return State
|
||||||
|
|
||||||
|
|
||||||
|
def _make_step_node(n: int) -> Any:
|
||||||
|
def node(state: dict) -> dict:
|
||||||
|
return {"messages": [AIMessage(content=f"step-{n}", id=f"ai-{n}")]}
|
||||||
|
|
||||||
|
return node
|
||||||
|
|
||||||
|
|
||||||
|
def _build_graph(saver: Any, steps: int = STEPS) -> Any:
|
||||||
|
builder = StateGraph(_state_schema())
|
||||||
|
for i in range(steps):
|
||||||
|
builder.add_node(f"step{i}", _make_step_node(i))
|
||||||
|
builder.set_entry_point("step0")
|
||||||
|
for i in range(steps - 1):
|
||||||
|
builder.add_edge(f"step{i}", f"step{i + 1}")
|
||||||
|
builder.set_finish_point(f"step{steps - 1}")
|
||||||
|
return builder.compile(checkpointer=saver)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_interrupt_graph(saver: Any) -> Any:
|
||||||
|
"""step0 -> step1 -> pause(interrupt) -> step2 -> step3."""
|
||||||
|
|
||||||
|
def pause(state: dict) -> dict:
|
||||||
|
answer = interrupt({"question": "continue?"})
|
||||||
|
return {"messages": [AIMessage(content=f"resumed:{answer}", id="ai-resume")]}
|
||||||
|
|
||||||
|
builder = StateGraph(_state_schema())
|
||||||
|
builder.add_node("step0", _make_step_node(0))
|
||||||
|
builder.add_node("step1", _make_step_node(1))
|
||||||
|
builder.add_node("pause", pause)
|
||||||
|
builder.add_node("step2", _make_step_node(2))
|
||||||
|
builder.add_node("step3", _make_step_node(3))
|
||||||
|
builder.set_entry_point("step0")
|
||||||
|
builder.add_edge("step0", "step1")
|
||||||
|
builder.add_edge("step1", "pause")
|
||||||
|
builder.add_edge("pause", "step2")
|
||||||
|
builder.add_edge("step2", "step3")
|
||||||
|
builder.set_finish_point("step3")
|
||||||
|
return builder.compile(checkpointer=saver)
|
||||||
|
|
||||||
|
|
||||||
|
def _config() -> dict[str, Any]:
|
||||||
|
return {"configurable": {"thread_id": f"cache-itest-{uuid4().hex}"}}
|
||||||
|
|
||||||
|
|
||||||
|
def _input() -> dict[str, Any]:
|
||||||
|
return {"messages": [HumanMessage(content="kickoff", id="h-0")]}
|
||||||
|
|
||||||
|
|
||||||
|
def _digest(values: dict[str, Any]) -> list[tuple[str, str, str | None]]:
|
||||||
|
return [(m.type, m.content, m.id) for m in values["messages"]]
|
||||||
|
|
||||||
|
|
||||||
|
def _history_digests(snapshots: list[Any]) -> list[list[tuple[str, str, str | None]]]:
|
||||||
|
return [_digest(s.values) for s in snapshots]
|
||||||
|
|
||||||
|
|
||||||
|
def _expected_final_digest(steps: int = STEPS) -> list[tuple[str, str, str | None]]:
|
||||||
|
return [("human", "kickoff", "h-0"), *[("ai", f"step-{n}", f"ai-{n}") for n in range(steps)]]
|
||||||
|
|
||||||
|
|
||||||
|
def _make_cached_stack(max_entries: int = 128) -> tuple[_CountingInMemorySaver, CachedHistorySaver, Any, CheckpointStateAccessor]:
|
||||||
|
inner = _CountingInMemorySaver()
|
||||||
|
saver = CachedHistorySaver(inner, MemoryCheckpointHistoryCache(max_entries), key_prefix=f"itest-{uuid4().hex}")
|
||||||
|
graph = _build_graph(saver)
|
||||||
|
accessor = CheckpointStateAccessor.bind(graph, saver, mode="delta")
|
||||||
|
return inner, saver, graph, accessor
|
||||||
|
|
||||||
|
|
||||||
|
def _make_oracle_stack() -> tuple[_CountingInMemorySaver, Any, CheckpointStateAccessor]:
|
||||||
|
inner = _CountingInMemorySaver()
|
||||||
|
graph = _build_graph(inner)
|
||||||
|
accessor = CheckpointStateAccessor.bind(graph, inner, mode="delta")
|
||||||
|
return inner, graph, accessor
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_sequential_run_composes_without_inner_walks() -> None:
|
||||||
|
"""A cached run must serve warm reads with strictly fewer inner history
|
||||||
|
walks than the identical uncached run, composing histories instead."""
|
||||||
|
inner, saver, graph, accessor = _make_cached_stack()
|
||||||
|
config = _config()
|
||||||
|
await graph.ainvoke(_input(), config)
|
||||||
|
|
||||||
|
final = await accessor.aget(config)
|
||||||
|
assert _digest(final.values) == _expected_final_digest()
|
||||||
|
|
||||||
|
# Cold pass: materialize every checkpoint in the thread.
|
||||||
|
cold = await accessor.ahistory(config)
|
||||||
|
cold_walks = inner.history_walks
|
||||||
|
|
||||||
|
# Warm pass: identical reads must be served entirely from the cache.
|
||||||
|
warm = await accessor.ahistory(config)
|
||||||
|
assert inner.history_walks == cold_walks, "warm re-read triggered an inner walk"
|
||||||
|
assert _history_digests(warm) == _history_digests(cold)
|
||||||
|
|
||||||
|
# Differential oracle: same run through the raw saver.
|
||||||
|
oracle_inner, oracle_graph, oracle_accessor = _make_oracle_stack()
|
||||||
|
oracle_config = _config()
|
||||||
|
await oracle_graph.ainvoke(_input(), oracle_config)
|
||||||
|
await oracle_accessor.aget(oracle_config)
|
||||||
|
oracle_history = await oracle_accessor.ahistory(oracle_config)
|
||||||
|
|
||||||
|
assert _history_digests(cold) == _history_digests(oracle_history)
|
||||||
|
assert oracle_inner.history_walks > 0
|
||||||
|
assert cold_walks < oracle_inner.history_walks
|
||||||
|
assert saver.stats()["compose_hits"] > 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_cache_disabled_parity() -> None:
|
||||||
|
"""A zero-entry cache must behave exactly like the raw inner saver."""
|
||||||
|
inner, saver, graph, accessor = _make_cached_stack(max_entries=0)
|
||||||
|
config = _config()
|
||||||
|
await graph.ainvoke(_input(), config)
|
||||||
|
final = await accessor.aget(config)
|
||||||
|
history = await accessor.ahistory(config)
|
||||||
|
|
||||||
|
oracle_inner, oracle_graph, oracle_accessor = _make_oracle_stack()
|
||||||
|
oracle_config = _config()
|
||||||
|
await oracle_graph.ainvoke(_input(), oracle_config)
|
||||||
|
oracle_final = await oracle_accessor.aget(oracle_config)
|
||||||
|
oracle_history = await oracle_accessor.ahistory(oracle_config)
|
||||||
|
|
||||||
|
assert _digest(final.values) == _expected_final_digest()
|
||||||
|
assert _digest(final.values) == _digest(oracle_final.values)
|
||||||
|
assert _history_digests(history) == _history_digests(oracle_history)
|
||||||
|
assert saver.stats()["hits"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
async def _run_branch_scenario(
|
||||||
|
graph: Any,
|
||||||
|
accessor: CheckpointStateAccessor,
|
||||||
|
config: dict[str, Any],
|
||||||
|
*,
|
||||||
|
fork_next: str,
|
||||||
|
as_node: str,
|
||||||
|
branch_id: str,
|
||||||
|
) -> dict[tuple[str, ...], Any]:
|
||||||
|
"""Run to completion, fork at the checkpoint whose next node is
|
||||||
|
``fork_next``, then resume the branch to completion (the branch head is
|
||||||
|
then the thread's latest checkpoint). Returns the original chain's
|
||||||
|
snapshots keyed by their ``next`` tuple for pinned re-reads.
|
||||||
|
|
||||||
|
``fork_next`` must name a NON-snapshot checkpoint: on langgraph 1.2.9 a
|
||||||
|
``aupdate_state`` fork at a snapshot checkpoint silently drops the update
|
||||||
|
(verified against a raw InMemorySaver - upstream behavior, not the cache).
|
||||||
|
"""
|
||||||
|
await graph.ainvoke(_input(), config)
|
||||||
|
history = await accessor.ahistory(config)
|
||||||
|
by_next = {s.next: s for s in history}
|
||||||
|
base = by_next[(fork_next,)]
|
||||||
|
branch_config = await accessor.aupdate(
|
||||||
|
base.config,
|
||||||
|
{"messages": [AIMessage(content=branch_id, id=f"ai-{branch_id}")]},
|
||||||
|
as_node=as_node,
|
||||||
|
)
|
||||||
|
await graph.ainvoke(None, branch_config)
|
||||||
|
return by_next
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_branch_divergence_no_cross_contamination() -> None:
|
||||||
|
"""A forked branch and the original head must each materialize their own
|
||||||
|
distinct history through the SAME cached saver."""
|
||||||
|
inner, saver, graph, accessor = _make_cached_stack()
|
||||||
|
config = _config()
|
||||||
|
by_next = await _run_branch_scenario(graph, accessor, config, fork_next="step2", as_node="step2", branch_id="branch")
|
||||||
|
|
||||||
|
# Oracle: identical branch scenario on the raw saver.
|
||||||
|
oracle_inner, oracle_graph, oracle_accessor = _make_oracle_stack()
|
||||||
|
oracle_config = _config()
|
||||||
|
oracle_by_next = await _run_branch_scenario(oracle_graph, oracle_accessor, oracle_config, fork_next="step2", as_node="step2", branch_id="branch")
|
||||||
|
|
||||||
|
# Branch head = thread's latest checkpoint after the forked resume.
|
||||||
|
branch_head = await accessor.aget(config)
|
||||||
|
oracle_branch_head = await oracle_accessor.aget(oracle_config)
|
||||||
|
expected_branch = [
|
||||||
|
("human", "kickoff", "h-0"),
|
||||||
|
("ai", "step-0", "ai-0"),
|
||||||
|
("ai", "step-1", "ai-1"),
|
||||||
|
("ai", "branch", "ai-branch"),
|
||||||
|
("ai", "step-3", "ai-3"),
|
||||||
|
("ai", "step-4", "ai-4"),
|
||||||
|
]
|
||||||
|
assert _digest(branch_head.values) == expected_branch
|
||||||
|
assert _digest(branch_head.values) == _digest(oracle_branch_head.values)
|
||||||
|
|
||||||
|
# The original chain still materializes its own un-branched history: the
|
||||||
|
# snapshot head and a non-snapshot mid checkpoint (cache-exercising read).
|
||||||
|
for next_key, expected_len in [((), 6), (("step4",), 5)]:
|
||||||
|
reread_original = await accessor.aget(by_next[next_key].config)
|
||||||
|
oracle_reread_original = await oracle_accessor.aget(oracle_by_next[next_key].config)
|
||||||
|
assert _digest(reread_original.values) == _expected_final_digest()[:expected_len]
|
||||||
|
assert _digest(reread_original.values) == _digest(oracle_reread_original.values)
|
||||||
|
assert _digest(reread_original.values) != _digest(branch_head.values)
|
||||||
|
|
||||||
|
|
||||||
|
async def _run_interrupt_scenario(graph: Any, accessor: CheckpointStateAccessor, config: dict[str, Any]) -> None:
|
||||||
|
result = await graph.ainvoke(_input(), config)
|
||||||
|
assert "__interrupt__" in result
|
||||||
|
await graph.ainvoke(Command(resume="yes"), config)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_interrupt_resume_appended_head_writes() -> None:
|
||||||
|
"""Resume appends writes under the interrupted head checkpoint; the cached
|
||||||
|
final state must equal the no-cache reference."""
|
||||||
|
inner = _CountingInMemorySaver()
|
||||||
|
saver = CachedHistorySaver(inner, MemoryCheckpointHistoryCache(128), key_prefix=f"itest-{uuid4().hex}")
|
||||||
|
graph = _build_interrupt_graph(saver)
|
||||||
|
accessor = CheckpointStateAccessor.bind(graph, saver, mode="delta")
|
||||||
|
config = _config()
|
||||||
|
await _run_interrupt_scenario(graph, accessor, config)
|
||||||
|
|
||||||
|
oracle_inner = _CountingInMemorySaver()
|
||||||
|
oracle_graph = _build_interrupt_graph(oracle_inner)
|
||||||
|
oracle_accessor = CheckpointStateAccessor.bind(oracle_graph, oracle_inner, mode="delta")
|
||||||
|
oracle_config = _config()
|
||||||
|
await _run_interrupt_scenario(oracle_graph, oracle_accessor, oracle_config)
|
||||||
|
|
||||||
|
final = await accessor.aget(config)
|
||||||
|
oracle_final = await oracle_accessor.aget(oracle_config)
|
||||||
|
expected = [
|
||||||
|
("human", "kickoff", "h-0"),
|
||||||
|
("ai", "step-0", "ai-0"),
|
||||||
|
("ai", "step-1", "ai-1"),
|
||||||
|
("ai", "resumed:yes", "ai-resume"),
|
||||||
|
("ai", "step-2", "ai-2"),
|
||||||
|
("ai", "step-3", "ai-3"),
|
||||||
|
]
|
||||||
|
assert _digest(final.values) == expected
|
||||||
|
assert _digest(final.values) == _digest(oracle_final.values)
|
||||||
|
|
||||||
|
# Every checkpoint along the resumed thread matches the oracle, including
|
||||||
|
# the interrupted head whose pending writes grew at resume time.
|
||||||
|
history = await accessor.ahistory(config)
|
||||||
|
oracle_history = await oracle_accessor.ahistory(oracle_config)
|
||||||
|
assert _history_digests(history) == _history_digests(oracle_history)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_eviction_only_costs_performance() -> None:
|
||||||
|
"""A 1-entry LRU thrashes on every read but must stay correct."""
|
||||||
|
inner, saver, graph, accessor = _make_cached_stack(max_entries=1)
|
||||||
|
config = _config()
|
||||||
|
await graph.ainvoke(_input(), config)
|
||||||
|
|
||||||
|
first_pass = await accessor.ahistory(config)
|
||||||
|
second_pass = await accessor.ahistory(config)
|
||||||
|
|
||||||
|
oracle_inner, oracle_graph, oracle_accessor = _make_oracle_stack()
|
||||||
|
oracle_config = _config()
|
||||||
|
await oracle_graph.ainvoke(_input(), oracle_config)
|
||||||
|
oracle_history = await oracle_accessor.ahistory(oracle_config)
|
||||||
|
|
||||||
|
assert _history_digests(first_pass) == _history_digests(oracle_history)
|
||||||
|
assert _history_digests(second_pass) == _history_digests(oracle_history)
|
||||||
|
assert saver.stats()["evictions"] > 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_rollback_supersede_does_not_pollute() -> None:
|
||||||
|
"""Re-running from an early checkpoint supersedes the head; the original
|
||||||
|
head's cached history must remain intact and re-readable."""
|
||||||
|
inner, saver, graph, accessor = _make_cached_stack()
|
||||||
|
config = _config()
|
||||||
|
by_next = await _run_branch_scenario(graph, accessor, config, fork_next="step0", as_node="step0", branch_id="rollback")
|
||||||
|
|
||||||
|
oracle_inner, oracle_graph, oracle_accessor = _make_oracle_stack()
|
||||||
|
oracle_config = _config()
|
||||||
|
oracle_by_next = await _run_branch_scenario(oracle_graph, oracle_accessor, oracle_config, fork_next="step0", as_node="step0", branch_id="rollback")
|
||||||
|
|
||||||
|
# New head equals the reference run's new head.
|
||||||
|
new_head = await accessor.aget(config)
|
||||||
|
oracle_new_head = await oracle_accessor.aget(oracle_config)
|
||||||
|
expected_new_head = [
|
||||||
|
("human", "kickoff", "h-0"),
|
||||||
|
("ai", "rollback", "ai-rollback"),
|
||||||
|
("ai", "step-1", "ai-1"),
|
||||||
|
("ai", "step-2", "ai-2"),
|
||||||
|
("ai", "step-3", "ai-3"),
|
||||||
|
("ai", "step-4", "ai-4"),
|
||||||
|
]
|
||||||
|
assert _digest(new_head.values) == expected_new_head
|
||||||
|
assert _digest(new_head.values) == _digest(oracle_new_head.values)
|
||||||
|
|
||||||
|
# Re-reading ORIGINAL chain checkpoints (pinned by checkpoint_id) returns
|
||||||
|
# their own original histories - their cached entries predate the fork and
|
||||||
|
# must be untouched by the superseding branch.
|
||||||
|
for next_key, expected_len in [((), 6), (("step4",), 5), (("step2",), 3)]:
|
||||||
|
reread = await accessor.aget(by_next[next_key].config)
|
||||||
|
oracle_reread = await oracle_accessor.aget(oracle_by_next[next_key].config)
|
||||||
|
assert _digest(reread.values) == _expected_final_digest()[:expected_len]
|
||||||
|
assert _digest(reread.values) == _digest(oracle_reread.values)
|
||||||
46
backend/tests/test_checkpoint_cache_config.py
Normal file
46
backend/tests/test_checkpoint_cache_config.py
Normal file
@ -0,0 +1,46 @@
|
|||||||
|
"""Config parsing for database.checkpoint_cache."""
|
||||||
|
|
||||||
|
from deerflow.config.database_config import CheckpointCacheConfig, DatabaseConfig
|
||||||
|
|
||||||
|
|
||||||
|
def test_checkpoint_cache_defaults():
|
||||||
|
cfg = DatabaseConfig()
|
||||||
|
assert cfg.checkpoint_cache.type == "memory"
|
||||||
|
assert cfg.checkpoint_cache.max_entries == 128
|
||||||
|
assert cfg.checkpoint_cache.redis_url is None
|
||||||
|
assert cfg.checkpoint_cache.ttl_seconds == 86400
|
||||||
|
assert cfg.checkpoint_cache.key_prefix == ""
|
||||||
|
|
||||||
|
|
||||||
|
def test_checkpoint_cache_from_dict_redis():
|
||||||
|
cfg = DatabaseConfig.model_validate(
|
||||||
|
{
|
||||||
|
"backend": "postgres",
|
||||||
|
"postgres_url": "postgresql://u:p@h/db",
|
||||||
|
"checkpoint_cache": {
|
||||||
|
"type": "redis",
|
||||||
|
"max_entries": 256,
|
||||||
|
"redis_url": "redis://localhost:6379/3",
|
||||||
|
"ttl_seconds": 3600,
|
||||||
|
"key_prefix": "prod:",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert cfg.checkpoint_cache.type == "redis"
|
||||||
|
assert cfg.checkpoint_cache.max_entries == 256
|
||||||
|
assert cfg.checkpoint_cache.redis_url == "redis://localhost:6379/3"
|
||||||
|
assert cfg.checkpoint_cache.ttl_seconds == 3600
|
||||||
|
assert cfg.checkpoint_cache.key_prefix == "prod:"
|
||||||
|
|
||||||
|
|
||||||
|
def test_checkpoint_cache_zero_max_entries_means_disabled():
|
||||||
|
cfg = CheckpointCacheConfig(max_entries=0)
|
||||||
|
assert cfg.max_entries == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_checkpoint_cache_rejects_negative_max_entries():
|
||||||
|
import pydantic
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
with pytest.raises(pydantic.ValidationError):
|
||||||
|
CheckpointCacheConfig(max_entries=-1)
|
||||||
130
backend/tests/test_checkpoint_cache_memory.py
Normal file
130
backend/tests/test_checkpoint_cache_memory.py
Normal file
@ -0,0 +1,130 @@
|
|||||||
|
"""Memory LRU backend for the checkpoint history cache."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from deerflow.runtime.checkpoint_cache.base import (
|
||||||
|
CACHE_FORMAT_VERSION,
|
||||||
|
CheckpointCacheStats,
|
||||||
|
make_history_key,
|
||||||
|
thread_key_stem,
|
||||||
|
)
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
|
||||||
|
|
||||||
|
def _entry(tag: str) -> dict:
|
||||||
|
return {"writes": [("task-1", "messages", tag)], "seed": f"seed-{tag}"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_make_history_key_is_stable_and_scoped():
|
||||||
|
k1 = make_history_key("ckpt-hist:v1:db0", "t1", "", "c1", "messages")
|
||||||
|
k2 = make_history_key("ckpt-hist:v1:db0", "t1", "", "c1", "messages")
|
||||||
|
assert k1 == k2
|
||||||
|
assert k1.startswith("ckpt-hist:v1:db0:t1:")
|
||||||
|
# ns / checkpoint / channel each change the key
|
||||||
|
assert k1 != make_history_key("ckpt-hist:v1:db0", "t1", "sub", "c1", "messages")
|
||||||
|
assert k1 != make_history_key("ckpt-hist:v1:db0", "t1", "", "c2", "messages")
|
||||||
|
assert k1 != make_history_key("ckpt-hist:v1:db0", "t1", "", "c1", "todos")
|
||||||
|
assert k1 != make_history_key("ckpt-hist:v1:db9", "t1", "", "c1", "messages")
|
||||||
|
assert CACHE_FORMAT_VERSION == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_many_miss_then_hit():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=4)
|
||||||
|
assert cache.get_many(["a"]) == {}
|
||||||
|
assert cache.stats().misses == 1
|
||||||
|
cache.set_many({"a": _entry("x")})
|
||||||
|
hit = cache.get_many(["a"])
|
||||||
|
assert hit["a"]["writes"] == [("task-1", "messages", "x")]
|
||||||
|
assert hit["a"]["seed"] == "seed-x"
|
||||||
|
assert cache.stats().hits == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_entry_without_seed_roundtrips_without_seed_key():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=4)
|
||||||
|
cache.set_many({"a": {"writes": []}})
|
||||||
|
hit = cache.get_many(["a"])
|
||||||
|
assert hit["a"] == {"writes": []}
|
||||||
|
assert "seed" not in hit["a"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_copy_on_read_returns_fresh_writes_list():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=4)
|
||||||
|
cache.set_many({"a": _entry("x")})
|
||||||
|
first = cache.get_many(["a"])["a"]
|
||||||
|
first["writes"].append(("task-2", "messages", "MUTATION"))
|
||||||
|
second = cache.get_many(["a"])["a"]
|
||||||
|
assert second["writes"] == [("task-1", "messages", "x")]
|
||||||
|
|
||||||
|
|
||||||
|
def test_caller_mutation_after_set_does_not_leak():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=4)
|
||||||
|
entry = _entry("x")
|
||||||
|
cache.set_many({"a": entry})
|
||||||
|
entry["writes"].append(("task-2", "messages", "MUTATION"))
|
||||||
|
assert cache.get_many(["a"])["a"]["writes"] == [("task-1", "messages", "x")]
|
||||||
|
|
||||||
|
|
||||||
|
def test_lru_evicts_oldest_and_counts():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=2)
|
||||||
|
cache.set_many({"a": _entry("a"), "b": _entry("b")})
|
||||||
|
cache.get_many(["a"]) # refresh a
|
||||||
|
cache.set_many({"c": _entry("c")}) # evicts b
|
||||||
|
assert cache.get_many(["b"]) == {}
|
||||||
|
assert cache.get_many(["a"]) != {}
|
||||||
|
assert cache.stats().evictions == 1
|
||||||
|
assert cache.stats().entries == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_zero_max_entries_disables():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=0)
|
||||||
|
assert cache.enabled is False
|
||||||
|
cache.set_many({"a": _entry("x")})
|
||||||
|
assert cache.get_many(["a"]) == {}
|
||||||
|
assert cache.stats().entries == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_delete_thread_purges_only_that_thread():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=16)
|
||||||
|
prefix = "ckpt-hist:v1:db0"
|
||||||
|
t1_keys = [make_history_key(prefix, "t1", "", f"c{i}", "messages") for i in range(3)]
|
||||||
|
t2_key = make_history_key(prefix, "t2", "", "c0", "messages")
|
||||||
|
# A thread_id that is a prefix of another must not over-match: the stem
|
||||||
|
# ends with ':' so "t1" never matches "t10"'s keys.
|
||||||
|
t10_key = make_history_key(prefix, "t10", "", "c0", "messages")
|
||||||
|
cache.set_many({k: _entry(k) for k in [*t1_keys, t2_key, t10_key]})
|
||||||
|
|
||||||
|
cache.delete_thread(prefix, "t1")
|
||||||
|
|
||||||
|
assert cache.stats().entries == 2
|
||||||
|
assert all(cache.get_many([k]) == {} for k in t1_keys)
|
||||||
|
assert cache.get_many([t2_key]) != {}
|
||||||
|
assert cache.get_many([t10_key]) != {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_adelete_thread_matches_sync():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=4)
|
||||||
|
prefix = "ckpt-hist:v1:db0"
|
||||||
|
key = make_history_key(prefix, "t1", "", "c0", "messages")
|
||||||
|
await cache.aset_many({key: _entry("x")})
|
||||||
|
await cache.adelete_thread(prefix, "t1")
|
||||||
|
assert cache.get_many([key]) == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_thread_key_stem_matches_make_history_key_layout():
|
||||||
|
key = make_history_key("p", "t1", "ns", "c1", "messages")
|
||||||
|
assert key.startswith(thread_key_stem("p", "t1"))
|
||||||
|
assert not key.startswith(thread_key_stem("p", "t"))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_async_protocol_matches_sync():
|
||||||
|
cache = MemoryCheckpointHistoryCache(max_entries=4)
|
||||||
|
await cache.aset_many({"a": _entry("x")})
|
||||||
|
hit = await cache.aget_many(["a"])
|
||||||
|
assert hit["a"]["seed"] == "seed-x"
|
||||||
|
stats = cache.stats()
|
||||||
|
assert isinstance(stats, CheckpointCacheStats)
|
||||||
|
assert stats.as_dict()["hits"] == 1
|
||||||
|
await cache.aclose()
|
||||||
|
assert cache.get_many(["a"]) == {}
|
||||||
96
backend/tests/test_checkpoint_cache_provider.py
Normal file
96
backend/tests/test_checkpoint_cache_provider.py
Normal file
@ -0,0 +1,96 @@
|
|||||||
|
"""Provider wiring: mode-gated wrapping in async and sync checkpointer factories."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from deerflow.config.app_config import AppConfig, set_app_config
|
||||||
|
from deerflow.runtime.checkpoint_mode import freeze_checkpoint_channel_mode
|
||||||
|
from deerflow.runtime.checkpointer.async_provider import make_checkpointer
|
||||||
|
from deerflow.runtime.checkpointer.cached_saver import CachedHistorySaver
|
||||||
|
from deerflow.runtime.checkpointer.provider import checkpointer_context, reset_checkpointer
|
||||||
|
|
||||||
|
|
||||||
|
# AppConfig requires the sandbox section (no default); the rest of the config
|
||||||
|
# is optional. Mirrors test_checkpoint_cache_redis.py's construction pattern.
|
||||||
|
def _app_config(mode: str, cache: dict | None = None) -> AppConfig:
|
||||||
|
database: dict = {"backend": "memory", "checkpoint_channel_mode": mode}
|
||||||
|
if cache is not None:
|
||||||
|
database["checkpoint_cache"] = cache
|
||||||
|
return AppConfig.model_validate(
|
||||||
|
{
|
||||||
|
"sandbox": {"use": "deerflow.sandbox.local.provider:LocalSandboxProvider"},
|
||||||
|
"database": database,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_delta_mode_wraps_with_cached_saver():
|
||||||
|
set_app_config(_app_config("delta"))
|
||||||
|
freeze_checkpoint_channel_mode("delta")
|
||||||
|
async with make_checkpointer() as saver:
|
||||||
|
assert isinstance(saver, CachedHistorySaver)
|
||||||
|
assert saver.stats()["entries"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_full_mode_yields_raw_saver():
|
||||||
|
set_app_config(_app_config("full"))
|
||||||
|
freeze_checkpoint_channel_mode("full")
|
||||||
|
async with make_checkpointer() as saver:
|
||||||
|
assert not isinstance(saver, CachedHistorySaver)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_zero_max_entries_disables_but_still_wraps():
|
||||||
|
set_app_config(_app_config("delta", {"max_entries": 0}))
|
||||||
|
freeze_checkpoint_channel_mode("delta")
|
||||||
|
async with make_checkpointer() as saver:
|
||||||
|
assert isinstance(saver, CachedHistorySaver)
|
||||||
|
# Disabled cache -> every history call is a full walk on the inner saver.
|
||||||
|
assert saver._cache.enabled is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_delta_mode_wraps_memory():
|
||||||
|
set_app_config(_app_config("delta"))
|
||||||
|
freeze_checkpoint_channel_mode("delta")
|
||||||
|
reset_checkpointer()
|
||||||
|
with checkpointer_context() as saver:
|
||||||
|
assert isinstance(saver, CachedHistorySaver)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_redis_cache_type_is_config_error():
|
||||||
|
set_app_config(_app_config("delta", {"type": "redis"}))
|
||||||
|
freeze_checkpoint_channel_mode("delta")
|
||||||
|
reset_checkpointer()
|
||||||
|
with pytest.raises(ValueError, match="redis"):
|
||||||
|
with checkpointer_context():
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_full_mode_unwrapped():
|
||||||
|
set_app_config(_app_config("full"))
|
||||||
|
freeze_checkpoint_channel_mode("full")
|
||||||
|
reset_checkpointer()
|
||||||
|
with checkpointer_context() as saver:
|
||||||
|
assert not isinstance(saver, CachedHistorySaver)
|
||||||
|
|
||||||
|
|
||||||
|
def test_sync_cache_recreated_when_key_prefix_changes():
|
||||||
|
"""The singleton must not outlive its namespace: a prefix change without
|
||||||
|
a process restart leaves old-prefix entries unreachable and unpurgeable."""
|
||||||
|
reset_checkpointer()
|
||||||
|
set_app_config(_app_config("delta", {"key_prefix": "ns-a"}))
|
||||||
|
freeze_checkpoint_channel_mode("delta")
|
||||||
|
with checkpointer_context() as saver:
|
||||||
|
saver._cache.set_many({"ns-a:t1:x": {"writes": []}})
|
||||||
|
first_cache = saver._cache
|
||||||
|
# Same prefix: singleton is reused (warm across wrappers).
|
||||||
|
with checkpointer_context() as saver:
|
||||||
|
assert saver._cache is first_cache
|
||||||
|
assert saver._cache.stats().entries == 1
|
||||||
|
# Prefix change: fresh cache, stale namespace gone with the old instance.
|
||||||
|
set_app_config(_app_config("delta", {"key_prefix": "ns-b"}))
|
||||||
|
with checkpointer_context() as saver:
|
||||||
|
assert saver._cache is not first_cache
|
||||||
|
assert saver._cache.stats().entries == 0
|
||||||
|
reset_checkpointer()
|
||||||
245
backend/tests/test_checkpoint_cache_redis.py
Normal file
245
backend/tests/test_checkpoint_cache_redis.py
Normal file
@ -0,0 +1,245 @@
|
|||||||
|
"""Redis backend and provider factory for the checkpoint history cache."""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
|
||||||
|
|
||||||
|
from deerflow.config.app_config import AppConfig
|
||||||
|
from deerflow.runtime.checkpoint_cache.provider import (
|
||||||
|
checkpoint_cache_db_hash,
|
||||||
|
checkpoint_cache_key_prefix,
|
||||||
|
make_checkpoint_cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeRedis:
|
||||||
|
"""Minimal async redis stand-in: mget / set / pipeline / scan / unlink."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.store: dict[str, bytes] = {}
|
||||||
|
self.ttls: dict[str, int | None] = {}
|
||||||
|
self.unlinked: list[tuple[str, ...]] = []
|
||||||
|
|
||||||
|
async def mget(self, keys: list[str]) -> list[bytes | None]:
|
||||||
|
return [self.store.get(k) for k in keys]
|
||||||
|
|
||||||
|
def set(self, key: str, value: bytes, ex: int | None = None) -> None:
|
||||||
|
self.store[key] = value
|
||||||
|
self.ttls[key] = ex
|
||||||
|
|
||||||
|
async def scan(self, cursor: int = 0, match: str | None = None, count: int = 500) -> tuple[int, list[str]]:
|
||||||
|
import fnmatch
|
||||||
|
|
||||||
|
keys = sorted(self.store)
|
||||||
|
batch = keys[cursor : cursor + count]
|
||||||
|
if match is not None:
|
||||||
|
batch = [k for k in batch if fnmatch.fnmatchcase(k, match)]
|
||||||
|
next_cursor = cursor + count
|
||||||
|
return (0 if next_cursor >= len(keys) else next_cursor), batch
|
||||||
|
|
||||||
|
async def unlink(self, *keys: str) -> int:
|
||||||
|
self.unlinked.append(tuple(keys))
|
||||||
|
removed = 0
|
||||||
|
for key in keys:
|
||||||
|
removed += self.store.pop(key, None) is not None
|
||||||
|
return removed
|
||||||
|
|
||||||
|
def pipeline(self, transaction: bool = False) -> "_FakePipeline":
|
||||||
|
return _FakePipeline(self)
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _FakePipeline:
|
||||||
|
def __init__(self, client: _FakeRedis) -> None:
|
||||||
|
self._client = client
|
||||||
|
|
||||||
|
def set(self, key: str, value: bytes, ex: int | None = None) -> "_FakePipeline":
|
||||||
|
self._client.set(key, value, ex=ex)
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def execute(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class _FailingRedis(_FakeRedis):
|
||||||
|
"""Simulates a redis outage: every operation raises RedisError."""
|
||||||
|
|
||||||
|
async def mget(self, keys: list[str]) -> list[bytes | None]:
|
||||||
|
from redis.exceptions import RedisError
|
||||||
|
|
||||||
|
raise RedisError("connection refused")
|
||||||
|
|
||||||
|
async def scan(self, cursor: int = 0, match: str | None = None, count: int = 500) -> tuple[int, list[str]]:
|
||||||
|
from redis.exceptions import RedisError
|
||||||
|
|
||||||
|
raise RedisError("connection refused")
|
||||||
|
|
||||||
|
def pipeline(self, transaction: bool = False) -> "_FakePipeline":
|
||||||
|
from redis.exceptions import RedisError
|
||||||
|
|
||||||
|
raise RedisError("connection refused")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_cache(monkeypatch: pytest.MonkeyPatch, fake: _FakeRedis, ttl_seconds: int = 60, **kwargs: Any):
|
||||||
|
import deerflow.runtime.checkpoint_cache.redis as redis_mod
|
||||||
|
|
||||||
|
monkeypatch.setattr(redis_mod, "_create_client", lambda *a, **k: fake)
|
||||||
|
return redis_mod.RedisCheckpointHistoryCache("redis://unused", serde=JsonPlusSerializer(), ttl_seconds=ttl_seconds, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def _entry(i: int) -> dict:
|
||||||
|
# Real message-like payloads to prove serde fidelity beyond plain dicts.
|
||||||
|
from langchain_core.messages import AIMessage
|
||||||
|
|
||||||
|
return {"writes": [("task-1", "messages", AIMessage(content=f"m{i}", id=f"ai-{i}"))], "seed": [AIMessage(content="s", id="ai-s")]}
|
||||||
|
|
||||||
|
|
||||||
|
# AppConfig requires the sandbox section (no default); the rest of the config
|
||||||
|
# is optional. Mirrors test_checkpoint_mode.py's construction pattern.
|
||||||
|
def _app_config(database: dict) -> AppConfig:
|
||||||
|
return AppConfig.model_validate(
|
||||||
|
{
|
||||||
|
"sandbox": {"use": "deerflow.sandbox.local.provider:LocalSandboxProvider"},
|
||||||
|
"database": database,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_redis_roundtrip_preserves_types(monkeypatch: pytest.MonkeyPatch):
|
||||||
|
fake = _FakeRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake)
|
||||||
|
await cache.aset_many({"k1": _entry(1), "k2": {"writes": []}})
|
||||||
|
hit = await cache.aget_many(["k1", "k2", "k3"])
|
||||||
|
assert set(hit) == {"k1", "k2"}
|
||||||
|
msg = hit["k1"]["writes"][0][2]
|
||||||
|
assert msg.content == "m1" and msg.id == "ai-1" and msg.type == "ai"
|
||||||
|
assert "seed" not in hit["k2"]
|
||||||
|
assert cache.stats().hits == 2 and cache.stats().misses == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_redis_keys_land_verbatim_and_ttl_set(monkeypatch: pytest.MonkeyPatch):
|
||||||
|
fake = _FakeRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake)
|
||||||
|
await cache.aset_many({"k1": _entry(1)})
|
||||||
|
assert list(fake.store) == ["k1"]
|
||||||
|
assert fake.ttls["k1"] == 60
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_redis_outage_degrades_to_all_miss(monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture):
|
||||||
|
fake = _FailingRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake)
|
||||||
|
with caplog.at_level("WARNING"):
|
||||||
|
assert await cache.aget_many(["k1", "k2"]) == {}
|
||||||
|
assert cache.stats().misses == 2 and cache.stats().hits == 0
|
||||||
|
assert "mget failed" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_redis_outage_skips_write(monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture):
|
||||||
|
fake = _FailingRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake)
|
||||||
|
with caplog.at_level("WARNING"):
|
||||||
|
await cache.aset_many({"k1": _entry(1)}) # must not raise
|
||||||
|
assert fake.store == {}
|
||||||
|
assert "write failed" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_zero_ttl_disables_expiry_explicitly(monkeypatch: pytest.MonkeyPatch):
|
||||||
|
fake = _FakeRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake, ttl_seconds=0)
|
||||||
|
await cache.aset_many({"k1": _entry(1)})
|
||||||
|
assert cache._ttl is None
|
||||||
|
assert fake.ttls["k1"] is None # SET without EX: redis maxmemory policy only
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_adelete_thread_purges_matching_keys_only(monkeypatch: pytest.MonkeyPatch):
|
||||||
|
fake = _FakeRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake)
|
||||||
|
prefix = "ckpt-hist:v1:db0"
|
||||||
|
await cache.aset_many(
|
||||||
|
{
|
||||||
|
f"{prefix}:t1:aaa": _entry(1),
|
||||||
|
f"{prefix}:t1:bbb": _entry(2),
|
||||||
|
f"{prefix}:t10:ccc": _entry(3), # 't1' stem must not over-match 't10'
|
||||||
|
f"{prefix}:t2:ddd": _entry(4),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
await cache.adelete_thread(prefix, "t1")
|
||||||
|
|
||||||
|
assert sorted(fake.store) == [f"{prefix}:t10:ccc", f"{prefix}:t2:ddd"]
|
||||||
|
assert fake.unlinked # UNLINK, not DEL: non-blocking on big histories
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_adelete_thread_outage_degrades_without_raising(monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture):
|
||||||
|
fake = _FailingRedis()
|
||||||
|
cache = _make_cache(monkeypatch, fake)
|
||||||
|
with caplog.at_level("WARNING"):
|
||||||
|
await cache.adelete_thread("p", "t1") # must not raise
|
||||||
|
assert "thread purge failed" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_provider_memory_default():
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
|
||||||
|
async with make_checkpoint_cache(_app_config({"backend": "sqlite"}), serde=JsonPlusSerializer()) as cache:
|
||||||
|
assert isinstance(cache, MemoryCheckpointHistoryCache)
|
||||||
|
assert cache.enabled is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_provider_zero_max_entries_disables_any_type():
|
||||||
|
app_config = _app_config({"backend": "sqlite", "checkpoint_cache": {"type": "redis", "max_entries": 0}})
|
||||||
|
from deerflow.runtime.checkpoint_cache.memory import MemoryCheckpointHistoryCache
|
||||||
|
|
||||||
|
async with make_checkpoint_cache(app_config, serde=JsonPlusSerializer()) as cache:
|
||||||
|
assert isinstance(cache, MemoryCheckpointHistoryCache)
|
||||||
|
assert cache.enabled is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_db_hash_distinguishes_backends_and_targets():
|
||||||
|
from deerflow.config.database_config import DatabaseConfig
|
||||||
|
|
||||||
|
sqlite_cfg = DatabaseConfig.model_validate({"backend": "sqlite", "sqlite_dir": "/tmp/a"})
|
||||||
|
pg_cfg = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "postgresql://u:p@h/db"})
|
||||||
|
pg_cfg2 = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "postgresql://u:p@h/other"})
|
||||||
|
assert checkpoint_cache_db_hash(sqlite_cfg) != checkpoint_cache_db_hash(pg_cfg)
|
||||||
|
assert checkpoint_cache_db_hash(pg_cfg) != checkpoint_cache_db_hash(pg_cfg2)
|
||||||
|
assert len(checkpoint_cache_db_hash(pg_cfg)) == 12
|
||||||
|
|
||||||
|
|
||||||
|
def test_db_hash_stable_across_credential_rotation():
|
||||||
|
"""Same database, rotated user/password -> same cache namespace."""
|
||||||
|
from deerflow.config.database_config import DatabaseConfig
|
||||||
|
|
||||||
|
before = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "postgresql://alice:secret1@pg.internal:5432/deerflow"})
|
||||||
|
rotated = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "postgresql://bob:secret2@pg.internal:5432/deerflow"})
|
||||||
|
driver_suffix = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "postgresql+asyncpg://alice:secret1@pg.internal:5432/deerflow"})
|
||||||
|
other_db = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "postgresql://alice:secret1@pg.internal:5432/other"})
|
||||||
|
assert checkpoint_cache_db_hash(before) == checkpoint_cache_db_hash(rotated)
|
||||||
|
assert checkpoint_cache_db_hash(before) == checkpoint_cache_db_hash(driver_suffix)
|
||||||
|
assert checkpoint_cache_db_hash(before) != checkpoint_cache_db_hash(other_db)
|
||||||
|
|
||||||
|
|
||||||
|
def test_db_hash_unparseable_url_falls_back_to_raw():
|
||||||
|
from deerflow.config.database_config import DatabaseConfig
|
||||||
|
|
||||||
|
cfg = DatabaseConfig.model_validate({"backend": "postgres", "postgres_url": "not-a-url"})
|
||||||
|
assert len(checkpoint_cache_db_hash(cfg)) == 12 # stable, never raises
|
||||||
|
|
||||||
|
|
||||||
|
def test_key_prefix_override_wins():
|
||||||
|
app_config = _app_config({"backend": "sqlite", "checkpoint_cache": {"key_prefix": "custom:"}})
|
||||||
|
assert checkpoint_cache_key_prefix(app_config) == "custom:"
|
||||||
|
default = checkpoint_cache_key_prefix(_app_config({"backend": "sqlite"}))
|
||||||
|
assert default.startswith("ckpt-hist:v1:")
|
||||||
@ -15,7 +15,7 @@
|
|||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Bump this number when the config schema changes.
|
# Bump this number when the config schema changes.
|
||||||
# Run `make config-upgrade` to merge new fields into your local config.yaml.
|
# Run `make config-upgrade` to merge new fields into your local config.yaml.
|
||||||
config_version: 31
|
config_version: 32
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Logging
|
# Logging
|
||||||
@ -1978,6 +1978,16 @@ database:
|
|||||||
checkpoint_graph_cache:
|
checkpoint_graph_cache:
|
||||||
# Gateway thread-state accessor graphs (per assistant/mode/cadence).
|
# Gateway thread-state accessor graphs (per assistant/mode/cadence).
|
||||||
accessor_graph_max: 64
|
accessor_graph_max: 64
|
||||||
|
# Delta-mode checkpoint history cache (only used when checkpoint_channel_mode: delta).
|
||||||
|
# Pure performance policy: safe to differ across workers, never frozen.
|
||||||
|
# checkpoint_cache:
|
||||||
|
# type: memory # memory | redis (redis is Gateway/async only)
|
||||||
|
# max_entries: 128 # 0 disables the cache
|
||||||
|
# redis_url: null # or DEER_FLOW_CHECKPOINT_CACHE_REDIS_URL / REDIS_URL
|
||||||
|
# ttl_seconds: 86400 # redis leak safety net; bounds residual copies if a
|
||||||
|
# # thread-delete purge fails (purge itself is immediate).
|
||||||
|
# # 0 explicitly disables expiry (redis maxmemory only)
|
||||||
|
# key_prefix: "" # default: hash of the database identity
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Run Events Configuration
|
# Run Events Configuration
|
||||||
|
|||||||
@ -124,7 +124,7 @@ they resolve from the `secrets` map):
|
|||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
config: |
|
config: |
|
||||||
config_version: 31
|
config_version: 32
|
||||||
models:
|
models:
|
||||||
- name: gpt-4
|
- name: gpt-4
|
||||||
use: langchain_openai:ChatOpenAI
|
use: langchain_openai:ChatOpenAI
|
||||||
|
|||||||
@ -240,7 +240,7 @@ ingress:
|
|||||||
# -- DeerFlow config.yaml content. Secrets MUST stay as $VAR references — never
|
# -- DeerFlow config.yaml content. Secrets MUST stay as $VAR references — never
|
||||||
# inline literal secret values here. The default enables provisioner sandbox.
|
# inline literal secret values here. The default enables provisioner sandbox.
|
||||||
config: |
|
config: |
|
||||||
config_version: 31
|
config_version: 32
|
||||||
log_level: info
|
log_level: info
|
||||||
|
|
||||||
models: []
|
models: []
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user