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.
Files changed (20) hide show
  1. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/PKG-INFO +1 -1
  2. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/__init__.py +3 -1
  3. langgraph_runtime_inmem-0.37.0.dev2/langgraph_runtime_inmem/_persistence.py +88 -0
  4. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/checkpoint.py +13 -4
  5. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/database.py +97 -19
  6. langgraph_runtime_inmem-0.37.0.dev2/langgraph_runtime_inmem/encryption.py +232 -0
  7. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/ops.py +326 -163
  8. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/queue.py +13 -3
  9. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/store.py +8 -2
  10. langgraph_runtime_inmem-0.36.0rc2/langgraph_runtime_inmem/_persistence.py +0 -64
  11. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/.gitignore +0 -0
  12. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/Makefile +0 -0
  13. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/README.md +0 -0
  14. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/inmem_stream.py +0 -0
  15. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/lifespan.py +0 -0
  16. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/metrics.py +0 -0
  17. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/retry.py +0 -0
  18. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/langgraph_runtime_inmem/routes.py +0 -0
  19. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/pyproject.toml +0 -0
  20. {langgraph_runtime_inmem-0.36.0rc2 → langgraph_runtime_inmem-0.37.0.dev2}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: langgraph-runtime-inmem
3
- Version: 0.36.0rc2
3
+ Version: 0.37.0.dev2
4
4
  Summary: Inmem implementation for the LangGraph API server.
5
5
  Author-email: Will Fu-Hinthorn <will@langchain.dev>
6
6
  License: Elastic-2.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.36.0rc2"
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 not DISABLE_FILE_PERSISTENCE else defaultdict,
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
- __persistence_hook__=register_persistent_dict,
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
- if os.path.exists(OPS_FILENAME):
168
- try:
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