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