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 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)