langgraph-runtime-inmem 0.36.0rc2__tar.gz → 0.37.0.dev2__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/PKG-INFO +1 -1
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/__init__.py +3 -1
- langgraph_runtime_inmem-0.37.0.dev2/langgraph_runtime_inmem/_persistence.py +88 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/checkpoint.py +13 -4
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/database.py +97 -19
- langgraph_runtime_inmem-0.37.0.dev2/langgraph_runtime_inmem/encryption.py +232 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/ops.py +326 -163
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/queue.py +13 -3
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/store.py +8 -2
- langgraph_runtime_inmem-0.36.0rc2/langgraph_runtime_inmem/_persistence.py +0 -64
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/.gitignore +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/Makefile +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/README.md +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/inmem_stream.py +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/lifespan.py +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/metrics.py +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/retry.py +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/routes.py +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/pyproject.toml +0 -0
- {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/uv.lock +0 -0
|
@@ -1,6 +1,7 @@
|
|
|
1
1
|
from langgraph_runtime_inmem import (
|
|
2
2
|
checkpoint,
|
|
3
3
|
database,
|
|
4
|
+
encryption,
|
|
4
5
|
lifespan,
|
|
5
6
|
metrics,
|
|
6
7
|
ops,
|
|
@@ -10,11 +11,12 @@ from langgraph_runtime_inmem import (
|
|
|
10
11
|
store,
|
|
11
12
|
)
|
|
12
13
|
|
|
13
|
-
__version__ = "0.
|
|
14
|
+
__version__ = "0.37.0.dev2"
|
|
14
15
|
__all__ = [
|
|
15
16
|
"ops",
|
|
16
17
|
"database",
|
|
17
18
|
"checkpoint",
|
|
19
|
+
"encryption",
|
|
18
20
|
"lifespan",
|
|
19
21
|
"retry",
|
|
20
22
|
"store",
|
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
"""Periodic flushing for all PersistentDict stores."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import functools
|
|
6
|
+
import logging
|
|
7
|
+
import os
|
|
8
|
+
import threading
|
|
9
|
+
import weakref
|
|
10
|
+
|
|
11
|
+
from langgraph.checkpoint.memory import PersistentDict
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger(__name__)
|
|
14
|
+
|
|
15
|
+
_stores: dict[str, weakref.ref[PersistentDict]] = {}
|
|
16
|
+
# Held while changing _stores and while syncing one entry, so once
|
|
17
|
+
# unregister_persistent_dict() returns no sync of that dict is running or will start.
|
|
18
|
+
_stores_lock = threading.Lock()
|
|
19
|
+
_flush_thread: tuple[threading.Event, threading.Thread] | None = None
|
|
20
|
+
_flush_interval: int = 10
|
|
21
|
+
DISABLE_FILE_PERSISTENCE = (
|
|
22
|
+
os.getenv("LANGGRAPH_DISABLE_FILE_PERSISTENCE", "false").lower() == "true"
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def register_persistent_dict(d: PersistentDict) -> None:
|
|
27
|
+
"""Register a PersistentDict for periodic flushing."""
|
|
28
|
+
if DISABLE_FILE_PERSISTENCE:
|
|
29
|
+
return
|
|
30
|
+
global _flush_thread
|
|
31
|
+
with _stores_lock:
|
|
32
|
+
_stores[d.filename] = weakref.ref(d)
|
|
33
|
+
if _flush_thread is None:
|
|
34
|
+
logger.info("Starting dev persistence flush loop")
|
|
35
|
+
stop_event = threading.Event()
|
|
36
|
+
_flush_thread = (
|
|
37
|
+
stop_event,
|
|
38
|
+
threading.Thread(
|
|
39
|
+
target=functools.partial(_flush_loop, stop_event), daemon=True
|
|
40
|
+
),
|
|
41
|
+
)
|
|
42
|
+
_flush_thread[1].start()
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def unregister_persistent_dict(d: PersistentDict) -> None:
|
|
46
|
+
"""Stop periodic flushing of ``d`` if it's still the dict registered for its file."""
|
|
47
|
+
with _stores_lock:
|
|
48
|
+
ref = _stores.get(d.filename)
|
|
49
|
+
if ref is not None and ref() is d:
|
|
50
|
+
_stores.pop(d.filename, None)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def stop_flush_loop() -> None:
|
|
54
|
+
"""Stop the background flush thread."""
|
|
55
|
+
global _flush_thread
|
|
56
|
+
if _flush_thread is not None:
|
|
57
|
+
logger.info("Stopping dev persistence flush loop")
|
|
58
|
+
_flush_thread[0].set()
|
|
59
|
+
_flush_thread[1].join()
|
|
60
|
+
_flush_thread = None
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _flush_once() -> None:
|
|
64
|
+
with _stores_lock:
|
|
65
|
+
entries = list(_stores.items())
|
|
66
|
+
for filename, ref in entries:
|
|
67
|
+
with _stores_lock:
|
|
68
|
+
# Skip entries unregistered or replaced since the snapshot.
|
|
69
|
+
if _stores.get(filename) is not ref:
|
|
70
|
+
continue
|
|
71
|
+
# Compare to None: an empty dict is falsy but still needs flushing.
|
|
72
|
+
if (store := ref()) is None:
|
|
73
|
+
_stores.pop(filename, None)
|
|
74
|
+
logger.debug(
|
|
75
|
+
"Dropped garbage-collected dev persistence file %s", filename
|
|
76
|
+
)
|
|
77
|
+
continue
|
|
78
|
+
try:
|
|
79
|
+
store.sync()
|
|
80
|
+
except Exception:
|
|
81
|
+
# One bad file must not stop the others (or later retries) from flushing.
|
|
82
|
+
logger.exception("Failed to flush dev persistence file %s", filename)
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def _flush_loop(stop_event: threading.Event) -> None:
|
|
86
|
+
while not stop_event.wait(timeout=_flush_interval):
|
|
87
|
+
_flush_once()
|
|
88
|
+
logger.info("dev persistence flush loop exiting")
|
|
@@ -54,6 +54,7 @@ class InMemorySaver(InMemorySaverBase):
|
|
|
54
54
|
self,
|
|
55
55
|
*,
|
|
56
56
|
serde: SerializerProtocol | None = None,
|
|
57
|
+
persist: bool = True,
|
|
57
58
|
__persistence_hook__: Callable[[PersistentDict], None] | None = None,
|
|
58
59
|
) -> None:
|
|
59
60
|
self.filename = os.path.join(".langgraph_api", ".langgraph_checkpoint.")
|
|
@@ -68,8 +69,6 @@ class InMemorySaver(InMemorySaverBase):
|
|
|
68
69
|
os.mkdir(".langgraph_api")
|
|
69
70
|
thisfname = self.filename + str(i) + ".pckl"
|
|
70
71
|
d = PersistentDict(*args, filename=thisfname)
|
|
71
|
-
if __persistence_hook__:
|
|
72
|
-
__persistence_hook__(d)
|
|
73
72
|
|
|
74
73
|
try:
|
|
75
74
|
d.load()
|
|
@@ -94,13 +93,21 @@ class InMemorySaver(InMemorySaverBase):
|
|
|
94
93
|
os.remove(self.filename)
|
|
95
94
|
except Exception:
|
|
96
95
|
pass
|
|
96
|
+
# Register only once loaded, so a flush can't write the still-empty
|
|
97
|
+
# dict over the file.
|
|
98
|
+
if __persistence_hook__:
|
|
99
|
+
__persistence_hook__(d)
|
|
97
100
|
return d
|
|
98
101
|
|
|
99
102
|
from langgraph_api.serde import Serializer # noqa: PLC0415
|
|
100
103
|
|
|
104
|
+
# persist=False: the saver borrows another saver's dicts instead of owning
|
|
105
|
+
# files. Owning them would load the pickles and take over their flush
|
|
106
|
+
# registrations (keyed by filename).
|
|
107
|
+
use_files = persist and not DISABLE_FILE_PERSISTENCE
|
|
101
108
|
super().__init__(
|
|
102
109
|
serde=serde if serde is not None else Serializer(),
|
|
103
|
-
factory=factory if
|
|
110
|
+
factory=factory if use_files else defaultdict,
|
|
104
111
|
)
|
|
105
112
|
|
|
106
113
|
def put(
|
|
@@ -286,9 +293,11 @@ def Checkpointer(*args, unpack_hook=None, **kwargs):
|
|
|
286
293
|
else:
|
|
287
294
|
ext_hook = unpack_hook
|
|
288
295
|
|
|
296
|
+
# A separate saver for its JSON-friendly serde and per-request latest_iter;
|
|
297
|
+
# the data itself is MEMORY's.
|
|
289
298
|
saver = InMemorySaver(
|
|
290
299
|
serde=Serializer(__unpack_ext_hook__=ext_hook),
|
|
291
|
-
|
|
300
|
+
persist=False,
|
|
292
301
|
**kwargs,
|
|
293
302
|
)
|
|
294
303
|
saver.writes = MEMORY.writes
|
|
@@ -1,10 +1,13 @@
|
|
|
1
1
|
import asyncio
|
|
2
|
+
import contextlib
|
|
3
|
+
import glob
|
|
2
4
|
import os
|
|
5
|
+
import pickle
|
|
3
6
|
import uuid
|
|
4
7
|
from collections import defaultdict
|
|
5
8
|
from collections.abc import AsyncIterator, Awaitable, Callable
|
|
6
9
|
from contextlib import asynccontextmanager
|
|
7
|
-
from datetime import datetime
|
|
10
|
+
from datetime import UTC, datetime
|
|
8
11
|
from typing import TYPE_CHECKING, Any, NotRequired
|
|
9
12
|
from uuid import UUID
|
|
10
13
|
|
|
@@ -154,6 +157,96 @@ async def connect(
|
|
|
154
157
|
yield InMemConnectionProto()
|
|
155
158
|
|
|
156
159
|
|
|
160
|
+
def _load_ops_file() -> None:
|
|
161
|
+
"""Load the saved ops file into ``GLOBAL_STORE``, setting it aside if unloadable."""
|
|
162
|
+
if not os.path.exists(OPS_FILENAME):
|
|
163
|
+
return
|
|
164
|
+
|
|
165
|
+
# Not GLOBAL_STORE.load(): it treats a truncated file as empty.
|
|
166
|
+
try:
|
|
167
|
+
with open(OPS_FILENAME, "rb") as f:
|
|
168
|
+
data = pickle.load(f)
|
|
169
|
+
if not isinstance(data, dict):
|
|
170
|
+
raise TypeError(f"expected a dict, found {type(data).__name__}")
|
|
171
|
+
except Exception as e:
|
|
172
|
+
if os.path.getsize(OPS_FILENAME) == 0:
|
|
173
|
+
return
|
|
174
|
+
_set_aside_ops_file(e)
|
|
175
|
+
return
|
|
176
|
+
GLOBAL_STORE.update(data)
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
def _set_aside_ops_file(reason: Exception) -> None:
|
|
180
|
+
"""Rename an unloadable ops file to ``<file>.bak-<UTC timestamp>`` and log why."""
|
|
181
|
+
from langgraph_api.graph import graph_file_module_names # noqa: PLC0415
|
|
182
|
+
|
|
183
|
+
stamp = datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ")
|
|
184
|
+
backup = f"{OPS_FILENAME}.bak-{stamp}"
|
|
185
|
+
n = 1
|
|
186
|
+
while os.path.exists(backup):
|
|
187
|
+
backup = f"{OPS_FILENAME}.bak-{stamp}-{n}"
|
|
188
|
+
n += 1
|
|
189
|
+
# Don't catch: carrying on would let the next flush overwrite the file.
|
|
190
|
+
os.replace(OPS_FILENAME, backup)
|
|
191
|
+
with contextlib.suppress(FileNotFoundError):
|
|
192
|
+
os.remove(RETRY_COUNTER_FILENAME)
|
|
193
|
+
try:
|
|
194
|
+
with open(f"{backup}.reason", "w") as f:
|
|
195
|
+
f.write(f"{type(reason).__name__}: {reason}")
|
|
196
|
+
except OSError as e:
|
|
197
|
+
logger.warning("Could not save why %s was set aside: %s", backup, e)
|
|
198
|
+
|
|
199
|
+
message = (
|
|
200
|
+
f"Could not load saved threads, runs and crons from {OPS_FILENAME}\n"
|
|
201
|
+
f" Reason: {type(reason).__name__}: {reason}\n"
|
|
202
|
+
" The server is starting without them. The file was NOT deleted; it was moved to:\n"
|
|
203
|
+
f" {backup}\n"
|
|
204
|
+
)
|
|
205
|
+
if (
|
|
206
|
+
isinstance(reason, ModuleNotFoundError)
|
|
207
|
+
and reason.name in graph_file_module_names()
|
|
208
|
+
):
|
|
209
|
+
message += (
|
|
210
|
+
f" It references a class from your graph file (module '{reason.name}'), saved by an older\n"
|
|
211
|
+
" version of langgraph-api. It can't be restored; newer versions no longer save\n"
|
|
212
|
+
" data this way, so this won't happen again."
|
|
213
|
+
)
|
|
214
|
+
else:
|
|
215
|
+
message += (
|
|
216
|
+
" To restore it: fix the cause, stop the server, and rename that file back to\n"
|
|
217
|
+
f" {os.path.basename(OPS_FILENAME)}. Threads created since this start will be lost.\n"
|
|
218
|
+
" This usually means a class stored in thread state was renamed, moved, or\n"
|
|
219
|
+
" failed to import (e.g. a syntax error in its module)."
|
|
220
|
+
)
|
|
221
|
+
logger.error(message, error_type=type(reason).__name__, backup=backup)
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
def _warn_about_ops_backups() -> None:
|
|
225
|
+
"""Log each ops-file backup left by an earlier start, with why it was set aside."""
|
|
226
|
+
paths = glob.glob(f"{glob.escape(OPS_FILENAME)}.bak-*")
|
|
227
|
+
backups = sorted(p for p in paths if not p.endswith(".reason"))
|
|
228
|
+
for reason in paths:
|
|
229
|
+
if reason.endswith(".reason") and reason.removesuffix(".reason") not in backups:
|
|
230
|
+
os.remove(reason)
|
|
231
|
+
if not backups:
|
|
232
|
+
return
|
|
233
|
+
|
|
234
|
+
lines = []
|
|
235
|
+
for backup in backups:
|
|
236
|
+
try:
|
|
237
|
+
with open(f"{backup}.reason") as f:
|
|
238
|
+
cause = f.read()
|
|
239
|
+
except OSError:
|
|
240
|
+
cause = "unknown cause"
|
|
241
|
+
lines.append(f" {backup}: {cause}")
|
|
242
|
+
logger.warning(
|
|
243
|
+
f"{len(backups)} unloadable ops file backup(s). Rename one to "
|
|
244
|
+
f"{os.path.basename(OPS_FILENAME)} once its cause is fixed, or delete it:\n"
|
|
245
|
+
+ "\n".join(lines),
|
|
246
|
+
backups=backups,
|
|
247
|
+
)
|
|
248
|
+
|
|
249
|
+
|
|
157
250
|
async def start_pool() -> None:
|
|
158
251
|
if store._STORE_CONFIG is None:
|
|
159
252
|
from langgraph_api import config as langgraph_config # noqa: PLC0415
|
|
@@ -164,24 +257,8 @@ async def start_pool() -> None:
|
|
|
164
257
|
|
|
165
258
|
if not os.path.exists(".langgraph_api"):
|
|
166
259
|
os.mkdir(".langgraph_api")
|
|
167
|
-
|
|
168
|
-
|
|
169
|
-
GLOBAL_STORE.load()
|
|
170
|
-
except ModuleNotFoundError:
|
|
171
|
-
logger.error(
|
|
172
|
-
"Unable to load cached data - your code has changed in a way that's incompatible with the cache."
|
|
173
|
-
"\nThis usually happens when you've:"
|
|
174
|
-
"\n - Renamed or moved classes"
|
|
175
|
-
"\n - Changed class structures"
|
|
176
|
-
"\n - Pulled updates that modified class definitions in a way that's incompatible with the cache"
|
|
177
|
-
"\n\nRemoving invalid cache data stored at path: .langgraph_api"
|
|
178
|
-
)
|
|
179
|
-
await asyncio.to_thread(os.remove, OPS_FILENAME)
|
|
180
|
-
await asyncio.to_thread(os.remove, RETRY_COUNTER_FILENAME)
|
|
181
|
-
except Exception as e:
|
|
182
|
-
logger.error("Failed to load cached data: %s", str(e))
|
|
183
|
-
await asyncio.to_thread(os.remove, OPS_FILENAME)
|
|
184
|
-
await asyncio.to_thread(os.remove, RETRY_COUNTER_FILENAME)
|
|
260
|
+
_warn_about_ops_backups()
|
|
261
|
+
_load_ops_file()
|
|
185
262
|
for k in ["runs", "threads", "assistant_versions", "assistants"]:
|
|
186
263
|
if not GLOBAL_STORE.get(k):
|
|
187
264
|
GLOBAL_STORE[k] = []
|
|
@@ -201,6 +278,7 @@ async def start_pool() -> None:
|
|
|
201
278
|
|
|
202
279
|
|
|
203
280
|
async def stop_pool() -> None:
|
|
281
|
+
# Stop flushing before closing: close() clears the dicts, and a later flush would write {}.
|
|
204
282
|
stop_flush_loop()
|
|
205
283
|
await asyncio.to_thread(GLOBAL_STORE.close)
|
|
206
284
|
await asyncio.to_thread(GLOBAL_RETRY_COUNTER.close)
|
|
@@ -0,0 +1,232 @@
|
|
|
1
|
+
"""Structural JSON encryption owned by the Python in-memory runtime.
|
|
2
|
+
|
|
3
|
+
Unlike Postgres persistence, this remains here while local persistence is
|
|
4
|
+
Python-owned.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import asyncio
|
|
10
|
+
from typing import TYPE_CHECKING, Any
|
|
11
|
+
|
|
12
|
+
if TYPE_CHECKING:
|
|
13
|
+
from collections.abc import Mapping
|
|
14
|
+
|
|
15
|
+
from langgraph_api.encryption.aes_json import AesEncryptionInstance
|
|
16
|
+
from langgraph_api.encryption.custom import JsonEncryptionWrapper, ModelType
|
|
17
|
+
|
|
18
|
+
NESTED_ENCRYPTED_SUBFIELDS: dict[tuple[str, str], tuple[str, ...]] = {
|
|
19
|
+
("run", "kwargs"): ("input", "config", "context", "command"),
|
|
20
|
+
("run", "config"): ("configurable", "metadata"),
|
|
21
|
+
("cron", "payload"): ("metadata", "context", "input", "config"),
|
|
22
|
+
("cron", "config"): ("configurable", "metadata"),
|
|
23
|
+
("assistant", "config"): ("configurable",),
|
|
24
|
+
("thread", "config"): ("configurable",),
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
ENCRYPTION_FIELDS: dict[str, tuple[str, ...]] = {
|
|
28
|
+
"thread": ("metadata", "config", "values", "interrupts", "error"),
|
|
29
|
+
"run": ("metadata", "kwargs"),
|
|
30
|
+
"assistant": ("metadata", "config", "context"),
|
|
31
|
+
"cron": ("metadata", "payload"),
|
|
32
|
+
"store": ("value",),
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
NEVER_ENCRYPT_FIELDS = frozenset(
|
|
36
|
+
{
|
|
37
|
+
"thread_id",
|
|
38
|
+
"run_id",
|
|
39
|
+
"assistant_id",
|
|
40
|
+
"graph_id",
|
|
41
|
+
"checkpoint_id",
|
|
42
|
+
"task_id",
|
|
43
|
+
"__pregel_checkpointer",
|
|
44
|
+
"__pregel_resuming",
|
|
45
|
+
"__pregel_durability",
|
|
46
|
+
"__pregel_stream",
|
|
47
|
+
"__pregel_task_id",
|
|
48
|
+
"__pregel_checkpoint_ns",
|
|
49
|
+
"__after_seconds__",
|
|
50
|
+
"__request_start_time_ms__",
|
|
51
|
+
"__encryption_context__",
|
|
52
|
+
"__blob_encryption_context__",
|
|
53
|
+
"langgraph_version",
|
|
54
|
+
"langgraph_api_version",
|
|
55
|
+
"langgraph_plan",
|
|
56
|
+
"langgraph_host",
|
|
57
|
+
"langgraph_api_url",
|
|
58
|
+
"langgraph_request_id",
|
|
59
|
+
"langgraph_auth_user_id",
|
|
60
|
+
"langgraph_auth_permissions",
|
|
61
|
+
}
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
NEVER_ENCRYPT_PATHS = frozenset(
|
|
65
|
+
{
|
|
66
|
+
"run.kwargs.config.recursion_limit",
|
|
67
|
+
"run.kwargs.config.max_concurrency",
|
|
68
|
+
"run.kwargs.temporary",
|
|
69
|
+
"run.kwargs.config.configurable.ttl",
|
|
70
|
+
"checkpoint.metadata.source",
|
|
71
|
+
"checkpoint.metadata.step",
|
|
72
|
+
"checkpoint.metadata.parents",
|
|
73
|
+
"checkpoint.metadata.run_attempt",
|
|
74
|
+
"checkpoint.metadata.counters_since_delta_snapshot",
|
|
75
|
+
}
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def should_skip_encryption(key: str, path: str) -> bool:
|
|
80
|
+
return key in NEVER_ENCRYPT_FIELDS or f"{path}.{key}" in NEVER_ENCRYPT_PATHS
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def extract_blob_encryption_context(
|
|
84
|
+
data: dict[str, Any] | None,
|
|
85
|
+
) -> dict[str, Any] | None:
|
|
86
|
+
"""Return checkpoint blob context carried by an encrypted resource."""
|
|
87
|
+
if data is None:
|
|
88
|
+
return None
|
|
89
|
+
|
|
90
|
+
from langgraph_api.encryption.shared import ( # noqa: PLC0415
|
|
91
|
+
BLOB_ENCRYPTION_CONTEXT_KEY,
|
|
92
|
+
ENCRYPTION_CONTEXT_KEY,
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
return data.get(BLOB_ENCRYPTION_CONTEXT_KEY) or data.get(ENCRYPTION_CONTEXT_KEY)
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def encryption_fields(model_type: ModelType) -> list[str]:
|
|
99
|
+
"""Return fields structurally encrypted by the local runtime."""
|
|
100
|
+
return list(ENCRYPTION_FIELDS[model_type])
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
async def _decrypt_field(
|
|
104
|
+
obj: dict[str, Any],
|
|
105
|
+
field_name: str,
|
|
106
|
+
encryption_instance: JsonEncryptionWrapper | AesEncryptionInstance | None,
|
|
107
|
+
model_type: ModelType,
|
|
108
|
+
) -> tuple[str, Any]:
|
|
109
|
+
from langgraph_api.encryption.middleware import ( # noqa: PLC0415
|
|
110
|
+
decrypt_json_if_needed,
|
|
111
|
+
)
|
|
112
|
+
|
|
113
|
+
if not obj.get(field_name):
|
|
114
|
+
return field_name, obj.get(field_name)
|
|
115
|
+
|
|
116
|
+
decrypted = await decrypt_json_if_needed(
|
|
117
|
+
obj[field_name], encryption_instance, model_type, field=field_name
|
|
118
|
+
)
|
|
119
|
+
nested_fields = NESTED_ENCRYPTED_SUBFIELDS.get((model_type, field_name), ())
|
|
120
|
+
if decrypted is not None:
|
|
121
|
+
results = await asyncio.gather(
|
|
122
|
+
*(
|
|
123
|
+
_decrypt_field(decrypted, name, encryption_instance, model_type)
|
|
124
|
+
for name in nested_fields
|
|
125
|
+
if isinstance(decrypted.get(name), dict)
|
|
126
|
+
)
|
|
127
|
+
)
|
|
128
|
+
for name, value in results:
|
|
129
|
+
decrypted[name] = value
|
|
130
|
+
return field_name, decrypted
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
async def decrypt_fields(
|
|
134
|
+
obj: dict[str, Any],
|
|
135
|
+
model_type: ModelType,
|
|
136
|
+
fields: list[str],
|
|
137
|
+
encryption_instance: JsonEncryptionWrapper | AesEncryptionInstance | None,
|
|
138
|
+
) -> None:
|
|
139
|
+
"""Decrypt configured fields, including local-runtime nested structure."""
|
|
140
|
+
results = await asyncio.gather(
|
|
141
|
+
*(
|
|
142
|
+
_decrypt_field(obj, name, encryption_instance, model_type)
|
|
143
|
+
for name in fields
|
|
144
|
+
if name in obj
|
|
145
|
+
)
|
|
146
|
+
)
|
|
147
|
+
for name, value in results:
|
|
148
|
+
obj[name] = value
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
async def _encrypt_field(
|
|
152
|
+
data: Mapping[str, Any],
|
|
153
|
+
field_name: str,
|
|
154
|
+
encryption_instance: JsonEncryptionWrapper | AesEncryptionInstance | None,
|
|
155
|
+
model_type: ModelType,
|
|
156
|
+
path: str | None = None,
|
|
157
|
+
) -> tuple[str, Any]:
|
|
158
|
+
from langgraph_api.encryption.middleware import ( # noqa: PLC0415
|
|
159
|
+
encrypt_json_if_needed,
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
if field_name not in data or data[field_name] is None:
|
|
163
|
+
return field_name, data.get(field_name)
|
|
164
|
+
|
|
165
|
+
field_data = data[field_name]
|
|
166
|
+
if model_type == "thread" and field_name == "values":
|
|
167
|
+
from langgraph_api.serde import json_dumpb, json_loads # noqa: PLC0415
|
|
168
|
+
|
|
169
|
+
field_data = json_loads(json_dumpb(field_data))
|
|
170
|
+
current_path = f"{path}.{field_name}" if path else f"{model_type}.{field_name}"
|
|
171
|
+
nested_fields = NESTED_ENCRYPTED_SUBFIELDS.get((model_type, field_name), ())
|
|
172
|
+
extracted: dict[str, Any] = {}
|
|
173
|
+
if nested_fields:
|
|
174
|
+
if not isinstance(field_data, dict):
|
|
175
|
+
raise TypeError(
|
|
176
|
+
f"'{field_name}' must be a dict for encryption, "
|
|
177
|
+
f"got {type(field_data).__name__}"
|
|
178
|
+
)
|
|
179
|
+
for name in nested_fields:
|
|
180
|
+
value = field_data.get(name)
|
|
181
|
+
if isinstance(value, dict) and value:
|
|
182
|
+
extracted[name] = value
|
|
183
|
+
if extracted:
|
|
184
|
+
field_data = {k: v for k, v in field_data.items() if k not in extracted}
|
|
185
|
+
|
|
186
|
+
encrypted = await encrypt_json_if_needed(
|
|
187
|
+
field_data,
|
|
188
|
+
encryption_instance,
|
|
189
|
+
model_type,
|
|
190
|
+
field=field_name,
|
|
191
|
+
path=current_path,
|
|
192
|
+
should_skip=should_skip_encryption,
|
|
193
|
+
)
|
|
194
|
+
if extracted and isinstance(encrypted, dict):
|
|
195
|
+
results = await asyncio.gather(
|
|
196
|
+
*(
|
|
197
|
+
_encrypt_field(
|
|
198
|
+
{name: value}, name, encryption_instance, model_type, current_path
|
|
199
|
+
)
|
|
200
|
+
for name, value in extracted.items()
|
|
201
|
+
)
|
|
202
|
+
)
|
|
203
|
+
for name, value in results:
|
|
204
|
+
encrypted[name] = value
|
|
205
|
+
marker = value.get("__encryption_context__")
|
|
206
|
+
if (
|
|
207
|
+
"__encryption_context__" not in encrypted
|
|
208
|
+
and isinstance(marker, dict)
|
|
209
|
+
and marker.get("__langgraph_encryption_type__") == "aes"
|
|
210
|
+
):
|
|
211
|
+
encrypted["__encryption_context__"] = marker
|
|
212
|
+
return field_name, encrypted
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
async def encrypt_fields(
|
|
216
|
+
data: Mapping[str, Any],
|
|
217
|
+
model_type: ModelType,
|
|
218
|
+
fields: list[str],
|
|
219
|
+
encryption_instance: JsonEncryptionWrapper | AesEncryptionInstance | None,
|
|
220
|
+
) -> dict[str, Any]:
|
|
221
|
+
"""Encrypt configured fields, preserving local-runtime nested structure."""
|
|
222
|
+
result = dict(data)
|
|
223
|
+
encrypted_fields = await asyncio.gather(
|
|
224
|
+
*(
|
|
225
|
+
_encrypt_field(data, name, encryption_instance, model_type)
|
|
226
|
+
for name in fields
|
|
227
|
+
if name in data
|
|
228
|
+
)
|
|
229
|
+
)
|
|
230
|
+
for name, value in encrypted_fields:
|
|
231
|
+
result[name] = value
|
|
232
|
+
return result
|