Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
data.py483 linesDownload Raw Back to trees
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