deerflow-code/offline-backend-20260512/backend/app/gateway/thread_copy.py
2026-09-07 18:24:55 +08:00

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