Yash030/claude-code-proxy
2
1"""Tree data structures for message queue.2 3Contains MessageState, MessageNode, and MessageTree classes.4"""5 6import asyncio7from collections import deque8from contextlib import asynccontextmanager9from dataclasses import dataclass, field10from datetime import UTC, datetime11from enum import Enum12from typing import Any13 14from loguru import logger15 16from ..models import IncomingMessage17 18 19class _SnapshotQueue:20 """Queue with snapshot/remove helpers, backed by a deque and a set index."""21 22 def __init__(self) -> None:23 self._deque: deque[str] = deque()24 self._set: set[str] = set()25 26 async def put(self, item: str) -> None:27 self._deque.append(item)28 self._set.add(item)29 30 def put_nowait(self, item: str) -> None:31 self._deque.append(item)32 self._set.add(item)33 34 def get_nowait(self) -> str:35 if not self._deque:36 raise asyncio.QueueEmpty()37 item = self._deque.popleft()38 self._set.discard(item)39 return item40 41 def qsize(self) -> int:42 return len(self._deque)43 44 def get_snapshot(self) -> list[str]:45 """Return current queue contents in FIFO order (read-only copy)."""46 return list(self._deque)47 48 def remove_if_present(self, item: str) -> bool:49 """Remove item from queue if present (O(1) membership check). Returns True if removed."""50 if item not in self._set:51 return False52 self._set.discard(item)53 self._deque = deque(x for x in self._deque if x != item)54 return True55 56 57class MessageState(Enum):58 """State of a message node in the tree."""59 60 PENDING = "pending" # Queued, waiting to be processed61 IN_PROGRESS = "in_progress" # Currently being processed by Claude62 COMPLETED = "completed" # Processing finished successfully63 ERROR = "error" # Processing failed64 65 66@dataclass67class MessageNode:68 """69 A node in the message tree.70 71 Each node represents a single message and tracks:72 - Its relationship to parent/children73 - Its processing state74 - Claude session information75 """76 77 node_id: str # Unique ID (typically message_id)78 incoming: IncomingMessage # The original message79 status_message_id: str # Bot's status message ID80 state: MessageState = MessageState.PENDING81 parent_id: str | None = None # Parent node ID (None for root)82 session_id: str | None = None # Claude session ID (forked from parent)83 children_ids: list[str] = field(default_factory=list)84 created_at: datetime = field(default_factory=lambda: datetime.now(UTC))85 completed_at: datetime | None = None86 error_message: str | None = None87 context: Any = None # Additional context if needed88 89 def set_context(self, context: Any) -> None:90 self.context = context91 92 def to_dict(self) -> dict:93 """Convert to dictionary for JSON serialization."""94 return {95 "node_id": self.node_id,96 "incoming": {97 "text": self.incoming.text,98 "chat_id": self.incoming.chat_id,99 "user_id": self.incoming.user_id,100 "message_id": self.incoming.message_id,101 "platform": self.incoming.platform,102 "reply_to_message_id": self.incoming.reply_to_message_id,103 "message_thread_id": self.incoming.message_thread_id,104 "username": self.incoming.username,105 },106 "status_message_id": self.status_message_id,107 "state": self.state.value,108 "parent_id": self.parent_id,109 "session_id": self.session_id,110 "children_ids": self.children_ids,111 "created_at": self.created_at.isoformat(),112 "completed_at": self.completed_at.isoformat()113 if self.completed_at114 else None,115 "error_message": self.error_message,116 }117 118 @classmethod119 def from_dict(cls, data: dict) -> MessageNode:120 """Create from dictionary (JSON deserialization)."""121 incoming_data = data["incoming"]122 incoming = IncomingMessage(123 text=incoming_data["text"],124 chat_id=incoming_data["chat_id"],125 user_id=incoming_data["user_id"],126 message_id=incoming_data["message_id"],127 platform=incoming_data["platform"],128 reply_to_message_id=incoming_data.get("reply_to_message_id"),129 message_thread_id=incoming_data.get("message_thread_id"),130 username=incoming_data.get("username"),131 )132 return cls(133 node_id=data["node_id"],134 incoming=incoming,135 status_message_id=data["status_message_id"],136 state=MessageState(data["state"]),137 parent_id=data.get("parent_id"),138 session_id=data.get("session_id"),139 children_ids=data.get("children_ids", []),140 created_at=datetime.fromisoformat(data["created_at"]),141 completed_at=datetime.fromisoformat(data["completed_at"])142 if data.get("completed_at")143 else None,144 error_message=data.get("error_message"),145 )146 147 148class MessageTree:149 """150 A tree of message nodes with queue functionality.151 152 Provides:153 - O(1) node lookup via hashmap154 - Per-tree message queue155 - Thread-safe operations via asyncio.Lock156 """157 158 def __init__(self, root_node: MessageNode):159 """160 Initialize tree with a root node.161 162 Args:163 root_node: The root message node164 """165 self.root_id = root_node.node_id166 self._nodes: dict[str, MessageNode] = {root_node.node_id: root_node}167 self._status_to_node: dict[str, str] = {168 root_node.status_message_id: root_node.node_id169 }170 self._queue: _SnapshotQueue = _SnapshotQueue()171 self._lock = asyncio.Lock()172 self._is_processing = False173 self._current_node_id: str | None = None174 self._current_task: asyncio.Task | None = None175 176 logger.debug(f"Created MessageTree with root {self.root_id}")177 178 def set_current_task(self, task: asyncio.Task | None) -> None:179 """Set the current processing task. Caller must hold lock."""180 self._current_task = task181 182 @property183 def is_processing(self) -> bool:184 """Check if tree is currently processing a message."""185 return self._is_processing186 187 async def add_node(188 self,189 node_id: str,190 incoming: IncomingMessage,191 status_message_id: str,192 parent_id: str,193 ) -> MessageNode:194 """195 Add a child node to the tree.196 197 Args:198 node_id: Unique ID for the new node199 incoming: The incoming message200 status_message_id: Bot's status message ID201 parent_id: Parent node ID202 203 Returns:204 The created MessageNode205 """206 async with self._lock:207 if parent_id not in self._nodes:208 raise ValueError(f"Parent node {parent_id} not found in tree")209 210 node = MessageNode(211 node_id=node_id,212 incoming=incoming,213 status_message_id=status_message_id,214 parent_id=parent_id,215 state=MessageState.PENDING,216 )217 218 self._nodes[node_id] = node219 self._status_to_node[status_message_id] = node_id220 self._nodes[parent_id].children_ids.append(node_id)221 222 logger.debug(f"Added node {node_id} as child of {parent_id}")223 return node224 225 def get_node(self, node_id: str) -> MessageNode | None:226 """Get a node by ID (O(1) lookup)."""227 return self._nodes.get(node_id)228 229 def get_root(self) -> MessageNode:230 """Get the root node."""231 return self._nodes[self.root_id]232 233 def get_children(self, node_id: str) -> list[MessageNode]:234 """Get all child nodes of a given node."""235 node = self._nodes.get(node_id)236 if not node:237 return []238 return [self._nodes[cid] for cid in node.children_ids if cid in self._nodes]239 240 def get_parent(self, node_id: str) -> MessageNode | None:241 """Get the parent node."""242 node = self._nodes.get(node_id)243 if not node or not node.parent_id:244 return None245 return self._nodes.get(node.parent_id)246 247 def get_parent_session_id(self, node_id: str) -> str | None:248 """249 Get the parent's session ID for forking.250 251 Returns None for root nodes.252 """253 parent = self.get_parent(node_id)254 return parent.session_id if parent else None255 256 async def update_state(257 self,258 node_id: str,259 state: MessageState,260 session_id: str | None = None,261 error_message: str | None = None,262 ) -> None:263 """Update a node's state."""264 async with self._lock:265 node = self._nodes.get(node_id)266 if not node:267 logger.warning(f"Node {node_id} not found for state update")268 return269 270 node.state = state271 if session_id:272 node.session_id = session_id273 if error_message:274 node.error_message = error_message275 if state in (MessageState.COMPLETED, MessageState.ERROR):276 node.completed_at = datetime.now(UTC)277 278 logger.debug(f"Node {node_id} state -> {state.value}")279 280 async def enqueue(self, node_id: str) -> int:281 """282 Add a node to the processing queue.283 284 Returns:285 Queue position (1-indexed)286 """287 async with self._lock:288 await self._queue.put(node_id)289 position = self._queue.qsize()290 logger.debug(f"Enqueued node {node_id}, position {position}")291 return position292 293 async def dequeue(self) -> str | None:294 """295 Get the next node ID from the queue.296 297 Returns None if queue is empty.298 """299 try:300 return self._queue.get_nowait()301 except asyncio.QueueEmpty:302 return None303 304 async def get_queue_snapshot(self) -> list[str]:305 """306 Get a snapshot of the current queue order.307 308 Returns:309 List of node IDs in FIFO order.310 """311 async with self._lock:312 return self._queue.get_snapshot()313 314 def get_queue_size(self) -> int:315 """Get number of messages waiting in queue."""316 return self._queue.qsize()317 318 def remove_from_queue(self, node_id: str) -> bool:319 """320 Remove node_id from the internal queue if present.321 322 Caller must hold the tree lock (e.g. via with_lock).323 Returns True if node was removed, False if not in queue.324 """325 return self._queue.remove_if_present(node_id)326 327 @asynccontextmanager328 async def with_lock(self):329 """Async context manager for tree lock. Use when multiple operations need atomicity."""330 async with self._lock:331 yield332 333 def set_processing_state(self, node_id: str | None, is_processing: bool) -> None:334 """Set processing state. Caller must hold lock for consistency with queue operations."""335 self._is_processing = is_processing336 self._current_node_id = node_id if is_processing else None337 338 def clear_current_node(self) -> None:339 """Clear the currently processing node ID. Caller must hold lock."""340 self._current_node_id = None341 342 def is_current_node(self, node_id: str) -> bool:343 """Check if node_id is the currently processing node."""344 return self._current_node_id == node_id345 346 def put_queue_unlocked(self, node_id: str) -> None:347 """Add node to queue. Caller must hold lock (e.g. via with_lock)."""348 self._queue.put_nowait(node_id)349 350 def cancel_current_task(self) -> bool:351 """Cancel the currently running task. Returns True if a task was cancelled."""352 if self._current_task and not self._current_task.done():353 self._current_task.cancel()354 return True355 return False356 357 def set_node_error_sync(self, node: MessageNode, error_message: str) -> None:358 """Synchronously mark a node as ERROR. Caller must ensure no concurrent access."""359 node.state = MessageState.ERROR360 node.error_message = error_message361 node.completed_at = datetime.now(UTC)362 363 def drain_queue_and_mark_cancelled(364 self, error_message: str = "Cancelled by user"365 ) -> list[MessageNode]:366 """367 Drain the queue, mark each node as ERROR, and return affected nodes.368 Does not acquire lock; caller must ensure no concurrent queue access.369 """370 nodes: list[MessageNode] = []371 while True:372 try:373 node_id = self._queue.get_nowait()374 except asyncio.QueueEmpty:375 break376 node = self._nodes.get(node_id)377 if node:378 self.set_node_error_sync(node, error_message)379 nodes.append(node)380 return nodes381 382 def reset_processing_state(self) -> None:383 """Reset processing flags after cancel/cleanup."""384 self._is_processing = False385 self._current_node_id = None386 387 @property388 def current_node_id(self) -> str | None:389 """Get the ID of the node currently being processed."""390 return self._current_node_id391 392 def to_dict(self) -> dict:393 """Serialize tree to dictionary."""394 return {395 "root_id": self.root_id,396 "nodes": {nid: node.to_dict() for nid, node in self._nodes.items()},397 }398 399 def _add_node_from_dict(self, node: MessageNode) -> None:400 """Register a deserialized node into the tree's internal indices."""401 self._nodes[node.node_id] = node402 self._status_to_node[node.status_message_id] = node.node_id403 404 @classmethod405 def from_dict(cls, data: dict) -> MessageTree:406 """Deserialize tree from dictionary."""407 root_id = data["root_id"]408 nodes_data = data["nodes"]409 410 # Create root node first411 root_node = MessageNode.from_dict(nodes_data[root_id])412 tree = cls(root_node)413 414 # Add remaining nodes and build status->node index415 for node_id, node_data in nodes_data.items():416 if node_id != root_id:417 node = MessageNode.from_dict(node_data)418 tree._add_node_from_dict(node)419 420 return tree421 422 def all_nodes(self) -> list[MessageNode]:423 """Get all nodes in the tree."""424 return list(self._nodes.values())425 426 def has_node(self, node_id: str) -> bool:427 """Check if a node exists in this tree."""428 return node_id in self._nodes429 430 def find_node_by_status_message(self, status_msg_id: str) -> MessageNode | None:431 """Find the node that has this status message ID (O(1) lookup)."""432 node_id = self._status_to_node.get(status_msg_id)433 return self._nodes.get(node_id) if node_id else None434 435 def get_descendants(self, node_id: str) -> list[str]:436 """437 Get node_id and all descendant IDs (subtree).438 439 Returns:440 List of node IDs including the given node.441 """442 if node_id not in self._nodes:443 return []444 result: list[str] = []445 stack = [node_id]446 while stack:447 nid = stack.pop()448 result.append(nid)449 node = self._nodes.get(nid)450 if node:451 stack.extend(node.children_ids)452 return result453 454 def remove_branch(self, branch_root_id: str) -> list[MessageNode]:455 """456 Remove a subtree (branch_root and all descendants) from the tree.457 458 Updates parent's children_ids. Caller must hold lock for consistency.459 Does not acquire lock internally.460 461 Returns:462 List of removed nodes.463 """464 if branch_root_id not in self._nodes:465 return []466 467 parent = self.get_parent(branch_root_id)468 removed = []469 for nid in self.get_descendants(branch_root_id):470 node = self._nodes.get(nid)471 if node:472 removed.append(node)473 del self._nodes[nid]474 del self._status_to_node[node.status_message_id]475 476 if parent and branch_root_id in parent.children_ids:477 parent.children_ids = [478 c for c in parent.children_ids if c != branch_root_id479 ]480 481 logger.debug(f"Removed branch {branch_root_id} ({len(removed)} nodes)")482 return removed483 