From e25360e30ac07dcedf4d7c5d65a67e96dd825e4a Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Mon, 14 Sep 2026 07:58:20 -0400 Subject: [PATCH 01/14] Automatically background auth-like commands to allow users to enter passwords, fix auto-PTY on linux background commands --- cecli/tools/command.py | 68 ++++++++++++++++++++++++++++++++++++++---- 1 file changed, 62 insertions(+), 6 deletions(-) diff --git a/cecli/tools/command.py b/cecli/tools/command.py index 6fd9c3be2aa..1177850bd4b 100644 --- a/cecli/tools/command.py +++ b/cecli/tools/command.py @@ -2,6 +2,7 @@ import fnmatch import os import platform +import re # PTY support for interactive commands (avoids pipe buffering issues) try: @@ -22,6 +23,40 @@ from cecli.tools.utils.responses import ToolResponse from cecli.tools.validations import ToolValidations +# Commands an LLM is likely to run during development where the user must +# type input for the task to proceed (passwords, passphrases, host-key +# confirmations, credential logins, editor handoffs). Matching commands run +# with user_input_required=True so the user can respond and the command can +# complete. Patterns are matched case-insensitively against the full command. +# +# Long-running commands are already moved to the background by the timeout +# mechanism, and session-style tools (ssh, su, database REPLs) read from +# stdin when backgrounded, so they need no special handling here. +# +# Read-only viewers and live-monitoring tools (less, more, man, top, htop, +# watch, tail -f, docker logs -f) are intentionally excluded: an LLM would +# run those in the background to watch output rather than wait for typed +# input, so forcing interactivity would block the background use case. +INTERACTIVE_COMMAND_PATTERNS = [ + # Privilege escalation, user switching, and password entry + r"^\s*(sudo|doas|runas|passwd)\b", + # Remote access: passwords, key passphrases, host-key confirmations + r"^\s*(scp|rsync|ssh-keygen|ssh-add|ssh-copy-id)\b", + # Passphrase prompts (gpg is also used for commit signing) + r"^\s*(gpg|gpg2)\b", + r"^\s*openssl\s+(enc|pkcs12|pkey|genpkey|rsa|genrsa|req)\b", + # Interactive credential / login flows + r"^\s*(gh|docker|npm|yarn|pnpm|az|aws|gcloud|heroku|firebase|vercel|netlify)\s+(auth|login|logout|configure|sso)\b", + # Editors: the user must edit content for the task to proceed + r"^\s*(vi|vim|nvim|nano|emacs|pico)\b", + # Git flows that hand off to an editor or need hunk-by-hunk input + r"^\s*git\s+(add\s+(-p|--patch)|commit\s+(-e|--edit)|rebase\s+(-i|--interactive)|mergetool|config\s+(-e|--edit))\b", + # Windows credential / remote-execution tools + r"^\s*(net\s+use|get-credential|psexec|cmdkey)\b", + # Config editors that open an interactive editor + r"^\s*(crontab\s+-e|visudo)\b", +] + class Tool(BaseTool): NORM_NAME = "command" @@ -66,17 +101,19 @@ class Tool(BaseTool): "type": "string", "description": ( "Input to send. Use with background=True to send at " - "start time, or with background_key + action='stdin'." + "start time, or with background_key + action='stdin'. " + "End the input with a newline to submit a line to an " + "interactive prompt." ), }, "pty": { "type": "boolean", "description": ( - "Use a pseudo-terminal (PTY). Auto-enabled on Unix for " - "background commands. Useful for interactive programs " - "like 'vi' or 'top'." + "Use a pseudo-terminal (PTY). Auto-enabled on Unix " + "when omitted; set false to force pipe mode. A PTY lets " + "you send stdin to long-running background commands." ), - "default": False, + "default": None, }, "user_input_required": { "type": "boolean", @@ -120,7 +157,7 @@ async def execute( background_key=None, action=None, stdin=None, - pty=False, + pty=None, user_input_required=False, timeout=0, **kwargs, @@ -129,6 +166,8 @@ async def execute( Execute a shell command or interact with background processes. For new commands: provide 'command' (and optionally 'background', 'stdin', 'pty'). + PTY is auto-enabled on Unix when 'pty' is omitted, so long-running + backgrounded commands can receive input via background_key + action='stdin'. When 'user_input_required' is True, runs the command interactively using a pseudo-terminal (PTY), allowing the user to provide inputs like passwords or navigate terminal interfaces. @@ -182,6 +221,12 @@ async def execute( background = True command = command.strip()[:-1].strip() + # Force interactive handling for commands known to prompt for input + # (e.g. sudo, passphrase, credential, and editor prompts) so the user + # can respond and the command can complete. + if cls._requires_user_input(command): + user_input_required = True + # Get user confirmation confirmed = await cls._get_confirmation(coder, command, background) if not confirmed: @@ -662,6 +707,17 @@ async def _handle_errors(cls, coder, command_string, e): response.append_error(f"Error executing command: {str(e)}") return response + @classmethod + def _requires_user_input(cls, command_string): + """Return True if command matches a known interactive-input pattern.""" + if not command_string: + return False + + return any( + re.search(pattern, command_string, re.IGNORECASE) + for pattern in INTERACTIVE_COMMAND_PATTERNS + ) + @classmethod def format_output(cls, coder, mcp_server, tool_response): """Format output for Command tool.""" From e80051cce1ffbd15803ea1f7c7c6d1729be67e4c Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Mon, 14 Sep 2026 20:11:08 -0400 Subject: [PATCH 02/14] - Update background command injection to be page based with a faster debounce cycle than other post messages - Allow models to intentionally pull the latest command arguments - Automatically require user input for auth-like commands like `sudo` --- cecli/coders/agent_coder.py | 67 +++++-- cecli/helpers/background_commands.py | 127 ++++++++++++- cecli/helpers/conversation/integration.py | 80 ++++++++- cecli/tools/command.py | 126 +++++++++---- tests/basic/test_background_commands.py | 167 +++++++++++++++++ .../test_background_command_injection.py | 170 ++++++++++++++++++ tests/tools/test_command_timeout_paging.py | 104 +++++------ tests/tools/test_resource_manager_paging.py | 24 ++- 8 files changed, 742 insertions(+), 123 deletions(-) create mode 100644 tests/conversations/test_background_command_injection.py diff --git a/cecli/coders/agent_coder.py b/cecli/coders/agent_coder.py index 4384283dea3..9a5a9d04cb2 100644 --- a/cecli/coders/agent_coder.py +++ b/cecli/coders/agent_coder.py @@ -652,6 +652,10 @@ def format_chat_chunks(self): # Add post-message context blocks (priority 250 - between CUR and REMINDER) ConversationService.get_chunks(self).add_post_message_context_blocks() + # Background command output is debounced independently so it is not + # re-dumped every turn as its contents mutate slightly + ConversationService.get_chunks(self).add_background_command_output() + # Add sub-agent states context block (same priority as post-message blocks) ConversationService.get_chunks(self).add_sub_agent_states() @@ -1845,25 +1849,44 @@ def get_background_command_output(self): """ Get background command output to append after the main message. - Returns: - String containing formatted background command output, or empty string if none + Emits a roster of active commands (keeping command keys in context) + plus any new incremental output. Output that has been paged to disk is + advertised by command key and read on demand through ``ResourceManager`` + paging. """ - # Get output from all running background commands - bg_outputs = BackgroundCommandManager.get_all_command_outputs(clear=True) + command_info = BackgroundCommandManager.list_background_commands() - if not bg_outputs: + if not command_info: return "" - # Get command info to show actual command strings - command_info = BackgroundCommandManager.list_background_commands() + new_outputs = {} + for command_key in command_info: + output = BackgroundCommandManager.get_new_command_output(command_key) + if output.strip(): + new_outputs[command_key] = output - # Create formatted output for background commands output = "--- Background Commands Output ---\n" - for command_key, cmd_output in bg_outputs.items(): - if cmd_output.strip(): # Only add if there's output - # Get the actual command string if available - command_str = command_info.get(command_key, {}).get("command", command_key) - output += f"\n[bg: {command_str}]\n{cmd_output}\n" + output += "Commands:\n" + + paged_keys = [] + for command_key, info in sorted(command_info.items()): + status = "running" if info.get("running", False) else "finished" + pages = info.get("pages", 0) + page_note = f"pages 1-{pages}" if pages else "no pages yet" + output += ( + f"- {command_key} [{status}] `{info.get('command', command_key)}`" + f" — {info.get('total_chars', 0):,} chars, {page_note}\n" + ) + if pages: + paged_keys.append(command_key) + + for command_key, cmd_output in new_outputs.items(): + output += f"\nNew output ({command_key}):\n{cmd_output}\n" + + if paged_keys: + output += "\nPaged output is available via `ResourceManager` (up to 3 pages):\n" + for command_key in paged_keys: + output += f'{{"paging": [{{"target": "{command_key}", "page": 1}}]}}\n' # Clean up stale (finished) background commands after reading their output for command_key, info in command_info.items(): @@ -1872,6 +1895,24 @@ def get_background_command_output(self): return output + def get_background_command_state(self): + """ + Return a lightweight status snapshot of tracked background commands. + + Maps each command key to its running state and page count so the + injection layer can detect finish and page-flush transitions cheaply + without consuming any output. + """ + command_info = BackgroundCommandManager.list_background_commands() + + return { + key: { + "running": bool(info.get("running", False)), + "pages": info.get("pages", 0), + } + for key, info in command_info.items() + } + def get_git_status(self): """ Generate a git status context block for repository information. diff --git a/cecli/helpers/background_commands.py b/cecli/helpers/background_commands.py index a698ac1313c..1a7f846558c 100644 --- a/cecli/helpers/background_commands.py +++ b/cecli/helpers/background_commands.py @@ -106,6 +106,110 @@ def size(self) -> int: return len(self.buffer) +class PagedOutputBuffer: + """ + Thread-safe output window that spills full pages to disk. + + Output accumulates in memory until it reaches ``page_size`` characters. + Full pages are written to ``pages_dir`` as ``{n}.txt`` and dropped from + memory, so the in-memory window never grows beyond ``page_size``. Unlike + ``CircularBuffer``, ``total_added`` is a monotonic stream offset that is + never reset, so incremental readers can resume after content has been + flushed; anything older than the in-memory window lives on disk. + """ + + def __init__(self, page_size: int = 4096, pages_dir: Optional[str] = None): + self.page_size = max(1, int(page_size)) + self.pages_dir = pages_dir + self.buffer = deque() + self.lock = threading.Lock() + self.total_added = 0 + self.window_start = 0 + self.page_count = 0 + + def append(self, text: str) -> None: + """Append text, spilling full pages to disk once the window fills.""" + if not text: + return + + with self.lock: + self.buffer.extend(text) + self.total_added += len(text) + + if self.total_added - self.window_start >= self.page_size: + self._spill_locked() + + def get_all(self, clear: bool = False) -> str: + """Return the in-memory window (content not yet flushed to disk).""" + with self.lock: + result = "".join(self.buffer) + + if clear: + self.buffer.clear() + self.window_start = self.total_added + + return result + + def get_new_output(self, last_read_position: int) -> Tuple[str, int]: + """Return window content past ``last_read_position`` and the new offset. + + Content older than the window has already been paged to disk, so the + read position is clamped forward to the window start rather than + replaying bytes that are only available as pages. + """ + with self.lock: + if last_read_position >= self.total_added: + return "", self.total_added + + start = max(last_read_position, self.window_start) + new_output = "".join(self.buffer)[start - self.window_start :] + + return new_output, self.total_added + + def clear(self) -> None: + """Drop the in-memory window; flushed pages are unaffected.""" + with self.lock: + self.buffer.clear() + self.window_start = self.total_added + + def size(self) -> int: + """Get current buffer size in characters.""" + with self.lock: + return len(self.buffer) + + def _spill_locked(self) -> None: + content = "".join(self.buffer) + + if not self.pages_dir: + # Without a page directory, retain only the newest page in memory. + if len(content) > self.page_size: + dropped = len(content) - self.page_size + content = content[dropped:] + self.window_start += dropped + self.buffer = deque(content) + + return + + while len(content) >= self.page_size: + page = content[: self.page_size] + content = content[self.page_size :] + self.page_count += 1 + self._write_page_locked(self.page_count, page) + self.window_start += self.page_size + + self.buffer = deque(content) + + def _write_page_locked(self, page_number: int, content: str) -> None: + os.makedirs(self.pages_dir, exist_ok=True) + abs_path = os.path.join(self.pages_dir, f"{page_number}.txt") + tmp_path = f"{abs_path}.tmp" + + with safe_open(tmp_path, "w") as page_file: + page_file.write(content) + + os.replace(tmp_path, abs_path) + + class InputBuffer: """ Thread-safe buffer for queuing input to be sent to a process. @@ -442,6 +546,9 @@ def start_background_command( existing_input_buffer: Optional[InputBuffer] = None, use_pty: bool = False, master_fd: Optional[int] = None, + command_key: Optional[str] = None, + page_size: Optional[int] = None, + pages_dir: Optional[str] = None, ) -> str: """ Start a command in background. @@ -452,15 +559,23 @@ def start_background_command( cwd: Working directory for command max_buffer_size: Maximum buffer size for output existing_process: Optional existing subprocess.Popen to register - existing_buffer: Optional existing CircularBuffer to use + existing_buffer: Optional existing buffer to use (CircularBuffer or PagedOutputBuffer) persist: If True, output buffer won't be cleared when read + command_key: Optional pre-generated command key; generated when omitted + page_size: Characters per page when paging output to disk + pages_dir: Directory where full output pages are written Returns: Command key for future reference """ try: - # Use existing buffer or create new one - buffer = existing_buffer or CircularBuffer(max_size=max_buffer_size) + # Use existing buffer or create a paged/circular one + if existing_buffer is not None: + buffer = existing_buffer + elif page_size and pages_dir: + buffer = PagedOutputBuffer(page_size=page_size, pages_dir=pages_dir) + else: + buffer = CircularBuffer(max_size=max_buffer_size) # Use existing process or start new one # Use provided master_fd (e.g., from _execute_with_timeout) or default to None @@ -523,7 +638,7 @@ def start_background_command( ) # Generate unique key and store - command_key = cls._generate_command_key(command) + command_key = command_key or cls._generate_command_key(command) with cls._lock: cls._background_commands[command_key] = bg_process @@ -691,6 +806,10 @@ def list_background_commands(cls) -> Dict[str, Dict[str, any]]: "command": bg_process.command, "running": bg_process.is_alive(), "buffer_size": bg_process.buffer.size(), + "pages": getattr(bg_process.buffer, "page_count", 0), + "total_chars": getattr( + bg_process.buffer, "total_added", bg_process.buffer.size() + ), "start_time": bg_process.start_time, "end_time": bg_process.end_time, "duration": ( diff --git a/cecli/helpers/conversation/integration.py b/cecli/helpers/conversation/integration.py index 202b00b7665..0cf9073b6a0 100644 --- a/cecli/helpers/conversation/integration.py +++ b/cecli/helpers/conversation/integration.py @@ -930,7 +930,9 @@ def add_post_message_context_blocks(self) -> None: """ Add post-message context blocks to conversation (priority 250). - Post-message blocks include: tool_context/write_context, background_command_output + Post-message blocks include: todo_list, context_summary, tool_context, + and write_context. Background command output is injected separately via + ``add_background_command_output`` with its own debounce. """ coder = self.get_coder() if not coder: @@ -972,12 +974,6 @@ def add_post_message_context_blocks(self) -> None: if write_context: message_blocks["write_context"] = write_context - # Add background command output if any - if hasattr(coder, "get_background_command_output"): - bg_output = coder.get_background_command_output() - if bg_output: - message_blocks["background_command_output"] = bg_output - # Add post-message blocks to conversation manager with stable hash keys for block_type, block_content in message_blocks.items(): ConversationService.get_manager(coder).add_message( @@ -989,6 +985,56 @@ def add_post_message_context_blocks(self) -> None: force=True, ) + def add_background_command_output(self, frequency=5): + """ + Inject background command output at most once every ``frequency`` turns, + except when a command finishes or flushes a new page. Those transitions + bypass the debounce so important changes surface immediately. + + Debounced independently from the other post-message blocks: the injected + content mutates slightly as commands produce output, so re-adding it + every turn would churn the conversation tail without the usual hash-key + deduplication catching it. + """ + coder = self.get_coder() + if not coder: + return + + if not hasattr(coder, "use_enhanced_context") or not coder.use_enhanced_context: + return + + if not hasattr(coder, "get_background_command_output"): + return + + last_turn = self.message_tracker.get("background_command_output") + due = last_turn is None or coder.turn_count - last_turn >= frequency + + previous_state = self.message_tracker.get("background_command_state") + state = None + if hasattr(coder, "get_background_command_state"): + state = coder.get_background_command_state() + + significant = self._has_background_signal(previous_state, state) + if state is not None: + self.message_tracker["background_command_state"] = state + + if not significant and not due: + return + + bg_output = coder.get_background_command_output() + if not bg_output: + return + + self.message_tracker["background_command_output"] = coder.turn_count + ConversationService.get_manager(coder).add_message( + message_dict={"role": "user", "content": bg_output}, + tag=MessageTag.STATIC, + priority=DEFAULT_TAG_PRIORITY[MessageTag.REMINDER] + 25, + mark_for_delete=0, + hash_key=("post_message", "background_command_output"), + force=True, + ) + def add_sub_agent_states(self) -> None: """ Add sub-agent states context block to conversation (priority 250). @@ -1056,6 +1102,26 @@ def debounce_message_injection(self, coder, message_type="default", frequency=10 return not should_send + @staticmethod + def _has_background_signal(previous, current): + """Return True when a command finished or flushed a new page since ``previous``.""" + if not current or not previous: + return False + + for key, info in current.items(): + prior = previous.get(key) + + if prior is None: + return True + + if prior.get("running") and not info.get("running"): + return True + + if info.get("pages", 0) > prior.get("pages", 0): + return True + + return False + def _cancel_post_message_injections(self, modulus=10): coder = self.get_coder() if not coder: diff --git a/cecli/tools/command.py b/cecli/tools/command.py index 1177850bd4b..42a4e82ac78 100644 --- a/cecli/tools/command.py +++ b/cecli/tools/command.py @@ -86,15 +86,16 @@ class Tool(BaseTool): "type": "string", "description": ( "Key of an existing background command to interact with. " - "Use with 'action' (stdin/stop)." + "Use with 'action' (stdin/stop/tail)." ), }, "action": { "type": "string", - "enum": ["stdin", "stop"], + "enum": ["stdin", "stop", "tail"], "description": ( "Action on a background command. Requires background_key: " - "'stdin' to send input, 'stop' to terminate." + "'stdin' to send input, 'stop' to terminate, 'tail' to read " + "the latest output." ), }, "stdin": { @@ -171,7 +172,7 @@ async def execute( When 'user_input_required' is True, runs the command interactively using a pseudo-terminal (PTY), allowing the user to provide inputs like passwords or navigate terminal interfaces. - For background interactions: provide 'background_key' + 'action' (stdin/stop). + For background interactions: provide 'background_key' + 'action' (stdin/stop/tail). Commands run with timeout from agent_config['command_timeout'] (default: 30 seconds), """ @@ -201,8 +202,11 @@ async def execute( elif action == "stop": return await cls._stop_background_command(coder, background_key) + elif action == "tail": + return await cls._tail_background_command(coder, background_key) + else: - response.append_error(f"Unknown action '{action}'. Use one of: stdin, stop.") + response.append_error(f"Unknown action '{action}'. Use one of: stdin, stop, tail.") return response if not command: @@ -320,12 +324,17 @@ async def _execute_background(cls, coder, command_string, use_pty=None, stdin=No use_pty = platform.system() != "Windows" # Use static manager to start background command + command_key, page_size, pages_dir = cls._paging_config(coder, command_string) + command_key = BackgroundCommandManager.start_background_command( command_string, verbose=coder.verbose, cwd=coder.root, - max_buffer_size=4096, + max_buffer_size=page_size or 4096, use_pty=use_pty, + command_key=command_key, + page_size=page_size, + pages_dir=pages_dir, ) # Send stdin to the background command if provided @@ -352,7 +361,7 @@ async def _execute_with_timeout(cls, coder, command_string, timeout, use_pty=Non import asyncio import subprocess - from cecli.helpers.background_commands import CircularBuffer + from cecli.helpers.background_commands import PagedOutputBuffer response = ToolResponse(cls.NORM_NAME) @@ -364,8 +373,9 @@ async def _execute_with_timeout(cls, coder, command_string, timeout, use_pty=Non if use_pty is None: use_pty = platform.system() != "Windows" - # Create output buffer - buffer = CircularBuffer(max_size=4096) + # Create output buffer (paged when context management is enabled) + command_key, page_size, pages_dir = cls._paging_config(coder, command_string) + buffer = PagedOutputBuffer(page_size=page_size or 4096, pages_dir=pages_dir) # Decide whether to use PTY master_fd = None @@ -424,6 +434,7 @@ async def _execute_with_timeout(cls, coder, command_string, timeout, use_pty=Non existing_buffer=buffer, persist=True, master_fd=master_fd, + command_key=command_key, ) # Now monitor the process with an event-driven race instead of @@ -475,30 +486,10 @@ async def _execute_with_timeout(cls, coder, command_string, timeout, use_pty=Non command_completed = wait_task in done output_content = buffer.get_all(clear=command_completed) or "" - # Tokens are roughly 3-4 characters - output_limit = int(coder.large_file_token_threshold * 3.5) - - if coder.context_management_enabled and len(output_content) > output_limit * 1.25: - folder_path, file_list, alias_paths = ( - BackgroundCommandManager.save_paginated_output( - output=output_content, - command_key=command_key, - page_size=output_limit, - abs_root_path_func=coder.abs_root_path, - local_agent_folder_func=coder.local_agent_folder, - ) - ) - total_size = len(output_content) + pages_notice = cls._pages_notice(command_key, getattr(buffer, "page_count", 0)) + if pages_notice: output_content = ( - f"[Large Response ({total_size} characters). " - f"Output saved in {len(file_list)} pages.]\n" - f"Command key: {command_key}\n" - f"Pages: 1-{len(file_list)}\n" - "Use `ResourceManager` to view up to 3 pages at a time:\n" - f'{{"paging": [{{"target": "{command_key}", "page": 1}}]}}\n' - "Change page or add entries to read other pages (maximum 3 entries). " - "Do not use add, read_only, or standard CLI tools to view command output " - "files. Pages are returned directly, not added to file context." + f"{output_content}\n\n{pages_notice}" if output_content else pages_notice ) if command_completed: @@ -570,8 +561,8 @@ async def _execute_foreground(cls, coder, command_string): # Format the output for the result message output_content = combined_output or "" - output_limit = coder.large_file_token_threshold - if coder.context_management_enabled and len(output_content) > output_limit * 1.25: + output_limit = cls._page_size(coder) + if coder.context_management_enabled and len(output_content) > output_limit: # Generate a unique key for file naming fg_key = BackgroundCommandManager._generate_command_key(command_string) # Save full output to paginated files instead of truncating @@ -718,6 +709,73 @@ def _requires_user_input(cls, command_string): for pattern in INTERACTIVE_COMMAND_PATTERNS ) + @classmethod + def _page_size(cls, coder): + """Characters per output page (~3.5 characters per LLM token).""" + return max(1, int(getattr(coder, "large_file_token_threshold", 8192) * 3.5)) + + @classmethod + def _paging_config(cls, coder, command_string): + """Return (command_key, page_size, pages_dir) when output paging is enabled.""" + if not getattr(coder, "context_management_enabled", False): + return None, None, None + + page_size = cls._page_size(coder) + command_key = BackgroundCommandManager._generate_command_key(command_string) + pages_dir = coder.abs_root_path(coder.local_agent_folder(command_key)) + + return command_key, page_size, pages_dir + + @classmethod + async def _tail_background_command(cls, coder, command_key): + """Return the latest in-memory output and page roster for a background command.""" + command_info = BackgroundCommandManager.list_background_commands() + info = command_info.get(command_key) + + response = ToolResponse(cls.NORM_NAME) + if not info: + response.append_error(f"Background command {command_key} not found.") + return response + + status = "running" if info.get("running", False) else "finished" + output = BackgroundCommandManager.get_new_command_output(command_key) + pages = info.get("pages", 0) + + lines = [ + f"Background command {command_key} [{status}]: {info.get('command', command_key)}", + f"Output so far: {info.get('total_chars', 0):,} chars", + ] + if pages: + lines.append(f"Paged output: pages 1-{pages}. Read with ResourceManager paging.") + lines.append(f'{{"paging": [{{"target": "{command_key}", "page": 1}}]}}') + + if output.strip(): + lines.append("New output since last read:") + lines.append(output) + else: + lines.append("No new output since last read.") + + response.append_result("\n".join(lines)) + + return response + + @staticmethod + def _pages_notice(command_key, page_count): + """Guidance for reading command output that has been paged to disk.""" + if not page_count: + return "" + + return ( + f"[Output paged to disk: {page_count} page(s).]\n" + f"Command key: {command_key}\n" + f"Pages: 1-{page_count}\n" + "Use `ResourceManager` to view up to 3 pages at a time:\n" + f'{{"paging": [{{"target": "{command_key}", "page": 1}}]}}\n' + "Change the page number to read other pages (maximum 3 entries). " + "Do not use add, read_only, or standard CLI tools to view command output " + "files. Pages are returned directly, not added to file context." + ) + @classmethod def format_output(cls, coder, mcp_server, tool_response): """Format output for Command tool.""" diff --git a/tests/basic/test_background_commands.py b/tests/basic/test_background_commands.py index 5ca98746186..1f7f079a274 100644 --- a/tests/basic/test_background_commands.py +++ b/tests/basic/test_background_commands.py @@ -49,8 +49,10 @@ def readline(self): _install_stubs() from cecli.helpers.background_commands import ( # noqa: E402 + BackgroundCommandManager, BackgroundProcess, CircularBuffer, + PagedOutputBuffer, ) @@ -235,3 +237,168 @@ def readline(self): success, output, exit_code = bg_process.stop() assert success is True assert exit_code == -1 # terminate() sets returncode to -1 in MockProcess + + +def test_paged_output_buffer_spills_pages_to_disk(tmp_path): + """Full pages are flushed to disk and dropped from the in-memory window.""" + buffer = PagedOutputBuffer(page_size=5, pages_dir=str(tmp_path / "pages")) + + buffer.append("abc") + assert buffer.get_all() == "abc" + assert buffer.page_count == 0 + + buffer.append("de") + assert buffer.page_count == 1 + assert buffer.get_all() == "" + assert buffer.total_added == 5 + assert (tmp_path / "pages" / "1.txt").read_text(encoding="utf-8") == "abcde" + + buffer.append("fgh") + assert buffer.page_count == 1 + assert buffer.get_all() == "fgh" + + buffer.append("ij") + assert buffer.page_count == 2 + assert buffer.get_all() == "" + assert (tmp_path / "pages" / "2.txt").read_text(encoding="utf-8") == "fghij" + + # Atomic writes leave no temporary files behind + assert not list((tmp_path / "pages").glob("*.tmp")) + + +def test_paged_output_buffer_incremental_reads_clamp_to_window(tmp_path): + """Readers resume monotonically; already-paged content is not replayed.""" + buffer = PagedOutputBuffer(page_size=4, pages_dir=str(tmp_path / "p")) + + assert buffer.get_new_output(0) == ("", 0) + + buffer.append("abcd") + assert buffer.get_new_output(0) == ("", 4) + + buffer.append("ef") + assert buffer.get_new_output(4) == ("ef", 6) + assert buffer.get_new_output(6) == ("", 6) + + +def test_paged_output_buffer_without_pages_dir_is_bounded(): + """With no page directory the window simply keeps the newest page.""" + buffer = PagedOutputBuffer(page_size=5, pages_dir=None) + + buffer.append("abcdefgh") + + assert buffer.get_all() == "defgh" + assert buffer.page_count == 0 + + +def test_tail_background_command_reports_output_and_pages(monkeypatch): + """The tail action reports status, page roster, and new output.""" + import asyncio + + from cecli.tools.command import Tool as CommandTool + + monkeypatch.setattr( + BackgroundCommandManager, + "list_background_commands", + lambda: { + "bg_1_1234": { + "command": "pytest -q", + "running": True, + "pages": 3, + "total_chars": 42, + } + }, + ) + monkeypatch.setattr( + BackgroundCommandManager, "get_new_command_output", lambda key: "new line\n" + ) + + response = asyncio.run(CommandTool._tail_background_command(object(), "bg_1_1234")) + content = response.to_dict()["result"][0]["content"] + + assert "bg_1_1234" in content + assert "running" in content + assert "pages 1-3" in content + assert '{"paging": [{"target": "bg_1_1234", "page": 1}]}' in content + assert "new line" in content + + +def test_tail_background_command_missing_key(monkeypatch): + import asyncio + + from cecli.tools.command import Tool as CommandTool + + monkeypatch.setattr(BackgroundCommandManager, "list_background_commands", lambda: {}) + + response = asyncio.run(CommandTool._tail_background_command(object(), "bg_9_9999")) + + assert response.to_dict()["errors"] + + +def test_get_background_command_output_roster_incremental_and_pages(monkeypatch): + """Injection lists a stable roster, new output, and page guidance.""" + from cecli.coders.agent_coder import AgentCoder + + monkeypatch.setattr( + BackgroundCommandManager, + "list_background_commands", + lambda: { + "bg_1_1234": { + "command": "pytest -q", + "running": True, + "pages": 2, + "total_chars": 100, + }, + "bg_2_5678": { + "command": "npm run build", + "running": True, + "pages": 0, + "total_chars": 12, + }, + }, + ) + monkeypatch.setattr( + BackgroundCommandManager, + "get_new_command_output", + lambda key: f"out-{key}\n", + ) + stopped = [] + monkeypatch.setattr( + BackgroundCommandManager, "stop_background_command", lambda key: stopped.append(key) + ) + + output = AgentCoder.get_background_command_output(object()) + + assert "bg_1_1234" in output + assert "pages 1-2" in output + assert "no pages yet" in output + assert "out-bg_1_1234" in output + assert '{"paging": [{"target": "bg_1_1234", "page": 1}]}' in output + assert stopped == [] + + +def test_get_background_command_output_stops_finished_commands(monkeypatch): + """Finished commands are reported once and then removed from tracking.""" + from cecli.coders.agent_coder import AgentCoder + + monkeypatch.setattr( + BackgroundCommandManager, + "list_background_commands", + lambda: { + "bg_3_0001": { + "command": "true", + "running": False, + "pages": 0, + "total_chars": 0, + } + }, + ) + monkeypatch.setattr(BackgroundCommandManager, "get_new_command_output", lambda key: "") + stopped = [] + monkeypatch.setattr( + BackgroundCommandManager, "stop_background_command", lambda key: stopped.append(key) + ) + + output = AgentCoder.get_background_command_output(object()) + + assert "finished" in output + assert stopped == ["bg_3_0001"] diff --git a/tests/conversations/test_background_command_injection.py b/tests/conversations/test_background_command_injection.py new file mode 100644 index 00000000000..f154d57adaa --- /dev/null +++ b/tests/conversations/test_background_command_injection.py @@ -0,0 +1,170 @@ +"""Tests for the independently debounced background command output injection.""" + +import uuid + +from cecli.helpers.conversation import ConversationService + + +class MockCoder: + def __init__(self): + self.uuid = str(uuid.uuid4()) + self.use_enhanced_context = True + self.turn_count = 0 + self.output_calls = 0 + self.output = "roster" + + def get_background_command_output(self): + self.output_calls += 1 + return self.output + + +def _make_chunks(coder): + manager = ConversationService.get_manager(coder) + manager.reset() + chunks = ConversationService.get_chunks(coder) + chunks.message_tracker = {} + return chunks, manager + + +def test_background_command_output_debounced_to_every_five_turns(): + coder = MockCoder() + chunks, manager = _make_chunks(coder) + + for turn in range(7): + coder.turn_count = turn + chunks.add_background_command_output(frequency=5) + + # Injected once up front, then not again until turn 5 + assert coder.output_calls == 2 + assert [message["content"] for message in manager.get_messages_dict()] == ["roster"] + + +def test_background_command_output_replaces_same_hash_key(): + coder = MockCoder() + chunks, manager = _make_chunks(coder) + + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + + coder.output = "roster v2" + coder.turn_count = 5 + chunks.add_background_command_output(frequency=5) + + assert [message["content"] for message in manager.get_messages_dict()] == ["roster v2"] + + +def test_background_command_output_skipped_without_enhanced_context(): + coder = MockCoder() + coder.use_enhanced_context = False + chunks, manager = _make_chunks(coder) + + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + + assert coder.output_calls == 0 + assert manager.get_messages_dict() == [] + + +def test_empty_background_command_output_does_not_consume_window(): + coder = MockCoder() + coder.output = "" + chunks, manager = _make_chunks(coder) + + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + coder.turn_count = 1 + chunks.add_background_command_output(frequency=5) + + # Empty output never injects and never marks the tracker, so polling continues + assert coder.output_calls == 2 + assert manager.get_messages_dict() == [] + + coder.output = "roster" + coder.turn_count = 2 + chunks.add_background_command_output(frequency=5) + + assert coder.output_calls == 3 + assert [message["content"] for message in manager.get_messages_dict()] == ["roster"] + + +class SignalCoder(MockCoder): + def __init__(self): + super().__init__() + self.state = {} + + def get_background_command_state(self): + return self.state + + +def test_finish_transition_bypasses_debounce(): + coder = SignalCoder() + chunks, manager = _make_chunks(coder) + + coder.state = {"bg_1_0001": {"running": True, "pages": 0}} + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 1 + + # Still running, nothing changed, and not due -> skip + coder.turn_count = 1 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 1 + + # Finished -> inject immediately despite the frequency window + coder.state = {"bg_1_0001": {"running": False, "pages": 0}} + coder.turn_count = 2 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 2 + + +def test_new_page_flush_bypasses_debounce(): + coder = SignalCoder() + chunks, _ = _make_chunks(coder) + + coder.state = {"bg_1_0001": {"running": True, "pages": 0}} + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 1 + + coder.state = {"bg_1_0001": {"running": True, "pages": 1}} + coder.turn_count = 1 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 2 + + +def test_new_command_appearance_bypasses_debounce(): + coder = SignalCoder() + chunks, _ = _make_chunks(coder) + + coder.state = {"bg_1_0001": {"running": True, "pages": 0}} + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 1 + + coder.state = { + "bg_1_0001": {"running": True, "pages": 0}, + "bg_2_0002": {"running": True, "pages": 0}, + } + coder.turn_count = 1 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 2 + + +def test_running_command_without_change_respects_frequency(): + coder = SignalCoder() + chunks, _ = _make_chunks(coder) + + coder.state = {"bg_1_0001": {"running": True, "pages": 0}} + coder.turn_count = 0 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 1 + + for turn in (1, 2, 3, 4): + coder.turn_count = turn + chunks.add_background_command_output(frequency=5) + + assert coder.output_calls == 1 + + coder.turn_count = 5 + chunks.add_background_command_output(frequency=5) + assert coder.output_calls == 2 diff --git a/tests/tools/test_command_timeout_paging.py b/tests/tools/test_command_timeout_paging.py index d2cc2f8d37c..c4fb654a12f 100644 --- a/tests/tools/test_command_timeout_paging.py +++ b/tests/tools/test_command_timeout_paging.py @@ -33,35 +33,43 @@ async def elapsed_command(monkeypatch, tmp_path): context_blocks_cache={}, edit_allowed=False, interrupt_event=asyncio.Event(), + large_file_token_threshold=8, + context_management_enabled=True, ) manager = background_commands.BackgroundCommandManager target = "bg_1_1234" process = Mock() popen = Mock(return_value=process) - buffer = background_commands.CircularBuffer() - get_all = Mock(wraps=buffer.get_all) - monkeypatch.setattr(buffer, "get_all", get_all) - monkeypatch.setattr(background_commands, "CircularBuffer", Mock(return_value=buffer)) + state = {"output": "", "buffer": None} + + real_buffer_cls = background_commands.PagedOutputBuffer + + def make_buffer(page_size=4096, pages_dir=None): + buffer = real_buffer_cls(page_size=page_size, pages_dir=pages_dir) + state["buffer"] = buffer + if state["output"]: + buffer.append(state["output"]) + return buffer + + monkeypatch.setattr(background_commands, "PagedOutputBuffer", make_buffer) monkeypatch.setattr("subprocess.Popen", popen) + monkeypatch.setattr(manager, "_generate_command_key", Mock(return_value=target)) start = Mock(return_value=target) stop = Mock() - save = Mock(wraps=manager.save_paginated_output) monkeypatch.setattr(manager, "start_background_command", start) monkeypatch.setattr(manager, "stop_background_command", stop) - monkeypatch.setattr(manager, "save_paginated_output", save) pending_tasks = [] async def pending_wait(*args, **kwargs): pending_tasks.append(asyncio.current_task()) await asyncio.get_running_loop().create_future() - to_thread = Mock(side_effect=pending_wait) - monkeypatch.setattr(asyncio, "to_thread", to_thread) + monkeypatch.setattr(asyncio, "to_thread", Mock(side_effect=pending_wait)) async def execute(output, threshold=8, enabled=True): coder.large_file_token_threshold = threshold coder.context_management_enabled = enabled - buffer.append(output) + state["output"] = output response = await CommandTool._execute_with_timeout( coder, "pending command", 0.001, use_pty=False ) @@ -75,14 +83,12 @@ async def execute(output, threshold=8, enabled=True): assert len(pending_tasks) == 1 assert not pending_tasks[0].done() assert not coder.interrupt_event.is_set() - to_thread.assert_called_once_with(process.wait) popen.assert_called_once() start.assert_called_once() assert start.call_args.kwargs["existing_process"] is process - assert start.call_args.kwargs["existing_buffer"] is buffer + assert start.call_args.kwargs["existing_buffer"] is state["buffer"] assert start.call_args.kwargs["persist"] is True - get_all.assert_called_once_with(clear=False) - assert buffer.get_all() == output + assert start.call_args.kwargs["command_key"] == (target if enabled else None) stop.assert_not_called() process.wait.assert_not_called() process.terminate.assert_not_called() @@ -90,7 +96,7 @@ async def execute(output, threshold=8, enabled=True): return content try: - yield SimpleNamespace(execute=execute, coder=coder, target=target, save=save) + yield SimpleNamespace(execute=execute, coder=coder, target=target, state=state) finally: for task in pending_tasks: task.cancel() @@ -99,59 +105,38 @@ async def execute(output, threshold=8, enabled=True): @pytest.mark.asyncio -async def test_elapsed_timeout_saves_pages_readable_without_adding_context(elapsed_command): - output = "first line: café\nsecond line\n" * 3 - content = await elapsed_command.execute(output) +async def test_elapsed_timeout_pages_output_and_exposes_command_keys(elapsed_command): coder = elapsed_command.coder target = elapsed_command.target page_size = int(coder.large_file_token_threshold * 3.5) - expected_pages = [ - output[index : index + page_size] for index in range(0, len(output), page_size) - ] + output = "first line: café\nsecond line\n" * 3 + content = await elapsed_command.execute(output) + num_pages = len(output) // page_size + assert num_pages >= 1 assert output not in content - assert f"Large Response ({len(output)} characters)" in content - assert f"Output saved in {len(expected_pages)} pages." in content + assert f"Output paged to disk: {num_pages} page(s)." in content assert "ResourceManager" in content assert "not added to file context" in content assert "command_key::" not in content example = next(line for line in content.splitlines() if line.startswith('{"paging"')) assert json.loads(example) == {"paging": [{"target": target, "page": 1}]} - elapsed_command.save.assert_called_once_with( - output=output, - command_key=target, - page_size=page_size, - abs_root_path_func=coder.abs_root_path, - local_agent_folder_func=coder.local_agent_folder, - ) + folder = Path(coder.abs_root_path(coder.local_agent_folder(target))) assert {path.name for path in folder.iterdir()} == { - f"{page}.txt" for page in range(1, len(expected_pages) + 1) + f"{page}.txt" for page in range(1, num_pages + 1) } saved_pages = [ - (folder / f"{page}.txt").read_text(encoding="utf-8") - for page in range(1, len(expected_pages) + 1) + (folder / f"{page}.txt").read_text(encoding="utf-8") for page in range(1, num_pages + 1) ] - assert saved_pages == expected_pages - assert "".join(saved_pages) == output + assert "".join(saved_pages) == output[: num_pages * page_size] editable_before = coder.abs_fnames.copy() read_only_before = coder.abs_read_only_fnames.copy() - for index in range(0, len(expected_pages), 3): - batch = expected_pages[index : index + 3] - response = await ResourceManagerTool.execute( - coder, - paging=[ - {"target": target, "page": page} - for page in range(index + 1, index + len(batch) + 1) - ], - ) - result = response.to_dict() - assert result["errors"] == [] - assert len(result["result"]) == len(batch) - for item, expected in zip(result["result"], batch): - assert item["content"].endswith(expected) - + response = await ResourceManagerTool.execute(coder, paging=[{"target": target, "page": 1}]) + result = response.to_dict() + assert result["errors"] == [] + assert result["result"][0]["content"].endswith(output[:page_size]) assert coder.abs_fnames == editable_before assert coder.abs_read_only_fnames == read_only_before coder._add_file_to_context.assert_not_called() @@ -159,23 +144,27 @@ async def test_elapsed_timeout_saves_pages_readable_without_adding_context(elaps @pytest.mark.asyncio @pytest.mark.parametrize( - "threshold,length,paged", [(8, 35, False), (8, 36, True), (9, 38, False), (9, 39, True)] + "threshold,length,paged", [(8, 27, False), (8, 28, True), (9, 30, False), (9, 31, True)] ) -async def test_elapsed_timeout_paging_uses_strict_rounded_threshold( +async def test_elapsed_timeout_paging_triggers_at_page_size( elapsed_command, threshold, length, paged ): output = "x" * length content = await elapsed_command.execute(output, threshold=threshold) if paged: - elapsed_command.save.assert_called_once() - assert elapsed_command.save.call_args.kwargs["page_size"] == int(threshold * 3.5) - assert "Large Response" in content + page_size = int(threshold * 3.5) + assert "Output paged to disk: 1 page(s)." in content assert output not in content + folder = Path( + elapsed_command.coder.abs_root_path( + elapsed_command.coder.local_agent_folder(elapsed_command.target) + ) + ) + assert (folder / "1.txt").read_text(encoding="utf-8") == output[:page_size] else: - elapsed_command.save.assert_not_called() assert f"Output captured so far:\n{output}\n" in content - assert "Large Response" not in content + assert "Output paged to disk" not in content @pytest.mark.asyncio @@ -188,5 +177,4 @@ async def test_elapsed_timeout_keeps_small_empty_or_unmanaged_output_inline( content = await elapsed_command.execute(output, enabled=enabled) assert f"Output captured so far:\n{output}\n" in content - assert "Large Response" not in content - elapsed_command.save.assert_not_called() + assert "Output paged to disk" not in content diff --git a/tests/tools/test_resource_manager_paging.py b/tests/tools/test_resource_manager_paging.py index 04ecc1576b9..1237b29dc6e 100644 --- a/tests/tools/test_resource_manager_paging.py +++ b/tests/tools/test_resource_manager_paging.py @@ -284,9 +284,9 @@ async def test_command_large_output_guidance_uses_paging_array( manager = command.BackgroundCommandManager target = "bg_1_1234" output = "large command output\n" * 50 + monkeypatch.setattr(manager, "_generate_command_key", Mock(return_value=target)) save = Mock(return_value=("pages", ["1.txt", "2.txt"], ["command_key::old/1.txt"])) monkeypatch.setattr(manager, "save_paginated_output", save) - monkeypatch.setattr(manager, "_generate_command_key", Mock(return_value=target)) if execution_path == "foreground": monkeypatch.setattr(command, "run_cmd_subprocess", Mock(return_value=(0, output))) @@ -297,9 +297,15 @@ async def test_command_large_output_guidance_uses_paging_array( monkeypatch.setattr("subprocess.Popen", Mock(return_value=process)) monkeypatch.setattr(manager, "start_background_command", Mock(return_value=target)) monkeypatch.setattr(manager, "stop_background_command", Mock()) - buffer = Mock() - buffer.get_all.return_value = output - monkeypatch.setattr(background_commands, "CircularBuffer", Mock(return_value=buffer)) + + real_buffer_cls = background_commands.PagedOutputBuffer + + def make_buffer(page_size=4096, pages_dir=None): + buffer = real_buffer_cls(page_size=page_size, pages_dir=pages_dir) + buffer.append(output) + return buffer + + monkeypatch.setattr(background_commands, "PagedOutputBuffer", make_buffer) response = await command.Tool._execute_with_timeout(coder, "echo test", 30, use_pty=False) result = response.to_dict() @@ -310,6 +316,10 @@ async def test_command_large_output_guidance_uses_paging_array( assert "ResourceManager" in content assert "command_key::" not in content assert "not added to file context" in content - save.assert_called_once() - assert save.call_args.kwargs["output"] == output - assert save.call_args.kwargs["command_key"] == target + + if execution_path == "foreground": + save.assert_called_once() + assert save.call_args.kwargs["output"] == output + assert save.call_args.kwargs["command_key"] == target + else: + save.assert_not_called() From 42b3f06cfec2f81b7ba56e18226f0f3e1cb346ab Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Tue, 15 Sep 2026 02:18:28 -0400 Subject: [PATCH 03/14] Add MCP server default request timeout, make sure connect_all() errors can't block --- cecli/helpers/coroutines.py | 19 +++++++ cecli/mcp/manager.py | 69 +++++++++++++++--------- cecli/mcp/oauth.py | 91 +++++++++++++++++++++++++------- cecli/mcp/server.py | 51 ++++++++++++++++-- cecli/tui/worker.py | 38 +++++++++++-- cecli/website/docs/config/mcp.md | 31 +++++++++++ tests/mcp/test_manager_retry.py | 41 ++++++++++++-- 7 files changed, 284 insertions(+), 56 deletions(-) diff --git a/cecli/helpers/coroutines.py b/cecli/helpers/coroutines.py index 968410ecf48..0d2ae4c8294 100644 --- a/cecli/helpers/coroutines.py +++ b/cecli/helpers/coroutines.py @@ -109,3 +109,22 @@ async def interruptible(coroutine, interrupt_event): return main_task.result(), False except asyncio.CancelledError: return None, True + + +def task_is_cancelling() -> bool: + """Return True when the running asyncio task has a pending cancellation. + + Used to tell a genuine cancellation of the caller apart from cancellation + errors that transports (e.g. MCP's anyio TaskGroups) surface for ordinary + connection failures. + """ + task = asyncio.current_task() + if task is None: + return False + + cancelling = getattr(task, "cancelling", None) + if cancelling is not None: + return cancelling() > 0 + + # Python 3.10 has no Task.cancelling(); fall back to the private flag. + return bool(getattr(task, "_must_cancel", False)) diff --git a/cecli/mcp/manager.py b/cecli/mcp/manager.py index 9f3cb429208..a2216e0267b 100644 --- a/cecli/mcp/manager.py +++ b/cecli/mcp/manager.py @@ -1,6 +1,7 @@ import asyncio -from cecli.mcp.server import LocalServer, McpServer +from cecli.helpers.coroutines import task_is_cancelling +from cecli.mcp.server import DEFAULT_MCP_REQUEST_TIMEOUT, LocalServer, McpServer from cecli.tools.utils.registry import ToolRegistry @@ -196,9 +197,10 @@ async def connect_server(self, name: str) -> bool: return True # Retry with exponential backoff for transient connection failures. - # Note: This also fixes a latent bug where asyncio.CancelledError was - # silently caught and treated as a connection failure. CancelledError is - # now re-raised to properly propagate cancellation. + # A genuine cancellation must still propagate, but some transports + # (MCP's streamable-HTTP runs its writer inside an anyio TaskGroup) + # surface ordinary connection failures as CancelledError, so only a + # real cancellation of this task is re-raised. # When io is None (e.g., during from_servers before IO is assigned), # _log_warning and _log_error silently return — retries still happen # but with no user-visible feedback. This is intentional. @@ -206,38 +208,55 @@ async def connect_server(self, name: str) -> bool: delay = 1.0 backoff = 2.0 max_delay = 30.0 + # Bound each attempt so a server that accepts a connection but never + # speaks MCP cannot wedge startup. The SDK read timeout usually fires + # first; this wait_for is a backstop for hangs inside the transport. + try: + base_timeout = float(server._request_timeout_seconds()) + except (TypeError, ValueError): + base_timeout = DEFAULT_MCP_REQUEST_TIMEOUT + + if base_timeout <= 0: + base_timeout = DEFAULT_MCP_REQUEST_TIMEOUT + + attempt_timeout = base_timeout + 5 for attempt in range(1, max_retries + 1): + error = None + try: - session = await server.connect() - tools_result = await session.list_tools() + session = await asyncio.wait_for(server.connect(), timeout=attempt_timeout) + tools_result = await asyncio.wait_for(session.list_tools(), timeout=attempt_timeout) tools = _mcp_tools_to_openai_tools(tools_result.tools) self._server_tools[server.name] = tools self._connected_servers.add(server) self._log_verbose(f"Connected to MCP server: {name}") return True except asyncio.CancelledError: - raise + if task_is_cancelling(): + raise + error = "connection cancelled by transport" except Exception as e: - if attempt < max_retries and server.name != "unnamed-server": - self._log_warning( - f"Connection attempt {attempt} failed for {name}, " - f"retrying in {delay}s... ({e})" - ) + error = e - await asyncio.sleep(delay) - delay = min(delay * backoff, max_delay) - else: - if server.name != "unnamed-server": - self._log_error( - f"Failed to connect to MCP server {name} " - f"after {max_retries} attempts: {e}" - ) - if server.is_connected: - # Session was established but tool listing failed; tear - # it down so the transport/subprocess doesn't leak. - await server.disconnect() - return False + if attempt < max_retries and server.name != "unnamed-server": + self._log_warning( + f"Connection attempt {attempt} failed for {name}, " + f"retrying in {delay}s... ({error})" + ) + await asyncio.sleep(delay) + delay = min(delay * backoff, max_delay) + else: + if server.name != "unnamed-server": + self._log_error( + f"Failed to connect to MCP server {name} " + f"after {max_retries} attempts: {error}" + ) + if server.is_connected: + # Session was established but tool listing failed; tear + # it down so the transport/subprocess doesn't leak. + await server.disconnect() + return False async def disconnect_server(self, name: str) -> bool: """ diff --git a/cecli/mcp/oauth.py b/cecli/mcp/oauth.py index ec5390f77b5..31ff00f0475 100644 --- a/cecli/mcp/oauth.py +++ b/cecli/mcp/oauth.py @@ -15,20 +15,29 @@ from cecli.decoding import safe_open -def create_oauth_callback_server( - port, path="/callback" -) -> Tuple[Callable[[], Awaitable[Tuple[str, str]]], Callable[[], None]]: +def create_oauth_callback_server(port, path="/callback") -> Tuple[ + Callable[[], Awaitable[Tuple[str, str]]], + Callable[[], None], + Callable[[], None], +]: """ Create a local HTTP server to handle OAuth callback. + The listener is started lazily via the returned ``ensure_started`` callable + so servers that never actually trigger OAuth don't leave a daemon HTTP + server (and its bound port) running for the life of the process. + Returns: - Tuple of (async callback handler function, shutdown function) + Tuple of (async callback handler, shutdown function, start function) """ auth_code = None state = None server_error = None callback_received = threading.Event() server = None + server_thread = None + server_started = threading.Event() + start_lock = threading.Lock() class OAuthCallbackHandler(http.server.SimpleHTTPRequestHandler): def do_GET(self): @@ -78,27 +87,71 @@ def do_GET(self): def log_message(self, format, *args): pass - # Start server in a separate thread def start_server(): - nonlocal server + nonlocal server, server_error + srv = None try: - server = socketserver.TCPServer(("localhost", port), OAuthCallbackHandler) - server.serve_forever() + srv = socketserver.TCPServer(("localhost", port), OAuthCallbackHandler) + server = srv + server_started.set() + srv.serve_forever() except Exception as e: - server_error = f"Server error: {e}" # noqa + server_error = f"Server error: {e}" + server_started.set() callback_received.set() + finally: + if srv is not None: + try: + srv.server_close() + except Exception: + pass + + def ensure_started(): + """Start the callback listener once, blocking briefly for the bind.""" + nonlocal server_thread + with start_lock: + if server_started.is_set(): + return + + server_thread = threading.Thread( + target=start_server, daemon=True, name="oauth-callback-server" + ) + server_thread.start() + + # Wait for the bind to complete so the browser redirect cannot race the + # listener coming up. + server_started.wait(timeout=5) - server_thread = threading.Thread(target=start_server, daemon=True) - server_thread.start() - - # Shutdown function def shutdown(): + """Stop the callback listener. Idempotent and safe to call repeatedly. + + ``socketserver.shutdown()`` blocks until ``serve_forever`` exits; async + callers must run this in a thread (see ``asyncio.to_thread``) so a stuck + listener can never wedge the event loop. + """ nonlocal server - if server: - server.shutdown() + with start_lock: + srv = server server = None + if srv is None: + return + + try: + srv.shutdown() + except Exception: + pass + finally: + try: + srv.server_close() + except Exception: + pass + async def get_auth_code() -> Tuple[str, str]: + # Backstop for callers that didn't start the listener via the redirect + # handler first (e.g. a resumed flow). + ensure_started() + # Wait for callback to be received MINUTES = 5 timeout = MINUTES * 60 @@ -106,23 +159,23 @@ async def get_auth_code() -> Tuple[str, str]: start_time = time.time() while not callback_received.is_set(): if time.time() - start_time > timeout: - shutdown() + await asyncio.to_thread(shutdown) raise Exception(f"OAuth callback timed out after {MINUTES} minutes") # Small sleep to avoid busy waiting await asyncio.sleep(0.1) if server_error: - shutdown() + await asyncio.to_thread(shutdown) raise Exception(server_error) if not auth_code: - shutdown() + await asyncio.to_thread(shutdown) raise Exception("No authorization code received") return auth_code, state - return get_auth_code, shutdown + return get_auth_code, shutdown, ensure_started def get_token_file_path(): diff --git a/cecli/mcp/server.py b/cecli/mcp/server.py index a38ab2ebd12..a7bf9c172b5 100644 --- a/cecli/mcp/server.py +++ b/cecli/mcp/server.py @@ -5,6 +5,7 @@ import threading import webbrowser from contextlib import AsyncExitStack +from datetime import timedelta from enum import Enum, auto from urllib.parse import urlparse @@ -28,6 +29,10 @@ MIN_KEEPALIVE_INTERVAL = 5 MAX_KEEPALIVE_INTERVAL = 300 FAILED_PING_THRESHOLD = 3 +# Default per-request timeout (seconds) for MCP handshake/tool calls. Two +# minutes is generous enough for slow first-run bootstraps (uvx/npx/Docker) while +# still bounding a server that accepts a connection but never speaks MCP. +DEFAULT_MCP_REQUEST_TIMEOUT = 120 logger = logging.getLogger(__name__) @@ -246,7 +251,13 @@ async def _open_session(self): stdio_client(server_params, errlog=err_file) ) read, write = stdio_transport - session = await self.exit_stack.enter_async_context(ClientSession(read, write)) + session = await self.exit_stack.enter_async_context( + ClientSession( + read, + write, + read_timeout_seconds=timedelta(seconds=self._request_timeout_seconds()), + ) + ) await session.initialize() self.session = session @@ -323,6 +334,25 @@ async def _run_session(self): self.session = None self._connection_loop = None + def _request_timeout_seconds(self) -> float: + """Per-request timeout (seconds) for the MCP handshake and tool calls. + + A wedged transport would otherwise block startup forever; the timeout is + applied to the SDK ``ClientSession`` so a silent server raises (and can be + retried or reported as failed) instead of hanging. Overridable per server + via the ``timeout`` config key. + """ + raw = self.config.get("timeout") + if raw is not None: + try: + value = float(raw) + if value > 0: + return value + except (TypeError, ValueError): + pass + + return DEFAULT_MCP_REQUEST_TIMEOUT + class HttpBasedMcpServer(McpServer): """Base class for HTTP-based MCP servers (HTTP streaming and SSE).""" @@ -376,12 +406,16 @@ async def _create_oauth_provider(self): redirect_uri = f"http://localhost:{port}/callback" - get_auth_code, shutdown = create_oauth_callback_server(port) + get_auth_code, shutdown, ensure_callback_server = create_oauth_callback_server(port) # Store shutdown function for cleanup self._oauth_shutdown = shutdown async def handle_redirect(auth_url: str) -> None: + # Start the local listener before opening the browser so the OAuth + # redirect can never race the callback server binding. + ensure_callback_server() + if self.io: self.io.tool_output(f"\nAuthentication required for MCP server: {self.name}") self.io.tool_output("\nPlease open this URL in your browser to authenticate:") @@ -440,7 +474,13 @@ async def _open_session(self): read, write = _unpack_transport(transport) - session = await self.exit_stack.enter_async_context(ClientSession(read, write)) + session = await self.exit_stack.enter_async_context( + ClientSession( + read, + write, + read_timeout_seconds=timedelta(seconds=self._request_timeout_seconds()), + ) + ) await session.initialize() self.session = session @@ -580,7 +620,10 @@ async def _close_session(self, cancel_keepalive: bool = True): logger.info(f"Keepalive task stopped for {self.name}") if hasattr(self, "_oauth_shutdown"): - self._oauth_shutdown() + # Run the blocking socketserver shutdown off-loop so a stuck callback + # server can't wedge the MCP event loop (and so the surrounding + # wait_for timeout can still fire). + await asyncio.to_thread(self._oauth_shutdown) self._http_client = None diff --git a/cecli/tui/worker.py b/cecli/tui/worker.py index 9553bbe13d0..e5fa92fdd7e 100644 --- a/cecli/tui/worker.py +++ b/cecli/tui/worker.py @@ -10,6 +10,7 @@ from cecli.coders import Coder from cecli.commands import ReloadProgramSignal, SwitchCoderSignal from cecli.helpers.conversation import ConversationService, MessageTag +from cecli.helpers.coroutines import task_is_cancelling logger = logging.getLogger(__name__) # Suppress asyncio task destroyed warnings during shutdown @@ -60,10 +61,14 @@ def _run_thread(self): try: self.loop.run_until_complete(self._async_run()) - except BaseException: - # Catch anything that could bring down the thread, and just let it exit. - # This includes KeyboardInterrupt, SystemExit, etc. - pass + except BaseException as e: + # A normal stop() stops the loop, which makes run_until_complete + # raise RuntimeError; that and a cancellation after running=False + # are expected shutdown paths, not crashes. + graceful = not self.running and isinstance(e, (asyncio.CancelledError, RuntimeError)) + if not graceful: + logger.error("Coder worker thread stopped unexpectedly", exc_info=e) + self._notify_crash(e) finally: self._cleanup_loop() @@ -117,6 +122,14 @@ async def _async_run(self): if mcp_manager is not None: try: await mcp_manager.connect_all() + except asyncio.CancelledError: + # connect_all uses gather; a single transport (e.g. MCP's + # streamable-HTTP) can surface an ordinary connection failure + # as CancelledError and abort the whole gather. Only propagate + # a genuine cancellation of this worker. + if task_is_cancelling(): + raise + logger.warning("MCP connect_all was cancelled by a server transport; continuing") except Exception as e: logger.error("Failed to connect MCP servers in worker: %s", e, exc_info=True) @@ -285,6 +298,23 @@ def stop(self): if self.thread and self.thread.is_alive(): self.thread.join(timeout=2.0) + def _notify_crash(self, exc): + """Tell the TUI the worker died so it can surface the error and exit. + + Without this the TUI keeps running with a dead worker and appears hung. + """ + try: + self.output_queue.put( + { + "type": "error", + "message": f"Worker stopped unexpectedly: {exc!r}", + "coder_uuid": getattr(self.coder, "uuid", None), + } + ) + self.output_queue.put({"type": "exit"}) + except Exception: + pass + def _create_event_loop(self): """Create the event loop used by the coder worker thread. diff --git a/cecli/website/docs/config/mcp.md b/cecli/website/docs/config/mcp.md index 9bd5e516673..8af15affee2 100644 --- a/cecli/website/docs/config/mcp.md +++ b/cecli/website/docs/config/mcp.md @@ -30,6 +30,37 @@ mcp-servers: keepalive_interval: 60 # Send a heartbeat every 60 seconds ``` +### Request Timeout + +`timeout`: (Optional) The per-request timeout in **seconds** for a server's connection handshake and subsequent requests. This keeps a server that accepts a connection but never responds from blocking cecli indefinitely at startup. + +- `timeout`: (Optional) A positive number of seconds. + - If not provided, it defaults to **120** seconds (2 minutes) — generous enough for slow first-run bootstraps while still bounding a server that never responds. + - It bounds the `initialize` handshake, the initial tool listing (`list_tools`), and later requests made on that server's session. + +A server that stays silent past its timeout is marked as failed to connect (it is retried up to **3** times) instead of hanging cecli, so increase `timeout` for servers that legitimately take a while to start — for example a first-run `uvx`/`npx` download or a Docker image pull. + +Example with a longer timeout: + +```yaml +mcp-servers: + mcpServers: + serena: + transport: stdio + command: uvx + args: [ + "--from", + "git+https://github.com/oraios/serena", + "serena", + "start-mcp-server", + "--context", + "ide", + "--project", + "/path/to/project" + ] + timeout: 300 # Allow extra time for a first-run build +``` + You have two ways of sharing your MCP server configuration with cecli. diff --git a/tests/mcp/test_manager_retry.py b/tests/mcp/test_manager_retry.py index 850938702cd..734d4807245 100644 --- a/tests/mcp/test_manager_retry.py +++ b/tests/mcp/test_manager_retry.py @@ -323,15 +323,48 @@ async def test_connect_server_propagates_cancelled_error_during_retry(mock_serve @pytest.mark.asyncio -async def test_connect_server_propagates_cancelled_error_during_connect(mock_server, mock_io): - """TC-008: connect_server re-raises CancelledError when server.connect() raises it.""" +async def test_connect_server_treats_transport_cancellation_as_failure(mock_server, mock_io): + """TC-008: a transport-level CancelledError is a failed attempt, not a propagated cancel. + + MCP's streamable-HTTP transport surfaces an unreachable server as + CancelledError from its anyio TaskGroup; it must not abort startup or the + caller, so it is retried and reported like any other connection failure. + """ manager = McpServerManager(servers=[mock_server], io=mock_io) mock_server.connect.side_effect = asyncio.CancelledError() + with patch("asyncio.sleep"): + result = await manager.connect_server("test-server") + + assert result is False + assert mock_server.connect.call_count == 3 + assert mock_io.tool_error.call_count == 1 + + +# --------------------------------------------------------------------------- +# TC-008b: connect_server propagates cancellation of the calling task +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_propagates_genuine_cancellation(mock_server, mock_io): + """TC-008b: cancelling the calling task still propagates out of connect_server.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + + started = asyncio.Event() + + async def _slow_connect(): + started.set() + await asyncio.sleep(3600) + + mock_server.connect.side_effect = _slow_connect + task = asyncio.create_task(manager.connect_server("test-server")) + await started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): - await manager.connect_server("test-server") + await task - assert mock_server.connect.call_count == 1 mock_io.tool_error.assert_not_called() From aead8d1dd234cce963a35561547230fc9bc8accc Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Tue, 15 Sep 2026 02:33:40 -0400 Subject: [PATCH 04/14] Update cancellation expectations for python 3.10 --- cecli/helpers/coroutines.py | 7 ++++++- tests/mcp/test_manager_retry.py | 9 +++++++++ 2 files changed, 15 insertions(+), 1 deletion(-) diff --git a/cecli/helpers/coroutines.py b/cecli/helpers/coroutines.py index 0d2ae4c8294..3676d676b55 100644 --- a/cecli/helpers/coroutines.py +++ b/cecli/helpers/coroutines.py @@ -117,6 +117,11 @@ def task_is_cancelling() -> bool: Used to tell a genuine cancellation of the caller apart from cancellation errors that transports (e.g. MCP's anyio TaskGroups) surface for ordinary connection failures. + + Reliable on Python 3.11+ (``Task.cancelling()``). On 3.10 there is no public + signal: ``_must_cancel`` is already cleared by the time the CancelledError is + delivered, so this is best-effort and usually reports False. Callers must + therefore treat False as "not known to be cancelling" rather than proof. """ task = asyncio.current_task() if task is None: @@ -126,5 +131,5 @@ def task_is_cancelling() -> bool: if cancelling is not None: return cancelling() > 0 - # Python 3.10 has no Task.cancelling(); fall back to the private flag. + # Python 3.10 fallback: best-effort only (see docstring). return bool(getattr(task, "_must_cancel", False)) diff --git a/tests/mcp/test_manager_retry.py b/tests/mcp/test_manager_retry.py index 734d4807245..976a7eac2d2 100644 --- a/tests/mcp/test_manager_retry.py +++ b/tests/mcp/test_manager_retry.py @@ -15,6 +15,7 @@ """ import asyncio +import sys from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch @@ -346,6 +347,14 @@ async def test_connect_server_treats_transport_cancellation_as_failure(mock_serv # --------------------------------------------------------------------------- +# Python 3.10 cannot tell a real cancellation from a transport-level one: it has no +# Task.cancelling(), _must_cancel is already cleared when the handler runs, and the +# traceback loses the origin frame. task_is_cancelling() therefore reports False +# there, so this propagation guarantee only holds on 3.11+. +@pytest.mark.skipif( + sys.version_info < (3, 11), + reason="Task.cancelling() is required to distinguish a real cancel from a transport cancel", +) @pytest.mark.asyncio async def test_connect_server_propagates_genuine_cancellation(mock_server, mock_io): """TC-008b: cancelling the calling task still propagates out of connect_server.""" From 24cf3948ae71e006992888d9d547edd9588c417f Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Tue, 15 Sep 2026 02:59:18 -0400 Subject: [PATCH 05/14] Make MCP connection timeout changes more explicit/idiomatic --- cecli/helpers/coroutines.py | 6 +++--- cecli/mcp/manager.py | 14 ++++++++----- cecli/mcp/oauth.py | 35 ++++++++++++++++++--------------- cecli/mcp/server.py | 39 ++++++++++++++++++++----------------- cecli/tui/worker.py | 15 +++++++++----- 5 files changed, 62 insertions(+), 47 deletions(-) diff --git a/cecli/helpers/coroutines.py b/cecli/helpers/coroutines.py index 3676d676b55..3d4b4b46089 100644 --- a/cecli/helpers/coroutines.py +++ b/cecli/helpers/coroutines.py @@ -127,9 +127,9 @@ def task_is_cancelling() -> bool: if task is None: return False - cancelling = getattr(task, "cancelling", None) - if cancelling is not None: - return cancelling() > 0 + cancelling_fn = getattr(task, "cancelling", None) + if cancelling_fn is not None: + return cancelling_fn() > 0 # Python 3.10 fallback: best-effort only (see docstring). return bool(getattr(task, "_must_cancel", False)) diff --git a/cecli/mcp/manager.py b/cecli/mcp/manager.py index a2216e0267b..2063c09fc77 100644 --- a/cecli/mcp/manager.py +++ b/cecli/mcp/manager.py @@ -4,6 +4,10 @@ from cecli.mcp.server import DEFAULT_MCP_REQUEST_TIMEOUT, LocalServer, McpServer from cecli.tools.utils.registry import ToolRegistry +# Slack added on top of a server's request timeout for the connect/list_tools +# backstop, covering time the transport spends outside the SDK's read timeout. +CONNECT_BACKSTOP_GRACE_SECONDS = 5 + class McpServerManager: """ @@ -204,7 +208,8 @@ async def connect_server(self, name: str) -> bool: # When io is None (e.g., during from_servers before IO is assigned), # _log_warning and _log_error silently return — retries still happen # but with no user-visible feedback. This is intentional. - max_retries = 3 if server.name != "unnamed-server" else 1 + is_unnamed = server.name == "unnamed-server" + max_retries = 1 if is_unnamed else 3 delay = 1.0 backoff = 2.0 max_delay = 30.0 @@ -219,10 +224,9 @@ async def connect_server(self, name: str) -> bool: if base_timeout <= 0: base_timeout = DEFAULT_MCP_REQUEST_TIMEOUT - attempt_timeout = base_timeout + 5 + attempt_timeout = base_timeout + CONNECT_BACKSTOP_GRACE_SECONDS for attempt in range(1, max_retries + 1): - error = None try: session = await asyncio.wait_for(server.connect(), timeout=attempt_timeout) @@ -239,7 +243,7 @@ async def connect_server(self, name: str) -> bool: except Exception as e: error = e - if attempt < max_retries and server.name != "unnamed-server": + if attempt < max_retries: self._log_warning( f"Connection attempt {attempt} failed for {name}, " f"retrying in {delay}s... ({error})" @@ -247,7 +251,7 @@ async def connect_server(self, name: str) -> bool: await asyncio.sleep(delay) delay = min(delay * backoff, max_delay) else: - if server.name != "unnamed-server": + if not is_unnamed: self._log_error( f"Failed to connect to MCP server {name} " f"after {max_retries} attempts: {error}" diff --git a/cecli/mcp/oauth.py b/cecli/mcp/oauth.py index 31ff00f0475..5804e9bdd95 100644 --- a/cecli/mcp/oauth.py +++ b/cecli/mcp/oauth.py @@ -14,6 +14,10 @@ from cecli.decoding import safe_open +# How long ``ensure_started`` waits for the callback listener to bind before +# giving up so the browser redirect can't race the listener coming up. +CALLBACK_BIND_TIMEOUT_SECONDS = 5 + def create_oauth_callback_server(port, path="/callback") -> Tuple[ Callable[[], Awaitable[Tuple[str, str]]], @@ -35,7 +39,6 @@ def create_oauth_callback_server(port, path="/callback") -> Tuple[ server_error = None callback_received = threading.Event() server = None - server_thread = None server_started = threading.Event() start_lock = threading.Lock() @@ -100,27 +103,19 @@ def start_server(): server_started.set() callback_received.set() finally: - if srv is not None: - try: - srv.server_close() - except Exception: - pass + _close_quietly(srv) def ensure_started(): """Start the callback listener once, blocking briefly for the bind.""" - nonlocal server_thread with start_lock: if server_started.is_set(): return - server_thread = threading.Thread( - target=start_server, daemon=True, name="oauth-callback-server" - ) - server_thread.start() + threading.Thread(target=start_server, daemon=True, name="oauth-callback-server").start() # Wait for the bind to complete so the browser redirect cannot race the # listener coming up. - server_started.wait(timeout=5) + server_started.wait(timeout=CALLBACK_BIND_TIMEOUT_SECONDS) def shutdown(): """Stop the callback listener. Idempotent and safe to call repeatedly. @@ -142,10 +137,7 @@ def shutdown(): except Exception: pass finally: - try: - srv.server_close() - except Exception: - pass + _close_quietly(srv) async def get_auth_code() -> Tuple[str, str]: # Backstop for callers that didn't start the listener via the redirect @@ -278,3 +270,14 @@ async def set_client_info(self, client_info: OAuthClientInformationFull) -> None all_tokens[self.server_name]["client_info"] = json.loads(client_info.model_dump_json()) save_mcp_oauth_tokens(all_tokens) + + +def _close_quietly(server) -> None: + """Close an HTTP server socket, ignoring a missing server or close errors.""" + if server is None: + return + + try: + server.server_close() + except Exception: + pass diff --git a/cecli/mcp/server.py b/cecli/mcp/server.py index a7bf9c172b5..c3758e76a55 100644 --- a/cecli/mcp/server.py +++ b/cecli/mcp/server.py @@ -251,15 +251,7 @@ async def _open_session(self): stdio_client(server_params, errlog=err_file) ) read, write = stdio_transport - session = await self.exit_stack.enter_async_context( - ClientSession( - read, - write, - read_timeout_seconds=timedelta(seconds=self._request_timeout_seconds()), - ) - ) - await session.initialize() - self.session = session + session = await self._enter_client_session(read, write) return session @@ -334,6 +326,25 @@ async def _run_session(self): self.session = None self._connection_loop = None + async def _enter_client_session(self, read, write): + """Enter a client session on the shared exit stack and initialize it. + + Builds the SDK ``ClientSession`` with the server's configured request + timeout so the stdio and HTTP transports share identical timeout + behavior, then stores the initialized session on ``self.session``. + """ + session = await self.exit_stack.enter_async_context( + ClientSession( + read, + write, + read_timeout_seconds=timedelta(seconds=self._request_timeout_seconds()), + ) + ) + await session.initialize() + self.session = session + + return session + def _request_timeout_seconds(self) -> float: """Per-request timeout (seconds) for the MCP handshake and tool calls. @@ -474,15 +485,7 @@ async def _open_session(self): read, write = _unpack_transport(transport) - session = await self.exit_stack.enter_async_context( - ClientSession( - read, - write, - read_timeout_seconds=timedelta(seconds=self._request_timeout_seconds()), - ) - ) - await session.initialize() - self.session = session + session = await self._enter_client_session(read, write) await self.start_keepalive() diff --git a/cecli/tui/worker.py b/cecli/tui/worker.py index e5fa92fdd7e..b1602657e7f 100644 --- a/cecli/tui/worker.py +++ b/cecli/tui/worker.py @@ -62,11 +62,7 @@ def _run_thread(self): try: self.loop.run_until_complete(self._async_run()) except BaseException as e: - # A normal stop() stops the loop, which makes run_until_complete - # raise RuntimeError; that and a cancellation after running=False - # are expected shutdown paths, not crashes. - graceful = not self.running and isinstance(e, (asyncio.CancelledError, RuntimeError)) - if not graceful: + if not self._is_graceful_shutdown(e): logger.error("Coder worker thread stopped unexpectedly", exc_info=e) self._notify_crash(e) finally: @@ -298,6 +294,15 @@ def stop(self): if self.thread and self.thread.is_alive(): self.thread.join(timeout=2.0) + def _is_graceful_shutdown(self, exc) -> bool: + """Return True when a worker-loop exception is an expected shutdown artifact. + + A normal stop() stops the loop, which surfaces from run_until_complete + as RuntimeError; a cancellation after running goes False surfaces as + CancelledError. Neither is a crash. + """ + return not self.running and isinstance(exc, (asyncio.CancelledError, RuntimeError)) + def _notify_crash(self, exc): """Tell the TUI the worker died so it can surface the error and exit. From 51e92ff6518bd1f937c205973813822080109cc8 Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Thu, 17 Sep 2026 09:16:27 -0400 Subject: [PATCH 06/14] Omit temperature entirely for github copilot provider --- cecli/helpers/model_config/agent.py | 11 +++++++++-- cecli/models.py | 6 ++++-- tests/basic/test_models.py | 12 ++++++++++++ tests/helpers/test_model_config.py | 17 +++++++++++++++++ 4 files changed, 42 insertions(+), 4 deletions(-) diff --git a/cecli/helpers/model_config/agent.py b/cecli/helpers/model_config/agent.py index 863bf13c251..f5e9615f89b 100644 --- a/cecli/helpers/model_config/agent.py +++ b/cecli/helpers/model_config/agent.py @@ -9,7 +9,7 @@ from typing import Dict, Optional -from .identifiers import is_anthropic +from .identifiers import is_anthropic, is_github_copilot from .utils import supports_reasoning @@ -67,7 +67,14 @@ def derive_agent_config(provider: Optional[str], route: str, record: Optional[Di "uses_messages_api": uses_messages_api, } - if reasoning or record.get("supports_adaptive_thinking"): + if ( + reasoning + or record.get("supports_adaptive_thinking") + or is_github_copilot(provider, route, record) + ): + # Reasoning, adaptive and GitHub Copilot models do not take an explicit + # sampling temperature; omit the parameter entirely rather than sending + # the generic default of 0. agent["use_temperature"] = False return agent diff --git a/cecli/models.py b/cecli/models.py index 351d689f574..04d5e199474 100644 --- a/cecli/models.py +++ b/cecli/models.py @@ -1348,8 +1348,10 @@ async def send_completion( temperature = float(self.use_temperature) kwargs["temperature"] = temperature else: - if override_kwargs and override_kwargs.get("temperature", None): - override_kwargs.pop("temperature", None) + # Omit temperature entirely when the model does not use it; the + # key must be dropped even when its override value is falsy (0). + if override_kwargs and "temperature" in override_kwargs: + override_kwargs.pop("temperature") effective_tools = tools diff --git a/tests/basic/test_models.py b/tests/basic/test_models.py index 748f8b8072b..a7357abe027 100644 --- a/tests/basic/test_models.py +++ b/tests/basic/test_models.py @@ -511,6 +511,18 @@ def test_use_temperature_settings(self): model.use_temperature = 0.7 assert model.use_temperature == 0.7 + @patch("cecli.models.litellm.acompletion") + async def test_use_temperature_false_omits_temperature(self, mock_completion): + model = Model("github/o1-mini") + model.extra_params = {} + messages = [{"role": "user", "content": "Hello"}] + + await model.send_completion( + messages, functions=None, stream=False, override_kwargs={"temperature": 0} + ) + + assert "temperature" not in mock_completion.call_args.kwargs + @patch("cecli.models.litellm.acompletion") async def test_request_timeout_default(self, mock_completion): model = Model("gpt-4") diff --git a/tests/helpers/test_model_config.py b/tests/helpers/test_model_config.py index b228cae5a2a..80f2a90e37a 100644 --- a/tests/helpers/test_model_config.py +++ b/tests/helpers/test_model_config.py @@ -634,6 +634,23 @@ def test_non_gpt_github_copilot_stays_chat(): assert config["llm"]["mode"] == "chat" +def test_github_copilot_models_omit_temperature(): + """Copilot models never send a sampling temperature, even in chat mode. + + Copilot models are absent from the individual model configs, so the + metadata-derived agent block supplies the rule. + """ + record = _record( + litellm_provider="github_copilot", + supports_reasoning=False, + supported_endpoints=["/v1/chat/completions"], + ) + config = get_default_config("github_copilot/gpt-4o", [{"github_copilot/gpt-4o": record}]) + + assert config["llm"]["mode"] == "chat" + assert config["agent"]["use_temperature"] is False + + def test_adaptive_thinking_sets_use_temperature_false(): record = _record(supports_reasoning=False, supports_adaptive_thinking=True) config = get_default_config("adaptive", [{"adaptive": record}]) From cccbb1a30a08a86b9a5b780377ed194be0f7382b Mon Sep 17 00:00:00 2001 From: Philippe Back Date: Fri, 18 Sep 2026 01:10:31 +0200 Subject: [PATCH 07/14] fix: safely handle non-string message content in tokens command Co-authored-by: cecli (gemini/gemini-3.8-flash) --- cecli/commands/tokens.py | 94 +++++++++++++++------------- tests/basic/test_commands.py | 117 +++++++++++++++++++++++++++++++++++ 2 files changed, 168 insertions(+), 43 deletions(-) diff --git a/cecli/commands/tokens.py b/cecli/commands/tokens.py index 1ef30e6e825..349a89e00b2 100644 --- a/cecli/commands/tokens.py +++ b/cecli/commands/tokens.py @@ -1,4 +1,4 @@ -from typing import List +from typing import Any, Dict, List, Optional from cecli.commands.utils.base_command import BaseCommand from cecli.commands.utils.helpers import format_command_result @@ -9,6 +9,48 @@ class TokensCommand(BaseCommand): NORM_NAME = "tokens" DESCRIPTION = "Report on the number of tokens used by the current chat context" + @classmethod + def _extract_file_name(cls, msg: Dict[str, Any]) -> Optional[str]: + """Extract file name from a message dictionary.""" + if not isinstance(msg, dict): + return None + + # Check explicit image_file metadata first + fname = msg.get("image_file") + if fname: + return fname + + content = msg.get("content") + if isinstance(content, str): + if content.startswith(("Original File Contents For", "Current File Contents For")): + lines = content.split("\n", 3) + if len(lines) > 1: + return lines[1].strip() + elif content.startswith("Image file: "): + return content[len("Image file: "):].strip() + elif isinstance(content, list): + for part in content: + if isinstance(part, dict): + if part.get("image_file"): + return part.get("image_file") + text = part.get("text") + if isinstance(text, str): + if text.startswith(("Original File Contents For", "Current File Contents For")): + lines = text.split("\n", 3) + if len(lines) > 1: + return lines[1].strip() + elif text.startswith("Image file: "): + return text[len("Image file: "):].strip() + elif isinstance(part, str): + if part.startswith(("Original File Contents For", "Current File Contents For")): + lines = part.split("\n", 3) + if len(lines) > 1: + return lines[1].strip() + elif part.startswith("Image file: "): + return part[len("Image file: "):].strip() + + return None + @classmethod async def execute(cls, io, coder, args, **kwargs): res = [] @@ -123,27 +165,10 @@ async def execute(cls, io, coder, args, **kwargs): # Group messages by file (each file has user and assistant messages) file_tokens = {} for msg in readonly_msgs: - # Extract file name from message content - content = msg.get("content", "") - if content.startswith("Original File Contents For"): - # Extract file path from "File Contents {path}:" - lines = content.split("\n", 3) - if lines: - file_line = lines[1] - fname = file_line.strip() - # Calculate tokens for this message - tokens = coder.main_model.token_count([msg]) - if fname not in file_tokens: - file_tokens[fname] = 0 - file_tokens[fname] += tokens - elif "image_file" in msg: - # Handle image files - fname = msg.get("image_file") - if fname: - tokens = coder.main_model.token_count([msg]) - if fname not in file_tokens: - file_tokens[fname] = 0 - file_tokens[fname] += tokens + fname = cls._extract_file_name(msg) + if fname: + tokens = coder.main_model.token_count([msg]) + file_tokens[fname] = file_tokens.get(fname, 0) + tokens # Add to results for fname, tokens in file_tokens.items(): @@ -158,27 +183,10 @@ async def execute(cls, io, coder, args, **kwargs): msgs = ConversationService.get_manager(coder).get_messages_dict(tag=tag) if msgs: for msg in msgs: - # Extract file name from message content - content = msg.get("content", "") - if content.startswith("Original File Contents For"): - # Extract file path from "File Contents {path}:" - lines = content.split("\n", 3) - if lines: - file_line = lines[1] - fname = file_line.strip() - # Calculate tokens for this message - tokens = coder.main_model.token_count([msg]) - if fname not in editable_file_tokens: - editable_file_tokens[fname] = 0 - editable_file_tokens[fname] += tokens - elif "image_file" in msg: - # Handle image files - fname = msg.get("image_file") - if fname: - tokens = coder.main_model.token_count([msg]) - if fname not in editable_file_tokens: - editable_file_tokens[fname] = 0 - editable_file_tokens[fname] += tokens + fname = cls._extract_file_name(msg) + if fname: + tokens = coder.main_model.token_count([msg]) + editable_file_tokens[fname] = editable_file_tokens.get(fname, 0) + tokens # Add editable files to results for fname, tokens in editable_file_tokens.items(): diff --git a/tests/basic/test_commands.py b/tests/basic/test_commands.py index c367cafefd1..b20d723036f 100644 --- a/tests/basic/test_commands.py +++ b/tests/basic/test_commands.py @@ -420,6 +420,123 @@ async def test_cmd_tokens(self): self.assertIn("foo.txt", console_output) self.assertIn("bar.txt", console_output) + async def test_cmd_tokens_with_image_list_content(self): + from cecli.commands.tokens import TokensCommand + from cecli.helpers.conversation import ConversationService, MessageTag + + # Direct unit tests for _extract_file_name + self.assertIsNone(TokensCommand._extract_file_name(None)) + self.assertIsNone(TokensCommand._extract_file_name("not a dict")) + self.assertIsNone(TokensCommand._extract_file_name({})) + self.assertIsNone(TokensCommand._extract_file_name({"content": None})) + self.assertIsNone(TokensCommand._extract_file_name({"content": "Just a normal message"})) + + # String content + self.assertEqual( + TokensCommand._extract_file_name({ + "content": "Original File Contents For:\n/path/to/file.py\n\ncode..." + }), + "/path/to/file.py", + ) + self.assertEqual( + TokensCommand._extract_file_name({ + "content": "Current File Contents For:\n/path/to/file2.py\n\ncode..." + }), + "/path/to/file2.py", + ) + self.assertEqual( + TokensCommand._extract_file_name({"content": "Image file: photo.png"}), + "photo.png", + ) + + # image_file key + self.assertEqual( + TokensCommand._extract_file_name({"image_file": "photo.jpg"}), + "photo.jpg", + ) + self.assertEqual( + TokensCommand._extract_file_name({ + "image_file": "photo.jpg", + "content": [{"type": "image_url", "image_url": {}}], + }), + "photo.jpg", + ) + + # List content + self.assertEqual( + TokensCommand._extract_file_name({ + "content": [ + {"type": "text", "text": "Image file: nested.png"}, + {"type": "image_url", "image_url": {}}, + ] + }), + "nested.png", + ) + self.assertEqual( + TokensCommand._extract_file_name({ + "content": [ + {"type": "text", "text": "Original File Contents For:\nmodule.py\n\ncode"} + ] + }), + "module.py", + ) + self.assertEqual( + TokensCommand._extract_file_name({ + "content": [{"image_file": "from_part.png"}] + }), + "from_part.png", + ) + self.assertEqual( + TokensCommand._extract_file_name({ + "content": ["Image file: string_part.png"] + }), + "string_part.png", + ) + + # Integration test with Coder / Commands execution + io = InputOutput(pretty=False, fancy_input=False, yes=True) + coder = await Coder.create(self.GPT35, None, io) + commands = Commands(io, coder) + + manager = ConversationService.get_manager(coder) + img_msg = { + "role": "user", + "content": [ + {"type": "text", "text": "Image file: test_image.png"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,123"}}, + ], + "image_file": "test_image.png", + } + manager.add_message( + message_dict=img_msg, + tag=MessageTag.READONLY_FILES, + hash_key=("image_user", "test_image.png"), + ) + + img_msg_no_attr = { + "role": "user", + "content": [ + {"type": "text", "text": "Image file: another_image.png"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,456"}}, + ], + } + manager.add_message( + message_dict=img_msg_no_attr, + tag=MessageTag.CHAT_FILES, + hash_key=("image_user", "another_image.png"), + ) + + stdout = StringIO() + sys.stdout = stdout + try: + commands.execute("tokens", "") + finally: + sys.stdout = sys.__stdout__ + + console_output = stdout.getvalue() + self.assertIn("test_image.png", console_output) + self.assertIn("another_image.png", console_output) + async def test_cmd_add_from_subdir(self): repo = git.Repo.init() repo.config_writer().set_value("user", "name", "Test User").release() From 00f45cc721c1e022fca90ae83ce7d3f2c5117913 Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Fri, 18 Sep 2026 09:17:09 -0400 Subject: [PATCH 08/14] Add binary data estimation and tool schemas to `/tokens` output --- cecli/commands/tokens.py | 112 ++++++++++++++++++++++++++++++++++++--- 1 file changed, 106 insertions(+), 6 deletions(-) diff --git a/cecli/commands/tokens.py b/cecli/commands/tokens.py index 349a89e00b2..7d56c71422e 100644 --- a/cecli/commands/tokens.py +++ b/cecli/commands/tokens.py @@ -1,3 +1,4 @@ +import json from typing import Any, Dict, List, Optional from cecli.commands.utils.base_command import BaseCommand @@ -27,7 +28,7 @@ def _extract_file_name(cls, msg: Dict[str, Any]) -> Optional[str]: if len(lines) > 1: return lines[1].strip() elif content.startswith("Image file: "): - return content[len("Image file: "):].strip() + return content[len("Image file: ") :].strip() elif isinstance(content, list): for part in content: if isinstance(part, dict): @@ -35,22 +36,114 @@ def _extract_file_name(cls, msg: Dict[str, Any]) -> Optional[str]: return part.get("image_file") text = part.get("text") if isinstance(text, str): - if text.startswith(("Original File Contents For", "Current File Contents For")): + if text.startswith( + ("Original File Contents For", "Current File Contents For") + ): lines = text.split("\n", 3) if len(lines) > 1: return lines[1].strip() elif text.startswith("Image file: "): - return text[len("Image file: "):].strip() + return text[len("Image file: ") :].strip() elif isinstance(part, str): if part.startswith(("Original File Contents For", "Current File Contents For")): lines = part.split("\n", 3) if len(lines) > 1: return lines[1].strip() elif part.startswith("Image file: "): - return part[len("Image file: "):].strip() + return part[len("Image file: ") :].strip() return None + @staticmethod + def binary_token_count(data_url: str, bytes_per_pixel: float = 1.0) -> int: + """Estimate OpenAI vision tokens for a base64 media data URL. + + Assumes a 4:3 aspect ratio and 1 byte per pixel, which is a reasonable middle + ground for binary media when the true dimensions are unknown. The tile cost is + inflated by 1.33 and floored to keep the estimate on the pessimistic side. + """ + import math + + b64_data = data_url.split(",", 1)[1] if "," in data_url else data_url + padding = b64_data.count("=") + rough_bytes = max(0, (len(b64_data) * 3 // 4) - padding) + total_pixels = rough_bytes / bytes_per_pixel + height = math.sqrt((3 / 4) * total_pixels) + width = (4 / 3) * height + + if max(width, height) > 2048: + scale = 2048 / max(width, height) + width *= scale + height *= scale + + if min(width, height) > 768: + scale = 768 / min(width, height) + width *= scale + height *= scale + + tiles = math.ceil(width / 512) * math.ceil(height / 512) + + return math.floor((85 + (170 * tiles)) * 1.33) + + @classmethod + def _extract_key_paths(cls, obj: Any, prefix: str = "") -> List[str]: + """Return the dot-separated paths of every leaf in a nested structure. + + List indices are included as numeric segments, e.g. "0.image_url.url". + """ + if isinstance(obj, dict): + items = obj.items() + elif isinstance(obj, list): + items = enumerate(obj) + else: + return [prefix] if prefix else [] + + paths = [] + for key, value in items: + path = f"{prefix}.{key}" if prefix else str(key) + paths.extend(cls._extract_key_paths(value, path)) + + return paths + + @classmethod + def _get_path_value(cls, obj: Any, path: str) -> Any: + """Resolve a dot-separated path produced by _extract_key_paths.""" + for part in path.split("."): + obj = obj[int(part)] if isinstance(obj, list) else obj[part] + + return obj + + @classmethod + def _count_value_tokens(cls, coder, key: str, value: Any) -> int: + """Count tokens for an extracted value, estimating data URLs as binary media.""" + if key == "url" and isinstance(value, str) and value.startswith("data:"): + return cls.binary_token_count(value) + + return coder.main_model.token_count(value) + + @classmethod + def _count_message_tokens(cls, coder, msg: Dict[str, Any]) -> int: + """Count tokens for a message, estimating binary data URLs instead of raw base64.""" + content = msg.get("content") if isinstance(msg, dict) else None + + if not isinstance(content, (list, dict)): + return coder.main_model.token_count([msg]) + + total = 0 + + for path in cls._extract_key_paths(content): + key = path.rsplit(".", 1)[-1] + + if not key.isdigit() and key not in ("text", "content", "url"): + continue + + value = cls._get_path_value(content, path) + + if isinstance(value, str): + total += cls._count_value_tokens(coder, key, value) + + return total + @classmethod async def execute(cls, io, coder, args, **kwargs): res = [] @@ -97,6 +190,13 @@ async def execute(cls, io, coder, args, **kwargs): res.append((system_tokens, "system messages", "")) + # tool definitions + tool_list = coder.get_tool_list() + + if tool_list: + tokens_tools = coder.main_model.token_count(json.dumps(tool_list)) + res.append((tokens_tools, "tool schemas", "")) + # chat history msgs_done = ConversationService.get_manager(coder).get_messages_dict(tag=MessageTag.DONE) msgs_cur = ConversationService.get_manager(coder).get_messages_dict(tag=MessageTag.CUR) @@ -167,7 +267,7 @@ async def execute(cls, io, coder, args, **kwargs): for msg in readonly_msgs: fname = cls._extract_file_name(msg) if fname: - tokens = coder.main_model.token_count([msg]) + tokens = cls._count_message_tokens(coder, msg) file_tokens[fname] = file_tokens.get(fname, 0) + tokens # Add to results @@ -185,7 +285,7 @@ async def execute(cls, io, coder, args, **kwargs): for msg in msgs: fname = cls._extract_file_name(msg) if fname: - tokens = coder.main_model.token_count([msg]) + tokens = cls._count_message_tokens(coder, msg) editable_file_tokens[fname] = editable_file_tokens.get(fname, 0) + tokens # Add editable files to results From 37fb15c790684f237e3e7973bf8e3144e73f55d4 Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Fri, 18 Sep 2026 09:18:49 -0400 Subject: [PATCH 09/14] Fix formatting --- tests/basic/test_commands.py | 56 +++++++++++++++++++----------------- 1 file changed, 29 insertions(+), 27 deletions(-) diff --git a/tests/basic/test_commands.py b/tests/basic/test_commands.py index b20d723036f..a64d7277f23 100644 --- a/tests/basic/test_commands.py +++ b/tests/basic/test_commands.py @@ -433,15 +433,15 @@ async def test_cmd_tokens_with_image_list_content(self): # String content self.assertEqual( - TokensCommand._extract_file_name({ - "content": "Original File Contents For:\n/path/to/file.py\n\ncode..." - }), + TokensCommand._extract_file_name( + {"content": "Original File Contents For:\n/path/to/file.py\n\ncode..."} + ), "/path/to/file.py", ) self.assertEqual( - TokensCommand._extract_file_name({ - "content": "Current File Contents For:\n/path/to/file2.py\n\ncode..." - }), + TokensCommand._extract_file_name( + {"content": "Current File Contents For:\n/path/to/file2.py\n\ncode..."} + ), "/path/to/file2.py", ) self.assertEqual( @@ -455,41 +455,43 @@ async def test_cmd_tokens_with_image_list_content(self): "photo.jpg", ) self.assertEqual( - TokensCommand._extract_file_name({ - "image_file": "photo.jpg", - "content": [{"type": "image_url", "image_url": {}}], - }), + TokensCommand._extract_file_name( + { + "image_file": "photo.jpg", + "content": [{"type": "image_url", "image_url": {}}], + } + ), "photo.jpg", ) # List content self.assertEqual( - TokensCommand._extract_file_name({ - "content": [ - {"type": "text", "text": "Image file: nested.png"}, - {"type": "image_url", "image_url": {}}, - ] - }), + TokensCommand._extract_file_name( + { + "content": [ + {"type": "text", "text": "Image file: nested.png"}, + {"type": "image_url", "image_url": {}}, + ] + } + ), "nested.png", ) self.assertEqual( - TokensCommand._extract_file_name({ - "content": [ - {"type": "text", "text": "Original File Contents For:\nmodule.py\n\ncode"} - ] - }), + TokensCommand._extract_file_name( + { + "content": [ + {"type": "text", "text": "Original File Contents For:\nmodule.py\n\ncode"} + ] + } + ), "module.py", ) self.assertEqual( - TokensCommand._extract_file_name({ - "content": [{"image_file": "from_part.png"}] - }), + TokensCommand._extract_file_name({"content": [{"image_file": "from_part.png"}]}), "from_part.png", ) self.assertEqual( - TokensCommand._extract_file_name({ - "content": ["Image file: string_part.png"] - }), + TokensCommand._extract_file_name({"content": ["Image file: string_part.png"]}), "string_part.png", ) From f40f954f008c713c430c3eb29c53afab92dd9d92 Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Fri, 18 Sep 2026 09:34:57 -0400 Subject: [PATCH 10/14] Add `alt+v` for image pasting in terminal --- cecli/tui/app.py | 41 +++++++++++++++ cecli/tui/widgets/input_area.py | 5 ++ cecli/website/docs/config/tui.md | 8 ++- tests/tui/test_app.py | 86 ++++++++++++++++++++++++++++++++ 4 files changed, 138 insertions(+), 2 deletions(-) diff --git a/cecli/tui/app.py b/cecli/tui/app.py index 7e5dfd32af3..d20e8a4f311 100644 --- a/cecli/tui/app.py +++ b/cecli/tui/app.py @@ -201,6 +201,12 @@ def __init__(self, coder_worker, output_queue, input_queue, args): description="Record Voice", show=True, ) + self.bind( + self._encode_keys(self.get_keys_for("paste")), + "paste_clipboard", + description="Paste Clipboard", + show=True, + ) self.register_theme(BASE_THEME) self.theme = "cecli" @@ -284,6 +290,7 @@ def _get_config(self): "editor": "ctrl+o", "history": "alt+shift+h", "voice": "ctrl+r", + "paste": "alt+v", "focus": "ctrl+f", "cancel": "ctrl+c", "clear": "ctrl+l", @@ -829,6 +836,12 @@ def on_input_area_submit(self, message: InputArea.Submit): self.action_start_voice() return + if stripped == "/paste": + input_area = self.query_one("#input", InputArea) + input_area.value = "" + self.action_paste_clipboard() + return + # Intercept /editor and /edit commands to handle with TUI suspension if ( stripped in ("/editor", "/edit") @@ -1280,6 +1293,11 @@ def action_start_voice(self): self._voice_stopping = False self.run_worker(self._run_voice(coder, self._voice_stop_queue), group="voice") + def action_paste_clipboard(self): + """Paste an image or text from the system clipboard (keyboard shortcut).""" + coder = self._get_visible_coder() + self.run_worker(self._run_paste(coder), group="paste") + def action_open_editor(self): """Open an external editor to compose a prompt (keyboard shortcut).""" # Get current input text to use as initial content @@ -2124,6 +2142,29 @@ async def _run_voice(self, coder, stop_queue): self._voice_stopping = False self.update_key_hints(generating=self._currently_generating) + async def _run_paste(self, coder): + """Run the paste command on a worker so clipboard I/O doesn't block the UI.""" + from cecli.commands.paste import PasteCommand + + try: + await PasteCommand.execute(coder.io, coder, "") + self._refresh_file_list(coder) + except Exception as err: + self.show_error(f"Unable to paste clipboard content: {err}") + + def _refresh_file_list(self, coder): + """Refresh the file list and autocomplete after the coder's chat files change.""" + try: + input_area = self.query_one("#input", InputArea) + files = list(coder.get_addable_relative_files()) + commands = coder.commands.get_commands() if getattr(coder, "commands", None) else [] + input_area.update_autocomplete_data(files, commands) + + file_list = self.query_one("#file-list", FileList) + file_list.update_files() + except Exception: + pass + def patch_color_name_to_rgb(): """Inject Rich 256-color names into Textual's COLOR_NAME_TO_RGB dict. diff --git a/cecli/tui/widgets/input_area.py b/cecli/tui/widgets/input_area.py index f3c2867783e..eecfd66c3e5 100644 --- a/cecli/tui/widgets/input_area.py +++ b/cecli/tui/widgets/input_area.py @@ -247,6 +247,11 @@ def on_key(self, event) -> None: self.app.action_start_voice() return + if self.app.is_key_for("paste", event.key): + event.stop() + event.prevent_default() + self.app.action_paste_clipboard() + if self.app.is_key_for("history", event.key): event.stop() event.prevent_default() diff --git a/cecli/website/docs/config/tui.md b/cecli/website/docs/config/tui.md index 27445a30072..c444dd14bbf 100644 --- a/cecli/website/docs/config/tui.md +++ b/cecli/website/docs/config/tui.md @@ -54,7 +54,9 @@ tui-config: completion: "tab" stop: "escape" editor: "ctrl+o" - history: "ctrl+r" + history: "alt+shift+h" + voice: "ctrl+r" + paste: "alt+v" cycle_forward: "tab" cycle_backward: "shift+tab" input_start: "ctrl+home" @@ -82,7 +84,9 @@ The TUI provides customizable key bindings for all major actions. The default ke | Cancel | `ctrl+c` | Stop and stash current input prompt | | Stop | `escape` | Interrupt the current LLM response or task | | Editor | `ctrl+o` | Open up default terminal text editor for input | -| Search History | `ctrl+r` | Search through history for previous commands (requires fzf to be installed) | +| Search History | `alt+shift+h` | Search through history for previous commands (requires fzf to be installed) | +| Voice | `ctrl+r` | Record and transcribe voice input (dictation) into the chat | +| Paste Clipboard | `alt+v` | Paste image or text from the system clipboard into the chat | | Cycle Forward | `tab` | Cycle forward through completion suggestions | | Cycle Backward | `shift+tab` | Cycle backward through completion suggestions | | Input Start | `ctrl+home` | Move cursor to start of first line | diff --git a/tests/tui/test_app.py b/tests/tui/test_app.py index 89ee8259722..e651bb2862c 100644 --- a/tests/tui/test_app.py +++ b/tests/tui/test_app.py @@ -513,3 +513,89 @@ def test_submit_voice_intercepts_without_agent_queue(tui_instance, text): assert tui_instance.input_queue.empty() push_input.assert_not_called() wake_input.assert_not_called() + + +@pytest.mark.asyncio +async def test_run_paste_executes_paste_command(tui_instance): + """The paste hotkey runs the paste command on a worker without queuing input.""" + coder = MagicMock() + tui_instance._get_visible_coder = MagicMock(return_value=coder) + tui_instance.run_worker = MagicMock() + tui_instance._refresh_file_list = MagicMock() + + tui_instance.action_paste_clipboard() + + coroutine = tui_instance.run_worker.call_args.args[0] + tui_instance.run_worker.assert_called_once_with(coroutine, group="paste") + tui_instance._get_visible_coder.assert_called_once_with() + + with patch("cecli.commands.paste.PasteCommand.execute", new=AsyncMock()) as execute: + await coroutine + execute.assert_awaited_once_with(coder.io, coder, "") + tui_instance._refresh_file_list.assert_called_once_with(coder) + + +@pytest.mark.asyncio +async def test_run_paste_reports_errors(tui_instance): + """Clipboard failures surface through the status bar instead of crashing.""" + coder = MagicMock() + tui_instance.show_error = MagicMock() + + with patch( + "cecli.commands.paste.PasteCommand.execute", + new=AsyncMock(side_effect=RuntimeError("no clipboard")), + ): + await tui_instance._run_paste(coder) + + tui_instance.show_error.assert_called_once_with( + "Unable to paste clipboard content: no clipboard" + ) + + +@pytest.mark.parametrize("text", ["/paste", " /paste \n"]) +def test_submit_paste_intercepts_without_agent_queue(tui_instance, text): + import queue + + input_area = MagicMock(value=text) + tui_instance.query_one = MagicMock(return_value=input_area) + tui_instance.action_paste_clipboard = MagicMock() + tui_instance.add_user_message = MagicMock() + tui_instance.input_queue = queue.Queue() + + with ( + patch("cecli.tui.app.queues.push_coder_input") as push_input, + patch("cecli.tui.app.queues.wake_input_waiters") as wake_input, + ): + tui_instance.on_input_area_submit(MagicMock(value=text)) + + assert input_area.value == "" + tui_instance.action_paste_clipboard.assert_called_once_with() + input_area.save_to_history.assert_not_called() + tui_instance.add_user_message.assert_not_called() + assert tui_instance.input_queue.empty() + push_input.assert_not_called() + wake_input.assert_not_called() + + +def test_refresh_file_list_updates_autocomplete_and_file_list(tui_instance): + """The paste refresh updates both autocomplete data and the file list.""" + coder = MagicMock() + coder.get_addable_relative_files.return_value = ["a.py", "b.py"] + coder.commands.get_commands.return_value = ["/add"] + + input_area = MagicMock() + file_list = MagicMock() + + def mock_query_one(selector, *args): + if selector == "#input": + return input_area + if selector == "#file-list": + return file_list + raise AssertionError(f"unexpected selector: {selector}") + + tui_instance.query_one = mock_query_one + + tui_instance._refresh_file_list(coder) + + input_area.update_autocomplete_data.assert_called_once_with(["a.py", "b.py"], ["/add"]) + file_list.update_files.assert_called_once_with() From 4bc39b5ed4a245025f9bdbcfc7b58e4ef3a8b2be Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Sat, 19 Sep 2026 01:05:40 -0400 Subject: [PATCH 11/14] Explicitly mark commands finished after run in case of errors --- cecli/coders/base_coder.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/cecli/coders/base_coder.py b/cecli/coders/base_coder.py index 933301fc488..02a611ab587 100755 --- a/cecli/coders/base_coder.py +++ b/cecli/coders/base_coder.py @@ -1999,7 +1999,13 @@ async def preproc_user_input(self, inp): if self.commands.is_run_command(inp): self.commands.cmd_running_event.clear() # Command is running - return await self.commands.run(inp, coder=self, **run_kwargs) + try: + return await self.commands.run(inp, coder=self, **run_kwargs) + finally: + # Dispatch can return early (unknown/ambiguous command) or + # raise without ever running a command; the gate must still be + # reopened or the input/output loops park forever. + self.commands.cmd_running_event.set() await self.check_for_file_mentions(inp) inp = await self.check_for_urls(inp) From b3642792d49615de58529d6dba5bcfe48e729e83 Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Sat, 19 Sep 2026 01:06:23 -0400 Subject: [PATCH 12/14] Fix potential deadlock in git repo iterator --- cecli/repo.py | 61 ++++++++++++++++++++++++++++++++++++++++----------- 1 file changed, 48 insertions(+), 13 deletions(-) diff --git a/cecli/repo.py b/cecli/repo.py index 679ac42262d..cda1cdbda09 100644 --- a/cecli/repo.py +++ b/cecli/repo.py @@ -42,6 +42,11 @@ git.Git.USE_SHELL = False +# A malformed git tree can make ``next(iterator)`` raise IndexError without +# advancing. Cap the retries so traversal cannot spin forever while holding +# ``_git_lock``, which would wedge every other git operation. +MAX_TREE_INDEX_ERRORS = 3 + @contextlib.contextmanager def set_git_env(var_name, value, original_value): @@ -559,23 +564,28 @@ def get_tracked_files(self): else: try: iterator = commit.tree.traverse() - blob = None # Initialize blob + index_errors = 0 while True: try: blob = next(iterator) + except StopIteration: + break + except IndexError: + # next() cannot advance past the entry that + # raised, so retrying would spin forever while + # holding _git_lock. Give up after a few tries. + index_errors += 1 + if index_errors > MAX_TREE_INDEX_ERRORS: + self.io.tool_warning( + "GitRepo: Index error encountered while reading git tree" + " object. Skipping remaining entries." + ) + break + continue + else: if blob.type == "blob": # blob is a file # Use sys.intern() to deduplicate path strings in memory files.add(sys.intern(blob.path)) - except IndexError: - # Handle potential index error during tree traversal - # without relying on potentially unassigned 'blob' - self.io.tool_warning( - "GitRepo: Index error encountered while reading git tree object." - " Skipping." - ) - continue - except StopIteration: - break except ANY_GIT_ERROR as err: self.git_repo_error = err self.io.tool_error(f"Unable to list files in git repo: {err}") @@ -964,7 +974,9 @@ class GitRepoProxy: def __init__(self, target): self._target = target self._executor = concurrent.futures.ThreadPoolExecutor( - max_workers=1, thread_name_prefix="git-repo" + max_workers=1, + thread_name_prefix="git-repo", + initializer=_mark_git_executor_thread, ) _instances: dict[str, "GitRepoProxy"] = {} @@ -1048,8 +1060,21 @@ def unwrap(cls, repo): def __del__(self): try: - if hasattr(self, "_target") and self._target is not None: + if not (hasattr(self, "_target") and self._target is not None): + return + + # Finalization can run on the executor thread itself (e.g. when + # garbage collection happens inside a git call). Submitting would + # make the single worker wait on itself, so tear down inline there. + if getattr(_git_executor_thread, "active", False): + self._target.__del__() + return + + try: self._executor.submit(self._target.__del__).result(timeout=5) + except RuntimeError: + # Executor already shut down: no worker will pick this up. + self._target.__del__() except Exception: pass @@ -1129,3 +1154,13 @@ def __setattr__(self, name, value): super().__setattr__(name, value) else: setattr(self._target, name, value) + + +# Marks worker threads created by a ``GitRepoProxy`` executor. Re-entrant work +# (e.g. a finalizer running during garbage collection on the executor thread) +# must run inline: resubmitting to the single-worker pool would deadlock. +_git_executor_thread = threading.local() + + +def _mark_git_executor_thread(): + _git_executor_thread.active = True From ca072134c476688ea7dcada4901f1a8ba5e7047b Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Sat, 19 Sep 2026 03:05:58 -0400 Subject: [PATCH 13/14] Update interrupt handler strategy: - Propagate asyncio cancellation errors instead of KeyboardInterrupts - Make task cancellation a bounded best-effort wait to prevent deadlocks - Destroy background command infrastructure when interrupting a running command - Collect MCP connection errors individually - Harden the worker loop against escaped interrupts by resuming the same run task --- cecli/coders/agent_coder.py | 11 +- cecli/helpers/coroutines.py | 121 ++++++++++++++------- cecli/helpers/lock_detect.py | 76 ++++++++++++++ cecli/mcp/manager.py | 40 +++++-- cecli/models.py | 4 +- cecli/tools/command.py | 42 +++++++- cecli/tui/worker.py | 198 +++++++++++++++++++++++++++++++---- 7 files changed, 424 insertions(+), 68 deletions(-) diff --git a/cecli/coders/agent_coder.py b/cecli/coders/agent_coder.py index 9a5a9d04cb2..483c1e49f1c 100644 --- a/cecli/coders/agent_coder.py +++ b/cecli/coders/agent_coder.py @@ -939,7 +939,11 @@ async def gather_and_await(): lint_coro = self.lint_edited(edited, show_output=False) lint_errors, interrupted = await interruptible(lint_coro, self.interrupt_event) if interrupted: - raise KeyboardInterrupt("Interrupted during linting") + # Abort the turn with CancelledError, which the linear run loop + # handles by re-prompting for input. KeyboardInterrupt is a + # BaseException that would escape the worker's event loop and + # leave it unable to accept further prompts. + raise asyncio.CancelledError("Interrupted during linting") has_errors = False @@ -1078,7 +1082,10 @@ async def reply_completed(self): sleep_coro = asyncio.sleep(command_timeout / 2) _res, interrupted = await interruptible(sleep_coro, self.interrupt_event) if interrupted: - raise KeyboardInterrupt("Interrupted while waiting for background commands") + # Use CancelledError (not KeyboardInterrupt) so the interrupt + # stays inside the worker loop and re-prompts, matching the + # lint interrupt path above. + raise asyncio.CancelledError("Interrupted while waiting for background commands") return True # Check for recently finished commands that need reflection diff --git a/cecli/helpers/coroutines.py b/cecli/helpers/coroutines.py index 3d4b4b46089..6c8e590404e 100644 --- a/cecli/helpers/coroutines.py +++ b/cecli/helpers/coroutines.py @@ -9,16 +9,10 @@ # Strong reference pool to protect fire-and-forget tasks from garbage collection background_tasks: set = set() - -def _handle_result(task: asyncio.Task) -> None: - """Callback to clean up references and capture exceptions safely.""" - background_tasks.discard(task) - try: - task.result() - except asyncio.CancelledError: - pass - except Exception as e: - logger.error(f"Background task failed: {e}", exc_info=True) +# Grace period granted to a cancelled task to unwind. Transports that defer or +# swallow cancellation (AnyIO cancel scopes, shielded streams) would otherwise +# block the interrupting coroutine forever. +INTERRUPT_UNWIND_GRACE_SECONDS = 5.0 def fire_and_forget(coro) -> asyncio.Task: @@ -35,6 +29,7 @@ async def interruptible_async_generator(async_generator, interrupt_event): """ gen = async_generator.__aiter__() interrupt_task = asyncio.create_task(interrupt_event.wait()) + next_task = None try: while True: @@ -44,11 +39,7 @@ async def interruptible_async_generator(async_generator, interrupt_event): ) if interrupt_task in done: - next_task.cancel() - try: - await next_task - except asyncio.CancelledError: - pass + await cancel_and_abandon([next_task]) break if next_task in done: @@ -57,11 +48,9 @@ async def interruptible_async_generator(async_generator, interrupt_event): except StopAsyncIteration: break finally: - interrupt_task.cancel() - try: - await interrupt_task - except asyncio.CancelledError: - pass + # Cancel whichever task is still pending if the generator is closed early + # (e.g. GeneratorExit at the wait or yield), not just the interrupt waiter. + await cancel_and_abandon([task for task in (next_task, interrupt_task) if task is not None]) def is_active(task): @@ -83,6 +72,10 @@ async def interruptible(coroutine, interrupt_event): A tuple of (result, interrupted). - If not interrupted: (coroutine_result, False) - If interrupted: (None, True) + + A task that ignores cancellation is abandoned once the grace period + elapses (see cancel_and_abandon) so a wedged transport cannot block the + caller forever. """ if interrupt_event is None: interrupt_event = ThreadSafeEvent() @@ -90,25 +83,26 @@ async def interruptible(coroutine, interrupt_event): main_task = asyncio.create_task(coroutine) interrupt_task = asyncio.create_task(interrupt_event.wait()) - done, pending = await asyncio.wait( - {main_task, interrupt_task}, - return_when=asyncio.FIRST_COMPLETED, - ) + try: + done, pending = await asyncio.wait( + {main_task, interrupt_task}, + return_when=asyncio.FIRST_COMPLETED, + ) - for task in pending: - task.cancel() - try: - await task - except asyncio.CancelledError: - pass # Expected + await cancel_and_abandon(pending) - if interrupt_task in done: - return None, True + if interrupt_task in done: + return None, True - try: - return main_task.result(), False - except asyncio.CancelledError: - return None, True + try: + return main_task.result(), False + except asyncio.CancelledError: + return None, True + finally: + # When the caller is cancelled while waiting, neither branch above + # runs, so cancel any remaining tasks here instead of orphaning the + # inner coroutine (which would keep running detached). + await cancel_and_abandon([main_task, interrupt_task]) def task_is_cancelling() -> bool: @@ -133,3 +127,58 @@ def task_is_cancelling() -> bool: # Python 3.10 fallback: best-effort only (see docstring). return bool(getattr(task, "_must_cancel", False)) + + +def drain_task_results(tasks): + """Consume exceptions from finished tasks so they aren't reported as unretrieved.""" + for task in tasks: + if task.cancelled(): + continue + + try: + exc = task.exception() + except (asyncio.CancelledError, asyncio.InvalidStateError): + continue + + if exc is not None: + logger.warning(f"Task raised while unwinding after interrupt: {exc}") + + +async def cancel_and_abandon(tasks, grace=INTERRUPT_UNWIND_GRACE_SECONDS): + """Cancel tasks and wait up to ``grace`` seconds for them to unwind. + + Cancellation is cooperative, and some transports (AnyIO cancel scopes, + shielded streams) defer or swallow it. Waiting on them unconditionally + would block the caller forever, so anything still running after the grace + period is abandoned: it keeps a strong reference so it can finish on its + own without tripping the "Task was destroyed but it is pending" warning. + """ + pending_tasks = [task for task in tasks if not task.done()] + + if not pending_tasks: + return + + for task in pending_tasks: + task.cancel() + + done, still_running = await asyncio.wait(set(pending_tasks), timeout=grace) + drain_task_results(done) + + for task in still_running: + logger.warning( + f"Task ignored cancellation after {grace}s; abandoning it. It may still be" + " running and could touch shared state later." + ) + background_tasks.add(task) + task.add_done_callback(_handle_result) + + +def _handle_result(task: asyncio.Task) -> None: + """Callback to clean up references and capture exceptions safely.""" + background_tasks.discard(task) + try: + task.result() + except asyncio.CancelledError: + pass + except Exception as e: + logger.error(f"Background task failed: {e}", exc_info=True) diff --git a/cecli/helpers/lock_detect.py b/cecli/helpers/lock_detect.py index 21822523d77..0682f25b3b7 100644 --- a/cecli/helpers/lock_detect.py +++ b/cecli/helpers/lock_detect.py @@ -1,3 +1,4 @@ +import asyncio import logging import os import sys @@ -12,6 +13,13 @@ logging.getLogger("asyncio").setLevel(logging.DEBUG) +# Registered event loops whose asyncio tasks are included in each dump. +# Task state is only read from the owning loop (see _dump_async_state), since +# thread stacks alone cannot show coroutine/task state. +_async_loops: dict = {} +_async_state_providers: dict = {} + + def dump_stacks_to_file(filename=".cecli/logs/threads.log", interval=5, max_prints=10): """Periodically writes stack traces to a file, resetting it after max_prints.""" print_count = 0 @@ -46,12 +54,80 @@ def dump_stacks_to_file(filename=".cecli/logs/threads.log", interval=5, max_prin f.write("=" * 70 + "\n") + for name in list(_async_loops.keys()): + _schedule_async_dump(name, filename) + print_count += 1 # Increment after a successful write except Exception as e: print(f"Error writing stack dump to file: {e}", file=sys.stderr) +def register_async_loop(name, loop) -> None: + """Register an event loop to include asyncio task dumps for.""" + _async_loops[name] = loop + + +def register_async_state_provider(name, provider) -> None: + """Register a callable returning extra state text for a loop's dump.""" + _async_state_providers[name] = provider + + +def _format_async_tasks(loop, name) -> str: + """Return a readable dump of every task on ``loop`` (loop thread only).""" + lines = [f"\n------------------- ASYNCIO TASKS ({name}) -------------------\n"] + + try: + tasks = asyncio.all_tasks(loop) + except Exception as e: + return f"\n" + + for task in tasks: + lines.append(f"\n[{task.get_name()}] done={task.done()} cancelled={task.cancelled()}\n") + lines.append(f" {task.get_coro()!r}\n") + + for frame in task.get_stack(limit=10): + lines.append("".join(traceback.format_stack(frame, limit=1))) + + return "".join(lines) + + +def _dump_async_state(name, filename) -> None: + """Append asyncio task/state info for one loop. Runs on the loop thread.""" + loop = _async_loops.get(name) + if loop is None or loop.is_closed(): + return + + parts = [] + provider = _async_state_providers.get(name) + + if provider is not None: + try: + parts.append(provider()) + except Exception as e: + parts.append(f"\n") + + parts.append(_format_async_tasks(loop, name)) + + try: + with safe_open(filename, "a") as f: + f.write("".join(parts)) + except Exception: + pass + + +def _schedule_async_dump(name, filename) -> None: + """Schedule an async-state dump onto a registered loop, if it is running.""" + loop = _async_loops.get(name) + if loop is None or loop.is_closed() or not loop.is_running(): + return + + try: + loop.call_soon_threadsafe(_dump_async_state, name, filename) + except RuntimeError: + pass + + # Start the monitor in a background daemon thread monitor_thread = threading.Thread( target=dump_stacks_to_file, diff --git a/cecli/mcp/manager.py b/cecli/mcp/manager.py index 2063c09fc77..4209b1be1ff 100644 --- a/cecli/mcp/manager.py +++ b/cecli/mcp/manager.py @@ -8,6 +8,9 @@ # backstop, covering time the transport spends outside the SDK's read timeout. CONNECT_BACKSTOP_GRACE_SECONDS = 5 +# Servers whose failures are expected in normal operation and should not warn. +QUIET_SERVER_NAMES = {"unnamed-server", "local"} + class McpServerManager: """ @@ -387,14 +390,39 @@ async def _connect(server: McpServer) -> tuple[McpServer, bool, bool]: return (server, success, True) - results = await asyncio.gather(*(_connect(server) for server in self._servers)) + # return_exceptions=True keeps one bad server from aborting the batch and + # discarding every other server's result. A CancelledError that a + # transport raises for an ordinary connection failure lands in the + # results as a value; a genuine cancellation of this task still + # propagates through the gather itself, so it is not swallowed here. + outcomes = await asyncio.gather( + *(_connect(server) for server in self._servers), return_exceptions=True + ) + + results = [] + for server, outcome in zip(self._servers, outcomes): + if isinstance(outcome, BaseException): + if isinstance(outcome, asyncio.CancelledError) and task_is_cancelling(): + # A genuine cancellation of this task must not be downgraded + # to a failed server; let it propagate to the caller. + raise outcome + + # CancelledError discards its arguments, so fall back to the + # type name rather than reporting an empty cause. + cause = str(outcome) or type(outcome).__name__ + + if server.name.lower() not in QUIET_SERVER_NAMES: + self._log_warning(f"MCP server {server.name} failed to initialize: {cause}") + + # Reported here with its cause already, so flag it as + # not-attempted to keep the generic warning below from + # repeating the same failure. + results.append((server, False, False)) + else: + results.append(outcome) for server, did_connect, attempted in results: - if ( - attempted - and not did_connect - and server.name.lower() not in ["unnamed-server", "local"] - ): + if attempted and not did_connect and server.name.lower() not in QUIET_SERVER_NAMES: self._log_warning( f"MCP tool initialization failed after multiple retries: {server.name}" ) diff --git a/cecli/models.py b/cecli/models.py index 04d5e199474..827f433e35f 100644 --- a/cecli/models.py +++ b/cecli/models.py @@ -1487,7 +1487,7 @@ async def send_completion( completion_coro = litellm.acompletion(**kwargs) res, interrupted = await coroutines.interruptible(completion_coro, interrupt_event) if interrupted: - raise KeyboardInterrupt("Interrupted during acompletion") + raise asyncio.CancelledError("Interrupted during acompletion") return hash_object, res except litellm.ContextWindowExceededError as err: @@ -1532,7 +1532,7 @@ async def send_completion( asyncio.sleep(retry_delay), interrupt_event ) if interrupted: - raise KeyboardInterrupt("Interrupted during retry sleep") + raise asyncio.CancelledError("Interrupted during retry sleep") else: await asyncio.sleep(retry_delay) continue diff --git a/cecli/tools/command.py b/cecli/tools/command.py index 42a4e82ac78..ecc0488bd99 100644 --- a/cecli/tools/command.py +++ b/cecli/tools/command.py @@ -527,11 +527,26 @@ async def _execute_with_timeout(cls, coder, command_string, timeout, use_pty=Non f"Output captured so far:\n{output_content}\n" ) return response + except asyncio.CancelledError: + # The turn was cancelled (e.g. worker.interrupt) before the process + # finished. Terminate and unregister it so the child, its + # reader/writer threads, and the wait_task thread don't outlive the + # interrupted turn. + success, _, _ = BackgroundCommandManager.stop_background_command(command_key) + if not success: + cls._terminate_process(process) + + raise finally: interrupt_task.cancel() timeout_task.cancel() - if wait_task.done() and not wait_task.cancelled(): + if not wait_task.done() and process.returncode is not None: + # The process was terminated (interrupt/cancel) but wait_task + # is still pending; cancel the await. The blocked executor + # thread is freed once the process exits. + wait_task.cancel() + elif wait_task.done() and not wait_task.cancelled(): # Retrieve any exception to avoid "task exception was never # retrieved" warnings. On timeout the process continues in the # background, so wait_task may legitimately still be pending. @@ -844,3 +859,28 @@ def format_output(cls, coder, mcp_server, tool_response): # Output footer tool_footer(coder=coder, tool_response=tool_response, params=params) + + @classmethod + def _terminate_process(cls, process): + """Terminate a foreground command process, killing it if it lingers. + + Best-effort and never raises, so an interrupt handler can call it while + unwinding without masking the original exception. + """ + import subprocess + + try: + try: + process.terminate() + except ProcessLookupError: + return + + try: + process.wait(timeout=1) + except subprocess.TimeoutExpired: + try: + process.kill() + except ProcessLookupError: + pass + except Exception: + pass diff --git a/cecli/tui/worker.py b/cecli/tui/worker.py index b1602657e7f..80e4a8d1d39 100644 --- a/cecli/tui/worker.py +++ b/cecli/tui/worker.py @@ -2,15 +2,21 @@ import asyncio import logging +import os import sys import threading +import time +import traceback import warnings from typing import Optional from cecli.coders import Coder from cecli.commands import ReloadProgramSignal, SwitchCoderSignal from cecli.helpers.conversation import ConversationService, MessageTag -from cecli.helpers.coroutines import task_is_cancelling +from cecli.helpers.coroutines import ( + cancel_and_abandon, + task_is_cancelling, +) logger = logging.getLogger(__name__) # Suppress asyncio task destroyed warnings during shutdown @@ -20,10 +26,18 @@ warnings.filterwarnings("ignore", message=".*Task was destroyed.*") warnings.filterwarnings("ignore", message=".*coroutine.*was never awaited.*") +CRASH_LOG_DIR = ".cecli/logs" +CRASH_LOG_PATH = os.path.join(CRASH_LOG_DIR, "worker-crash.log") + class CoderWorker: """Runs Coder in a background thread with its own event loop.""" + # A KeyboardInterrupt raised inside a child task is re-raised out of the + # event loop by asyncio; give the loop a few chances to settle back into + # waiting for input before treating the repeated interrupt as a crash. + MAX_CONSECUTIVE_INTERRUPTS = 5 + def __init__(self, coder, output_queue, input_queue): """Initialize worker with coder instance and communication queues. @@ -59,12 +73,42 @@ def _run_thread(self): queues.set_input_loop(self.loop) + self._register_async_dump() + + consecutive_interrupts = 0 + run_task = None try: - self.loop.run_until_complete(self._async_run()) - except BaseException as e: - if not self._is_graceful_shutdown(e): - logger.error("Coder worker thread stopped unexpectedly", exc_info=e) - self._notify_crash(e) + while self.running: + if run_task is None or run_task.done(): + run_task = self.loop.create_task(self._async_run()) + + try: + self.loop.run_until_complete(run_task) + break + except KeyboardInterrupt as e: + # asyncio re-raises KeyboardInterrupt (and SystemExit) out + # of the event loop when a *child task* raises it, so it + # escapes run_until_complete instead of being caught by the + # coroutine awaiting that task. Treat it as a soft interrupt: + # resume the same task if it is still pending, or let the top + # of the loop start a fresh run loop if it already finished. + self._log_worker_error("KeyboardInterrupt escaped the worker loop", e) + consecutive_interrupts += 1 + if consecutive_interrupts > self.MAX_CONSECUTIVE_INTERRUPTS: + logger.error( + "Repeated KeyboardInterrupt escaped the worker loop", exc_info=e + ) + self._notify_crash(e) + break + if run_task.done() and not run_task.cancelled(): + run_task.exception() + time.sleep(0.05) + except BaseException as e: + if not self._is_graceful_shutdown(e): + self._log_worker_error("Coder worker thread stopped unexpectedly", e) + logger.error("Coder worker thread stopped unexpectedly", exc_info=e) + self._notify_crash(e) + break finally: self._cleanup_loop() @@ -92,14 +136,14 @@ def _cleanup_loop(self): for task in pending: task.cancel() - # Only try to gather if loop isn't stopped + # Let the cancelled tasks unwind, but only for a bounded time: + # a task that ignores cancellation (see coroutines.interruptible) + # must not wedge this thread, which would then never finish. if self.loop.is_running(): pass # Can't do much if loop is still running elif pending: try: - self.loop.run_until_complete( - asyncio.gather(*pending, return_exceptions=True) - ) + self.loop.run_until_complete(cancel_and_abandon(list(pending))) except RuntimeError: pass # Loop already stopped except KeyboardInterrupt: @@ -255,17 +299,22 @@ def interrupt(self): except Exception: pass - if target_coder and hasattr(target_coder, "io") and target_coder.io: - # Cancel the output task if it exists - if hasattr(target_coder.io, "output_task") and target_coder.io.output_task: - target_coder.io.output_task.cancel() - # Also set output_running to False to stop the output_task loop - if hasattr(target_coder, "output_running"): - target_coder.output_running = False + if not target_coder or not getattr(target_coder, "io", None): + return + + # The output task, output flag, and interrupt event live on the worker + # loop, but interrupt() runs on the TUI thread, so marshal the touch + # onto that loop instead of mutating loop-bound state directly. + loop = self.loop + if loop is not None and loop.is_running() and not loop.is_closed(): + try: + loop.call_soon_threadsafe(self._apply_interrupt, target_coder) + return + except RuntimeError: + pass - # Cancel any tracked generate task on the coder directly - if hasattr(target_coder, "interrupt_event") and target_coder.interrupt_event: - target_coder.interrupt_event.set() + # Loop is stopped or gone; apply directly so a stuck turn still stops. + self._apply_interrupt(target_coder) def stop(self): """Stop the worker thread gracefully.""" @@ -307,12 +356,16 @@ def _notify_crash(self, exc): """Tell the TUI the worker died so it can surface the error and exit. Without this the TUI keeps running with a dead worker and appears hung. + The full traceback is included so the failure is actionable instead of a + bare ``KeyboardInterrupt()``. """ + detail = self._format_exception(exc) + try: self.output_queue.put( { "type": "error", - "message": f"Worker stopped unexpectedly: {exc!r}", + "message": f"Worker stopped unexpectedly: {exc!r}\n{detail}", "coder_uuid": getattr(self.coder, "uuid", None), } ) @@ -335,3 +388,106 @@ def _create_event_loop(self): return asyncio.ProactorEventLoop() return asyncio.new_event_loop() + + def _apply_interrupt(self, target_coder): + """Stop the coder's output loop, cancel its output task, and signal its interrupt event. + + Runs on the worker loop (see interrupt) so the loop-bound task and + event are only touched from their owning loop. + """ + if hasattr(target_coder, "output_running"): + target_coder.output_running = False + + output_task = getattr(target_coder.io, "output_task", None) + if output_task: + output_task.cancel() + + interrupt_event = getattr(target_coder, "interrupt_event", None) + if interrupt_event: + interrupt_event.set() + + def _format_exception(self, exc): + """Return a readable traceback string for an exception.""" + if exc is None: + return "No exception details available." + + return "".join(traceback.format_exception(type(exc), exc, exc.__traceback__)) + + def _log_worker_error(self, message, exc): + """Append a worker error and its traceback to a log file. + + The TUI owns stdout/stderr while it runs, so a traceback printed there + is lost as soon as the program exits. Persisting it to disk lets the + user report the exact failure. + """ + try: + os.makedirs(CRASH_LOG_DIR, exist_ok=True) + + with open(CRASH_LOG_PATH, "a", encoding="utf-8", errors="replace") as f: + stamp = time.strftime("%Y-%m-%d %H:%M:%S") + f.write(f"\n===== {stamp} | {message} =====\n") + f.write(self._format_exception(exc)) + f.write("\n") + except Exception: + pass + + def _register_async_dump(self): + """Register the worker loop for thread dumps when debug tracing is on. + + The thread-dump monitor only runs in debug mode, so skip registration + otherwise; a normal run should not stream periodic task dumps to + ``.cecli/logs/threads.log``. + """ + if "cecli.helpers.lock_detect" not in sys.modules or not self._thread_dump_enabled(): + return + + from cecli.helpers import lock_detect + + lock_detect.register_async_loop("worker", self.loop) + lock_detect.register_async_state_provider("worker", self._format_async_state) + + def _thread_dump_enabled(self): + """Return True when debug tracing that starts the thread monitor is on.""" + env_override = os.getenv("CECLI_DEBUG_THREAD_LOG", "").lower() in ("1", "true", "yes") + env_debug = os.getenv("CECLI_DEBUG", "").lower() in ("1", "true", "yes") + args = getattr(getattr(self, "coder", None), "args", None) + + return bool(env_override or env_debug or getattr(args, "debug", False)) + + def _format_async_state(self): + """Return coder task/flag state for the lock_detect async dump.""" + from cecli.helpers.background_commands import BackgroundCommandManager + from cecli.helpers.coroutines import is_active + + lines = ["\n------------------- CODER STATE -------------------\n"] + coder = self.coder + + try: + io = getattr(coder, "io", None) + lines.append(f"input_task active: {is_active(getattr(io, 'input_task', None))}\n") + lines.append(f"output_task active: {is_active(getattr(io, 'output_task', None))}\n") + interrupt_event = getattr(coder, "interrupt_event", None) + lines.append( + f"interrupt_event set: {bool(interrupt_event and interrupt_event.is_set())}\n" + ) + lines.append(f"input_running: {getattr(coder, 'input_running', None)}\n") + lines.append(f"output_running: {getattr(coder, 'output_running', None)}\n") + + commands = getattr(coder, "commands", None) + lines.append( + f"cmd_running_event set: {bool(commands and commands.cmd_running_event.is_set())}\n" + ) + lines.append(f"worker running: {self.running}\n") + + background = BackgroundCommandManager.list_background_commands() + if not background: + lines.append(" (no background commands)\n") + + for key, info in sorted(background.items()): + lines.append( + f" - {key}: running={info.get('running')} pages={info.get('pages', 0)}\n" + ) + except Exception as e: + lines.append(f"\n") + + return "".join(lines) From ab974c9a64fec246ee8575ac95998ae3aa4c1775 Mon Sep 17 00:00:00 2001 From: Dustin Washington Date: Sat, 19 Sep 2026 17:49:43 -0400 Subject: [PATCH 14/14] #526, #594: Refactor session management system to work with sub agents --- cecli/coders/base_coder.py | 3 +- cecli/commands/list_sessions.py | 2 +- cecli/commands/load_session.py | 9 +- cecli/commands/save_session.py | 8 +- cecli/helpers/sessions/__init__.py | 20 ++ cecli/helpers/sessions/layout.py | 112 ++++++ cecli/helpers/sessions/manager.py | 325 +++++++++++++++++ cecli/helpers/sessions/payload.py | 287 +++++++++++++++ cecli/helpers/sessions/storage.py | 245 +++++++++++++ cecli/helpers/sessions/subagents.py | 152 ++++++++ cecli/main.py | 2 +- cecli/sessions.py | 505 --------------------------- tests/basic/test_sessions.py | 2 +- tests/basic/test_sessions_manager.py | 264 +++++++++++++- 14 files changed, 1417 insertions(+), 519 deletions(-) create mode 100644 cecli/helpers/sessions/__init__.py create mode 100644 cecli/helpers/sessions/layout.py create mode 100644 cecli/helpers/sessions/manager.py create mode 100644 cecli/helpers/sessions/payload.py create mode 100644 cecli/helpers/sessions/storage.py create mode 100644 cecli/helpers/sessions/subagents.py delete mode 100644 cecli/sessions.py diff --git a/cecli/coders/base_coder.py b/cecli/coders/base_coder.py index 02a611ab587..cc62c69493e 100755 --- a/cecli/coders/base_coder.py +++ b/cecli/coders/base_coder.py @@ -46,6 +46,7 @@ from cecli.helpers.memory_control import trim_memory from cecli.helpers.observations.service import ObservationService from cecli.helpers.profiler import TokenProfiler +from cecli.helpers.sessions import SessionManager from cecli.helpers.threading import ThreadSafeEvent from cecli.history import ChatSummary from cecli.hooks import HookIntegration @@ -64,7 +65,6 @@ from cecli.repomap import RepoMap from cecli.report import update_error_prefix from cecli.run_cmd import run_cmd_async -from cecli.sessions import SessionManager from cecli.tools.utils.output import print_tool_response from cecli.tools.utils.registry import ToolRegistry from cecli.utils import copy_tool_call, format_tokens, is_image_file @@ -5234,6 +5234,7 @@ async def auto_save_session(self, force=False): session_manager.save_session, getattr(self.args, "auto_save_session_name", "auto-save"), False, + True, ) except Exception: # Don't show errors for auto-save to avoid interrupting the user experience diff --git a/cecli/commands/list_sessions.py b/cecli/commands/list_sessions.py index b896b17e547..84a5f0da718 100644 --- a/cecli/commands/list_sessions.py +++ b/cecli/commands/list_sessions.py @@ -11,7 +11,7 @@ class ListSessionsCommand(BaseCommand): @classmethod async def execute(cls, io, coder, args, **kwargs): """Execute the list-sessions command with given parameters.""" - from cecli import sessions + from cecli.helpers import sessions session_manager = sessions.SessionManager(coder, io) sessions_list = session_manager.list_sessions() diff --git a/cecli/commands/load_session.py b/cecli/commands/load_session.py index f3c38396a8b..f9910845795 100644 --- a/cecli/commands/load_session.py +++ b/cecli/commands/load_session.py @@ -15,7 +15,7 @@ async def execute(cls, io, coder, args, **kwargs): io.tool_output("Usage: /load-session ") return format_command_result(io, "load-session", "No session name provided") - from cecli import sessions + from cecli.helpers import sessions session_manager = sessions.SessionManager(coder, io) await session_manager.load_session(args.strip()) @@ -26,7 +26,7 @@ async def execute(cls, io, coder, args, **kwargs): def get_completions(cls, io, coder, args) -> List[str]: """Get completion options for load-session command.""" # Return available session names for completion - from cecli import sessions + from cecli.helpers import sessions session_manager = sessions.SessionManager(coder, io) sessions_list = session_manager.list_sessions() @@ -41,7 +41,10 @@ def get_help(cls) -> str: help_text += "\nExamples:\n" help_text += " /load-session my-feature # Load session 'my-feature'\n" help_text += " /load-session bug-fix # Load session 'bug-fix'\n" - help_text += "\nSessions are loaded from the .cecli/sessions/ directory.\n" + help_text += ( + "\nSessions are loaded from the .cecli/sessions/ directory. Loading a session" + " also restores any sub-agents saved with it.\n" + ) help_text += ( "Use /list-sessions to see saved sessions and /save-session to save a session.\n" ) diff --git a/cecli/commands/save_session.py b/cecli/commands/save_session.py index 16a96d61453..cf32180ec96 100644 --- a/cecli/commands/save_session.py +++ b/cecli/commands/save_session.py @@ -15,7 +15,7 @@ async def execute(cls, io, coder, args, **kwargs): io.tool_error("Please provide a session name to save.") return format_command_result(io, "save-session", "No session name provided") - from cecli import sessions + from cecli.helpers import sessions session_manager = sessions.SessionManager(coder, io) session_manager.save_session(args.strip()) @@ -26,7 +26,7 @@ async def execute(cls, io, coder, args, **kwargs): def get_completions(cls, io, coder, args) -> List[str]: """Get completion options for save-session command.""" # Return existing session names for completion to prevent accidental overwrites - from cecli import sessions + from cecli.helpers import sessions session_manager = sessions.SessionManager(coder, io) sessions_list = session_manager.list_sessions() @@ -42,6 +42,10 @@ def get_help(cls) -> str: help_text += " /save-session my-feature # Save session as 'my-feature'\n" help_text += " /save-session bug-fix # Save session as 'bug-fix'\n" help_text += "\nSessions are saved in the .cecli/sessions/ directory as JSON files.\n" + help_text += ( + "When the session has sub-agents it is saved as a folder holding the primary" + " and each sub-agent payload.\n" + ) help_text += "Use /list-sessions to see saved sessions and /load-session to load them.\n" help_text += ( "\nNote: Existing session names will be shown for tab completion to help prevent" diff --git a/cecli/helpers/sessions/__init__.py b/cecli/helpers/sessions/__init__.py new file mode 100644 index 00000000000..5d50218726c --- /dev/null +++ b/cecli/helpers/sessions/__init__.py @@ -0,0 +1,20 @@ +"""Session persistence helpers. + +The session concerns are split into focused modules so each one stays readable: + +* ``layout`` - where session payloads live on disk (single files, folder + bundles, and the ``s/`` tree used for sub-agents). +* ``storage`` - reading/writing payload files and reference documents, including + optional encryption. +* ``payload`` - serialising a coder's state into a payload and applying a + payload back onto a coder. +* ``subagents`` - detecting, saving, and resolving the sub-agents that belong to + a session. +* ``manager`` - :class:`SessionManager`, the facade that ties the above together + for save/list/load. +""" + +from .layout import SESSION_REFERENCE_TYPE +from .manager import SessionManager + +__all__ = ["SessionManager", "SESSION_REFERENCE_TYPE"] diff --git a/cecli/helpers/sessions/layout.py b/cecli/helpers/sessions/layout.py new file mode 100644 index 00000000000..f3128621c6c --- /dev/null +++ b/cecli/helpers/sessions/layout.py @@ -0,0 +1,112 @@ +"""On-disk layout for saved sessions. + +Two shapes are supported: + +* a single ``{name}.json`` file (optionally a reference document pointing at an + auto-save payload held in a coder's agent folder), and +* a ``{name}/`` folder bundle holding ``primary.json`` plus an ``s/`` tree of + sub-agent ``agent.json`` payloads. + +Sub-agent payloads always live in an ``s/{child}/`` directory beside the primary +payload. That lets the same discovery logic cover auto-save payloads (which use +``s/{uuid}/{session}.json``) and folder bundles (``s/{child}/agent.json``). +""" + +from __future__ import annotations + +import os +from pathlib import Path +from typing import Dict, List, Tuple + +SESSION_REFERENCE_TYPE = "session-reference" +PRIMARY_PAYLOAD_NAME = "primary.json" +SUB_AGENT_PAYLOAD_NAME = "agent.json" +SUB_AGENT_DIR_NAME = "s" + + +def session_directory(coder) -> Path: + """Return (and create) the coder's ``.cecli/sessions`` directory.""" + session_dir = Path(coder.abs_root_path(".cecli/sessions")) + os.makedirs(session_dir, exist_ok=True) + + return session_dir + + +def entry_for(session_dir: Path, session_name: str) -> Path: + """Return the canonical ``sessions/{name}.json`` file path.""" + return session_dir / f"{session_name}.json" + + +def bundle_dir(session_dir: Path, session_name: str) -> Path: + """Return the ``sessions/{name}/`` folder bundle path.""" + return session_dir / session_name + + +def agent_payload_file(coder, session_name: str) -> Path: + """Return the auto-save payload path inside the coder's agent folder.""" + rel_path = coder.local_agent_folder(f"{session_name}.json") + + return Path(coder.abs_root_path(rel_path)) + + +def sub_agents_dir(primary_payload: Path) -> Path: + """Return the ``s/`` directory that owns a primary payload's sub-agents.""" + return primary_payload.parent / SUB_AGENT_DIR_NAME + + +def primary_payload_in(directory: Path) -> Path | None: + """Return the primary payload inside a folder bundle, if present.""" + candidate = directory / PRIMARY_PAYLOAD_NAME + + return candidate if candidate.is_file() else None + + +def payload_in_sub_dir(sub_dir: Path, primary_name: str) -> Path | None: + """Return the payload held by one sub-agent directory. + + Folder bundles name it ``agent.json``; auto-save payloads reuse the primary + payload's filename. + """ + candidate = sub_dir / SUB_AGENT_PAYLOAD_NAME + if candidate.is_file(): + return candidate + + candidate = sub_dir / primary_name + if candidate.is_file(): + return candidate + + return None + + +def sub_agent_payloads(primary_payload: Path) -> List[Path]: + """Return every sub-agent payload stored beside ``primary_payload``.""" + subs_dir = sub_agents_dir(primary_payload) + if not subs_dir.is_dir(): + return [] + + payloads = [] + for sub_dir in sorted(subs_dir.iterdir()): + if not sub_dir.is_dir(): + continue + + payload = payload_in_sub_dir(sub_dir, primary_payload.name) + if payload is not None: + payloads.append(payload) + + return payloads + + +def discover_entries(session_dir: Path) -> List[Tuple[str, Path]]: + """Return ``(name, locator)`` for every saved session in ``session_dir``. + + A locator is either a ``{name}.json`` file or a ``{name}/`` folder bundle; + callers resolve it to the payload file they need. + """ + entries: Dict[str, Path] = {} + for path in session_dir.iterdir(): + if path.is_file() and path.suffix == ".json": + entries[path.stem] = path + elif path.is_dir() and primary_payload_in(path) is not None: + entries[path.name] = path + + return sorted(entries.items()) diff --git a/cecli/helpers/sessions/manager.py b/cecli/helpers/sessions/manager.py new file mode 100644 index 00000000000..70ddb7a0628 --- /dev/null +++ b/cecli/helpers/sessions/manager.py @@ -0,0 +1,325 @@ +"""``SessionManager`` - the facade over the session helper modules.""" + +from __future__ import annotations + +import logging +import shutil +from pathlib import Path +from typing import Dict, List, Optional + +from . import layout, subagents +from .layout import SESSION_REFERENCE_TYPE +from .payload import apply_payload, build_payload +from .storage import ( + encrypt_settings, + read_payload, + resolve_payload_file, + write_payload, + write_reference, +) + +logger = logging.getLogger(__name__) + + +class SessionManager: + """Manages chat session saving, listing, and loading.""" + + SESSION_REFERENCE_TYPE = SESSION_REFERENCE_TYPE + + def __init__(self, coder, io): + self.coder = coder + self.io = io + + # ------------------------------------------------------------------ # + # Saving + # ------------------------------------------------------------------ # + + def save_session(self, session_name: str, output=True, to_agent_folder=False) -> bool: + """Save the current chat session. + + ``to_agent_folder`` selects the auto-save shape: the payload is written + into the coder's local agent folder and a reference document is written + to the sessions directory (primary agent only). Otherwise the session is + saved as a self-contained ``sessions/{name}.json`` file, or a + ``sessions/{name}/`` folder bundle when the coder has sub-agents. + """ + if not session_name: + if output: + self.io.tool_error("Please provide a session name.") + + return False + + session_name = session_name.replace(".json", "") + + try: + if to_agent_folder: + return self._save_to_agent_folder(session_name, output) + + return self._save_named(session_name, output) + + except Exception as e: + self.io.tool_error(f"Error saving session: {e}") + + return False + + def _save_to_agent_folder(self, session_name: str, output: bool) -> bool: + """Auto-save: payload in the agent folder, reference in sessions dir.""" + session_dir = layout.session_directory(self.coder) + session_file = layout.entry_for(session_dir, session_name) + if session_file.exists() and output: + self.io.tool_warning(f"Session '{session_name}' already exists. Overwriting.") + + data_file = layout.agent_payload_file(self.coder, session_name) + session_data = build_payload(self.coder, self.io, session_name, self._sub_agent_name()) + + if not write_payload(self.coder, self.io, data_file, session_data): + return False + + # Only the primary agent owns the sessions-directory reference; + # sub-agents stop at their own payload so they cannot clobber the + # pointer that auto-load resolves. + if not session_data.get("agent_name") and not write_reference( + self.coder, self.io, session_file, session_name, data_file + ): + return False + + if output: + suffix = " (encrypted)" if encrypt_settings(self.coder)[0] else "" + self.io.tool_output(f"Session saved: {data_file}{suffix}") + + return True + + def _save_named(self, session_name: str, output: bool) -> bool: + """Explicit save: a folder bundle when sub-agents exist, else a file.""" + session_dir = layout.session_directory(self.coder) + bundle = layout.bundle_dir(session_dir, session_name) + + if output and (layout.entry_for(session_dir, session_name).exists() or bundle.exists()): + self.io.tool_warning(f"Session '{session_name}' already exists. Overwriting.") + + if subagents.live_sub_agents(self.coder): + return self._save_bundle(session_dir, session_name, output) + + return self._save_file(session_dir, session_name, output) + + def _save_file(self, session_dir: Path, session_name: str, output: bool) -> bool: + """Save a session without sub-agents as a single ``{name}.json`` file.""" + entry = layout.entry_for(session_dir, session_name) + session_data = build_payload(self.coder, self.io, session_name, self._sub_agent_name()) + + if not write_payload(self.coder, self.io, entry, session_data): + return False + + # Drop a stale bundle saved under the same name so one entry stays canonical. + shutil.rmtree(layout.bundle_dir(session_dir, session_name), ignore_errors=True) + + if output: + suffix = " (encrypted)" if encrypt_settings(self.coder)[0] else "" + self.io.tool_output(f"Session saved: {entry}{suffix}") + + return True + + def _save_bundle(self, session_dir: Path, session_name: str, output: bool) -> bool: + """Save a session with sub-agents as a ``{name}/`` folder bundle.""" + bundle = layout.bundle_dir(session_dir, session_name) + primary = bundle / layout.PRIMARY_PAYLOAD_NAME + session_data = build_payload(self.coder, self.io, session_name, self._sub_agent_name()) + + if not write_payload(self.coder, self.io, primary, session_data): + return False + + subagents.save_sub_agents(self.coder, self.io, session_name, layout.sub_agents_dir(primary)) + + # Drop a stale single-file entry saved under the same name. + layout.entry_for(session_dir, session_name).unlink(missing_ok=True) + + if output: + suffix = " (encrypted)" if encrypt_settings(self.coder)[0] else "" + self.io.tool_output(f"Session saved: {bundle}{suffix}") + + return True + + # ------------------------------------------------------------------ # + # Listing + # ------------------------------------------------------------------ # + + def list_sessions(self) -> List[Dict]: + """List all saved sessions with metadata.""" + from .storage import describe_session + + session_dir = layout.session_directory(self.coder) + entries = layout.discover_entries(session_dir) + + if not entries: + self.io.tool_output("No saved sessions found.") + + return [] + + sessions = [] + ordered = sorted(entries, key=lambda item: item[1].stat().st_mtime, reverse=True) + for name, locator in ordered: + try: + info = describe_session(self.coder, self.io, name, locator) + + except Exception as e: + self.io.tool_output(f" {name} [error reading: {e}]") + + continue + + if info is not None: + sessions.append(info) + + return sessions + + # ------------------------------------------------------------------ # + # Loading + # ------------------------------------------------------------------ # + + async def load_session(self, session_identifier: str, switch=True, quiet: bool = False) -> bool: + """Load a saved session by name or file path.""" + if not session_identifier: + self.io.tool_error("Please provide a session name or file path.") + + return False + + session_file = self._find_session_file(session_identifier) + if not session_file: + return False + + data_file = resolve_payload_file(self.coder, self.io, session_file, quiet=quiet) + if data_file is None: + return False + + session_data = read_payload(self.coder, self.io, data_file, quiet=quiet) + if session_data is None: + return False + + if not isinstance(session_data, dict) or "version" not in session_data: + if not quiet: + self.io.tool_error("Invalid session format.") + + return False + + # Apply session data + applied, loaded_edit_format = await self._apply_session_data(session_data, session_file) + + # A primary session may have sub-agents saved alongside it; restore the + # whole agent tree rather than just the primary conversation. + if applied and not session_data.get("agent_name"): + await self._reload_sub_agents(data_file) + + if applied and switch: + from cecli.commands import SwitchCoderSignal + + edit_format_to_switch_to = self.coder.edit_format + if loaded_edit_format: + edit_format_to_switch_to = loaded_edit_format + self.coder.edit_format = loaded_edit_format + + raise SwitchCoderSignal( + edit_format=edit_format_to_switch_to, + from_coder=self.coder, + summarize_from_coder=False, + show_announcements=True, + ) + + return applied + + # ------------------------------------------------------------------ # + # Sub-agent restore + # ------------------------------------------------------------------ # + + async def _reload_sub_agents(self, data_file: Path) -> None: + """Rebuild every sub-agent saved alongside a primary session payload.""" + payloads = layout.sub_agent_payloads(data_file) + if not payloads: + return + + from cecli.helpers.agents.service import AgentService + + service = AgentService.get_instance(self.coder) + + for sub_file in payloads: + try: + await self._reload_sub_agent(service, sub_file) + + except Exception as e: + logger.warning("Failed to restore sub-agent from %s: %s", sub_file, e) + self.io.tool_warning(f"Could not restore sub-agent from {sub_file}: {e}") + + async def _reload_sub_agent(self, service, sub_file: Path) -> None: + """Spawn and restore a single sub-agent from its saved payload.""" + sub_data = self._read_session_payload(sub_file, quiet=True) + if not isinstance(sub_data, dict): + return + + name = subagents.resolve_reload_agent_name( + sub_data.get("agent_name") or "worker", sub_data.get("agent_root") + ) + if not name: + return + + new_coder, _info = await service.spawn( + name, parent=self.coder, auto_reap=False, independent=True + ) + + sub_manager = SessionManager(new_coder, self.io) + applied, _edit_format = await sub_manager._apply_session_data( + sub_data, sub_file, sub_agent=True + ) + if not applied: + logger.warning("Restored sub-agent '%s' but could not apply its saved state", name) + + # ------------------------------------------------------------------ # + # Internal helpers + # ------------------------------------------------------------------ # + + def _find_session_file(self, session_identifier: str) -> Optional[Path]: + """Find a session locator (single file or folder bundle) by name or path.""" + session_file = Path(session_identifier) + if session_file.exists(): + return session_file + + session_dir = layout.session_directory(self.coder) + + if not session_identifier.endswith(".json"): + session_file = layout.entry_for(session_dir, session_identifier) + if session_file.exists(): + return session_file + + session_file = layout.bundle_dir(session_dir, session_identifier) + if session_file.exists(): + return session_file + + session_file = session_dir / session_identifier + if session_file.exists(): + return session_file + + self.io.tool_error(f"Session not found: {session_identifier}") + self.io.tool_output("Use /list-sessions to see available sessions.") + + return None + + def _read_session_file(self, session_file: Path, quiet: bool = False) -> dict | None: + """Resolve and read a session payload, following reference documents.""" + data_file = resolve_payload_file(self.coder, self.io, session_file, quiet=quiet) + if data_file is None: + return None + + return read_payload(self.coder, self.io, data_file, quiet=quiet) + + def _read_session_payload(self, data_file: Path, quiet: bool = False) -> dict | None: + """Read a session payload file, decrypting it when necessary.""" + return read_payload(self.coder, self.io, data_file, quiet=quiet) + + def _sub_agent_name(self) -> Optional[str]: + """Return this coder's sub-agent type, or ``None`` for the primary agent.""" + return subagents.detect_agent_name(self.coder) + + async def _apply_session_data( + self, session_data: Dict, session_file: Path, sub_agent: bool = False + ) -> tuple[bool, Optional[str]]: + """Apply a session payload to this manager's coder.""" + return await apply_payload( + self.coder, self.io, session_data, session_file, sub_agent=sub_agent + ) diff --git a/cecli/helpers/sessions/payload.py b/cecli/helpers/sessions/payload.py new file mode 100644 index 00000000000..ce64fdbffd9 --- /dev/null +++ b/cecli/helpers/sessions/payload.py @@ -0,0 +1,287 @@ +"""Serialising a coder's state to and from a session payload.""" + +from __future__ import annotations + +import os +from typing import Dict, Optional + +from cecli import models +from cecli.helpers.conversation import ConversationService, MessageTag + + +def build_payload(coder, io, session_name: str, agent_name: Optional[str] = None) -> Dict: + """Build a session payload dictionary from a coder's current state. + + ``agent_name`` is the sub-agent type when saving a sub-agent and ``None`` + for the primary agent. Workspace agents (``ws:*``) also record their root so + a later load can rebuild them at the right path. + """ + editable_files = [coder.get_rel_fname(abs_fname) for abs_fname in coder.abs_fnames] + read_only_files = [coder.get_rel_fname(abs_fname) for abs_fname in coder.abs_read_only_fnames] + read_only_stubs_files = [ + coder.get_rel_fname(abs_fname) for abs_fname in coder.abs_read_only_stubs_fnames + ] + + # Flush any queued messages so the saved chat history is complete + ConversationService.get_manager(coder).flush_queue() + + return { + "version": 1, + "session_name": session_name, + "agent_name": agent_name, + "agent_root": str(coder.root) if agent_name and agent_name.startswith("ws:") else None, + "model": coder.main_model.name, + "weak_model": coder.main_model.weak_model.name, + "editor_model": coder.main_model.editor_model.name, + "agent_model": coder.main_model.agent_model.name, + "editor_edit_format": coder.main_model.editor_edit_format, + "edit_format": coder.edit_format, + "chat_history": { + "done_messages": ( + ConversationService.get_manager(coder).get_messages_dict(MessageTag.DONE) + ), + "cur_messages": ( + ConversationService.get_manager(coder).get_messages_dict(MessageTag.CUR) + ), + }, + "files": { + "editable": editable_files, + "read_only": read_only_files, + "read_only_stubs": read_only_stubs_files, + }, + "settings": { + "auto_commits": coder.auto_commits, + "auto_lint": coder.auto_lint, + "auto_test": coder.auto_test, + }, + "todo_list": _read_todo(coder, io), + "mcps": _connected_mcps(coder), + "skills": _skills_data(coder), + "tools": _agent_config_data(coder), + "usage": { + "total_tokens_sent": coder.total_tokens_sent, + "total_tokens_received": coder.total_tokens_received, + "total_cached_tokens": coder.total_cached_tokens, + "total_cost": coder.total_cost, + }, + } + + +async def apply_payload( + coder, io, session_data: Dict, session_file, sub_agent: bool = False +) -> tuple[bool, Optional[str]]: + """Apply a session payload to a coder's state. + + ``sub_agent`` marks a sub-agent restore: progress output is suppressed and + global environment registries (MCP servers, skills, tools) are left untouched + so restoring one sub-agent cannot disturb the primary session's shared state. + + Returns: + A tuple of (success, edit_format). + """ + try: + # Clear current state + coder.abs_fnames = set() + coder.abs_read_only_fnames = set() + coder.abs_read_only_stubs_fnames = set() + + # Load files + files = session_data.get("files", {}) + _restore_files(coder, io, coder.abs_fnames, files.get("editable", [])) + _restore_files(coder, io, coder.abs_read_only_fnames, files.get("read_only", [])) + _restore_files( + coder, io, coder.abs_read_only_stubs_fnames, files.get("read_only_stubs", []) + ) + + # Load usage stats + usage = session_data.get("usage", {}) + coder.total_tokens_sent = usage.get("total_tokens_sent", 0) + coder.total_tokens_received = usage.get("total_tokens_received", 0) + coder.total_cached_tokens = usage.get("total_cached_tokens", 0) + coder.total_cost = usage.get("total_cost", 0.0) + # Loading a session seeds the cumulative counters but does not represent + # recent API usage, so clear the rolling token-rate buffer. + coder._reset_token_usage() + if session_data.get("model"): + coder.main_model = models.Model( + session_data.get("model", coder.args.model), + weak_model=session_data.get("weak_model", coder.args.weak_model), + editor_model=session_data.get("editor_model", coder.args.editor_model), + agent_model=session_data.get("agent_model", coder.args.agent_model), + editor_edit_format=session_data.get( + "editor_edit_format", coder.args.editor_edit_format + ), + io=io, + verbose=coder.args.verbose, + retries=coder.main_model.retries, + debug=coder.main_model.debug, + ) + + # Load settings + settings = session_data.get("settings", {}) + if "auto_commits" in settings: + coder.auto_commits = settings["auto_commits"] + if "auto_lint" in settings: + coder.auto_lint = settings["auto_lint"] + if "auto_test" in settings: + coder.auto_test = settings["auto_test"] + + _restore_todo(coder, io, session_data) + + # Clear CUR and DONE messages from ConversationManager + ConversationService.get_manager(coder).reset() + ConversationService.get_files(coder).reset() + coder.format_chat_chunks() + + # Load chat history + chat_history = session_data.get("chat_history", {}) + for msg in chat_history.get("done_messages", []): + ConversationService.get_manager(coder).add_message( + message_dict=msg, + tag=MessageTag.DONE, + ) + for msg in chat_history.get("cur_messages", []): + ConversationService.get_manager(coder).add_message( + message_dict=msg, + tag=MessageTag.CUR, + ) + + if not sub_agent: + io.tool_output(f"Session loaded: {session_data.get('session_name', session_file.stem)}") + io.tool_output( + f"Model: {session_data.get('model', 'unknown')}, Edit format:" + f" {session_data.get('edit_format', 'unknown')}" + ) + + # Show summary + num_messages = len(coder.done_messages) + len(coder.cur_messages) + num_files = ( + len(coder.abs_fnames) + + len(coder.abs_read_only_fnames) + + len(coder.abs_read_only_stubs_fnames) + ) + if not sub_agent: + io.tool_output(f"Loaded {num_messages} messages and {num_files} files") + + # Load MCPs + if not sub_agent and getattr(coder, "mcp_manager", None): + await _restore_mcps(coder, session_data.get("mcps", [])) + + # Load skills + skills_data = session_data.get("skills") + if not sub_agent and skills_data and getattr(coder, "skills_manager", None): + coder.skills_manager.directory_paths = skills_data.get("skills_paths", []) + coder.skills_manager.include_list = set(skills_data.get("skills_includelist", [])) + coder.skills_manager.exclude_list = set(skills_data.get("skills_excludelist", [])) + + # Load tools config + agent_config_data = session_data.get("tools") + if not sub_agent and agent_config_data and hasattr(coder, "agent_config"): + coder.agent_config.update(agent_config_data) + from cecli.tools.utils.registry import ToolRegistry + + ToolRegistry.build_registry(agent_config=coder.agent_config) + coder.loaded_custom_tools = ToolRegistry.loaded_custom_tools + + # Return True and the edit format so the Coder can be switched + edit_format = session_data.get("edit_format") + + return True, edit_format + + except Exception as e: + io.tool_error(f"Error applying session data: {e}") + + return False, None + + +def _read_todo(coder, io) -> Optional[str]: + """Read the agent's todo file so it can be restored with the session.""" + try: + todo_path = coder.abs_root_path(coder.local_agent_folder("todo.txt")) + if os.path.isfile(todo_path): + todo_content = io.read_text(todo_path) + + return todo_content if todo_content is not None else "" + + except Exception as e: + io.tool_warning(f"Could not read todo list file: {e}") + + return None + + +def _connected_mcps(coder) -> list: + """Return the names of the coder's connected MCP servers.""" + if getattr(coder, "mcp_manager", None): + return [server.name for server in coder.mcp_manager.connected_servers] + + return [] + + +def _skills_data(coder) -> Optional[Dict]: + """Return the coder's skill path/filter configuration.""" + manager = getattr(coder, "skills_manager", None) + if not manager: + return None + + return { + "skills_paths": [str(p) for p in manager.directory_paths], + "skills_includelist": ( + list(manager.include_list) if manager.include_list is not None else [] + ), + "skills_excludelist": ( + list(manager.exclude_list) if manager.exclude_list is not None else [] + ), + } + + +def _agent_config_data(coder) -> Optional[Dict]: + """Return the coder's tool configuration.""" + if not hasattr(coder, "agent_config"): + return None + + return { + "tools_paths": coder.agent_config.get("tools_paths", []), + "tools_includelist": coder.agent_config.get("tools_includelist", []), + "tools_excludelist": coder.agent_config.get("tools_excludelist", []), + } + + +def _restore_files(coder, io, target: set, rel_fnames: list) -> None: + """Add every existing ``rel_fnames`` entry to ``target``.""" + for rel_fname in rel_fnames: + abs_fname = coder.abs_root_path(rel_fname) + if os.path.exists(abs_fname): + target.add(abs_fname) + else: + io.tool_warning(f"File not found, skipping: {rel_fname}") + + +def _restore_todo(coder, io, session_data: Dict) -> None: + """Restore (or clear) the agent's todo file from the session.""" + if "todo_list" not in session_data: + return + + todo_path = coder.abs_root_path(coder.local_agent_folder("todo.txt")) + todo_content = session_data.get("todo_list") + + try: + if todo_content is None: + if os.path.exists(todo_path): + os.remove(todo_path) + else: + io.write_text(todo_path, todo_content) + + except Exception as e: + io.tool_warning(f"Could not restore todo list: {e}") + + +async def _restore_mcps(coder, saved_mcps: list) -> None: + """Reconcile the coder's MCP connections with the saved list.""" + current_mcps = {server.name for server in coder.mcp_manager.connected_servers} + saved_mcps_set = set(saved_mcps) + + for mcp_name in current_mcps - saved_mcps_set: + await coder.mcp_manager.disconnect_server(mcp_name) + + for mcp_name in saved_mcps_set - current_mcps: + await coder.mcp_manager.connect_server(mcp_name) diff --git a/cecli/helpers/sessions/storage.py b/cecli/helpers/sessions/storage.py new file mode 100644 index 00000000000..76e68742a0c --- /dev/null +++ b/cecli/helpers/sessions/storage.py @@ -0,0 +1,245 @@ +"""Reading and writing session payloads and reference documents.""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Dict, Optional + +from cecli.decoding import safe_open +from cecli.helpers import crypto as session_crypto + +from .layout import PRIMARY_PAYLOAD_NAME, SESSION_REFERENCE_TYPE, sub_agent_payloads + + +def encrypt_settings(coder) -> tuple[bool, bytes | None]: + """Return whether session encryption is on and the resolved key.""" + args = getattr(coder, "args", None) + if not args or not getattr(args, "session_encrypt", False): + return False, None + + key_file = getattr(args, "session_key_file", None) + + return True, session_crypto.resolve_key(key_file=key_file) + + +def resolve_payload_file(coder, io, session_file: Path, quiet: bool = False) -> Path | None: + """Resolve a session locator to the file holding its payload. + + A locator may be a folder bundle (``{name}/``), a single payload file, or a + reference document pointing at an auto-save payload in a coder's agent + folder. Returns ``None`` when the payload cannot be located. + """ + if session_file.is_dir(): + payload = session_file / PRIMARY_PAYLOAD_NAME + if payload.is_file(): + return payload + + if not quiet: + io.tool_error(f"Session folder is missing {PRIMARY_PAYLOAD_NAME}: {session_file}") + + return None + + try: + raw = session_file.read_bytes() + + except OSError as e: + if not quiet: + io.tool_error(f"Error reading session: {e}") + + return None + + if session_crypto.is_encrypted_payload(raw): + return session_file + + try: + parsed = json.loads(raw.decode("utf-8")) + + except (UnicodeDecodeError, json.JSONDecodeError): + return session_file + + if not isinstance(parsed, dict) or parsed.get("type") != SESSION_REFERENCE_TYPE: + return session_file + + reference = parsed.get("path") + if not reference: + if not quiet: + io.tool_error("Session reference is missing a path.") + + return None + + target = Path(reference) + if not target.is_absolute(): + target = Path(coder.abs_root_path(reference)) + + if not target.exists(): + if not quiet: + io.tool_error(f"Referenced session file not found: {reference}") + + return None + + return target + + +def read_payload(coder, io, data_file: Path, quiet: bool = False) -> dict | None: + """Read a session payload file, decrypting it when necessary.""" + try: + data = data_file.read_bytes() + + except OSError as e: + if not quiet: + io.tool_error(f"Error reading session: {e}") + + return None + + try: + if session_crypto.is_encrypted_payload(data): + args = getattr(coder, "args", None) + key_file = getattr(args, "session_key_file", None) if args else None + key = session_crypto.resolve_key(key_file=key_file) + if not key: + if not quiet: + io.tool_error( + "Session is encrypted but no key is configured " + f"({session_crypto.KEY_ENV} or --session-key-file)." + ) + + return None + + return session_crypto.decrypt_session_bytes(data, key) + + parsed = json.loads(data.decode("utf-8")) + if not isinstance(parsed, dict): + if not quiet: + io.tool_error("Invalid session format.") + + return None + + return parsed + + except session_crypto.SessionCryptoError as e: + if not quiet: + io.tool_error(str(e)) + + return None + + except (UnicodeDecodeError, json.JSONDecodeError) as e: + if not quiet: + io.tool_error(f"Error loading session: {e}") + + return None + + +def write_payload(coder, io, data_file: Path, session_data: dict) -> bool: + """Write a session payload, encrypting it when configured.""" + encrypt_enabled, key = encrypt_settings(coder) + + try: + data_file.parent.mkdir(parents=True, exist_ok=True) + + if encrypt_enabled: + if not key: + io.tool_error( + "Session encryption is enabled but no key is configured " + f"({session_crypto.KEY_ENV} or --session-key-file)." + ) + + return False + + data_file.write_bytes(session_crypto.encrypt_session_dict(session_data, key)) + else: + with safe_open(data_file, "w") as f: + json.dump(session_data, f, indent=2) + + return True + + except session_crypto.SessionCryptoError as e: + io.tool_error(str(e)) + + return False + + except OSError as e: + io.tool_error(f"Error saving session: {e}") + + return False + + +def write_reference(coder, io, reference_file: Path, session_name: str, data_file: Path) -> bool: + """Write a sessions-directory pointer to a payload stored elsewhere.""" + reference = { + "version": 1, + "type": SESSION_REFERENCE_TYPE, + "session_name": session_name, + "path": _relative_to_root(coder, data_file), + } + + try: + reference_file.parent.mkdir(parents=True, exist_ok=True) + with safe_open(reference_file, "w") as f: + json.dump(reference, f, indent=2) + + return True + + except OSError as e: + io.tool_error(f"Error saving session reference: {e}") + + return False + + +def describe_session(coder, io, name: str, locator: Path) -> Optional[Dict]: + """Build a list-row description for a saved session locator.""" + data_file = resolve_payload_file(coder, io, locator, quiet=True) + if data_file is None: + return None + + raw = data_file.read_bytes() + encrypted = session_crypto.is_encrypted_payload(raw) + + if encrypted: + _, key = encrypt_settings(coder) + if not key: + return { + "name": name, + "file": locator, + "model": "encrypted", + "edit_format": "—", + "num_messages": 0, + "num_files": 0, + "num_sub_agents": 0, + "encrypted": True, + } + + session_data = session_crypto.decrypt_session_bytes(raw, key) + else: + session_data = json.loads(raw.decode("utf-8")) + if not isinstance(session_data, dict): + raise ValueError("not a session object") + + chat_history = session_data.get("chat_history", {}) + files = session_data.get("files", {}) + + return { + "name": name, + "file": locator, + "model": session_data.get("model", "unknown"), + "edit_format": session_data.get("edit_format", "unknown"), + "num_messages": ( + len(chat_history.get("done_messages", [])) + len(chat_history.get("cur_messages", [])) + ), + "num_files": ( + len(files.get("editable", [])) + + len(files.get("read_only", [])) + + len(files.get("read_only_stubs", [])) + ), + "num_sub_agents": len(sub_agent_payloads(data_file)), + "encrypted": encrypted, + } + + +def _relative_to_root(coder, data_file: Path) -> str: + """Store payload paths relative to the coder root when possible.""" + try: + return str(data_file.relative_to(coder.abs_root_path(""))) + + except ValueError: + return str(data_file) diff --git a/cecli/helpers/sessions/subagents.py b/cecli/helpers/sessions/subagents.py new file mode 100644 index 00000000000..f77daa6daf7 --- /dev/null +++ b/cecli/helpers/sessions/subagents.py @@ -0,0 +1,152 @@ +"""Detecting, saving, and resolving the sub-agents that belong to a session.""" + +from __future__ import annotations + +import logging +from pathlib import Path +from typing import Dict, List, Optional, Set, Tuple + +from .layout import SUB_AGENT_PAYLOAD_NAME +from .payload import build_payload +from .storage import write_payload + +logger = logging.getLogger(__name__) + + +def detect_agent_name(coder) -> Optional[str]: + """Return a coder's sub-agent type, or ``None`` for the primary agent.""" + coder_uuid = getattr(coder, "uuid", None) + parent_uuid = getattr(coder, "parent_uuid", None) + if not isinstance(coder_uuid, str) or not coder_uuid: + return None + + if not isinstance(parent_uuid, str) or not parent_uuid: + return None + + try: + from cecli.helpers.agents.service import AgentService + + return AgentService.get_instance(coder).get_agent_name(coder) + + except Exception: + return None + + +def live_sub_agents(coder) -> List[Tuple[str, object]]: + """Return ``(agent_name, sub_coder)`` for every live sub-agent of ``coder``. + + Descendants at any depth are included so a saved session captures the whole + delegation tree, not just the direct children. + """ + try: + from cecli.helpers.agents.service import AgentService + + service = AgentService.get_instance(coder) + + except Exception: + return [] + + infos: Dict[str, object] = {} + for info in list(service.sub_agents.values()): + sub_coder = getattr(info, "coder", None) + sub_uuid = getattr(sub_coder, "uuid", None) if sub_coder is not None else None + if sub_coder is not None and sub_uuid is not None: + infos[str(sub_uuid)] = info + + root_uuid = str(getattr(coder, "uuid", "")) + + agents = [] + for info in infos.values(): + agent_name = getattr(info, "name", None) + if agent_name and _descends_from(info, root_uuid, infos): + agents.append((agent_name, info.coder)) + + return agents + + +def save_sub_agents(coder, io, session_name: str, subs_dir: Path) -> int: + """Write each descendant sub-agent payload into ``subs_dir/{child}/agent.json``. + + Returns the number of sub-agents saved. + """ + subs_dir.mkdir(parents=True, exist_ok=True) + + used: Set[str] = set() + written = 0 + for agent_name, sub_coder in live_sub_agents(coder): + child_dir = subs_dir / _unique_child_name(agent_name, used) + payload = build_payload(sub_coder, io, session_name, agent_name=agent_name) + + if write_payload(sub_coder, io, child_dir / SUB_AGENT_PAYLOAD_NAME, payload): + written += 1 + + return written + + +def resolve_reload_agent_name(name: str, root: Optional[str]) -> Optional[str]: + """Resolve a stored agent type to a registered name, falling back to ``worker``.""" + from cecli.helpers.agents.service import AgentService + + registry = AgentService.get_registry() + + if name not in registry and name.startswith("ws:") and root: + ensure_ws_agent_registered(name, str(root)) + + registry = AgentService.get_registry() + + if name in registry: + return name + + if "worker" in registry: + return "worker" + + return None + + +def ensure_ws_agent_registered(name: str, root: str) -> None: + """Register a ``ws:`` sub-agent for a stored root when it is missing.""" + from cecli.helpers.workspaces.subagents import register_workspace_subagents + + project_name = name[3:] + register_workspace_subagents( + { + "name": project_name, + "projects": [{"name": project_name, "path": root}], + } + ) + + +def _descends_from(info, root_uuid: str, infos: Dict[str, object]) -> bool: + """Return True when ``info`` sits anywhere below ``root_uuid`` in the tree.""" + seen: Set[str] = set() + + current = info + while current is not None: + parent_uuid = getattr(current, "parent_uuid", None) + if not parent_uuid: + return False + + if parent_uuid == root_uuid: + return True + + if parent_uuid in seen: + return False + + seen.add(parent_uuid) + current = infos.get(parent_uuid) + + return False + + +def _unique_child_name(agent_name: str, used: Set[str]) -> str: + """Return a filesystem-safe, unique directory name for a sub-agent.""" + safe = "".join(c if c.isalnum() or c in "-_" else "_" for c in agent_name) or "agent" + candidate = safe + suffix = 2 + while candidate in used: + candidate = f"{safe}-{suffix}" + suffix += 1 + + used.add(candidate) + + return candidate diff --git a/cecli/main.py b/cecli/main.py index e200a1121d4..b59db284c83 100644 --- a/cecli/main.py +++ b/cecli/main.py @@ -1488,7 +1488,7 @@ def get_io(pretty): explicit_yes_required=True, ): try: - from cecli.sessions import SessionManager + from cecli.helpers.sessions import SessionManager session_manager = SessionManager(coder, io) await session_manager.load_session( diff --git a/cecli/sessions.py b/cecli/sessions.py deleted file mode 100644 index 01a2d0b3a0b..00000000000 --- a/cecli/sessions.py +++ /dev/null @@ -1,505 +0,0 @@ -"""Session management utilities for cecli.""" - -import json -import os -from pathlib import Path -from typing import Dict, List, Optional - -from cecli import models -from cecli.decoding import safe_open -from cecli.helpers import crypto as session_crypto -from cecli.helpers.conversation import ConversationService, MessageTag - - -class SessionManager: - """Manages chat session saving, listing, and loading.""" - - def __init__(self, coder, io): - self.coder = coder - self.io = io - - def save_session(self, session_name: str, output=True) -> bool: - """Save the current chat session to a named file.""" - if not session_name: - if output: - self.io.tool_error("Please provide a session name.") - return False - - session_name = session_name.replace(".json", "") - session_dir = self._get_session_directory() - session_file = session_dir / f"{session_name}.json" - - if session_file.exists(): - if output: - self.io.tool_warning(f"Session '{session_name}' already exists. Overwriting.") - - try: - session_data = self._build_session_data(session_name) - if not self._write_session_file(session_file, session_data): - return False - - if output: - suffix = " (encrypted)" if self._session_encrypt_settings()[0] else "" - self.io.tool_output(f"Session saved: {session_file}{suffix}") - - return True - - except Exception as e: - self.io.tool_error(f"Error saving session: {e}") - return False - - def list_sessions(self) -> List[Dict]: - """List all saved sessions with metadata.""" - session_dir = self._get_session_directory() - session_files = list(session_dir.glob("*.json")) - - if not session_files: - self.io.tool_output("No saved sessions found.") - return [] - - sessions = [] - for session_file in sorted(session_files, key=lambda x: x.stat().st_mtime, reverse=True): - try: - raw = session_file.read_bytes() - if session_crypto.is_encrypted_payload(raw): - _, key = self._session_encrypt_settings() - if not key: - sessions.append( - { - "name": session_file.stem, - "file": session_file, - "model": "encrypted", - "edit_format": "—", - "num_messages": 0, - "num_files": 0, - "encrypted": True, - } - ) - continue - session_data = session_crypto.decrypt_session_bytes(raw, key) - else: - session_data = json.loads(raw.decode("utf-8")) - if not isinstance(session_data, dict): - raise ValueError("not a session object") - - session_info = { - "name": session_file.stem, - "file": session_file, - "model": session_data.get("model", "unknown"), - "edit_format": session_data.get("edit_format", "unknown"), - "num_messages": ( - len(session_data.get("chat_history", {}).get("done_messages", [])) - + len(session_data.get("chat_history", {}).get("cur_messages", [])) - ), - "num_files": ( - len(session_data.get("files", {}).get("editable", [])) - + len(session_data.get("files", {}).get("read_only", [])) - + len(session_data.get("files", {}).get("read_only_stubs", [])) - ), - "encrypted": session_crypto.is_encrypted_payload(raw), - } - sessions.append(session_info) - - except Exception as e: - self.io.tool_output(f" {session_file.stem} [error reading: {e}]") - - return sessions - - async def load_session(self, session_identifier: str, switch=True, quiet: bool = False) -> bool: - """Load a saved session by name or file path.""" - if not session_identifier: - self.io.tool_error("Please provide a session name or file path.") - return False - - # Try to find the session file - session_file = self._find_session_file(session_identifier) - if not session_file: - return False - - session_data = self._read_session_file(session_file, quiet=quiet) - if session_data is None: - return False - - if not isinstance(session_data, dict) or "version" not in session_data: - if not quiet: - self.io.tool_error("Invalid session format.") - return False - - # Apply session data - applied, loaded_edit_format = await self._apply_session_data(session_data, session_file) - if applied and switch: - from cecli.commands import SwitchCoderSignal - - edit_format_to_switch_to = self.coder.edit_format - if loaded_edit_format: - edit_format_to_switch_to = loaded_edit_format - self.coder.edit_format = loaded_edit_format - - raise SwitchCoderSignal( - edit_format=edit_format_to_switch_to, - from_coder=self.coder, - summarize_from_coder=False, - show_announcements=True, - ) - return applied - - def _get_session_directory(self) -> Path: - """Get the session directory, creating it if necessary.""" - session_dir = Path(self.coder.abs_root_path(".cecli/sessions")) - os.makedirs(session_dir, exist_ok=True) - return session_dir - - def _session_encrypt_settings(self) -> tuple[bool, bytes | None]: - args = getattr(self.coder, "args", None) - if not args or not getattr(args, "session_encrypt", False): - return False, None - key_file = getattr(args, "session_key_file", None) - return True, session_crypto.resolve_key(key_file=key_file) - - def _read_session_file(self, session_file: Path, quiet: bool = False) -> dict | None: - try: - data = session_file.read_bytes() - except OSError as e: - if not quiet: - self.io.tool_error(f"Error reading session: {e}") - return None - try: - if session_crypto.is_encrypted_payload(data): - args = getattr(self.coder, "args", None) - key_file = getattr(args, "session_key_file", None) if args else None - key = session_crypto.resolve_key(key_file=key_file) - if not key: - if not quiet: - self.io.tool_error( - "Session is encrypted but no key is configured " - f"({session_crypto.KEY_ENV} or --session-key-file)." - ) - return None - return session_crypto.decrypt_session_bytes(data, key) - parsed = json.loads(data.decode("utf-8")) - if not isinstance(parsed, dict): - if not quiet: - self.io.tool_error("Invalid session format.") - return None - return parsed - except session_crypto.SessionCryptoError as e: - if not quiet: - self.io.tool_error(str(e)) - return None - except (UnicodeDecodeError, json.JSONDecodeError) as e: - if not quiet: - self.io.tool_error(f"Error loading session: {e}") - return None - - def _write_session_file(self, session_file: Path, session_data: dict) -> bool: - encrypt_enabled, key = self._session_encrypt_settings() - try: - if encrypt_enabled: - if not key: - self.io.tool_error( - "Session encryption is enabled but no key is configured " - f"({session_crypto.KEY_ENV} or --session-key-file)." - ) - return False - session_file.write_bytes(session_crypto.encrypt_session_dict(session_data, key)) - else: - with safe_open(session_file, "w") as f: - json.dump(session_data, f, indent=2) - return True - except session_crypto.SessionCryptoError as e: - self.io.tool_error(str(e)) - return False - except OSError as e: - self.io.tool_error(f"Error saving session: {e}") - return False - - def _build_session_data(self, session_name) -> Dict: - """Build session data dictionary from current coder state.""" - # Get relative paths for all files - editable_files = [ - self.coder.get_rel_fname(abs_fname) for abs_fname in self.coder.abs_fnames - ] - read_only_files = [ - self.coder.get_rel_fname(abs_fname) for abs_fname in self.coder.abs_read_only_fnames - ] - read_only_stubs_files = [ - self.coder.get_rel_fname(abs_fname) - for abs_fname in self.coder.abs_read_only_stubs_fnames - ] - - # Capture todo list content so it can be restored with the session - todo_content = None - try: - todo_path = self.coder.abs_root_path(self.coder.local_agent_folder("todo.txt")) - if os.path.isfile(todo_path): - todo_content = self.io.read_text(todo_path) - if todo_content is None: - todo_content = "" - except Exception as e: - self.io.tool_warning(f"Could not read todo list file: {e}") - - # Get CUR and DONE messages from ConversationManager - connected_mcps = [] - if hasattr(self.coder, "mcp_manager") and self.coder.mcp_manager: - connected_mcps = [server.name for server in self.coder.mcp_manager.connected_servers] - - # Get CUR and DONE messages from ConversationManager - connected_mcps = [] - if hasattr(self.coder, "mcp_manager") and self.coder.mcp_manager: - connected_mcps = [server.name for server in self.coder.mcp_manager.connected_servers] - - skills_data = None - if hasattr(self.coder, "skills_manager") and self.coder.skills_manager: - skills_data = { - "skills_paths": [str(p) for p in self.coder.skills_manager.directory_paths], - "skills_includelist": ( - list(self.coder.skills_manager.include_list) - if self.coder.skills_manager.include_list is not None - else [] - ), - "skills_excludelist": ( - list(self.coder.skills_manager.exclude_list) - if self.coder.skills_manager.exclude_list is not None - else [] - ), - } - - agent_config_data = None - if hasattr(self.coder, "agent_config"): - agent_config_data = { - "tools_paths": self.coder.agent_config.get("tools_paths", []), - "tools_includelist": self.coder.agent_config.get("tools_includelist", []), - "tools_excludelist": self.coder.agent_config.get("tools_excludelist", []), - } - - # Flush any queued messages so the saved chat history is complete - ConversationService.get_manager(self.coder).flush_queue() - - return { - "version": 1, - "session_name": session_name, - "model": self.coder.main_model.name, - "weak_model": self.coder.main_model.weak_model.name, - "editor_model": self.coder.main_model.editor_model.name, - "agent_model": self.coder.main_model.agent_model.name, - "editor_edit_format": self.coder.main_model.editor_edit_format, - "edit_format": self.coder.edit_format, - "chat_history": { - "done_messages": ( - ConversationService.get_manager(self.coder).get_messages_dict(MessageTag.DONE) - ), - "cur_messages": ( - ConversationService.get_manager(self.coder).get_messages_dict(MessageTag.CUR) - ), - }, - "files": { - "editable": editable_files, - "read_only": read_only_files, - "read_only_stubs": read_only_stubs_files, - }, - "settings": { - "auto_commits": self.coder.auto_commits, - "auto_lint": self.coder.auto_lint, - "auto_test": self.coder.auto_test, - }, - "todo_list": todo_content, - "mcps": connected_mcps, - "skills": skills_data, - "tools": agent_config_data, - "usage": { - "total_tokens_sent": self.coder.total_tokens_sent, - "total_tokens_received": self.coder.total_tokens_received, - "total_cached_tokens": self.coder.total_cached_tokens, - "total_cost": self.coder.total_cost, - }, - } - - def _find_session_file(self, session_identifier: str) -> Optional[Path]: - """Find session file by name or path.""" - # Check if it's a direct file path - session_file = Path(session_identifier) - if session_file.exists(): - return session_file - - # Check if it's a session name in the sessions directory - session_dir = self._get_session_directory() - - # Try with .json extension - if not session_identifier.endswith(".json"): - session_file = session_dir / f"{session_identifier}.json" - if session_file.exists(): - return session_file - - session_file = session_dir / f"{session_identifier}" - if session_file.exists(): - return session_file - - self.io.tool_error(f"Session not found: {session_identifier}") - self.io.tool_output("Use /list-sessions to see available sessions.") - return None - - async def _apply_session_data( - self, session_data: Dict, session_file: Path - ) -> (bool, Optional[str]): - """Apply session data to current coder state. - - Returns: - A tuple of (success, edit_format) - """ - try: - # Clear current state - self.coder.abs_fnames = set() - self.coder.abs_read_only_fnames = set() - self.coder.abs_read_only_stubs_fnames = set() - - # Load files - files = session_data.get("files", {}) - for rel_fname in files.get("editable", []): - abs_fname = self.coder.abs_root_path(rel_fname) - if os.path.exists(abs_fname): - self.coder.abs_fnames.add(abs_fname) - else: - self.io.tool_warning(f"File not found, skipping: {rel_fname}") - - for rel_fname in files.get("read_only", []): - abs_fname = self.coder.abs_root_path(rel_fname) - if os.path.exists(abs_fname): - self.coder.abs_read_only_fnames.add(abs_fname) - else: - self.io.tool_warning(f"File not found, skipping: {rel_fname}") - - for rel_fname in files.get("read_only_stubs", []): - abs_fname = self.coder.abs_root_path(rel_fname) - if os.path.exists(abs_fname): - self.coder.abs_read_only_stubs_fnames.add(abs_fname) - else: - self.io.tool_warning(f"File not found, skipping: {rel_fname}") - - # Load usage stats - usage = session_data.get("usage", {}) - self.coder.total_tokens_sent = usage.get("total_tokens_sent", 0) - self.coder.total_tokens_received = usage.get("total_tokens_received", 0) - self.coder.total_cached_tokens = usage.get("total_cached_tokens", 0) - self.coder.total_cost = usage.get("total_cost", 0.0) - # Loading a session seeds the cumulative counters but does not represent - # recent API usage, so clear the rolling token-rate buffer. - self.coder._reset_token_usage() - if session_data.get("model"): - self.coder.main_model = models.Model( - session_data.get("model", self.coder.args.model), - weak_model=session_data.get("weak_model", self.coder.args.weak_model), - editor_model=session_data.get("editor_model", self.coder.args.editor_model), - agent_model=session_data.get("agent_model", self.coder.args.agent_model), - editor_edit_format=session_data.get( - "editor_edit_format", self.coder.args.editor_edit_format - ), - io=self.io, - verbose=self.coder.args.verbose, - retries=self.coder.main_model.retries, - debug=self.coder.main_model.debug, - ) - - # Load settings - settings = session_data.get("settings", {}) - if "auto_commits" in settings: - self.coder.auto_commits = settings["auto_commits"] - if "auto_lint" in settings: - self.coder.auto_lint = settings["auto_lint"] - if "auto_test" in settings: - self.coder.auto_test = settings["auto_test"] - - # Restore todo list content if present in the session - if "todo_list" in session_data: - todo_path = self.coder.abs_root_path(self.coder.local_agent_folder("todo.txt")) - todo_content = session_data.get("todo_list") - try: - if todo_content is None: - if os.path.exists(todo_path): - os.remove(todo_path) - else: - self.io.write_text(todo_path, todo_content) - except Exception as e: - self.io.tool_warning(f"Could not restore todo list: {e}") - - # Clear CUR and DONE messages from ConversationManager - ConversationService.get_manager(self.coder).reset() - ConversationService.get_files(self.coder).reset() - self.coder.format_chat_chunks() - - # Load chat history - chat_history = session_data.get("chat_history", {}) - done_messages = chat_history.get("done_messages", []) - cur_messages = chat_history.get("cur_messages", []) - - # Add messages to ConversationManager (source of truth) - # Add done messages - for msg in done_messages: - ConversationService.get_manager(self.coder).add_message( - message_dict=msg, - tag=MessageTag.DONE, - ) - # Add current messages - for msg in cur_messages: - ConversationService.get_manager(self.coder).add_message( - message_dict=msg, - tag=MessageTag.CUR, - ) - - self.io.tool_output( - f"Session loaded: {session_data.get('session_name', session_file.stem)}" - ) - self.io.tool_output( - f"Model: {session_data.get('model', 'unknown')}, Edit format:" - f" {session_data.get('edit_format', 'unknown')}" - ) - - # Show summary - num_messages = len(self.coder.done_messages) + len(self.coder.cur_messages) - num_files = ( - len(self.coder.abs_fnames) - + len(self.coder.abs_read_only_fnames) - + len(self.coder.abs_read_only_stubs_fnames) - ) - self.io.tool_output(f"Loaded {num_messages} messages and {num_files} files") - - # Load MCPs - saved_mcps = session_data.get("mcps", []) - if hasattr(self.coder, "mcp_manager") and self.coder.mcp_manager: - current_mcps = {server.name for server in self.coder.mcp_manager.connected_servers} - saved_mcps_set = set(saved_mcps) - - to_disconnect = current_mcps - saved_mcps_set - for mcp_name in to_disconnect: - await self.coder.mcp_manager.disconnect_server(mcp_name) - - to_connect = saved_mcps_set - current_mcps - for mcp_name in to_connect: - await self.coder.mcp_manager.connect_server(mcp_name) - - # Load skills - skills_data = session_data.get("skills") - if skills_data and hasattr(self.coder, "skills_manager") and self.coder.skills_manager: - self.coder.skills_manager.directory_paths = skills_data.get("skills_paths", []) - self.coder.skills_manager.include_list = set( - skills_data.get("skills_includelist", []) - ) - self.coder.skills_manager.exclude_list = set( - skills_data.get("skills_excludelist", []) - ) - - # Load tools config - agent_config_data = session_data.get("tools") - if agent_config_data and hasattr(self.coder, "agent_config"): - self.coder.agent_config.update(agent_config_data) - from cecli.tools.utils.registry import ToolRegistry - - ToolRegistry.build_registry(agent_config=self.coder.agent_config) - self.coder.loaded_custom_tools = ToolRegistry.loaded_custom_tools - - # Return True and the edit format so the Coder can be switched - edit_format = session_data.get("edit_format") - return True, edit_format - - except Exception as e: - self.io.tool_error(f"Error applying session data: {e}") - return False, None diff --git a/tests/basic/test_sessions.py b/tests/basic/test_sessions.py index 7fafd220963..116880439b9 100644 --- a/tests/basic/test_sessions.py +++ b/tests/basic/test_sessions.py @@ -5,8 +5,8 @@ import pytest +from cecli.helpers.sessions import SessionManager from cecli.io import InputOutput -from cecli.sessions import SessionManager @pytest.fixture diff --git a/tests/basic/test_sessions_manager.py b/tests/basic/test_sessions_manager.py index 3e2769e31e0..47c1e7476e5 100644 --- a/tests/basic/test_sessions_manager.py +++ b/tests/basic/test_sessions_manager.py @@ -6,13 +6,13 @@ import os from pathlib import Path from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest from cecli.helpers import crypto as session_crypto +from cecli.helpers.sessions import SessionManager from cecli.io import InputOutput -from cecli.sessions import SessionManager def _prepare_workspace(coder, tmp_path) -> Path: @@ -38,15 +38,15 @@ def mock_coder(monkeypatch): conv_manager.get_messages_dict.return_value = [] files_manager = MagicMock() monkeypatch.setattr( - "cecli.sessions.ConversationService.get_manager", + "cecli.helpers.sessions.payload.ConversationService.get_manager", lambda _coder: conv_manager, ) monkeypatch.setattr( - "cecli.sessions.ConversationService.get_files", + "cecli.helpers.sessions.payload.ConversationService.get_files", lambda _coder: files_manager, ) monkeypatch.setattr( - "cecli.sessions.models.Model", + "cecli.helpers.sessions.payload.models.Model", lambda *args, **kwargs: main_model, ) @@ -217,3 +217,257 @@ async def test_load_encrypted_using_env_key_only(encrypt_coder, session_key_env, assert await manager.load_session(str(path), switch=False) is True loaded = session_crypto.decrypt_session_bytes(path.read_bytes(), session_key_env) assert loaded["edit_format"] == "architect" + + +def test_save_agent_folder_writes_reference(session_manager, mock_coder, tmp_path): + """Agent-folder saves keep the payload in the agent folder and a pointer in sessions/.""" + root = _prepare_workspace(mock_coder, tmp_path) + assert session_manager.save_session("auto-save", output=False, to_agent_folder=True) + + reference_file = root / ".cecli" / "sessions" / "auto-save.json" + reference = json.loads(reference_file.read_text(encoding="utf-8")) + assert reference["type"] == SessionManager.SESSION_REFERENCE_TYPE + assert reference["session_name"] == "auto-save" + + payload_file = root / reference["path"] + assert payload_file != reference_file + assert payload_file.exists() + assert json.loads(payload_file.read_text(encoding="utf-8"))["session_name"] == "auto-save" + + +def test_list_resolves_agent_folder_reference(session_manager, mock_coder, tmp_path): + _prepare_workspace(mock_coder, tmp_path) + assert session_manager.save_session("auto-save", output=False, to_agent_folder=True) + + rows = session_manager.list_sessions() + assert len(rows) == 1 + assert rows[0]["name"] == "auto-save" + assert rows[0]["model"] == "test_model" + + +@pytest.mark.asyncio +async def test_load_agent_folder_reference(session_manager, mock_coder, tmp_path): + root = _prepare_workspace(mock_coder, tmp_path) + assert session_manager.save_session("auto-save", output=False, to_agent_folder=True) + + reference_file = root / ".cecli" / "sessions" / "auto-save.json" + assert await session_manager.load_session(str(reference_file), switch=False) is True + + +def test_sub_agent_save_keeps_primary_reference(mock_coder, monkeypatch, tmp_path): + """A sub-agent auto-save must not clobber the primary session reference.""" + root = _prepare_workspace(mock_coder, tmp_path) + manager = SessionManager(mock_coder, mock_coder.io) + assert manager.save_session("auto-save", output=False, to_agent_folder=True) + + reference_file = root / ".cecli" / "sessions" / "auto-save.json" + original_reference = reference_file.read_text(encoding="utf-8") + + # Sub-agents write to their own agent folder; the reference stays put. + sub_dir = root / ".cecli" / "agents" / "sub456" + sub_dir.mkdir(parents=True, exist_ok=True) + mock_coder.local_agent_folder.side_effect = lambda x: f".cecli/agents/sub456/{x}" + monkeypatch.setattr(manager, "_sub_agent_name", lambda: "worker") + + assert manager.save_session("auto-save", output=False, to_agent_folder=True) + + assert reference_file.read_text(encoding="utf-8") == original_reference + + payload_file = sub_dir / "auto-save.json" + assert payload_file.exists() + data = json.loads(payload_file.read_text(encoding="utf-8")) + assert data["agent_name"] == "worker" + assert data["agent_root"] is None + + +def test_sub_agent_save_records_workspace_root(mock_coder, monkeypatch, tmp_path): + """Workspace sub-agents persist their type and root for later restoration.""" + root = _prepare_workspace(mock_coder, tmp_path) + sub_dir = root / ".cecli" / "agents" / "sub789" + sub_dir.mkdir(parents=True, exist_ok=True) + mock_coder.local_agent_folder.side_effect = lambda x: f".cecli/agents/sub789/{x}" + mock_coder.root = "/workspace/app" + + manager = SessionManager(mock_coder, mock_coder.io) + monkeypatch.setattr(manager, "_sub_agent_name", lambda: "ws:app") + + assert manager.save_session("auto-save", output=False, to_agent_folder=True) + + data = json.loads((sub_dir / "auto-save.json").read_text(encoding="utf-8")) + assert data["agent_name"] == "ws:app" + assert data["agent_root"] == "/workspace/app" + assert not (root / ".cecli" / "sessions" / "auto-save.json").exists() + + +@pytest.mark.asyncio +async def test_load_primary_reference_reloads_sub_agents(mock_coder, monkeypatch, tmp_path): + """Loading a primary session rebuilds the sub-agents saved beside it.""" + root = _prepare_workspace(mock_coder, tmp_path) + manager = SessionManager(mock_coder, mock_coder.io) + assert manager.save_session("auto-save", output=False, to_agent_folder=True) + + # Sub-agent payloads live in the "s/{uuid}/" directory next to the primary. + sub_dir = root / ".cecli" / "s" / "sub123" + sub_dir.mkdir(parents=True, exist_ok=True) + (sub_dir / "auto-save.json").write_text( + json.dumps({"version": 1, "session_name": "auto-save", "agent_name": "worker"}), + encoding="utf-8", + ) + + fake_service = MagicMock() + fake_service.spawn = AsyncMock(return_value=(MagicMock(), MagicMock())) + monkeypatch.setattr( + "cecli.helpers.agents.service.AgentService.get_instance", + classmethod(lambda cls, coder: fake_service), + ) + monkeypatch.setattr( + "cecli.helpers.agents.service.AgentService.get_registry", + classmethod(lambda cls: {"worker": object()}), + ) + monkeypatch.setattr(SessionManager, "_apply_session_data", AsyncMock(return_value=(True, None))) + + reference_file = root / ".cecli" / "sessions" / "auto-save.json" + assert await manager.load_session(str(reference_file), switch=False) is True + + fake_service.spawn.assert_awaited_once_with( + "worker", parent=mock_coder, auto_reap=False, independent=True + ) + + +def test_resolve_reload_agent_name_registers_workspace_agent(mock_coder, tmp_path): + """A stored ``ws:`` agent is re-registered from its persisted root.""" + from cecli.helpers.agents.service import AgentService + from cecli.helpers.sessions import subagents + from cecli.utils import make_repo + + project = tmp_path / "app" + project.mkdir(parents=True, exist_ok=True) + make_repo(project) + + registry = AgentService.get_registry() + registry.pop("ws:app", None) + + try: + name = subagents.resolve_reload_agent_name("ws:app", str(project)) + + assert name == "ws:app" + assert AgentService.get_registry()["ws:app"].metadata["root"] == str(project.resolve()) + finally: + registry.pop("ws:app", None) + + +def _fake_sub_coder(): + coder = MagicMock() + coder.args = SimpleNamespace(session_encrypt=False, session_key_file=None) + + return coder + + +def _fake_sub_services(monkeypatch, sub_agents): + monkeypatch.setattr( + "cecli.helpers.sessions.subagents.live_sub_agents", lambda coder: sub_agents + ) + + +def test_save_session_with_sub_agents_writes_bundle(mock_coder, monkeypatch, tmp_path): + """An explicit save with sub-agents writes a folder bundle of payloads.""" + root = _prepare_workspace(mock_coder, tmp_path) + manager = SessionManager(mock_coder, mock_coder.io) + _fake_sub_services(monkeypatch, [("worker", _fake_sub_coder()), ("ws:app", _fake_sub_coder())]) + monkeypatch.setattr( + "cecli.helpers.sessions.subagents.build_payload", + lambda coder, io, name, agent_name=None: { + "version": 1, + "session_name": name, + "agent_name": agent_name, + }, + ) + + assert manager.save_session("team", output=False) + + bundle = root / ".cecli" / "sessions" / "team" + assert (bundle / "primary.json").is_file() + assert (bundle / "s" / "worker" / "agent.json").is_file() + assert (bundle / "s" / "ws_app" / "agent.json").is_file() + assert not (root / ".cecli" / "sessions" / "team.json").exists() + + data = json.loads((bundle / "s" / "ws_app" / "agent.json").read_text(encoding="utf-8")) + assert data["agent_name"] == "ws:app" + + +def test_save_session_without_sub_agents_writes_file(mock_coder, monkeypatch, tmp_path): + """An explicit save without sub-agents keeps the single-file shape.""" + root = _prepare_workspace(mock_coder, tmp_path) + manager = SessionManager(mock_coder, mock_coder.io) + _fake_sub_services(monkeypatch, []) + + assert manager.save_session("solo", output=False) + + assert (root / ".cecli" / "sessions" / "solo.json").is_file() + assert not (root / ".cecli" / "sessions" / "solo").exists() + + +@pytest.mark.asyncio +async def test_load_session_bundle_reloads_sub_agents(mock_coder, monkeypatch, tmp_path): + """Loading a folder-bundle session rebuilds its sub-agents.""" + root = _prepare_workspace(mock_coder, tmp_path) + bundle = root / ".cecli" / "sessions" / "team" + sub_dir = bundle / "s" / "worker" + sub_dir.mkdir(parents=True, exist_ok=True) + (bundle / "primary.json").write_text( + json.dumps({"version": 1, "session_name": "team"}), encoding="utf-8" + ) + (sub_dir / "agent.json").write_text( + json.dumps({"version": 1, "session_name": "team", "agent_name": "worker"}), + encoding="utf-8", + ) + + fake_service = MagicMock() + fake_service.spawn = AsyncMock(return_value=(MagicMock(), MagicMock())) + monkeypatch.setattr( + "cecli.helpers.agents.service.AgentService.get_instance", + classmethod(lambda cls, coder: fake_service), + ) + monkeypatch.setattr( + "cecli.helpers.agents.service.AgentService.get_registry", + classmethod(lambda cls: {"worker": object()}), + ) + monkeypatch.setattr(SessionManager, "_apply_session_data", AsyncMock(return_value=(True, None))) + + manager = SessionManager(mock_coder, mock_coder.io) + assert await manager.load_session("team", switch=False) is True + + fake_service.spawn.assert_awaited_once_with( + "worker", parent=mock_coder, auto_reap=False, independent=True + ) + + +def test_list_sessions_discovers_bundle(mock_coder, tmp_path): + """Folder-bundle sessions show up in listings with a sub-agent count.""" + root = _prepare_workspace(mock_coder, tmp_path) + bundle = root / ".cecli" / "sessions" / "team" + (bundle / "s" / "worker").mkdir(parents=True, exist_ok=True) + (bundle / "primary.json").write_text( + json.dumps( + { + "version": 1, + "session_name": "team", + "model": "test_model", + "edit_format": "diff", + "chat_history": {"done_messages": [{}], "cur_messages": []}, + "files": {"editable": ["file1.py"]}, + } + ), + encoding="utf-8", + ) + (bundle / "s" / "worker" / "agent.json").write_text( + json.dumps({"version": 1, "agent_name": "worker"}), encoding="utf-8" + ) + + manager = SessionManager(mock_coder, mock_coder.io) + rows = manager.list_sessions() + + assert len(rows) == 1 + assert rows[0]["name"] == "team" + assert rows[0]["model"] == "test_model" + assert rows[0]["num_sub_agents"] == 1