Yash030/claude-code-proxy
2
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 