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