Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
queue_manager.py750 linesDownload Raw Back to trees
1"""Tree-based message queue: index, async node processor, and public manager API."""2 3import asyncio4from collections.abc import Awaitable, Callable5 6from loguru import logger7 8from config.settings import get_settings9from core.anthropic import get_user_facing_error_message10 11from ..models import IncomingMessage12from ..safe_diagnostics import format_exception_for_log13from .data import MessageNode, MessageState, MessageTree14 15 16class TreeRepository:17    """18    In-memory index of trees and node-to-root mappings.19 20    Used only by :class:`TreeQueueManager`; kept as a named type for tests.21    """22 23    def __init__(self) -> None:24        self._trees: dict[str, MessageTree] = {}  # root_id -> tree25        self._node_to_tree: dict[str, str] = {}  # node_id -> root_id26 27    def get_tree(self, root_id: str) -> MessageTree | None:28        """Get a tree by its root ID."""29        return self._trees.get(root_id)30 31    def get_tree_for_node(self, node_id: str) -> MessageTree | None:32        """Get the tree containing a given node."""33        root_id = self._node_to_tree.get(node_id)34        if not root_id:35            return None36        return self._trees.get(root_id)37 38    def get_node(self, node_id: str) -> MessageNode | None:39        """Get a node from any tree."""40        tree = self.get_tree_for_node(node_id)41        return tree.get_node(node_id) if tree else None42 43    def add_tree(self, root_id: str, tree: MessageTree) -> None:44        """Add a new tree to the repository."""45        self._trees[root_id] = tree46        self._node_to_tree[root_id] = root_id47        logger.debug("TREE_REPO: add_tree root_id={}", root_id)48 49    def register_node(self, node_id: str, root_id: str) -> None:50        """Register a node ID to a tree."""51        self._node_to_tree[node_id] = root_id52        logger.debug("TREE_REPO: register_node node_id={} root_id={}", node_id, root_id)53 54    def has_node(self, node_id: str) -> bool:55        """Check if a node is registered in any tree."""56        return node_id in self._node_to_tree57 58    def tree_count(self) -> int:59        """Get the number of trees in the repository."""60        return len(self._trees)61 62    def is_tree_busy(self, root_id: str) -> bool:63        """Check if a tree is currently processing."""64        tree = self._trees.get(root_id)65        return tree.is_processing if tree else False66 67    def is_node_tree_busy(self, node_id: str) -> bool:68        """Check if the tree containing a node is busy."""69        tree = self.get_tree_for_node(node_id)70        return tree.is_processing if tree else False71 72    def get_queue_size(self, node_id: str) -> int:73        """Get queue size for the tree containing a node."""74        tree = self.get_tree_for_node(node_id)75        return tree.get_queue_size() if tree else 076 77    def resolve_parent_node_id(self, msg_id: str) -> str | None:78        """79        Resolve a message ID to the actual parent node ID.80 81        Handles the case where msg_id is a status message ID82        (which maps to the tree but isn't an actual node).83 84        Returns:85            The node_id to use as parent, or None if not found86        """87        tree = self.get_tree_for_node(msg_id)88        if not tree:89            return None90 91        if tree.has_node(msg_id):92            return msg_id93 94        node = tree.find_node_by_status_message(msg_id)95        if node:96            return node.node_id97 98        return None99 100    def get_pending_children(self, node_id: str) -> list[MessageNode]:101        """102        Get all pending child nodes (recursively) of a given node.103 104        Used for error propagation - when a node fails, its pending105        children should also be marked as failed.106        """107        tree = self.get_tree_for_node(node_id)108        if not tree:109            return []110 111        pending: list[MessageNode] = []112        stack = [node_id]113 114        while stack:115            current_id = stack.pop()116            node = tree.get_node(current_id)117            if not node:118                continue119            for child_id in node.children_ids:120                child = tree.get_node(child_id)121                if child and child.state == MessageState.PENDING:122                    pending.append(child)123                    stack.append(child_id)124 125        return pending126 127    def all_trees(self) -> list[MessageTree]:128        """Get all trees in the repository."""129        return list(self._trees.values())130 131    def tree_ids(self) -> list[str]:132        """Get all tree root IDs."""133        return list(self._trees.keys())134 135    def unregister_nodes(self, node_ids: list[str]) -> None:136        """Remove node IDs from the node-to-tree mapping."""137        for nid in node_ids:138            self._node_to_tree.pop(nid, None)139 140    def remove_tree(self, root_id: str) -> MessageTree | None:141        """142        Remove a tree and all its node mappings from the repository.143 144        Returns:145            The removed tree, or None if not found.146        """147        tree = self._trees.pop(root_id, None)148        if not tree:149            return None150        for node in tree.all_nodes():151            self._node_to_tree.pop(node.node_id, None)152        logger.debug("TREE_REPO: remove_tree root_id={}", root_id)153        return tree154 155    def get_message_ids_for_chat(self, platform: str, chat_id: str) -> set[str]:156        """Get all message IDs (incoming + status) for a given platform/chat."""157        msg_ids: set[str] = set()158        for tree in self._trees.values():159            for node in tree.all_nodes():160                if str(node.incoming.platform) == str(platform) and str(161                    node.incoming.chat_id162                ) == str(chat_id):163                    if node.incoming.message_id is not None:164                        msg_ids.add(str(node.incoming.message_id))165                    if node.status_message_id:166                        msg_ids.add(str(node.status_message_id))167        return msg_ids168 169    def to_dict(self) -> dict:170        """Serialize all trees."""171        return {172            "trees": {rid: tree.to_dict() for rid, tree in self._trees.items()},173            "node_to_tree": self._node_to_tree.copy(),174        }175 176    @classmethod177    def from_dict(cls, data: dict) -> TreeRepository:178        """Deserialize from dictionary."""179        repo = cls()180        for root_id, tree_data in data.get("trees", {}).items():181            repo._trees[root_id] = MessageTree.from_dict(tree_data)182        repo._node_to_tree = data.get("node_to_tree", {})183        return repo184 185 186class TreeQueueProcessor:187    """188    Per-tree async queue processing (one manager owns one processor instance).189    """190 191    def __init__(192        self,193        queue_update_callback: Callable[[MessageTree], Awaitable[None]] | None = None,194        node_started_callback: Callable[[MessageTree, str], Awaitable[None]]195        | None = None,196    ) -> None:197        self._queue_update_callback = queue_update_callback198        self._node_started_callback = node_started_callback199 200    def set_queue_update_callback(201        self,202        queue_update_callback: Callable[[MessageTree], Awaitable[None]] | None,203    ) -> None:204        """Update the callback used to refresh queue positions."""205        self._queue_update_callback = queue_update_callback206 207    def set_node_started_callback(208        self,209        node_started_callback: Callable[[MessageTree, str], Awaitable[None]] | None,210    ) -> None:211        """Update the callback used when a queued node starts processing."""212        self._node_started_callback = node_started_callback213 214    async def _notify_queue_updated(self, tree: MessageTree) -> None:215        """Invoke queue update callback if set."""216        if not self._queue_update_callback:217            return218        try:219            await self._queue_update_callback(tree)220        except Exception as e:221            d = get_settings().log_messaging_error_details222            logger.warning(223                "Queue update callback failed: {}",224                format_exception_for_log(e, log_full_message=d),225            )226 227    async def _notify_node_started(self, tree: MessageTree, node_id: str) -> None:228        """Invoke node started callback if set."""229        if not self._node_started_callback:230            return231        try:232            await self._node_started_callback(tree, node_id)233        except Exception as e:234            d = get_settings().log_messaging_error_details235            logger.warning(236                "Node started callback failed: {}",237                format_exception_for_log(e, log_full_message=d),238            )239 240    async def process_node(241        self,242        tree: MessageTree,243        node: MessageNode,244        processor: Callable[[str, MessageNode], Awaitable[None]],245    ) -> None:246        """Process a single node and then check the queue."""247        if node.state == MessageState.ERROR:248            logger.info(249                f"Skipping node {node.node_id} as it is already in state {node.state}"250            )251            await self._process_next(tree, processor)252            return253 254        try:255            await processor(node.node_id, node)256        except asyncio.CancelledError:257            logger.info(f"Task for node {node.node_id} was cancelled")258            raise259        except Exception as e:260            d = get_settings().log_messaging_error_details261            logger.error(262                "Error processing node {}: {}",263                node.node_id,264                format_exception_for_log(e, log_full_message=d),265            )266            await tree.update_state(267                node.node_id,268                MessageState.ERROR,269                error_message=get_user_facing_error_message(e),270            )271        finally:272            async with tree.with_lock():273                tree.clear_current_node()274            await self._process_next(tree, processor)275 276    async def _process_next(277        self,278        tree: MessageTree,279        processor: Callable[[str, MessageNode], Awaitable[None]],280    ) -> None:281        """Process the next message in queue, if any."""282        next_node_id = None283        async with tree.with_lock():284            next_node_id = await tree.dequeue()285 286            if not next_node_id:287                tree.set_processing_state(None, False)288                logger.debug(f"Tree {tree.root_id} queue empty, marking as free")289                return290 291            tree.set_processing_state(next_node_id, True)292            logger.info(f"Processing next queued node {next_node_id}")293 294            node = tree.get_node(next_node_id)295            if node:296                tree.set_current_task(297                    asyncio.create_task(self.process_node(tree, node, processor))298                )299 300        if next_node_id:301            await self._notify_node_started(tree, next_node_id)302            await self._notify_queue_updated(tree)303 304    async def enqueue_and_start(305        self,306        tree: MessageTree,307        node_id: str,308        processor: Callable[[str, MessageNode], Awaitable[None]],309    ) -> bool:310        """311        Enqueue a node or start processing immediately.312 313        Returns:314            True if queued, False if processing immediately315        """316        async with tree.with_lock():317            if tree.is_processing:318                tree.put_queue_unlocked(node_id)319                queue_size = tree.get_queue_size()320                logger.info(f"Queued node {node_id}, position {queue_size}")321                return True322            else:323                tree.set_processing_state(node_id, True)324 325                node = tree.get_node(node_id)326                if node:327                    tree.set_current_task(328                        asyncio.create_task(self.process_node(tree, node, processor))329                    )330                return False331 332    def cancel_current(self, tree: MessageTree) -> bool:333        """Cancel the currently running task in a tree."""334        return tree.cancel_current_task()335 336 337class TreeQueueManager:338    """339    Manages multiple message trees: index + async processing.340 341    Each new conversation creates a new tree.342    Replies to existing messages add nodes to existing trees.343    """344 345    def __init__(346        self,347        queue_update_callback: Callable[[MessageTree], Awaitable[None]] | None = None,348        node_started_callback: Callable[[MessageTree, str], Awaitable[None]]349        | None = None,350        _repository: TreeRepository | None = None,351    ) -> None:352        self._repository = _repository or TreeRepository()353        self._processor = TreeQueueProcessor(354            queue_update_callback=queue_update_callback,355            node_started_callback=node_started_callback,356        )357        self._lock = asyncio.Lock()358 359        logger.info("TreeQueueManager initialized")360 361    async def create_tree(362        self,363        node_id: str,364        incoming: IncomingMessage,365        status_message_id: str,366    ) -> MessageTree:367        """368        Create a new tree with a root node.369 370        Args:371            node_id: ID for the root node372            incoming: The incoming message373            status_message_id: Bot's status message ID374 375        Returns:376            The created MessageTree377        """378        async with self._lock:379            root_node = MessageNode(380                node_id=node_id,381                incoming=incoming,382                status_message_id=status_message_id,383                state=MessageState.PENDING,384            )385 386            tree = MessageTree(root_node)387            self._repository.add_tree(node_id, tree)388 389            logger.info(f"Created new tree with root {node_id}")390            return tree391 392    async def add_to_tree(393        self,394        parent_node_id: str,395        node_id: str,396        incoming: IncomingMessage,397        status_message_id: str,398    ) -> tuple[MessageTree, MessageNode]:399        """400        Add a reply as a child node to an existing tree.401 402        Args:403            parent_node_id: ID of the parent message404            node_id: ID for the new node405            incoming: The incoming reply message406            status_message_id: Bot's status message ID407 408        Returns:409            Tuple of (tree, new_node)410        """411        async with self._lock:412            if not self._repository.has_node(parent_node_id):413                raise ValueError(f"Parent node {parent_node_id} not found in any tree")414 415            tree = self._repository.get_tree_for_node(parent_node_id)416            if not tree:417                raise ValueError(f"Parent node {parent_node_id} not found in any tree")418 419        node = await tree.add_node(420            node_id=node_id,421            incoming=incoming,422            status_message_id=status_message_id,423            parent_id=parent_node_id,424        )425 426        async with self._lock:427            self._repository.register_node(node_id, tree.root_id)428 429        logger.info(f"Added node {node_id} to tree {tree.root_id}")430        return tree, node431 432    def get_tree(self, root_id: str) -> MessageTree | None:433        """Get a tree by its root ID."""434        return self._repository.get_tree(root_id)435 436    def get_tree_for_node(self, node_id: str) -> MessageTree | None:437        """Get the tree containing a given node."""438        return self._repository.get_tree_for_node(node_id)439 440    def get_node(self, node_id: str) -> MessageNode | None:441        """Get a node from any tree."""442        return self._repository.get_node(node_id)443 444    def resolve_parent_node_id(self, msg_id: str) -> str | None:445        """Resolve a message ID to the actual parent node ID."""446        return self._repository.resolve_parent_node_id(msg_id)447 448    def is_tree_busy(self, root_id: str) -> bool:449        """Check if a tree is currently processing."""450        return self._repository.is_tree_busy(root_id)451 452    def is_node_tree_busy(self, node_id: str) -> bool:453        """Check if the tree containing a node is busy."""454        return self._repository.is_node_tree_busy(node_id)455 456    async def enqueue(457        self,458        node_id: str,459        processor: Callable[[str, MessageNode], Awaitable[None]],460    ) -> bool:461        """462        Enqueue a node for processing.463 464        If the tree is not busy, processing starts immediately.465        If busy, the message is queued.466 467        Args:468            node_id: Node to process469            processor: Async function to process the node470 471        Returns:472            True if queued, False if processing immediately473        """474        tree = self._repository.get_tree_for_node(node_id)475        if not tree:476            logger.error(f"No tree found for node {node_id}")477            return False478 479        return await self._processor.enqueue_and_start(tree, node_id, processor)480 481    def get_queue_size(self, node_id: str) -> int:482        """Get queue size for the tree containing a node."""483        return self._repository.get_queue_size(node_id)484 485    def get_pending_children(self, node_id: str) -> list[MessageNode]:486        """Get all pending child nodes (recursively) of a given node."""487        return self._repository.get_pending_children(node_id)488 489    async def mark_node_error(490        self,491        node_id: str,492        error_message: str,493        propagate_to_children: bool = True,494    ) -> list[MessageNode]:495        """496        Mark a node as ERROR and optionally propagate to pending children.497 498        Args:499            node_id: The node to mark as error500            error_message: Error description501            propagate_to_children: If True, also mark pending children as error502 503        Returns:504            List of all nodes marked as error (including children)505        """506        tree = self._repository.get_tree_for_node(node_id)507        if not tree:508            return []509 510        affected = []511        node = tree.get_node(node_id)512        if node:513            await tree.update_state(514                node_id, MessageState.ERROR, error_message=error_message515            )516            affected.append(node)517 518        if propagate_to_children:519            pending_children = self._repository.get_pending_children(node_id)520            for child in pending_children:521                await tree.update_state(522                    child.node_id,523                    MessageState.ERROR,524                    error_message=f"Parent failed: {error_message}",525                )526                affected.append(child)527 528        return affected529 530    async def cancel_tree(self, root_id: str) -> list[MessageNode]:531        """532        Cancel all queued and in-progress messages in a tree.533 534        Updates node states to ERROR and returns list of affected nodes535        that were actually active or in the current processing queue.536        """537        tree = self._repository.get_tree(root_id)538        if not tree:539            return []540 541        cancelled_nodes = []542 543        cleanup_count = 0544        async with tree.with_lock():545            if tree.cancel_current_task():546                current_id = tree.current_node_id547                if current_id:548                    node = tree.get_node(current_id)549                    if node and node.state not in (550                        MessageState.COMPLETED,551                        MessageState.ERROR,552                    ):553                        tree.set_node_error_sync(node, "Cancelled by user")554                        cancelled_nodes.append(node)555 556            queue_nodes = tree.drain_queue_and_mark_cancelled()557            cancelled_nodes.extend(queue_nodes)558            cancelled_ids = {n.node_id for n in cancelled_nodes}559 560            for node in tree.all_nodes():561                if (562                    node.state in (MessageState.PENDING, MessageState.IN_PROGRESS)563                    and node.node_id not in cancelled_ids564                ):565                    tree.set_node_error_sync(node, "Stale task cleaned up")566                    cleanup_count += 1567 568            tree.reset_processing_state()569 570        if cancelled_nodes:571            logger.info(572                f"Cancelled {len(cancelled_nodes)} active nodes in tree {root_id}"573            )574        if cleanup_count:575            logger.info(f"Cleaned up {cleanup_count} stale nodes in tree {root_id}")576 577        return cancelled_nodes578 579    async def cancel_node(self, node_id: str) -> list[MessageNode]:580        """581        Cancel a single node (queued or in-progress) without affecting other nodes.582 583        Returns:584            List containing the cancelled node if it was cancellable, else empty list.585        """586        tree = self._repository.get_tree_for_node(node_id)587        if not tree:588            return []589 590        async with tree.with_lock():591            node = tree.get_node(node_id)592            if not node:593                return []594 595            if node.state in (MessageState.COMPLETED, MessageState.ERROR):596                return []597 598            if tree.is_current_node(node_id):599                self._processor.cancel_current(tree)600 601            try:602                tree.remove_from_queue(node_id)603            except Exception:604                logger.debug(605                    "Failed to remove node from queue; will rely on state=ERROR"606                )607 608            tree.set_node_error_sync(node, "Cancelled by user")609 610            return [node]611 612    async def cancel_all(self) -> list[MessageNode]:613        """Cancel all messages in all trees."""614        async with self._lock:615            root_ids = list(self._repository.tree_ids())616            all_cancelled: list[MessageNode] = []617            for root_id in root_ids:618                all_cancelled.extend(await self.cancel_tree(root_id))619            return all_cancelled620 621    def cleanup_stale_nodes(self) -> int:622        """623        Mark any PENDING or IN_PROGRESS nodes in all trees as ERROR.624        Used on startup to reconcile restored state.625        """626        count = 0627        for tree in self._repository.all_trees():628            for node in tree.all_nodes():629                if node.state in (MessageState.PENDING, MessageState.IN_PROGRESS):630                    tree.set_node_error_sync(node, "Lost during server restart")631                    count += 1632        if count:633            logger.info(f"Cleaned up {count} stale nodes during startup")634        return count635 636    def get_tree_count(self) -> int:637        """Get the number of active message trees."""638        return self._repository.tree_count()639 640    def set_queue_update_callback(641        self,642        queue_update_callback: Callable[[MessageTree], Awaitable[None]] | None,643    ) -> None:644        """Set callback for queue position updates."""645        self._processor.set_queue_update_callback(queue_update_callback)646 647    def set_node_started_callback(648        self,649        node_started_callback: Callable[[MessageTree, str], Awaitable[None]] | None,650    ) -> None:651        """Set callback for when a queued node starts processing."""652        self._processor.set_node_started_callback(node_started_callback)653 654    def register_node(self, node_id: str, root_id: str) -> None:655        """Register a node ID to a tree (for external mapping)."""656        self._repository.register_node(node_id, root_id)657 658    async def cancel_branch(self, branch_root_id: str) -> list[MessageNode]:659        """660        Cancel all PENDING/IN_PROGRESS nodes in the subtree (branch_root + descendants).661        """662        tree = self._repository.get_tree_for_node(branch_root_id)663        if not tree:664            return []665 666        branch_ids = set(tree.get_descendants(branch_root_id))667        cancelled: list[MessageNode] = []668 669        async with tree.with_lock():670            for nid in branch_ids:671                node = tree.get_node(nid)672                if not node or node.state in (673                    MessageState.COMPLETED,674                    MessageState.ERROR,675                ):676                    continue677 678                if tree.is_current_node(nid):679                    self._processor.cancel_current(tree)680                    tree.set_node_error_sync(node, "Cancelled by user")681                    cancelled.append(node)682                else:683                    tree.remove_from_queue(nid)684                    tree.set_node_error_sync(node, "Cancelled by user")685                    cancelled.append(node)686 687        if cancelled:688            logger.info(f"Cancelled {len(cancelled)} nodes in branch {branch_root_id}")689        return cancelled690 691    async def remove_branch(692        self, branch_root_id: str693    ) -> tuple[list[MessageNode], str, bool]:694        """695        Remove a branch (subtree) from the tree.696 697        If branch_root is the tree root, removes the entire tree.698 699        Returns:700            (removed_nodes, root_id, removed_entire_tree)701        """702        tree = self._repository.get_tree_for_node(branch_root_id)703        if not tree:704            return ([], "", False)705 706        root_id = tree.root_id707 708        if branch_root_id == root_id:709            cancelled = await self.cancel_tree(root_id)710            removed_tree = self._repository.remove_tree(root_id)711            if removed_tree:712                return (removed_tree.all_nodes(), root_id, True)713            return (cancelled, root_id, True)714 715        async with tree.with_lock():716            removed = tree.remove_branch(branch_root_id)717 718        self._repository.unregister_nodes([n.node_id for n in removed])719        return (removed, root_id, False)720 721    def get_message_ids_for_chat(self, platform: str, chat_id: str) -> set[str]:722        """Get all message IDs for a given platform/chat."""723        return self._repository.get_message_ids_for_chat(platform, chat_id)724 725    def to_dict(self) -> dict:726        """Serialize all trees."""727        return self._repository.to_dict()728 729    @classmethod730    def from_dict(731        cls,732        data: dict,733        queue_update_callback: Callable[[MessageTree], Awaitable[None]] | None = None,734        node_started_callback: Callable[[MessageTree, str], Awaitable[None]]735        | None = None,736    ) -> TreeQueueManager:737        """Deserialize from dictionary."""738        return cls(739            queue_update_callback=queue_update_callback,740            node_started_callback=node_started_callback,741            _repository=TreeRepository.from_dict(data),742        )743 744 745__all__ = [746    "TreeQueueManager",747    "TreeQueueProcessor",748    "TreeRepository",749]750