Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
session.py306 linesDownload Raw Back to messaging
1"""2Session Store for Messaging Platforms3 4Provides persistent storage for mapping platform messages to Claude CLI session IDs5and message trees for conversation continuation.6"""7 8import contextlib9import json10import os11import tempfile12import threading13from datetime import UTC, datetime14from typing import Any15 16from loguru import logger17 18 19class SessionStore:20    """21    Persistent storage for message ↔ Claude session mappings and message trees.22 23    Uses a JSON file for storage with thread-safe operations.24    Platform-agnostic: works with any messaging platform.25    """26 27    def __init__(28        self,29        storage_path: str = "sessions.json",30        *,31        message_log_cap: int | None = None,32    ):33        self.storage_path = storage_path34        self._lock = threading.Lock()35        self._trees: dict[str, dict] = {}  # root_id -> tree data36        self._node_to_tree: dict[str, str] = {}  # node_id -> root_id37        # Per-chat message ID log used to support best-effort UI clearing (/clear).38        # Key: "{platform}:{chat_id}" -> list of records39        self._message_log: dict[str, list[dict[str, Any]]] = {}40        self._message_log_ids: dict[str, set[str]] = {}41        self._dirty = False42        self._save_timer: threading.Timer | None = None43        self._save_debounce_secs = 0.544        self._message_log_cap: int | None = message_log_cap45        self._load()46 47    def _make_chat_key(self, platform: str, chat_id: str) -> str:48        return f"{platform}:{chat_id}"49 50    def _load(self) -> None:51        """Load sessions and trees from disk."""52        if not os.path.exists(self.storage_path):53            return54 55        try:56            with open(self.storage_path, encoding="utf-8") as f:57                data = json.load(f)58 59            # Load trees60            self._trees = data.get("trees", {})61            self._node_to_tree = data.get("node_to_tree", {})62 63            # Load message log (optional/backward compatible)64            raw_log = data.get("message_log", {}) or {}65            if isinstance(raw_log, dict):66                self._message_log = {}67                self._message_log_ids = {}68                for chat_key, items in raw_log.items():69                    if not isinstance(chat_key, str) or not isinstance(items, list):70                        continue71                    cleaned: list[dict[str, Any]] = []72                    seen: set[str] = set()73                    for it in items:74                        if not isinstance(it, dict):75                            continue76                        mid = it.get("message_id")77                        if mid is None:78                            continue79                        mid_s = str(mid)80                        if mid_s in seen:81                            continue82                        seen.add(mid_s)83                        cleaned.append(84                            {85                                "message_id": mid_s,86                                "ts": str(it.get("ts") or ""),87                                "direction": str(it.get("direction") or ""),88                                "kind": str(it.get("kind") or ""),89                            }90                        )91                    self._message_log[chat_key] = cleaned92                    self._message_log_ids[chat_key] = seen93 94            logger.info(95                f"Loaded {len(self._trees)} trees and "96                f"{sum(len(v) for v in self._message_log.values())} msg_ids from {self.storage_path}"97            )98        except Exception as e:99            logger.error(f"Failed to load sessions: {e}")100 101    def _snapshot(self) -> dict:102        """Snapshot current state for serialization. Caller must hold self._lock."""103        return {104            "trees": dict(self._trees),105            "node_to_tree": dict(self._node_to_tree),106            "message_log": {k: list(v) for k, v in self._message_log.items()},107        }108 109    def _write_data(self, data: dict) -> None:110        """Atomically write data dict to disk. Must be called WITHOUT holding self._lock."""111        abs_target = os.path.abspath(self.storage_path)112        dir_name = os.path.dirname(abs_target) or "."113        fd, tmp_path = tempfile.mkstemp(114            dir=dir_name, prefix=".sessions.", suffix=".tmp.json"115        )116        try:117            with os.fdopen(fd, "w", encoding="utf-8") as f:118                json.dump(data, f, indent=2)119                f.flush()120                os.fsync(f.fileno())121            os.replace(tmp_path, abs_target)122        except BaseException:123            with contextlib.suppress(OSError):124                os.unlink(tmp_path)125            raise126 127    def _schedule_save(self) -> None:128        """Schedule a debounced save. Caller must hold self._lock."""129        self._dirty = True130        if self._save_timer is not None:131            self._save_timer.cancel()132            self._save_timer = None133        self._save_timer = threading.Timer(134            self._save_debounce_secs, self._save_from_timer135        )136        self._save_timer.daemon = True137        self._save_timer.start()138 139    def _save_from_timer(self) -> None:140        """Timer callback: save if dirty. Runs in timer thread."""141        with self._lock:142            if not self._dirty:143                self._save_timer = None144                return145            snapshot = self._snapshot()146            self._dirty = False147            self._save_timer = None148        try:149            self._write_data(snapshot)150        except Exception as e:151            logger.error(f"Failed to save sessions: {e}")152            with self._lock:153                self._dirty = True154 155    def _flush_save(self) -> dict:156        """Cancel pending timer and snapshot current state. Caller must hold self._lock.157        Returns snapshot dict; caller must call _write_data(snapshot) outside the lock."""158        if self._save_timer is not None:159            self._save_timer.cancel()160            self._save_timer = None161        self._dirty = False162        return self._snapshot()163 164    def flush_pending_save(self) -> None:165        """Flush any pending debounced save. Call on shutdown to avoid losing data."""166        with self._lock:167            snapshot = self._flush_save()168        try:169            self._write_data(snapshot)170        except Exception as e:171            logger.error(f"Failed to save sessions: {e}")172            with self._lock:173                self._dirty = True174 175    def record_message_id(176        self,177        platform: str,178        chat_id: str,179        message_id: str,180        direction: str,181        kind: str,182    ) -> None:183        """Record a message_id for later best-effort deletion (/clear)."""184        if message_id is None:185            return186 187        chat_key = self._make_chat_key(str(platform), str(chat_id))188        mid = str(message_id)189 190        with self._lock:191            seen = self._message_log_ids.setdefault(chat_key, set())192            if mid in seen:193                return194 195            rec = {196                "message_id": mid,197                "ts": datetime.now(UTC).isoformat(),198                "direction": str(direction),199                "kind": str(kind),200            }201            self._message_log.setdefault(chat_key, []).append(rec)202            seen.add(mid)203 204            # Optional cap to prevent unbounded growth if configured.205            if self._message_log_cap is not None and self._message_log_cap > 0:206                items = self._message_log.get(chat_key, [])207                if len(items) > self._message_log_cap:208                    self._message_log[chat_key] = items[-self._message_log_cap :]209                    self._message_log_ids[chat_key] = {210                        str(x.get("message_id")) for x in self._message_log[chat_key]211                    }212 213            self._schedule_save()214 215    def get_message_ids_for_chat(self, platform: str, chat_id: str) -> list[str]:216        """Get all recorded message IDs for a chat (in insertion order)."""217        chat_key = self._make_chat_key(str(platform), str(chat_id))218        with self._lock:219            items = self._message_log.get(chat_key, [])220            return [221                str(x.get("message_id"))222                for x in items223                if x.get("message_id") is not None224            ]225 226    def clear_all(self) -> None:227        """Clear all stored sessions/trees/mappings and persist an empty store."""228        with self._lock:229            self._trees.clear()230            self._node_to_tree.clear()231            self._message_log.clear()232            self._message_log_ids.clear()233            snapshot = self._flush_save()234        try:235            self._write_data(snapshot)236        except Exception as e:237            logger.error(f"Failed to save sessions: {e}")238            with self._lock:239                self._dirty = True240 241    # ==================== Tree Methods ====================242 243    def save_tree(self, root_id: str, tree_data: dict) -> None:244        """245        Save a message tree.246 247        Args:248            root_id: Root node ID of the tree249            tree_data: Serialized tree data from tree.to_dict()250        """251        with self._lock:252            self._trees[root_id] = tree_data253 254            # Update node-to-tree mapping255            for node_id in tree_data.get("nodes", {}):256                self._node_to_tree[node_id] = root_id257 258            self._schedule_save()259            logger.debug(f"Saved tree {root_id}")260 261    def get_tree(self, root_id: str) -> dict | None:262        """Get a tree by its root ID."""263        with self._lock:264            return self._trees.get(root_id)265 266    def register_node(self, node_id: str, root_id: str) -> None:267        """Register a node ID to a tree root."""268        with self._lock:269            self._node_to_tree[node_id] = root_id270            self._schedule_save()271 272    def remove_node_mappings(self, node_ids: list[str]) -> None:273        """Remove node IDs from the node-to-tree mapping."""274        with self._lock:275            for nid in node_ids:276                self._node_to_tree.pop(nid, None)277            self._schedule_save()278 279    def remove_tree(self, root_id: str) -> None:280        """Remove a tree and all its node mappings from the store."""281        with self._lock:282            tree_data = self._trees.pop(root_id, None)283            if tree_data:284                for node_id in tree_data.get("nodes", {}):285                    self._node_to_tree.pop(node_id, None)286                self._schedule_save()287 288    def get_all_trees(self) -> dict[str, dict]:289        """Get all stored trees (public accessor)."""290        with self._lock:291            return dict(self._trees)292 293    def get_node_mapping(self) -> dict[str, str]:294        """Get the node-to-tree mapping (public accessor)."""295        with self._lock:296            return dict(self._node_to_tree)297 298    def sync_from_tree_data(299        self, trees: dict[str, dict], node_to_tree: dict[str, str]300    ) -> None:301        """Sync internal tree state from external data and persist."""302        with self._lock:303            self._trees = trees304            self._node_to_tree = node_to_tree305            self._schedule_save()306