riftpoint 1.0.0__py3-none-any.whl
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.
- riftpoint/__init__.py +85 -0
- riftpoint/checkpointer/__init__.py +7 -0
- riftpoint/checkpointer/async_saver.py +483 -0
- riftpoint/checkpointer/base.py +571 -0
- riftpoint/checkpointer/sync_saver.py +313 -0
- riftpoint/collapse/__init__.py +22 -0
- riftpoint/collapse/evaluators.py +362 -0
- riftpoint/collapse/resolver.py +249 -0
- riftpoint/decorators.py +370 -0
- riftpoint/logger.py +28 -0
- riftpoint/py.typed +1 -0
- riftpoint/runner/__init__.py +22 -0
- riftpoint/runner/beam.py +452 -0
- riftpoint/runner/branch.py +187 -0
- riftpoint/runner/multiverse.py +213 -0
- riftpoint/serde/__init__.py +21 -0
- riftpoint/serde/codec.py +261 -0
- riftpoint/visualization/__init__.py +9 -0
- riftpoint/visualization/mermaid.py +33 -0
- riftpoint/visualization/plot.py +47 -0
- riftpoint-1.0.0.dist-info/METADATA +144 -0
- riftpoint-1.0.0.dist-info/RECORD +23 -0
- riftpoint-1.0.0.dist-info/WHEEL +4 -0
riftpoint/__init__.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
"""RiftPoint package."""
|
|
2
|
+
|
|
3
|
+
from riftpoint.checkpointer import (
|
|
4
|
+
AsyncRiftCheckpointSaver,
|
|
5
|
+
BaseRiftSaver,
|
|
6
|
+
RiftCheckpointSaver,
|
|
7
|
+
)
|
|
8
|
+
from riftpoint.collapse import (
|
|
9
|
+
BaseEvaluator,
|
|
10
|
+
CollapseResult,
|
|
11
|
+
ConsensusEvaluator,
|
|
12
|
+
EvaluationResult,
|
|
13
|
+
HeuristicEvaluator,
|
|
14
|
+
JSONSchemaEvaluator,
|
|
15
|
+
LLMJudgeEvaluator,
|
|
16
|
+
MultiverseResolver,
|
|
17
|
+
)
|
|
18
|
+
from riftpoint.decorators import (
|
|
19
|
+
SpeculativeNodeConfig,
|
|
20
|
+
SpeculativeRaceMeta,
|
|
21
|
+
speculative_node,
|
|
22
|
+
)
|
|
23
|
+
from riftpoint.logger import logger
|
|
24
|
+
from riftpoint.runner import (
|
|
25
|
+
BeamSearchConfig,
|
|
26
|
+
BeamSearchResult,
|
|
27
|
+
BeamSearchRunner,
|
|
28
|
+
BranchInfo,
|
|
29
|
+
BranchManager,
|
|
30
|
+
BranchResult,
|
|
31
|
+
BranchSpec,
|
|
32
|
+
DepthSummary,
|
|
33
|
+
RiftRunner,
|
|
34
|
+
)
|
|
35
|
+
from riftpoint.serde import (
|
|
36
|
+
SessionSnapshot,
|
|
37
|
+
export_session_bytes,
|
|
38
|
+
export_session_json,
|
|
39
|
+
export_session_snapshot,
|
|
40
|
+
import_session_bytes,
|
|
41
|
+
import_session_json,
|
|
42
|
+
import_session_snapshot,
|
|
43
|
+
)
|
|
44
|
+
from riftpoint.visualization import (
|
|
45
|
+
plot_multiverse,
|
|
46
|
+
render_mermaid,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
__version__ = "1.0.0"
|
|
50
|
+
__all__ = [
|
|
51
|
+
"AsyncRiftCheckpointSaver",
|
|
52
|
+
"BaseEvaluator",
|
|
53
|
+
"BaseRiftSaver",
|
|
54
|
+
"BeamSearchConfig",
|
|
55
|
+
"BeamSearchResult",
|
|
56
|
+
"BeamSearchRunner",
|
|
57
|
+
"BranchInfo",
|
|
58
|
+
"BranchManager",
|
|
59
|
+
"BranchResult",
|
|
60
|
+
"BranchSpec",
|
|
61
|
+
"CollapseResult",
|
|
62
|
+
"ConsensusEvaluator",
|
|
63
|
+
"DepthSummary",
|
|
64
|
+
"EvaluationResult",
|
|
65
|
+
"HeuristicEvaluator",
|
|
66
|
+
"JSONSchemaEvaluator",
|
|
67
|
+
"LLMJudgeEvaluator",
|
|
68
|
+
"MultiverseResolver",
|
|
69
|
+
"RiftCheckpointSaver",
|
|
70
|
+
"RiftRunner",
|
|
71
|
+
"SessionSnapshot",
|
|
72
|
+
"SpeculativeNodeConfig",
|
|
73
|
+
"SpeculativeRaceMeta",
|
|
74
|
+
"__version__",
|
|
75
|
+
"export_session_bytes",
|
|
76
|
+
"export_session_json",
|
|
77
|
+
"export_session_snapshot",
|
|
78
|
+
"import_session_bytes",
|
|
79
|
+
"import_session_json",
|
|
80
|
+
"import_session_snapshot",
|
|
81
|
+
"logger",
|
|
82
|
+
"plot_multiverse",
|
|
83
|
+
"render_mermaid",
|
|
84
|
+
"speculative_node",
|
|
85
|
+
]
|
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
"""Checkpointer module for RiftPoint."""
|
|
2
|
+
|
|
3
|
+
from riftpoint.checkpointer.async_saver import AsyncRiftCheckpointSaver
|
|
4
|
+
from riftpoint.checkpointer.base import BaseRiftSaver
|
|
5
|
+
from riftpoint.checkpointer.sync_saver import RiftCheckpointSaver
|
|
6
|
+
|
|
7
|
+
__all__ = ["AsyncRiftCheckpointSaver", "BaseRiftSaver", "RiftCheckpointSaver"]
|
|
@@ -0,0 +1,483 @@
|
|
|
1
|
+
"""Asynchronous Janus-backed CheckpointSaver for LangGraph."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
from typing import TYPE_CHECKING, Any, cast
|
|
7
|
+
|
|
8
|
+
from langgraph.checkpoint.base import (
|
|
9
|
+
WRITES_IDX_MAP,
|
|
10
|
+
BaseCheckpointSaver,
|
|
11
|
+
get_checkpoint_metadata,
|
|
12
|
+
)
|
|
13
|
+
|
|
14
|
+
from riftpoint.checkpointer.base import BaseRiftSaver
|
|
15
|
+
|
|
16
|
+
if TYPE_CHECKING:
|
|
17
|
+
from collections.abc import AsyncIterator, Iterator, Sequence
|
|
18
|
+
|
|
19
|
+
from langchain_core.runnables import RunnableConfig
|
|
20
|
+
from langgraph.checkpoint.base import (
|
|
21
|
+
ChannelVersions,
|
|
22
|
+
Checkpoint,
|
|
23
|
+
CheckpointMetadata,
|
|
24
|
+
CheckpointTuple,
|
|
25
|
+
)
|
|
26
|
+
from langgraph.checkpoint.serde.base import SerializerProtocol
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class AsyncRiftCheckpointSaver(BaseCheckpointSaver[str], BaseRiftSaver):
|
|
30
|
+
"""Async multiversal checkpointer for LangGraph powered by Janus-Tachyon-RS."""
|
|
31
|
+
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
*,
|
|
35
|
+
serde: SerializerProtocol | None = None,
|
|
36
|
+
) -> None:
|
|
37
|
+
"""Initialize a new AsyncRiftCheckpointSaver.
|
|
38
|
+
|
|
39
|
+
Args:
|
|
40
|
+
serde: Optional serializer for state payloads.
|
|
41
|
+
Defaults to JsonPlusSerializer.
|
|
42
|
+
"""
|
|
43
|
+
BaseCheckpointSaver.__init__(self, serde=serde)
|
|
44
|
+
BaseRiftSaver.__init__(self, serde=serde)
|
|
45
|
+
self._locks: dict[str, asyncio.Lock] = {}
|
|
46
|
+
|
|
47
|
+
def _get_lock(self, thread_id: str) -> asyncio.Lock:
|
|
48
|
+
"""Retrieve or create an asyncio.Lock for the specified thread ID."""
|
|
49
|
+
if thread_id not in self._locks:
|
|
50
|
+
self._locks[thread_id] = asyncio.Lock()
|
|
51
|
+
return self._locks[thread_id]
|
|
52
|
+
|
|
53
|
+
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
|
54
|
+
"""Asynchronously retrieve a checkpoint tuple for the given configuration.
|
|
55
|
+
|
|
56
|
+
Args:
|
|
57
|
+
config: Configuration with thread_id, checkpoint_ns, and checkpoint_id.
|
|
58
|
+
|
|
59
|
+
Returns:
|
|
60
|
+
CheckpointTuple if found, None otherwise.
|
|
61
|
+
"""
|
|
62
|
+
thread_id = config["configurable"]["thread_id"]
|
|
63
|
+
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
64
|
+
raw_chk_id = get_checkpoint_metadata(config, {}).get("checkpoint_id") or config[
|
|
65
|
+
"configurable"
|
|
66
|
+
].get("checkpoint_id")
|
|
67
|
+
checkpoint_id = str(raw_chk_id) if raw_chk_id is not None else None
|
|
68
|
+
|
|
69
|
+
thread_storage = self._storage.get(thread_id, {}).get(checkpoint_ns, {})
|
|
70
|
+
if not thread_storage:
|
|
71
|
+
return None
|
|
72
|
+
|
|
73
|
+
target_id: str | None = checkpoint_id
|
|
74
|
+
if not target_id:
|
|
75
|
+
target_id = list(thread_storage.keys())[-1]
|
|
76
|
+
|
|
77
|
+
if target_id not in thread_storage:
|
|
78
|
+
return None
|
|
79
|
+
|
|
80
|
+
storage_entry = thread_storage[target_id]
|
|
81
|
+
return self._build_checkpoint_tuple(
|
|
82
|
+
thread_id,
|
|
83
|
+
checkpoint_ns,
|
|
84
|
+
target_id,
|
|
85
|
+
storage_entry,
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
async def _iter_ns_checkpoints_async(
|
|
89
|
+
self,
|
|
90
|
+
thread_id: str,
|
|
91
|
+
namespace: str,
|
|
92
|
+
filter_dict: dict[str, Any] | None,
|
|
93
|
+
before_id: str | None,
|
|
94
|
+
) -> AsyncIterator[CheckpointTuple]:
|
|
95
|
+
"""Traverse checkpoints for a namespace in reverse chronological order."""
|
|
96
|
+
checkpoints = self._storage.get(thread_id, {}).get(namespace, {})
|
|
97
|
+
for chk_id in reversed(list(checkpoints.keys())):
|
|
98
|
+
if before_id and chk_id >= before_id:
|
|
99
|
+
continue
|
|
100
|
+
|
|
101
|
+
storage_entry = checkpoints[chk_id]
|
|
102
|
+
_, ser_meta, _ = storage_entry
|
|
103
|
+
meta_dict = cast(
|
|
104
|
+
"CheckpointMetadata",
|
|
105
|
+
self.serde.loads_typed(ser_meta),
|
|
106
|
+
)
|
|
107
|
+
if not self._matches_filter(meta_dict, filter_dict):
|
|
108
|
+
continue
|
|
109
|
+
|
|
110
|
+
yield self._build_checkpoint_tuple(
|
|
111
|
+
thread_id,
|
|
112
|
+
namespace,
|
|
113
|
+
chk_id,
|
|
114
|
+
storage_entry,
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
async def alist(
|
|
118
|
+
self,
|
|
119
|
+
config: RunnableConfig | None,
|
|
120
|
+
*,
|
|
121
|
+
filter: dict[str, Any] | None = None, # noqa: A002
|
|
122
|
+
before: RunnableConfig | None = None,
|
|
123
|
+
limit: int | None = None,
|
|
124
|
+
) -> AsyncIterator[CheckpointTuple]:
|
|
125
|
+
"""Asynchronously list checkpoint tuples matching the specified criteria.
|
|
126
|
+
|
|
127
|
+
Args:
|
|
128
|
+
config: Optional configuration specifying thread_id and checkpoint_ns.
|
|
129
|
+
filter: Optional metadata filter key-value pairs.
|
|
130
|
+
before: Optional configuration specifying the upper bound checkpoint.
|
|
131
|
+
limit: Maximum number of checkpoints to return.
|
|
132
|
+
|
|
133
|
+
Yields:
|
|
134
|
+
Matching CheckpointTuple instances in reverse chronological order.
|
|
135
|
+
"""
|
|
136
|
+
thread_ids = (
|
|
137
|
+
[config["configurable"]["thread_id"]]
|
|
138
|
+
if config
|
|
139
|
+
else list(self._storage.keys())
|
|
140
|
+
)
|
|
141
|
+
yielded = 0
|
|
142
|
+
before_id = before["configurable"].get("checkpoint_id") if before else None
|
|
143
|
+
|
|
144
|
+
for thread_id in thread_ids:
|
|
145
|
+
ns_dict = self._storage.get(thread_id, {})
|
|
146
|
+
checkpoint_ns = (
|
|
147
|
+
config["configurable"].get("checkpoint_ns", "") if config else None
|
|
148
|
+
)
|
|
149
|
+
namespaces = (
|
|
150
|
+
[checkpoint_ns] if checkpoint_ns is not None else list(ns_dict.keys())
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
for ns in namespaces:
|
|
154
|
+
async for tuple_item in self._iter_ns_checkpoints_async(
|
|
155
|
+
thread_id, ns, filter, before_id
|
|
156
|
+
):
|
|
157
|
+
if limit is not None and yielded >= limit:
|
|
158
|
+
return
|
|
159
|
+
yield tuple_item
|
|
160
|
+
yielded += 1
|
|
161
|
+
|
|
162
|
+
async def aput(
|
|
163
|
+
self,
|
|
164
|
+
config: RunnableConfig,
|
|
165
|
+
checkpoint: Checkpoint,
|
|
166
|
+
metadata: CheckpointMetadata,
|
|
167
|
+
new_versions: ChannelVersions,
|
|
168
|
+
) -> RunnableConfig:
|
|
169
|
+
"""Asynchronously store a checkpoint snapshot into the Janus multiverse.
|
|
170
|
+
|
|
171
|
+
Args:
|
|
172
|
+
config: Configuration specifying thread_id and checkpoint_ns.
|
|
173
|
+
checkpoint: The checkpoint payload to save.
|
|
174
|
+
metadata: Associated checkpoint metadata.
|
|
175
|
+
new_versions: Channel versions for updated channels.
|
|
176
|
+
|
|
177
|
+
Returns:
|
|
178
|
+
Updated RunnableConfig pointing to the new checkpoint_id.
|
|
179
|
+
"""
|
|
180
|
+
thread_id = config["configurable"]["thread_id"]
|
|
181
|
+
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
182
|
+
checkpoint_id = checkpoint["id"]
|
|
183
|
+
parent_checkpoint_id = config["configurable"].get("checkpoint_id")
|
|
184
|
+
|
|
185
|
+
async with self._get_lock(thread_id):
|
|
186
|
+
c_dict = dict(checkpoint)
|
|
187
|
+
values = cast("dict[str, Any]", c_dict.pop("channel_values", {}))
|
|
188
|
+
self._store_blobs(thread_id, checkpoint_ns, values, new_versions)
|
|
189
|
+
|
|
190
|
+
stored_metadata = get_checkpoint_metadata(config, metadata)
|
|
191
|
+
stored_chk = cast("Checkpoint", c_dict)
|
|
192
|
+
self._storage[thread_id][checkpoint_ns][checkpoint_id] = (
|
|
193
|
+
self.serde.dumps_typed(stored_chk),
|
|
194
|
+
self.serde.dumps_typed(stored_metadata),
|
|
195
|
+
parent_checkpoint_id,
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
# Record moment in Janus Multiverse DAG
|
|
199
|
+
self._record_janus_checkpoint(
|
|
200
|
+
thread_id,
|
|
201
|
+
checkpoint_id,
|
|
202
|
+
stored_metadata,
|
|
203
|
+
)
|
|
204
|
+
|
|
205
|
+
return {
|
|
206
|
+
"configurable": {
|
|
207
|
+
"thread_id": thread_id,
|
|
208
|
+
"checkpoint_ns": checkpoint_ns,
|
|
209
|
+
"checkpoint_id": checkpoint_id,
|
|
210
|
+
}
|
|
211
|
+
}
|
|
212
|
+
|
|
213
|
+
async def aput_writes(
|
|
214
|
+
self,
|
|
215
|
+
config: RunnableConfig,
|
|
216
|
+
writes: Sequence[tuple[str, Any]],
|
|
217
|
+
task_id: str,
|
|
218
|
+
task_path: str = "",
|
|
219
|
+
) -> None:
|
|
220
|
+
"""Asynchronously store intermediate task writes for a checkpoint.
|
|
221
|
+
|
|
222
|
+
Args:
|
|
223
|
+
config: Configuration specifying thread_id, ns, and checkpoint_id.
|
|
224
|
+
writes: Sequence of (channel, value) writes.
|
|
225
|
+
task_id: Unique task identifier.
|
|
226
|
+
task_path: Optional execution path of the task.
|
|
227
|
+
"""
|
|
228
|
+
thread_id = config["configurable"]["thread_id"]
|
|
229
|
+
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
230
|
+
checkpoint_id = config["configurable"]["checkpoint_id"]
|
|
231
|
+
|
|
232
|
+
async with self._get_lock(thread_id):
|
|
233
|
+
outer_key = (thread_id, checkpoint_ns, checkpoint_id)
|
|
234
|
+
existing_writes = self._writes[outer_key]
|
|
235
|
+
|
|
236
|
+
for idx, (channel, value) in enumerate(writes):
|
|
237
|
+
inner_key = (task_id, WRITES_IDX_MAP.get(channel, idx))
|
|
238
|
+
if inner_key[1] >= 0 and inner_key in existing_writes:
|
|
239
|
+
continue
|
|
240
|
+
|
|
241
|
+
existing_writes[inner_key] = (
|
|
242
|
+
task_id,
|
|
243
|
+
channel,
|
|
244
|
+
self.serde.dumps_typed(value),
|
|
245
|
+
task_path,
|
|
246
|
+
)
|
|
247
|
+
|
|
248
|
+
async def adelete_thread(self, thread_id: str) -> None:
|
|
249
|
+
"""Asynchronously delete all checkpoints, writes, and Janus state for a thread.
|
|
250
|
+
|
|
251
|
+
Args:
|
|
252
|
+
thread_id: The thread ID to purge.
|
|
253
|
+
"""
|
|
254
|
+
async with self._get_lock(thread_id):
|
|
255
|
+
self._storage.pop(thread_id, None)
|
|
256
|
+
self._multiverses.pop(thread_id, None)
|
|
257
|
+
|
|
258
|
+
for write_key in list(self._writes.keys()):
|
|
259
|
+
if write_key[0] == thread_id:
|
|
260
|
+
del self._writes[write_key]
|
|
261
|
+
|
|
262
|
+
for blob_key in list(self._blobs.keys()):
|
|
263
|
+
if blob_key[0] == thread_id:
|
|
264
|
+
del self._blobs[blob_key]
|
|
265
|
+
|
|
266
|
+
self._locks.pop(thread_id, None)
|
|
267
|
+
|
|
268
|
+
def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
|
269
|
+
"""Synchronous wrapper for get_tuple."""
|
|
270
|
+
thread_id = config["configurable"]["thread_id"]
|
|
271
|
+
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
272
|
+
raw_chk_id = get_checkpoint_metadata(config, {}).get("checkpoint_id") or config[
|
|
273
|
+
"configurable"
|
|
274
|
+
].get("checkpoint_id")
|
|
275
|
+
checkpoint_id = str(raw_chk_id) if raw_chk_id is not None else None
|
|
276
|
+
|
|
277
|
+
thread_storage = self._storage.get(thread_id, {}).get(checkpoint_ns, {})
|
|
278
|
+
if not thread_storage:
|
|
279
|
+
return None
|
|
280
|
+
|
|
281
|
+
target_id = checkpoint_id or list(thread_storage.keys())[-1]
|
|
282
|
+
if target_id not in thread_storage:
|
|
283
|
+
return None
|
|
284
|
+
|
|
285
|
+
storage_entry = thread_storage[target_id]
|
|
286
|
+
return self._build_checkpoint_tuple(
|
|
287
|
+
thread_id,
|
|
288
|
+
checkpoint_ns,
|
|
289
|
+
target_id,
|
|
290
|
+
storage_entry,
|
|
291
|
+
)
|
|
292
|
+
|
|
293
|
+
def list(
|
|
294
|
+
self,
|
|
295
|
+
config: RunnableConfig | None,
|
|
296
|
+
*,
|
|
297
|
+
filter: dict[str, Any] | None = None, # noqa: A002
|
|
298
|
+
before: RunnableConfig | None = None,
|
|
299
|
+
limit: int | None = None,
|
|
300
|
+
) -> Iterator[CheckpointTuple]:
|
|
301
|
+
"""Synchronous wrapper for list."""
|
|
302
|
+
thread_ids = (
|
|
303
|
+
[config["configurable"]["thread_id"]]
|
|
304
|
+
if config
|
|
305
|
+
else list(self._storage.keys())
|
|
306
|
+
)
|
|
307
|
+
yielded = 0
|
|
308
|
+
before_id = before["configurable"].get("checkpoint_id") if before else None
|
|
309
|
+
|
|
310
|
+
for thread_id in thread_ids:
|
|
311
|
+
ns_dict = self._storage.get(thread_id, {})
|
|
312
|
+
checkpoint_ns = (
|
|
313
|
+
config["configurable"].get("checkpoint_ns", "") if config else None
|
|
314
|
+
)
|
|
315
|
+
namespaces = (
|
|
316
|
+
[checkpoint_ns] if checkpoint_ns is not None else list(ns_dict.keys())
|
|
317
|
+
)
|
|
318
|
+
|
|
319
|
+
for ns in namespaces:
|
|
320
|
+
for tuple_item in self._iter_ns_checkpoints(
|
|
321
|
+
thread_id, ns, filter, before_id
|
|
322
|
+
):
|
|
323
|
+
if limit is not None and yielded >= limit:
|
|
324
|
+
return
|
|
325
|
+
yield tuple_item
|
|
326
|
+
yielded += 1
|
|
327
|
+
|
|
328
|
+
def put(
|
|
329
|
+
self,
|
|
330
|
+
config: RunnableConfig,
|
|
331
|
+
checkpoint: Checkpoint,
|
|
332
|
+
metadata: CheckpointMetadata,
|
|
333
|
+
new_versions: ChannelVersions,
|
|
334
|
+
) -> RunnableConfig:
|
|
335
|
+
"""Synchronous wrapper for put."""
|
|
336
|
+
thread_id = config["configurable"]["thread_id"]
|
|
337
|
+
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
338
|
+
checkpoint_id = checkpoint["id"]
|
|
339
|
+
parent_checkpoint_id = config["configurable"].get("checkpoint_id")
|
|
340
|
+
|
|
341
|
+
c_dict = dict(checkpoint)
|
|
342
|
+
values = cast("dict[str, Any]", c_dict.pop("channel_values", {}))
|
|
343
|
+
self._store_blobs(thread_id, checkpoint_ns, values, new_versions)
|
|
344
|
+
|
|
345
|
+
stored_metadata = get_checkpoint_metadata(config, metadata)
|
|
346
|
+
stored_chk = cast("Checkpoint", c_dict)
|
|
347
|
+
self._storage[thread_id][checkpoint_ns][checkpoint_id] = (
|
|
348
|
+
self.serde.dumps_typed(stored_chk),
|
|
349
|
+
self.serde.dumps_typed(stored_metadata),
|
|
350
|
+
parent_checkpoint_id,
|
|
351
|
+
)
|
|
352
|
+
|
|
353
|
+
self._record_janus_checkpoint(
|
|
354
|
+
thread_id,
|
|
355
|
+
checkpoint_id,
|
|
356
|
+
stored_metadata,
|
|
357
|
+
)
|
|
358
|
+
|
|
359
|
+
return {
|
|
360
|
+
"configurable": {
|
|
361
|
+
"thread_id": thread_id,
|
|
362
|
+
"checkpoint_ns": checkpoint_ns,
|
|
363
|
+
"checkpoint_id": checkpoint_id,
|
|
364
|
+
}
|
|
365
|
+
}
|
|
366
|
+
|
|
367
|
+
def put_writes(
|
|
368
|
+
self,
|
|
369
|
+
config: RunnableConfig,
|
|
370
|
+
writes: Sequence[tuple[str, Any]],
|
|
371
|
+
task_id: str,
|
|
372
|
+
task_path: str = "",
|
|
373
|
+
) -> None:
|
|
374
|
+
"""Synchronous wrapper for put_writes."""
|
|
375
|
+
thread_id = config["configurable"]["thread_id"]
|
|
376
|
+
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
377
|
+
checkpoint_id = config["configurable"]["checkpoint_id"]
|
|
378
|
+
|
|
379
|
+
outer_key = (thread_id, checkpoint_ns, checkpoint_id)
|
|
380
|
+
existing_writes = self._writes[outer_key]
|
|
381
|
+
|
|
382
|
+
for idx, (channel, value) in enumerate(writes):
|
|
383
|
+
inner_key = (task_id, WRITES_IDX_MAP.get(channel, idx))
|
|
384
|
+
if inner_key[1] >= 0 and inner_key in existing_writes:
|
|
385
|
+
continue
|
|
386
|
+
|
|
387
|
+
existing_writes[inner_key] = (
|
|
388
|
+
task_id,
|
|
389
|
+
channel,
|
|
390
|
+
self.serde.dumps_typed(value),
|
|
391
|
+
task_path,
|
|
392
|
+
)
|
|
393
|
+
|
|
394
|
+
def delete_thread(self, thread_id: str) -> None:
|
|
395
|
+
"""Synchronous wrapper for delete_thread."""
|
|
396
|
+
self._storage.pop(thread_id, None)
|
|
397
|
+
self._multiverses.pop(thread_id, None)
|
|
398
|
+
|
|
399
|
+
for write_key in list(self._writes.keys()):
|
|
400
|
+
if write_key[0] == thread_id:
|
|
401
|
+
del self._writes[write_key]
|
|
402
|
+
|
|
403
|
+
for blob_key in list(self._blobs.keys()):
|
|
404
|
+
if blob_key[0] == thread_id:
|
|
405
|
+
del self._blobs[blob_key]
|
|
406
|
+
|
|
407
|
+
self._locks.pop(thread_id, None)
|
|
408
|
+
|
|
409
|
+
def visualize(self, thread_id: str) -> str:
|
|
410
|
+
"""Generate a Mermaid diagram representing the checkpoint history DAG.
|
|
411
|
+
|
|
412
|
+
Args:
|
|
413
|
+
thread_id: Thread ID to visualize.
|
|
414
|
+
|
|
415
|
+
Returns:
|
|
416
|
+
Mermaid formatted markdown string of the timeline DAG.
|
|
417
|
+
"""
|
|
418
|
+
lines = ["```mermaid", "graph TD"]
|
|
419
|
+
ns_dict = self._storage.get(thread_id, {})
|
|
420
|
+
for checkpoints in ns_dict.values():
|
|
421
|
+
for chk_id, (_, ser_meta, parent_id) in checkpoints.items():
|
|
422
|
+
meta = cast(
|
|
423
|
+
"CheckpointMetadata",
|
|
424
|
+
self.serde.loads_typed(ser_meta),
|
|
425
|
+
)
|
|
426
|
+
step = meta.get("step", "")
|
|
427
|
+
node_label = f'{chk_id}["{chk_id} (step: {step})"]'
|
|
428
|
+
lines.append(f" {node_label}")
|
|
429
|
+
if parent_id:
|
|
430
|
+
lines.append(f" {parent_id} --> {chk_id}")
|
|
431
|
+
lines.append("```")
|
|
432
|
+
return "\n".join(lines)
|
|
433
|
+
|
|
434
|
+
async def aplot(
|
|
435
|
+
self,
|
|
436
|
+
thread_id: str,
|
|
437
|
+
*,
|
|
438
|
+
output_path: str | None = None,
|
|
439
|
+
**kwargs: Any,
|
|
440
|
+
) -> Any:
|
|
441
|
+
"""Asynchronously render a plot of the multiversal DAG.
|
|
442
|
+
|
|
443
|
+
Args:
|
|
444
|
+
thread_id: Primary session thread ID.
|
|
445
|
+
output_path: Optional path to save figure.
|
|
446
|
+
**kwargs: Extra plotting options forwarded to Janus engine.
|
|
447
|
+
|
|
448
|
+
Returns:
|
|
449
|
+
Plot object.
|
|
450
|
+
"""
|
|
451
|
+
return self.plot(thread_id, output_path=output_path, **kwargs)
|
|
452
|
+
|
|
453
|
+
async def aexport_session(
|
|
454
|
+
self,
|
|
455
|
+
thread_id: str,
|
|
456
|
+
file_path: str | None = None,
|
|
457
|
+
) -> str:
|
|
458
|
+
"""Asynchronously export session state as JSON string or file.
|
|
459
|
+
|
|
460
|
+
Args:
|
|
461
|
+
thread_id: Session thread ID.
|
|
462
|
+
file_path: Optional destination file path.
|
|
463
|
+
|
|
464
|
+
Returns:
|
|
465
|
+
JSON string representation of session snapshot.
|
|
466
|
+
"""
|
|
467
|
+
lock = self._get_lock(thread_id)
|
|
468
|
+
async with lock:
|
|
469
|
+
return self.export_session(thread_id, file_path=file_path)
|
|
470
|
+
|
|
471
|
+
async def aimport_session(
|
|
472
|
+
self,
|
|
473
|
+
data: str,
|
|
474
|
+
) -> str:
|
|
475
|
+
"""Asynchronously import and restore session state from JSON.
|
|
476
|
+
|
|
477
|
+
Args:
|
|
478
|
+
data: JSON string content or path to file.
|
|
479
|
+
|
|
480
|
+
Returns:
|
|
481
|
+
Restored thread ID.
|
|
482
|
+
"""
|
|
483
|
+
return self.import_session(data)
|