68 lines
3.0 KiB
Python
68 lines
3.0 KiB
Python
"""Copy a conversation's LangGraph checkpoint state from one thread to another.
|
|
|
|
Used by the conversation-sharing feature: redeeming a share link duplicates the
|
|
sharer's threads into the redeemer's account. Each duplicated thread needs its
|
|
full checkpoint history copied under a fresh ``thread_id``.
|
|
|
|
The copy goes through the checkpointer's public API (``alist`` / ``aput`` /
|
|
``aput_writes``) rather than raw SQL, so it is independent of the checkpoint
|
|
table schema and works for every backend (SQLite, Postgres, in-memory).
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
async def copy_thread_state(checkpointer, src_thread_id: str, dst_thread_id: str) -> bool:
|
|
"""Duplicate every checkpoint of ``src_thread_id`` under ``dst_thread_id``.
|
|
|
|
Returns True when at least one checkpoint was copied, False when the source
|
|
thread has no checkpoints (e.g. a thread that was never run).
|
|
|
|
The destination keeps the source's checkpoint ids — they only need to be
|
|
unique within a ``thread_id``, and reusing them preserves the parent chain.
|
|
"""
|
|
src_config = {"configurable": {"thread_id": src_thread_id}}
|
|
# alist yields newest-first; reverse so parents are written before children.
|
|
tuples = [t async for t in checkpointer.alist(src_config)]
|
|
if not tuples:
|
|
return False
|
|
|
|
copied = 0
|
|
for tup in reversed(tuples):
|
|
configurable = (tup.config or {}).get("configurable", {})
|
|
checkpoint_ns = configurable.get("checkpoint_ns", "") or ""
|
|
parent_id = ((tup.parent_config or {}).get("configurable", {}) or {}).get("checkpoint_id")
|
|
|
|
put_config: dict = {"configurable": {"thread_id": dst_thread_id, "checkpoint_ns": checkpoint_ns}}
|
|
if parent_id:
|
|
# aput reads the parent checkpoint id from the incoming config.
|
|
put_config["configurable"]["checkpoint_id"] = parent_id
|
|
|
|
new_versions = tup.checkpoint.get("channel_versions", {}) or {}
|
|
await checkpointer.aput(put_config, tup.checkpoint, tup.metadata or {}, new_versions)
|
|
|
|
# pending_writes is a list of (task_id, channel, value); aput_writes
|
|
# wants them grouped per task as (channel, value) pairs.
|
|
if tup.pending_writes:
|
|
by_task: dict[str, list] = {}
|
|
for task_id, channel, value in tup.pending_writes:
|
|
by_task.setdefault(task_id, []).append((channel, value))
|
|
checkpoint_id = tup.checkpoint.get("id")
|
|
for task_id, writes in by_task.items():
|
|
writes_config = {
|
|
"configurable": {
|
|
"thread_id": dst_thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": checkpoint_id,
|
|
}
|
|
}
|
|
await checkpointer.aput_writes(writes_config, writes, task_id)
|
|
copied += 1
|
|
|
|
logger.info("Copied %d checkpoint(s) from thread %s to %s", copied, src_thread_id, dst_thread_id)
|
|
return copied > 0
|