diff --git a/cecli/args.py b/cecli/args.py index a5224cfd9ec..2f3efc0120d 100644 --- a/cecli/args.py +++ b/cecli/args.py @@ -61,6 +61,7 @@ "security_config", "retries", "custom", + "system_one", "tui_config", } ) @@ -360,7 +361,7 @@ def get_parser(default_config_files, git_root): metavar="RETRIES_JSON", help=( 'Specify LLM retry configuration as a JSON/YAML string (e.g., \'{"retry_on_empty": ' - "true}')" + 'true, "retry-on-unauthorized": false}\')' ), default=None, ) @@ -427,6 +428,16 @@ def get_parser(default_config_files, git_root): help="Specify hooks configuration as a JSON string", default=None, ) + group.add_argument( + "--system-one", + metavar="SYSTEM_ONE_JSON", + help=( + "Specify the System One decision endpoint as a JSON/YAML string (e.g.," + ' \'{"api_base": "http://127.0.0.1:8000", "api_key_env": ["SYSTEM_ONE_API_KEY"],' + ' "model_name": "von-latest"}\')' + ), + default=None, + ) group.add_argument( "--agent-model", metavar="AGENT_MODEL", diff --git a/cecli/args_formatter.py b/cecli/args_formatter.py index aaa9463c3b3..07cd7078cab 100644 --- a/cecli/args_formatter.py +++ b/cecli/args_formatter.py @@ -138,6 +138,7 @@ def _format_action(self, action): parts.append("# retry-timeout: 60") parts.append("# retry-backoff-factor: 2.0") parts.append("# retry-on-unavailable: true") + parts.append("# retry-on-unauthorized: false") parts.append("# retry-on-empty: false") parts.append("") return "\n".join(parts) diff --git a/cecli/coders/base_coder.py b/cecli/coders/base_coder.py index cc62c69493e..126781384fd 100755 --- a/cecli/coders/base_coder.py +++ b/cecli/coders/base_coder.py @@ -54,7 +54,6 @@ from cecli.linter import Linter from cecli.llm import litellm from cecli.mcp import LocalServer -from cecli.models import RETRY_TIMEOUT from cecli.reasoning_tags import ( REASONING_TAG, format_reasoning_content, @@ -67,7 +66,7 @@ from cecli.run_cmd import run_cmd_async 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 +from cecli.utils import copy_tool_call, format_tokens, is_image_file, safe_relpath from ..dump import dump # noqa: F401 from ..prompts.utils.registry import PromptObject, PromptRegistry @@ -2829,24 +2828,14 @@ async def format_in_executor(): except EmptyResponseError: self.io.tool_warning(self.empty_llm_tool_warning()) - retry_on_empty = False - retries_config = self.get_active_model().retries - if isinstance(retries_config, str): - try: - retries_config = json.loads(retries_config) - except json.JSONDecodeError: - self.io.tool_warning( - f"Could not parse retries config: {retries_config}" - ) - retries_config = {} - if isinstance(retries_config, dict): - retry_on_empty = retries_config.get("retry_on_empty", False) + retry_config = models.parse_retry_config(self.get_active_model().retries) + retry_on_empty = retry_config["retry_on_empty"] if not retry_on_empty: break - retry_delay *= 2 - if retry_delay > RETRY_TIMEOUT: + retry_delay *= retry_config["retry_backoff_factor"] + if retry_delay > retry_config["retry_timeout"]: self.io.tool_error("Retry timeout exceeded on empty response.") break @@ -2866,10 +2855,13 @@ async def format_in_executor(): exhausted = True break - should_retry = ex_info.retry + retry_config = models.parse_retry_config(self.get_active_model().retries) + + should_retry, _ = models.parse_model_error(retry_config, err) + if should_retry: - retry_delay *= 2 - if retry_delay > RETRY_TIMEOUT: + retry_delay *= retry_config["retry_backoff_factor"] + if retry_delay > retry_config["retry_timeout"]: should_retry = False if not should_retry: @@ -3193,7 +3185,8 @@ def _find_mcp_server_for_tool(self, tool_call): tool_name_from_schema = nested.getter(tool, "function.name") if ( tool_name_from_schema - and tool_name_from_schema.lower() == unprefixed_tool_name.lower() + and responses.sanitize_tool_name(tool_name_from_schema).lower() + == responses.sanitize_tool_name(unprefixed_tool_name).lower() ): # Find the McpServer instance that will be used for communication for server in self.mcp_manager: @@ -3415,6 +3408,10 @@ async def call_mcp_tool_from_session(self, session, tool_call): # with a "missing required parameter" error. arguments = responses.coerce_tool_structure(arguments) + # Providers require a provider-safe tool name, so the name coming back + # from the model may not be the name the MCP server advertises. + name = responses.original_tool_name(name) + return await session.call_tool(name=name, arguments=arguments) async def process_tool_calls(self, tool_call_response): @@ -3504,6 +3501,7 @@ def mcp_tools(self, value): def get_tool_list(self): """Get a flattened list of all MCP tools with server prefixes, filtered by registered_servers.""" + responses.register_tool_names(self.mcp_tools) tool_list = [] if self.mcp_tools: for server_name, server_tools in self.mcp_tools: @@ -4877,7 +4875,7 @@ async def allowed_to_edit(self, path): return if not Path(full_path).exists(): - rel_path = os.path.relpath(full_path) + rel_path = safe_relpath(full_path) if not await self.io.confirm_ask(f"Create new file? ({rel_path})", subject=path): self.io.tool_output(f"Skipping edits to {path}") return diff --git a/cecli/commands/add.py b/cecli/commands/add.py index bc783805923..f28400fe1d5 100644 --- a/cecli/commands/add.py +++ b/cecli/commands/add.py @@ -10,7 +10,7 @@ parse_quoted_filenames, quote_filename, ) -from cecli.utils import is_image_file, run_fzf +from cecli.utils import is_image_file, run_fzf, safe_relpath class AddCommand(BaseCommand): @@ -67,7 +67,7 @@ async def execute(cls, io, coder, args, **kwargs): io.tool_output(f"You can add to git with: /git add {fname}") continue - confirm_fname = os.path.relpath(fname) + confirm_fname = safe_relpath(fname) if len(confirm_fname) > 64: confirm_fname = f".../{os.path.basename(confirm_fname)}" diff --git a/cecli/commands/core.py b/cecli/commands/core.py index 2977ad172ed..312a1802625 100644 --- a/cecli/commands/core.py +++ b/cecli/commands/core.py @@ -92,57 +92,6 @@ def __init__( # Commands that should NOT trigger auto-processing of the queue self._MANAGEMENT_COMMANDS = {"queue", "list-queue", "remove-queue"} - # ── Queue Management Methods (CLI-33) ────────────────────────────── # - # - # The prompt queue itself lives on the coder (Coder.prompt_queue) and - # is managed by cecli.helpers.command_queue. These thin wrappers keep - # the /queue, /list-queue and /remove-queue command implementations - # stable while operating on the coder that owns this Commands - # instance, so each sub-agent's commands manage that sub-agent's own - # queue. - - @property - def prompt_queue(self): - """Proxy to the owning coder's prompt queue.""" - coder = self.coder - return coder.prompt_queue if coder is not None else [] - - def _insert_prompt(self, text: str, index: int) -> dict: - """Insert a prompt at the given index in the active coder's queue.""" - from cecli.helpers import command_queue - - return command_queue.insert_prompt(self._active_coder(), text, index) - - def _enqueue_prompt(self, text: str) -> dict: - """Add a prompt to the owning coder's queue.""" - from cecli.helpers import command_queue - - return command_queue.enqueue_prompt(self.coder, text) - - def _dequeue_prompt(self) -> dict | None: - """Remove and return the first item from the owning coder's queue.""" - from cecli.helpers import command_queue - - return command_queue.dequeue_prompt(self.coder) - - def _get_queue_length(self) -> int: - """Return the current number of items in the owning coder's queue.""" - from cecli.helpers import command_queue - - return command_queue.get_queue_length(self.coder) - - def _remove_from_queue(self, index: int) -> dict | None: - """Remove and return the item at the given index from the owning coder's queue.""" - from cecli.helpers import command_queue - - return command_queue.remove_from_queue(self.coder, index) - - def _clear_queue(self) -> list: - """Remove all items from the owning coder's queue and return them.""" - from cecli.helpers import command_queue - - return command_queue.clear_queue(self.coder) - def _load_custom_commands(self, custom_commands): """ Load custom commands from plugin paths. diff --git a/cecli/commands/insert_queue.py b/cecli/commands/insert_queue.py index fec5b0cda17..f4cc1922c9f 100644 --- a/cecli/commands/insert_queue.py +++ b/cecli/commands/insert_queue.py @@ -4,6 +4,7 @@ from cecli.commands.utils.base_command import BaseCommand from cecli.commands.utils.helpers import format_command_result +from cecli.helpers import command_queue class InsertQueueCommand(BaseCommand): @@ -57,7 +58,7 @@ async def execute(cls, io, coder, args, **kwargs): # Happy path: insert the prompt try: - item = coder.commands._insert_prompt(prompt_text, index) + item = command_queue.insert_prompt(coder, prompt_text, index) io.tool_output(f"Prompt inserted at position {index + 1} (id: {item['id']})") return f"Successfully executed {cls.NORM_NAME}." except ValueError as e: diff --git a/cecli/commands/list_queue.py b/cecli/commands/list_queue.py index 85ab6d9d1f8..f699efbf939 100644 --- a/cecli/commands/list_queue.py +++ b/cecli/commands/list_queue.py @@ -5,6 +5,7 @@ from cecli.commands.utils.base_command import BaseCommand from cecli.commands.utils.helpers import format_command_result +from cecli.helpers import command_queue class ListQueueCommand(BaseCommand): @@ -30,7 +31,7 @@ async def execute(cls, io, coder, args, **kwargs): io, cls.NORM_NAME, "", error="Command system not available. Cannot list queue." ) - queue = coder.commands.prompt_queue + queue = command_queue.list_queue(coder) # Sad path: empty queue if not queue: diff --git a/cecli/commands/paste.py b/cecli/commands/paste.py index fd7fc4306b8..b24462b6a27 100644 --- a/cecli/commands/paste.py +++ b/cecli/commands/paste.py @@ -1,5 +1,9 @@ +import base64 import os +import shutil +import subprocess import tempfile +from io import BytesIO from pathlib import Path from typing import List @@ -20,8 +24,24 @@ class PasteCommand(BaseCommand): @classmethod async def execute(cls, io, coder, args, **kwargs): try: - # Check for image first - image = ImageGrab.grabclipboard() + image = cls._grab_clipboard_image(io) + + if not isinstance(image, Image.Image): + # Fall back to text when there is no image + text = cls._read_clipboard_text(io) + if text: + if coder.tui and coder.tui(): + coder.tui().set_input_value(text) + else: + coder.io.set_placeholder(text) + + return format_command_result(io, "paste", "Pasted text from clipboard") + + # On WSL the clipboard is owned by Windows and images are not always + # bridged to X11/Wayland, so probe the Windows clipboard directly. + if cls._is_wsl(): + image = cls._grab_clipboard_image_windows(io) + if isinstance(image, Image.Image): if args.strip(): filename = args.strip() @@ -36,6 +56,9 @@ async def execute(cls, io, coder, args, **kwargs): temp_dir = tempfile.mkdtemp() temp_file_path = os.path.join(temp_dir, basename) image_format = "PNG" if basename.lower().endswith(".png") else "JPEG" + if image_format == "JPEG" and image.mode not in ("RGB", "L", "CMYK"): + image = image.convert("RGB") + image.save(temp_file_path, image_format) abs_file_path = Path(temp_file_path).resolve() @@ -54,16 +77,6 @@ async def execute(cls, io, coder, args, **kwargs): return format_command_result(io, "paste", f"Added clipboard image: {abs_file_path}") - # If not an image, try to get text - text = pyperclip.paste() - if text: - if coder.tui and coder.tui(): - coder.tui().set_input_value(text) - else: - coder.io.set_placeholder(text) - - return format_command_result(io, "paste", "Pasted text from clipboard") - io.tool_error("No image or text content found in clipboard.") return format_command_result( io, "paste", "No content found in clipboard", Exception("No content") @@ -93,3 +106,101 @@ def get_help(cls) -> str: ) help_text += "If text is in the clipboard, it will be displayed in the chat.\n" return help_text + + @classmethod + def _grab_clipboard_image(cls, io): + """Return an image from the clipboard, or None if it cannot be read.""" + try: + return ImageGrab.grabclipboard() + except Exception as e: + # WAYLAND_DISPLAY can be set with no live compositor (e.g. WSLg or a + # stale SSH env), which makes the Wayland probe fail even though the + # bridged X11 clipboard works. Retry once over X11. + if os.environ.get("WAYLAND_DISPLAY"): + cls._log_clipboard_error(io, f"Wayland image probe failed: {e}") + return cls._grab_clipboard_image_x11(io) + + # The image probe can fail even when the text clipboard works, so keep + # it quiet unless the user asked for detail. + cls._log_clipboard_error(io, f"Clipboard image read failed: {e}") + return None + + @classmethod + def _grab_clipboard_image_x11(cls, io): + """Retry the image probe with the Wayland display disabled.""" + wayland_display = os.environ.pop("WAYLAND_DISPLAY", None) + try: + return ImageGrab.grabclipboard() + except Exception as e: + cls._log_clipboard_error(io, f"Clipboard image read failed over X11: {e}") + return None + finally: + if wayland_display is not None: + os.environ["WAYLAND_DISPLAY"] = wayland_display + + @classmethod + def _grab_clipboard_image_windows(cls, io): + """Read an image from the Windows clipboard via PowerShell (WSL only).""" + powershell = shutil.which("powershell.exe") + if not powershell: + return None + + script = ( + "Add-Type -AssemblyName System.Windows.Forms,System.Drawing\n" + "$img = [System.Windows.Forms.Clipboard]::GetImage()\n" + "if ($img -eq $null) { exit 1 }\n" + "$ms = New-Object System.IO.MemoryStream\n" + "$img.Save($ms, [System.Drawing.Imaging.ImageFormat]::Png)\n" + "[Convert]::ToBase64String($ms.ToArray())\n" + ) + try: + result = subprocess.run( + [powershell, "-NoProfile", "-STA", "-Command", script], + capture_output=True, + timeout=30, + ) + except Exception as e: + cls._log_clipboard_error(io, f"Windows clipboard image read failed: {e}") + return None + + if result.returncode != 0 or not result.stdout.strip(): + detail = result.stderr.decode("utf-8", "replace").strip() + cls._log_clipboard_error(io, f"Windows clipboard image read failed: {detail}") + return None + + try: + image = Image.open(BytesIO(base64.b64decode(result.stdout))) + image.load() + return image + except Exception as e: + cls._log_clipboard_error(io, f"Windows clipboard image decode failed: {e}") + return None + + @classmethod + def _is_wsl(cls): + """Detect whether the process is running inside WSL.""" + if os.environ.get("WSL_DISTRO_NAME") or os.environ.get("WSL_INTEROP"): + return True + + try: + with open("/proc/version", "r", encoding="utf-8", errors="ignore") as f: + return "microsoft" in f.read().lower() + except OSError: + return False + + @classmethod + def _read_clipboard_text(cls, io): + """Return clipboard text, or an empty string if it is empty/unavailable.""" + try: + return pyperclip.paste() or "" + except Exception as e: + # pyperclip's WSL backend raises on an empty Windows clipboard + # (GetBytes(null)), which is not an error worth surfacing. + cls._log_clipboard_error(io, f"Clipboard text read failed: {e}") + return "" + + @classmethod + def _log_clipboard_error(cls, io, message): + """Surface clipboard probe errors only in verbose mode.""" + if getattr(io, "verbose", False): + io.tool_error(message) diff --git a/cecli/commands/queue.py b/cecli/commands/queue.py index 0d3113b4d97..9c4a858377c 100644 --- a/cecli/commands/queue.py +++ b/cecli/commands/queue.py @@ -4,6 +4,7 @@ from cecli.commands.utils.base_command import BaseCommand from cecli.commands.utils.helpers import format_command_result +from cecli.helpers import command_queue class QueueCommand(BaseCommand): @@ -55,8 +56,8 @@ async def execute(cls, io, coder, args, **kwargs): # Happy path: enqueue the prompt try: - item = coder.commands._enqueue_prompt(prompt_text) - position = len(coder.commands.prompt_queue) + item = command_queue.enqueue_prompt(coder, prompt_text) + position = command_queue.get_queue_length(coder) io.tool_output(f"Prompt queued at position {position} (id: {item['id']})") return f"Successfully executed {cls.NORM_NAME}." except ValueError as e: diff --git a/cecli/commands/remove_queue.py b/cecli/commands/remove_queue.py index 4ef3c9c6b68..99a5970ddd5 100644 --- a/cecli/commands/remove_queue.py +++ b/cecli/commands/remove_queue.py @@ -4,6 +4,7 @@ from cecli.commands.utils.base_command import BaseCommand from cecli.commands.utils.helpers import format_command_result +from cecli.helpers import command_queue class RemoveQueueCommand(BaseCommand): @@ -33,14 +34,14 @@ async def execute(cls, io, coder, args, **kwargs): ) # Sad path: empty queue - if coder.commands._get_queue_length() == 0: + if command_queue.get_queue_length(coder) == 0: return format_command_result( io, cls.NORM_NAME, "", error="Queue is empty. Nothing to remove." ) # Handle wildcard: clear entire queue if args and args.strip() == "*": - items = coder.commands._clear_queue() + items = command_queue.clear_queue(coder) count = len(items) io.tool_output(f"Removed all {count} queued prompt(s).") return f"Successfully executed {cls.NORM_NAME}." @@ -57,9 +58,9 @@ async def execute(cls, io, coder, args, **kwargs): error=f"Invalid index: '{args.strip()}'. Please provide a number or '*'.", ) - item = coder.commands._remove_from_queue(index) + item = command_queue.remove_from_queue(coder, index) if item is None: - queue_len = coder.commands._get_queue_length() + queue_len = command_queue.get_queue_length(coder) return format_command_result( io, cls.NORM_NAME, @@ -71,7 +72,7 @@ async def execute(cls, io, coder, args, **kwargs): return f"Successfully executed {cls.NORM_NAME}." # Interactive mode: no args provided - queue = coder.commands.prompt_queue + queue = command_queue.list_queue(coder) io.tool_output("Queued prompts:") for i, item in enumerate(queue, 1): text = item["text"][:80] @@ -91,7 +92,7 @@ def get_completions(cls, io, coder, args) -> List[str]: if not coder.commands: return [] - queue_len = coder.commands._get_queue_length() + queue_len = command_queue.get_queue_length(coder) completions = [str(i) for i in range(1, queue_len + 1)] completions.append("*") return completions diff --git a/cecli/helpers/background_commands.py b/cecli/helpers/background_commands.py index 1a7f846558c..52f30ffeded 100644 --- a/cecli/helpers/background_commands.py +++ b/cecli/helpers/background_commands.py @@ -604,6 +604,9 @@ def start_background_command( close_fds=True, text=True, bufsize=1, + # Detach from the TUI's controlling terminal so the child cannot + # steal keystrokes or reset terminal modes that Textual relies on. + start_new_session=True, universal_newlines=True, ) os.close(slave_fd) diff --git a/cecli/helpers/config_utils.py b/cecli/helpers/config_utils.py index 14f375e50e8..0504932ff40 100644 --- a/cecli/helpers/config_utils.py +++ b/cecli/helpers/config_utils.py @@ -226,11 +226,22 @@ def read_and_merge_all_configs( return merged +def _safe_deepcopy(val, memo=None): + try: + return copy.deepcopy(val, memo) + except (TypeError, copy.Error): + if isinstance(val, dict): + return {k: _safe_deepcopy(v, memo) for k, v in val.items()} + elif isinstance(val, list): + return [_safe_deepcopy(item, memo) for item in val] + return val + + def deep_merge(dict1: dict, dict2: dict, deep_merge_arrays: bool = True) -> dict: """ Recursively merges dict2 into dict1. """ - merged = copy.deepcopy(dict1) + merged = _safe_deepcopy(dict1) for key, value in dict2.items(): if deep_merge_arrays and value is None: @@ -248,7 +259,7 @@ def deep_merge(dict1: dict, dict2: dict, deep_merge_arrays: bool = True) -> dict merged[key] = _deduplicate_list(merged[key], value) else: - merged[key] = copy.deepcopy(value) + merged[key] = _safe_deepcopy(value) return merged diff --git a/cecli/helpers/llms/constants.py b/cecli/helpers/llms/constants.py new file mode 100644 index 00000000000..64f861f987a --- /dev/null +++ b/cecli/helpers/llms/constants.py @@ -0,0 +1,34 @@ +"""Shared constant values for the llms package. + +Only values that more than one module needs live here; single-use constants stay +next to the code that reads them. +""" + +from __future__ import annotations + +from typing import FrozenSet + +#: litellm/cecli runtime kwargs consumed by the shim or the pipeline (routing, +#: auth, timeouts, cache plumbing) that must never be sent as request params. +CONTROL_KWARGS: FrozenSet[str] = frozenset( + { + "model", + "messages", + "stream", + "tools", + "functions", + "api_base", + "base_url", + "api_key", + "extra_headers", + "headers", + "timeout", + "drop_params", + "allowed_openai_params", + "custom_llm_provider", + "cache_control_injection_points", + "logger_fn", + } +) + +__all__ = ["CONTROL_KWARGS"] diff --git a/cecli/helpers/llms/domains/chat.py b/cecli/helpers/llms/domains/chat.py index 8757ff8f50f..05470f446ef 100644 --- a/cecli/helpers/llms/domains/chat.py +++ b/cecli/helpers/llms/domains/chat.py @@ -106,10 +106,20 @@ async def chat_complete( headers: Dict[str, str], kwargs: Dict[str, Any], ) -> CompletionResponse: - env = _openai_env_override() + provider = resolved.get("_provider") + + if provider is None or getattr(provider, "honors_openai_env_override", True): + env = _openai_env_override() + else: + env = None + base = env[0] if env else resolved["api_base"] - url = f"{base}/chat/completions" - payload = chat_payload(resolved, messages, tools, False, kwargs) + url = provider.chat_url(resolved, base) if provider else f"{base}/chat/completions" + payload = ( + provider.chat_payload(resolved, messages, tools, False, kwargs) + if provider + else chat_payload(resolved, messages, tools, False, kwargs) + ) hdrs = {"Content-Type": "application/json", **headers} body: Optional[bytes] = None signer = resolved.get("_signer") @@ -134,6 +144,9 @@ async def chat_complete( resp.raise_for_status() data = resp.json() + if provider: + return provider.parse_chat_response(data, resolved) + return normalize_chat_response(data, resolved["model"]) @@ -145,10 +158,20 @@ async def chat_stream( headers: Dict[str, str], kwargs: Dict[str, Any], ) -> AsyncIterator[CompletionChunk]: - env = _openai_env_override() + provider = resolved.get("_provider") + + if provider is None or getattr(provider, "honors_openai_env_override", True): + env = _openai_env_override() + else: + env = None + base = env[0] if env else resolved["api_base"] - url = f"{base}/chat/completions" - payload = chat_payload(resolved, messages, tools, True, kwargs) + url = provider.chat_url(resolved, base) if provider else f"{base}/chat/completions" + payload = ( + provider.chat_payload(resolved, messages, tools, True, kwargs) + if provider + else chat_payload(resolved, messages, tools, True, kwargs) + ) hdrs = {"Content-Type": "application/json", **headers} body: Optional[bytes] = None signer = resolved.get("_signer") @@ -173,8 +196,12 @@ async def chat_stream( resp.raise_for_status() last_finish_reason = None - async for json_obj in sse_json_lines(resp): - chunk = parse_chat_chunk(json_obj) + lines = provider.chat_stream_json(resp) if provider else sse_json_lines(resp) + + async for json_obj in lines: + chunk = ( + provider.parse_chat_chunk(json_obj) if provider else parse_chat_chunk(json_obj) + ) if not chunk: continue diff --git a/cecli/helpers/llms/identifiers.py b/cecli/helpers/llms/identifiers.py index cd6f655ee03..097fe5a60e5 100644 --- a/cecli/helpers/llms/identifiers.py +++ b/cecli/helpers/llms/identifiers.py @@ -54,6 +54,13 @@ def is_openrouter(provider: Optional[str], route: str, record: Optional[Dict[str return provider == "openrouter" or record_provider == "openrouter" +def is_ollama(provider: Optional[str], route: str, record: Optional[Dict[str, Any]]) -> bool: + """True when the model is served through an Ollama provider slug.""" + provider = (provider or "").lower() + record_provider = ((record or {}).get("litellm_provider") or "").lower() + return provider in ("ollama", "ollama_chat") or record_provider in ("ollama", "ollama_chat") + + def is_claude_5_plus(provider: Optional[str], route: str, record: Optional[Dict[str, Any]]) -> bool: """True for Claude 5+ models, which use adaptive thinking + output_config. @@ -88,6 +95,7 @@ def gpt_version(route: str) -> float: "is_github_copilot", "is_meta", "is_openrouter", + "is_ollama", "is_claude_5_plus", "gpt_version", ] diff --git a/cecli/helpers/llms/litellm_compat.py b/cecli/helpers/llms/litellm_compat.py index 979af930115..4086f141ab0 100644 --- a/cecli/helpers/llms/litellm_compat.py +++ b/cecli/helpers/llms/litellm_compat.py @@ -29,6 +29,8 @@ from cecli.dump import dump # noqa: F401 from cecli.http import httpx +from .constants import CONTROL_KWARGS +from .identifiers import is_ollama from .runtime import log_error_response warnings.filterwarnings("ignore", category=UserWarning, module="pydantic") @@ -400,6 +402,16 @@ def __init__( for k, v in kwargs.items(): setattr(self, k, v) + def __str__(self) -> str: + s = super().__str__() + if s: + return s + cause = getattr(self, "__cause__", None) + if cause is not None: + cause_str = str(cause) or repr(cause) + return f"{self.__class__.__name__}: {cause_str}" + return self.__class__.__name__ + class APIConnectionError(_FacadeException): pass @@ -507,7 +519,7 @@ def _translate_http_error(err: httpx.HTTPStatusError) -> _FacadeException: except Exception: text = "" body = text.lower() - message = text or str(err) + message = text or str(err) or repr(err) if status == 400 and any( token in body for token in ("context", "context_length", "maximum context") @@ -801,16 +813,32 @@ async def acompletion(self, **kwargs: Any) -> Any: if headers: extra_headers = {**headers, **extra_headers} + # Ollama's native runner options (num_ctx, keep_alive, top_p, ...) live + # outside the OpenAI parameter set, so forward every request-level kwarg + # the caller supplied except those this shim or the pipeline consume + # internally. Every other provider keeps the narrow whitelist below so + # unrelated runtime kwargs are not leaked into its request body. + provider, _, route = (model or "").partition("/") passthrough: Dict[str, Any] = {} - for key in ( - "temperature", - "tool_choice", - "extra_body", - "prompt_cache_key", - "stream_options", - ): - if kwargs.get(key) is not None: - passthrough[key] = kwargs[key] + + if is_ollama(provider, route, None): + for key, value in kwargs.items(): + if key in CONTROL_KWARGS or value is None: + continue + + passthrough[key] = value + + passthrough.pop("max_completion_tokens", None) + else: + for key in ( + "temperature", + "tool_choice", + "extra_body", + "prompt_cache_key", + "stream_options", + ): + if kwargs.get(key) is not None: + passthrough[key] = kwargs[key] # The model-config pipeline formatters (helpers.format_reasoning / # helpers.format_thinking) lift reasoning_effort/thinking OUT of @@ -856,11 +884,11 @@ async def acompletion(self, **kwargs: Any) -> Any: **passthrough, ) except httpx.TimeoutException as err: - raise Timeout(str(err)) from err + raise Timeout(str(err) or repr(err)) from err except httpx.HTTPStatusError as err: raise _translate_http_error(err) from err except httpx.HTTPError as err: - raise APIConnectionError(str(err)) from err + raise APIConnectionError(str(err) or repr(err)) from err return _response_shim(resp, model) @@ -874,11 +902,11 @@ async def _stream_with_errors(self, gen: Any, model: Optional[str]) -> Any: async for chunk in gen: yield _chunk_shim(chunk, model) except httpx.TimeoutException as err: - raise Timeout(str(err)) from err + raise Timeout(str(err) or repr(err)) from err except httpx.HTTPStatusError as err: raise _translate_http_error(err) from err except httpx.HTTPError as err: - raise APIConnectionError(str(err)) from err + raise APIConnectionError(str(err) or repr(err)) from err def stream_chunk_builder( self, chunks: List[Any], messages: Optional[Any] = None, **kwargs: Any diff --git a/cecli/helpers/llms/pipeline.py b/cecli/helpers/llms/pipeline.py index 17ddaf57fc2..5cebf225e5f 100644 --- a/cecli/helpers/llms/pipeline.py +++ b/cecli/helpers/llms/pipeline.py @@ -56,6 +56,12 @@ async def acompletion( # payload (e.g. Bedrock Mantle's SigV4 path). resolved["_signer"] = getattr(provider, "sign_request", None) + # The chat domain reads the provider back off ``resolved`` so it can + # delegate endpoint/payload/parsing to providers that override the shared + # chat wire (e.g. Ollama's native /api/chat) without threading a new + # argument through every family entry point. + resolved["_provider"] = provider + headers = dict(resolved.get("extra_headers") or {}) headers.update(extra_headers or {}) diff --git a/cecli/helpers/llms/providers/base.py b/cecli/helpers/llms/providers/base.py index cc874fdc9c2..ae13c308ff5 100644 --- a/cecli/helpers/llms/providers/base.py +++ b/cecli/helpers/llms/providers/base.py @@ -24,6 +24,9 @@ class ProviderAdapter: normalize fields a stricter provider rejects (e.g. Mistral rejects ``reasoning_content`` / ``provider_specific_fields`` / ``function_call`` and a null tool-call ``index``). + - :meth:`chat_url` / :meth:`chat_payload` / :meth:`chat_stream_json` / + :meth:`parse_chat_response` / :meth:`parse_chat_chunk` - override the + shared chat-family wire (e.g. Ollama's native ``/api/chat``). - :meth:`normalize` - post-process a family-normalized response (e.g. meta encrypted-reasoning marker). """ @@ -36,6 +39,10 @@ class ProviderAdapter: #: (Mistral) set this False so the chat payload's coercer skips them. echoes_reasoning_content: bool = True + #: Whether OPENAI_API_BASE/OPENAI_API_KEY may redirect this provider's chat + #: request (see domains/chat.py). Native wires (Ollama) opt out. + honors_openai_env_override: bool = True + def resolve_api_base(self, resolved: Dict[str, Any]) -> str: """Return the api_base for a resolved config (default: as resolved).""" return resolved["api_base"] @@ -73,6 +80,41 @@ def transform_messages(self, messages: List[Dict[str, Any]]) -> List[Dict[str, A """ return messages + def chat_url(self, resolved: Dict[str, Any], base: str) -> str: + """Return the chat endpoint URL (default: OpenAI ``/chat/completions``).""" + return f"{base}/chat/completions" + + def chat_payload( + self, + resolved: Dict[str, Any], + messages: List[Dict[str, Any]], + tools: Optional[List[Dict[str, Any]]], + stream: bool, + kwargs: Dict[str, Any], + ) -> Dict[str, Any]: + """Build the chat request body (default: OpenAI-compatible payload).""" + from ..domains.chat import chat_payload + + return chat_payload(resolved, messages, tools, stream, kwargs) + + def chat_stream_json(self, resp: Any) -> Any: + """Yield parsed JSON objects from a chat stream (default: SSE lines).""" + from ..utils import sse_json_lines + + return sse_json_lines(resp) + + def parse_chat_response(self, data: Dict[str, Any], resolved: Dict[str, Any]) -> Any: + """Normalize a chat response (default: OpenAI-compatible parser).""" + from ..domains.chat import normalize_chat_response + + return normalize_chat_response(data, resolved["model"]) + + def parse_chat_chunk(self, data: Dict[str, Any]) -> Any: + """Normalize one streamed chat chunk (default: OpenAI-compatible parser).""" + from ..domains.chat import parse_chat_chunk + + return parse_chat_chunk(data) + def normalize( self, family: str, diff --git a/cecli/helpers/llms/providers/ollama.py b/cecli/helpers/llms/providers/ollama.py new file mode 100644 index 00000000000..2fcf6c25ce8 --- /dev/null +++ b/cecli/helpers/llms/providers/ollama.py @@ -0,0 +1,471 @@ +"""Ollama provider adapter for the llms package. + +Ollama's OpenAI-compatible ``/v1/chat/completions`` endpoint ignores runner +options such as ``num_ctx``; Ollama's own docs direct users to a Modelfile or the +native wire to change the context window. This adapter keeps the shared ``chat`` +family but overrides its chat hooks (endpoint, payload, stream parsing) to speak +Ollama's native ``/api/chat`` wire, where api-block parameters such as ``num_ctx`` +and ``keep_alive`` live under ``options``. +""" + +from __future__ import annotations + +import json +from typing import Any, AsyncIterator, Dict, List, Optional + +from ..constants import CONTROL_KWARGS +from ..types import ( + Choice, + CompletionChunk, + CompletionResponse, + Part, + PartsMessage, + ReasoningPart, + TextPart, + ToolCall, + ToolCallPart, + Usage, + parts_message_to_message, +) +from ..utils import extract_reasoning, split_data_url +from .base import ProviderAdapter + +#: Native-body keys the payload builder consumes itself, so they must not also +#: be forwarded as runner options (on top of the shared shim/pipeline kwargs). +_NATIVE_CONTROL_KEYS = frozenset( + { + "stream_options", + "prompt_cache_key", + "tool_choice", + "extra_body", + "reasoning_effort", + "thinking", + "parallel_tool_calls", + } +) + +#: Request keys that must never be forwarded as runner options. +_CONTROL_KEYS = CONTROL_KWARGS | _NATIVE_CONTROL_KEYS + + +#: Native request-body keys forwarded at the top level (not under ``options``). +_TOP_LEVEL_KEYS = frozenset({"keep_alive", "format", "think"}) + +#: api-block parameter names mapped onto their native ``options`` equivalent. +_OPTION_ALIASES = {"max_tokens": "num_predict", "max_completion_tokens": "num_predict"} + +#: Reasoning hints that explicitly turn thinking off. +_THINK_DISABLED = frozenset({"none", "disabled", "off", "false"}) + +#: cecli effort levels mapped onto Ollama's native ``think`` levels. +_THINK_LEVELS = {"low": "low", "medium": "medium", "high": "high", "max": "high"} + +#: Ollama's default native origin when no base URL resolves. +_DEFAULT_BASE = "http://localhost:11434" + + +class _OllamaWire: + """Shared native ``/api/chat`` chat hooks for the Ollama provider slugs.""" + + #: num_ctx and friends only work on the native wire; never let + #: OPENAI_API_BASE hijack an Ollama request. + honors_openai_env_override: bool = False + + def __init__(self) -> None: + self._tool_call_offset = 0 + + def chat_url(self, resolved: Dict[str, Any], base: str) -> str: + return f"{ollama_native_base(base)}/api/chat" + + def chat_payload( + self, + resolved: Dict[str, Any], + messages: List[Dict[str, Any]], + tools: Optional[List[Dict[str, Any]]], + stream: bool, + kwargs: Dict[str, Any], + ) -> Dict[str, Any]: + return ollama_payload(resolved, messages, tools, stream, kwargs) + + def chat_stream_json(self, resp: Any) -> Any: + self._tool_call_offset = 0 + + return _ndjson_lines(resp) + + def parse_chat_response(self, data: Dict[str, Any], resolved: Dict[str, Any]) -> Any: + return normalize_ollama_response(data, resolved["model"]) + + def parse_chat_chunk(self, data: Dict[str, Any]) -> Any: + """Normalize one chunk, numbering tool calls across the whole stream. + + Ollama streams each tool call whole (no id) with a per-chunk index that + restarts at 0, so a running stream offset is applied to keep parallel + calls distinct in the aggregators (which key by id, else index). + """ + chunk = parse_ollama_chunk(data, self._tool_call_offset) + + if chunk and chunk.tool_calls: + self._tool_call_offset += len(chunk.tool_calls) + + return chunk + + +class OllamaProvider(_OllamaWire, ProviderAdapter): + """Ollama ``ollama/`` slug: native /api/chat wire.""" + + provider: str = "ollama" + + +class OllamaChatProvider(_OllamaWire, ProviderAdapter): + """Ollama ``ollama_chat/`` slug: native /api/chat wire.""" + + provider: str = "ollama_chat" + + +def ollama_native_base(api_base: Optional[str]) -> str: + """Strip an OpenAI-compat ``/v1`` (or ``/api``) suffix to reach Ollama.""" + base = (api_base or _DEFAULT_BASE).rstrip("/") + + for suffix in ("/v1", "/api"): + if base.endswith(suffix): + base = base[: -len(suffix)] + + return base or _DEFAULT_BASE + + +def ollama_payload( + resolved: Dict[str, Any], + messages: List[Dict[str, Any]], + tools: Optional[List[Dict[str, Any]]], + stream: bool, + kwargs: Dict[str, Any], +) -> Dict[str, Any]: + """Build the native ``/api/chat`` payload from api-block/extra params.""" + body = _extra_body(resolved, kwargs) + payload: Dict[str, Any] = { + "model": resolved["route"], + "messages": _native_messages(messages), + "stream": stream, + } + + if tools: + payload["tools"] = tools + + options = _options(kwargs, body) + + if options: + payload["options"] = options + + keep_alive = kwargs.get("keep_alive") + + if keep_alive is None: + keep_alive = body.get("keep_alive") + + if keep_alive is not None: + payload["keep_alive"] = keep_alive + + think = _think_value(kwargs, body) + + if think is not None: + payload["think"] = think + + format_value = kwargs.get("format") + + if format_value is None: + format_value = body.get("format") + + if format_value is not None: + payload["format"] = format_value + + return payload + + +def normalize_ollama_response(data: Dict[str, Any], model: str) -> CompletionResponse: + """Convert a native ``/api/chat`` response into a normalized response.""" + message = data.get("message") or {} + parts: List[Part] = [] + content = message.get("content") + + if isinstance(content, str) and content: + parts.append(TextPart(text=content)) + + reasoning = message.get("thinking") or extract_reasoning(message) + + if reasoning: + parts.append(ReasoningPart(text=reasoning)) + + for tc in message.get("tool_calls") or []: + function = tc.get("function") or {} + parts.append( + ToolCallPart( + name=function.get("name", ""), + arguments=_parse_arguments(function.get("arguments")), + tool_call_id=tc.get("id"), + ) + ) + + pm = PartsMessage(role=message.get("role", "assistant"), parts=parts) + finish_reason = _map_finish_reason(data.get("done_reason") or "stop") + + if any(isinstance(part, ToolCallPart) for part in parts): + finish_reason = "tool_calls" + + return CompletionResponse( + id=data.get("id"), + model=model, + choices=[ + Choice( + index=0, + message=parts_message_to_message(pm), + finish_reason=finish_reason, + ) + ], + usage=_usage_from_native(data), + ) + + +def parse_ollama_chunk(data: Dict[str, Any], start_index: int = 0) -> Optional[CompletionChunk]: + """Convert one native NDJSON chunk into a normalized stream chunk. + + ``start_index`` offsets the tool-call index/id so a caller consuming a whole + stream keeps parallel calls distinct: Ollama streams each call whole with a + per-chunk index that restarts at 0, so the streaming hook passes a running + offset (see :meth:`_OllamaWire.parse_chat_chunk`). + """ + message = data.get("message") or {} + text = message.get("content") or "" + reasoning = message.get("thinking") or "" + tool_calls: List[ToolCall] = [] + + for offset, tc in enumerate(message.get("tool_calls") or []): + index = start_index + offset + function = tc.get("function") or {} + tool_calls.append( + ToolCall( + id=tc.get("id") or f"call_{index}", + name=function.get("name", ""), + arguments=_parse_arguments(function.get("arguments")), + index=index, + ) + ) + + done = bool(data.get("done")) + finish_reason = None + usage = None + + if done: + finish_reason = ( + "tool_calls" if tool_calls else _map_finish_reason(data.get("done_reason") or "stop") + ) + usage = _usage_from_native(data) + + if not text and not reasoning and not tool_calls and not done: + return None + + return CompletionChunk( + text=text, + reasoning=reasoning, + tool_calls=tool_calls, + finish_reason=finish_reason, + usage=usage, + ) + + +def _extra_body(resolved: Dict[str, Any], kwargs: Dict[str, Any]) -> Dict[str, Any]: + body = dict(resolved.get("extra_body") or {}) + body.update(kwargs.get("extra_body") or {}) + return body + + +def _options(kwargs: Dict[str, Any], body: Dict[str, Any]) -> Dict[str, Any]: + options: Dict[str, Any] = {} + + for source in (kwargs, body): + for key, value in source.items(): + if value is None or key in _CONTROL_KEYS or key in _TOP_LEVEL_KEYS: + continue + + options[_OPTION_ALIASES.get(key, key)] = value + + return options + + +def _think_value(kwargs: Dict[str, Any], body: Dict[str, Any]) -> Optional[Any]: + for source in (kwargs, body): + thinking = source.get("thinking") + + if thinking is not None: + return _normalize_think(thinking) + + for source in (kwargs, body): + effort = source.get("reasoning_effort") + + if effort: + return _normalize_think(effort) + + return None + + +def _normalize_think(value: Any) -> Optional[Any]: + """Map a generic thinking/effort hint onto Ollama's ``think`` field. + + Ollama accepts a boolean or a ``low``/``medium``/``high`` level: explicit + disables become ``False``, recognized levels are forwarded (``max`` -> + ``high``), and an unrecognized value returns ``None`` so the field is omitted + and Ollama's default applies rather than the level silently disabling + thinking. + """ + if isinstance(value, bool): + return value + + if isinstance(value, dict): + return value.get("type") != "disabled" + + if not isinstance(value, str): + return None + + level = value.strip().lower() + + if level in _THINK_DISABLED: + return False + + return _THINK_LEVELS.get(level) + + +def _native_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + out: List[Dict[str, Any]] = [] + + for msg in messages: + role = msg.get("role") + content, images = _native_content(msg.get("content")) + native: Dict[str, Any] = {"role": role} + + if content is not None: + native["content"] = content + + if images: + native["images"] = images + + reasoning = msg.get("reasoning_content") or extract_reasoning(msg) + + if role == "assistant" and reasoning: + native["thinking"] = reasoning + + tool_calls = msg.get("tool_calls") + + if tool_calls: + native["tool_calls"] = [_native_tool_call(tc) for tc in tool_calls] + + tool_call_id = msg.get("tool_call_id") + + if tool_call_id is not None: + native["tool_call_id"] = tool_call_id + + out.append(native) + + return out + + +def _native_content(content: Any) -> Any: + if isinstance(content, str): + return content, [] + + if not isinstance(content, list): + return (None if content is None else str(content)), [] + + texts: List[str] = [] + images: List[str] = [] + + for part in content: + if not isinstance(part, dict): + continue + + part_type = part.get("type") + + if part_type == "text" and isinstance(part.get("text"), str): + texts.append(part["text"]) + elif part_type == "image_url": + url = part.get("image_url") + + if isinstance(url, dict): + url = url.get("url") + + split = split_data_url(url) + + if split: + images.append(split[1]) + + text = "\n".join(item for item in texts if item) + return (text or None), images + + +def _native_tool_call(tc: Dict[str, Any]) -> Dict[str, Any]: + function = tc.get("function") or {} + return { + "function": { + "name": function.get("name", ""), + "arguments": _parse_arguments(function.get("arguments")), + } + } + + +def _parse_arguments(raw: Any) -> Dict[str, Any]: + if isinstance(raw, dict): + return raw + + if not isinstance(raw, str) or not raw.strip(): + return {} + + try: + parsed = json.loads(raw) + except json.JSONDecodeError: + return {"_raw": raw} + + return parsed if isinstance(parsed, dict) else {"_value": parsed} + + +def _usage_from_native(data: Dict[str, Any]) -> Optional[Usage]: + prompt_tokens = data.get("prompt_eval_count") + completion_tokens = data.get("eval_count") + + if prompt_tokens is None and completion_tokens is None: + return None + + return Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=(prompt_tokens or 0) + (completion_tokens or 0), + ) + + +def _map_finish_reason(reason: str) -> str: + reason = (reason or "").lower() + + if reason in ("stop", "length", "tool_calls"): + return reason + + return "stop" + + +async def _ndjson_lines(resp: Any) -> AsyncIterator[Dict[str, Any]]: + """Yield parsed JSON objects from Ollama's newline-delimited stream.""" + async for raw in resp.aiter_lines(): + line = raw.strip() + + if not line: + continue + + try: + yield json.loads(line) + except json.JSONDecodeError: + continue + + +__all__ = [ + "OllamaProvider", + "OllamaChatProvider", + "ollama_payload", + "ollama_native_base", + "normalize_ollama_response", + "parse_ollama_chunk", +] diff --git a/cecli/helpers/responses.py b/cecli/helpers/responses.py index ba1155a7ad9..2ac7a9d431e 100644 --- a/cecli/helpers/responses.py +++ b/cecli/helpers/responses.py @@ -247,6 +247,56 @@ def extract_tools_from_pseudo_json(content: str) -> Optional[List[ChatCompletion return None +_tool_name_to_sanitized: dict[str, str] = {} +_sanitized_to_tool_name: dict[str, str] = {} + + +def sanitize_tool_name(name: str) -> str: + """Make a name acceptable in OpenAI-style ``tools[].function.name``. + + Providers enforce ``^[A-Za-z0-9_-]{1,64}$`` on the function name, which + rejects names an MCP server may legitimately advertise (``browser.fetch``). + Idempotent, so it is also safe on an incoming tool call before the name is + compared against, or mapped back to, the server's own name. + """ + cached = _tool_name_to_sanitized.get(name) + + if cached is not None: + return cached + + sanitized = re.sub(r"[^A-Za-z0-9_-]", "_", name)[:64] + _tool_name_to_sanitized[name] = sanitized + + return sanitized + + +def register_tool_names(server_tools) -> None: + """Cache the original/sanitized MCP tool name pairs. + + ``server_tools`` is an iterable of ``(server_name, tools)`` pairs, i.e. the + shape of ``coder.mcp_tools``. Registering when the tool list is built means + mapping a call back to the server's advertised name is a dict lookup rather + than a scan over every tool on every call. + """ + for _server_name, tools in server_tools or []: + for tool in tools: + original = nested.getter(tool, "function.name", "") + + if not original: + continue + + _sanitized_to_tool_name[sanitize_tool_name(original)] = original + + +def original_tool_name(name: str) -> str: + """Map a sanitized tool name back to the name the MCP server advertises. + + Returns ``name`` unchanged when nothing matches, so an unsanitized name + still reaches the server as-is. + """ + return _sanitized_to_tool_name.get(name, name) + + def prefix_tool_name(server_name: str, tool_name: str) -> str: """ Prefix a tool name with the server name. @@ -256,9 +306,9 @@ def prefix_tool_name(server_name: str, tool_name: str) -> str: tool_name: Original tool name Returns: - Prefixed tool name in format "{server_name}--{tool_name}" + Prefixed, provider-safe tool name in format "{server_name}--{tool_name}" """ - return f"{server_name}--{tool_name}" + return sanitize_tool_name(f"{server_name}--{tool_name}") def unprefix_tool_name(prefixed_name: str) -> tuple[str, str]: diff --git a/cecli/helpers/sessions/manager.py b/cecli/helpers/sessions/manager.py index 70ddb7a0628..c5e29fe310d 100644 --- a/cecli/helpers/sessions/manager.py +++ b/cecli/helpers/sessions/manager.py @@ -248,19 +248,30 @@ async def _reload_sub_agents(self, data_file: Path) -> None: 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.""" + """Spawn and restore a single sub-agent from its saved payload. + + Only sub-agents that were still relevant when the session was saved are + rebuilt: independent agents are always restored, while a dependent agent + is restored only if it had not finished or errored. The transient + ``memorizer`` helper is never restored. + """ sub_data = self._read_session_payload(sub_file, quiet=True) if not isinstance(sub_data, dict): return + if not subagents.should_restore_sub_agent(sub_data): + return + name = subagents.resolve_reload_agent_name( sub_data.get("agent_name") or "worker", sub_data.get("agent_root") ) if not name: return + independent = bool(sub_data.get("independent", True)) + new_coder, _info = await service.spawn( - name, parent=self.coder, auto_reap=False, independent=True + name, parent=self.coder, auto_reap=False, independent=independent ) sub_manager = SessionManager(new_coder, self.io) diff --git a/cecli/helpers/sessions/payload.py b/cecli/helpers/sessions/payload.py index ce64fdbffd9..7d08d9e9df1 100644 --- a/cecli/helpers/sessions/payload.py +++ b/cecli/helpers/sessions/payload.py @@ -14,8 +14,19 @@ def build_payload(coder, io, session_name: str, agent_name: Optional[str] = None ``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. + a later load can rebuild them at the right path. Sub-agent payloads also + record their independence and lifecycle status so a reload can decide + whether they should be rebuilt. """ + # Sub-agent lifecycle flags come from the AgentService; the primary agent + # has none and stores ``None`` for both. + if agent_name: + from .subagents import sub_agent_state + + independent, status = sub_agent_state(coder) + else: + independent, status = None, None + 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 = [ @@ -25,11 +36,13 @@ def build_payload(coder, io, session_name: str, agent_name: Optional[str] = None # Flush any queued messages so the saved chat history is complete ConversationService.get_manager(coder).flush_queue() - return { + payload = { "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, + "independent": independent, + "status": status, "model": coder.main_model.name, "weak_model": coder.main_model.weak_model.name, "editor_model": coder.main_model.editor_model.name, @@ -66,6 +79,8 @@ def build_payload(coder, io, session_name: str, agent_name: Optional[str] = None }, } + return payload + async def apply_payload( coder, io, session_data: Dict, session_file, sub_agent: bool = False diff --git a/cecli/helpers/sessions/subagents.py b/cecli/helpers/sessions/subagents.py index f77daa6daf7..1824c384762 100644 --- a/cecli/helpers/sessions/subagents.py +++ b/cecli/helpers/sessions/subagents.py @@ -12,6 +12,10 @@ logger = logging.getLogger(__name__) +# The memorizer is a transient helper that is re-spawned on demand; a saved +# session must never try to resurrect it. +MEMORIZER_AGENT_NAME = "memorizer" + def detect_agent_name(coder) -> Optional[str]: """Return a coder's sub-agent type, or ``None`` for the primary agent.""" @@ -64,6 +68,33 @@ def live_sub_agents(coder) -> List[Tuple[str, object]]: return agents +def sub_agent_state(coder) -> Tuple[bool, Optional[str]]: + """Return ``(independent, status)`` for a tracked sub-agent coder. + + Falls back to ``(False, None)`` when the coder is not (or no longer) tracked + by the :class:`AgentService`, so payload building stays safe for callers with + no live service. + """ + from cecli.helpers.agents.service import AgentService + + coder_uuid = getattr(coder, "uuid", None) + if not isinstance(coder_uuid, str) or not coder_uuid: + return False, None + + try: + service = AgentService.get_instance(coder) + except Exception: + return False, None + + info = service.sub_agents.get(coder_uuid) + if info is None: + return False, None + + status = getattr(getattr(info, "status", None), "value", None) + + return bool(getattr(info, "independent", False)), status + + 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``. @@ -83,6 +114,27 @@ def save_sub_agents(coder, io, session_name: str, subs_dir: Path) -> int: return written +def should_restore_sub_agent(sub_data: Dict) -> bool: + """Return whether a saved sub-agent payload should be rebuilt on load. + + Independent agents are always restored regardless of how they finished. A + dependent agent is restored only while it is still unfinished and unerrored. + The transient ``memorizer`` helper is never restored. + """ + from cecli.helpers.agents.service import SubAgentStatus + + agent_name = sub_data.get("agent_name") + if not agent_name or agent_name == MEMORIZER_AGENT_NAME: + return False + + if sub_data.get("independent", True): + return True + + status = sub_data.get("status") + + return status not in (SubAgentStatus.FINISHED.value, SubAgentStatus.ERROR.value) + + 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 diff --git a/cecli/helpers/system_one/__init__.py b/cecli/helpers/system_one/__init__.py new file mode 100644 index 00000000000..93d5be027b6 --- /dev/null +++ b/cecli/helpers/system_one/__init__.py @@ -0,0 +1,86 @@ +"""System One decision endpoints, speaking the ``/v1/systemone`` wire format. + +Three layers live here, from the wire up: + +- :mod:`cecli.helpers.system_one.types` - the question/answer vocabulary + (``noul`` / ``choice`` / ``score``) and the parsing of response bodies. +- :mod:`cecli.helpers.system_one.api` - one call per decision (:func:`decide`, + :func:`judge`, :func:`rate`) plus :func:`system_one` / :func:`ask` for + several questions over one state. +- :mod:`cecli.helpers.system_one.integrations` - the shorthand forms other cecli + subsystems (hooks today) use, so their call sites stay one-liners. + +Endpoint settings come from ``--system-one`` (or the ``system-one`` config-file +key), installed with :func:`configure`, or from ``SYSTEM_ONE_*`` defaults. +""" + +from cecli.helpers.system_one.api import ( + DECIDE_QUESTION_ID, + JUDGE_QUESTION_ID, + RATE_QUESTION_ID, + ask, + decide, + judge, + rate, + system_one, +) +from cecli.helpers.system_one.client import SystemOneClient, SystemOneError +from cecli.helpers.system_one.config import ( + DEFAULT_API_BASE, + ENDPOINT_PATH, + SystemOneConfig, + configure, + get_config, + is_configured, + reset, +) +from cecli.helpers.system_one.types import ( + CHOICE, + NOUL, + QUESTION_TYPES, + SCORE, + ChoiceAnswer, + NoulAnswer, + ScoreAnswer, + SystemOneResponse, + build_payload, + choice, + noul, + parse_answer, + parse_response, + score, +) + +__all__ = [ + "CHOICE", + "DEFAULT_API_BASE", + "DECIDE_QUESTION_ID", + "JUDGE_QUESTION_ID", + "RATE_QUESTION_ID", + "ENDPOINT_PATH", + "NOUL", + "QUESTION_TYPES", + "SCORE", + "ChoiceAnswer", + "NoulAnswer", + "ScoreAnswer", + "SystemOneClient", + "SystemOneConfig", + "SystemOneError", + "SystemOneResponse", + "ask", + "build_payload", + "choice", + "configure", + "decide", + "get_config", + "is_configured", + "judge", + "noul", + "parse_answer", + "parse_response", + "rate", + "reset", + "score", + "system_one", +] diff --git a/cecli/helpers/system_one/api.py b/cecli/helpers/system_one/api.py new file mode 100644 index 00000000000..1efafe1e61c --- /dev/null +++ b/cecli/helpers/system_one/api.py @@ -0,0 +1,186 @@ +"""Application-level API over the System One endpoint. + +Mirrors the shape of the ``von`` SDK so a decision reads like a single call +rather than a hand-built question map:: + + from cecli.helpers import system_one + + r = await system_one.decide( + state="Payment gateway timeouts on charge authorizations.", + choices={"billing": "Payments and refunds", "infra": "Bugs and outages"}, + instructions="Which team owns this?", + ) + r.choice, r.confidence, r.probabilities + +``decide`` wraps a Choice question, ``judge`` a Noul (yes/no) question and +``rate`` a Score question; :func:`system_one` (alias :func:`ask`) covers the +multi-question case where several questions share one pass over the state. +""" + +from __future__ import annotations + +from typing import Any, Dict, Iterable, List, Mapping, Optional, Union + +from cecli.helpers.system_one.client import SystemOneClient +from cecli.helpers.system_one.types import ( + ChoiceAnswer, + NoulAnswer, + ScoreAnswer, + SystemOneResponse, + choice, + noul, + score, +) + +#: Question ids used when a single-question helper builds its own question map. +DECIDE_QUESTION_ID = "decision" +JUDGE_QUESTION_ID = "judgment" +RATE_QUESTION_ID = "rating" + + +async def ask( + state: Any, + questions: Mapping[str, Dict[str, Any]], + model: Optional[str] = None, + client: Optional[SystemOneClient] = None, +) -> SystemOneResponse: + """Evaluate ``state`` against a map of typed questions (one request).""" + return await (client or SystemOneClient()).evaluate(state, questions, model=model) + + +async def system_one( + state: Any, + questions: Mapping[str, Dict[str, Any]], + model: Optional[str] = None, + client: Optional[SystemOneClient] = None, +) -> SystemOneResponse: + """Evaluate ``state`` against several questions in one request. + + Name-for-name parity with the ``von`` SDK's ``system_one()``; :func:`ask` + is the same call under a name that reads better when the module itself is + already imported as ``system_one``. + """ + return await ask(state, questions, model=model, client=client) + + +async def decide( + state: Any, + choices: Union[Mapping[str, Any], Iterable[str]], + instructions: str = "Which option best describes the state?", + model: Optional[str] = None, + client: Optional[SystemOneClient] = None, +) -> ChoiceAnswer: + """Pick one option from ``choices`` for ``state``. + + ``choices`` maps an option to its rubric description, or is a bare + iterable of option names when no rubric text is needed. + """ + response = await ask( + state, + {DECIDE_QUESTION_ID: choice(instructions, choices)}, + model=model, + client=client, + ) + + return _unwrap_choice(response) + + +async def judge( + state: Any, + instructions: str = "Is the statement in the instructions true of the state?", + criteria: Optional[Mapping[str, Any]] = None, + model: Optional[str] = None, + client: Optional[SystemOneClient] = None, +) -> NoulAnswer: + """Answer a yes/no question about ``state`` as a probability of yes. + + ``criteria`` describes what yes and no mean (``{"true": ..., "false": + ...}``); omitting it is the weakest path, so prefer passing it or phrasing + the decision as a described Choice. + """ + if not isinstance(criteria, Mapping): + criteria = _true_false_criteria(criteria) or {} + + response = await ask( + state, + { + JUDGE_QUESTION_ID: noul( + instructions, true=criteria.get("true"), false=criteria.get("false") + ) + }, + model=model, + client=client, + ) + + return _unwrap_noul(response) + + +async def rate( + state: Any, + levels: List[Any], + instructions: str = "Rate the state against the rubric levels.", + model: Optional[str] = None, + client: Optional[SystemOneClient] = None, +) -> ScoreAnswer: + """Rate ``state`` along an ordered rubric of ``levels``.""" + response = await ask( + state, + {RATE_QUESTION_ID: score(instructions, levels)}, + model=model, + client=client, + ) + + return _unwrap_score(response) + + +def _true_false_criteria(criteria: Any) -> Optional[Dict[str, Any]]: + """Accept ``criteria`` given as a mapping or a (true, false) pair.""" + if criteria is None: + return None + + if isinstance(criteria, Mapping): + return { + "true": criteria.get("true"), + "false": criteria.get("false"), + } + + if isinstance(criteria, (list, tuple)) and len(criteria) == 2: + return {"true": criteria[0], "false": criteria[1]} + + return None + + +def _unwrap_choice(response: SystemOneResponse) -> ChoiceAnswer: + answer = _single(response) + + if not isinstance(answer, ChoiceAnswer): + raise TypeError(f"Expected a choice answer, got {type(answer).__name__}") + + return answer + + +def _unwrap_noul(response: SystemOneResponse) -> NoulAnswer: + answer = _single(response) + + if not isinstance(answer, NoulAnswer): + raise TypeError(f"Expected a noul answer, got {type(answer).__name__}") + + return answer + + +def _unwrap_score(response: SystemOneResponse) -> ScoreAnswer: + answer = _single(response) + + if not isinstance(answer, ScoreAnswer): + raise TypeError(f"Expected a score answer, got {type(answer).__name__}") + + return answer + + +def _single(response: SystemOneResponse) -> Any: + answers = list(response.answers.values()) + + if not answers: + raise ValueError("System One response carried no answers") + + return answers[0] diff --git a/cecli/helpers/system_one/client.py b/cecli/helpers/system_one/client.py new file mode 100644 index 00000000000..61162fe6552 --- /dev/null +++ b/cecli/helpers/system_one/client.py @@ -0,0 +1,162 @@ +"""HTTP transport for the System One evaluation endpoint. + +Sends ``POST {api_base}/v1/systemone`` requests, the wire format shared by +System One decision servers: any endpoint implementing it works, so pointing +:attr:`SystemOneConfig.api_base` at a locally served model (``von serve``) or a +gateway is just a config change. + +Retries follow the documented guidance for ``429``/``529`` responses: +exponential backoff, bounded by the config's ``max_retries``. +""" + +from __future__ import annotations + +import asyncio +import logging +import random +from typing import Any, Dict, Mapping, Optional + +from cecli.helpers.system_one.config import SystemOneConfig, get_config +from cecli.helpers.system_one.types import ( + SystemOneResponse, + build_payload, + parse_response, +) + +logger = logging.getLogger(__name__) + +#: HTTP statuses worth retrying (rate limited / overloaded). +RETRY_STATUS = frozenset({429, 529}) + +#: Statuses that mean the request was rejected for authentication. +AUTH_STATUS = frozenset({401, 403}) + +#: Base seconds for the exponential backoff sequence. +BACKOFF_BASE_SECONDS = 0.5 + + +class SystemOneError(RuntimeError): + """Raised when an evaluation request fails after retries.""" + + def __init__(self, message: str, status_code: Optional[int] = None) -> None: + super().__init__(message) + self.status_code = status_code + + +class SystemOneClient: + """Async client for one System One endpoint.""" + + def __init__(self, config: Optional[SystemOneConfig] = None) -> None: + self.config = config or get_config() + + def questions_for( + self, + state: Any, + questions: Mapping[str, Dict[str, Any]], + model: Optional[str] = None, + ) -> Dict[str, Any]: + """Return the request payload that :meth:`evaluate` would send.""" + return build_payload( + state, + questions, + model=model, + default_model=self.config.model_name, + ) + + async def evaluate( + self, + state: Any, + questions: Mapping[str, Dict[str, Any]], + model: Optional[str] = None, + ) -> SystemOneResponse: + """Evaluate ``state`` against typed ``questions`` and parse the answer map.""" + payload = self.questions_for(state, questions, model=model) + body = await self._post(payload) + + return parse_response(body) + + async def raw_evaluate(self, payload: Mapping[str, Any]) -> Dict[str, Any]: + """Send an already-built payload and return the decoded JSON body.""" + return await self._post(dict(payload)) + + async def _post(self, payload: Dict[str, Any]) -> Dict[str, Any]: + from cecli.http import httpx + + headers = self.config.build_headers() + url = self.config.endpoint_url + last_error: Optional[str] = None + last_status: Optional[int] = None + + for attempt in range(self.config.max_retries + 1): + if attempt: + await asyncio.sleep(_backoff_delay(attempt - 1)) + + try: + async with httpx.AsyncClient(timeout=self.config.timeout) as client: + response = await client.post(url, json=payload, headers=headers) + except Exception as exc: + last_error = f"{type(exc).__name__}: {exc}" + last_status = None + logger.warning("System One request to %s failed: %s", url, last_error) + continue + + if response.status_code < 400: + return _decode_body(response) + + last_status = response.status_code + last_error = _error_text(response, self.config) + + if response.status_code not in RETRY_STATUS or attempt >= self.config.max_retries: + break + + raise SystemOneError( + f"System One request to {url} failed ({last_status or 'transport error'}): " + f"{last_error}", + status_code=last_status, + ) + + +def _decode_body(response: Any) -> Dict[str, Any]: + try: + body = response.json() + except Exception as exc: + raise SystemOneError(f"System One response was not valid JSON: {exc}") from exc + + if not isinstance(body, Mapping): + raise SystemOneError( + f"System One response must be a JSON object, got {type(body).__name__}" + ) + + return dict(body) + + +def _error_text(response: Any, config: Optional[SystemOneConfig] = None) -> str: + """Describe a failed response, naming unset key variables on auth errors. + + An endpoint that needs a bearer token answers 401/403 whether the variable + in ``api_key_env`` is missing or simply empty, which otherwise reads as an + opaque server-side rejection. + """ + try: + text = response.text[:500] + except Exception: + text = "" + + auth_error = getattr(response, "status_code", 0) in (401, 403) + + if config is not None and auth_error and not config.resolve_api_key(): + names = ", ".join(config.api_key_env) or "(none configured)" + text += ( + f"\nNo API key was sent: none of {names} is set. Export one of them" + " or point api_key_env at the variable that holds your key." + ) + + return text + try: + return response.text[:500] + except Exception: + return "" + + +def _backoff_delay(retry_index: int) -> float: + return BACKOFF_BASE_SECONDS * (2**retry_index) * random.uniform(0.8, 1.2) diff --git a/cecli/helpers/system_one/config.py b/cecli/helpers/system_one/config.py new file mode 100644 index 00000000000..3e7eb206ed5 --- /dev/null +++ b/cecli/helpers/system_one/config.py @@ -0,0 +1,217 @@ +"""Endpoint configuration for System One decision calls. + +An endpoint is described with the same keys cecli uses for custom model +providers (``api_base`` / ``api_key_env`` / ``extra_headers``), so +``--system-one`` and the ``system-one`` config-file key read like a slimmed +down ``model-providers`` entry:: + + system-one: + api_base: "http://127.0.0.1:8000" + api_key_env: ["LITELLM_API_KEY"] + model_name: "von-latest" + extra_headers: + x-trace-id: "cecli" + +With no configuration the defaults are local and vendor neutral: only +``SYSTEM_ONE_*`` environment variables are probed, so pointing the helper at a +hosted service or naming a different API-key variable is always an explicit +configuration choice rather than an implicit fallback. +""" + +from __future__ import annotations + +import os +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional + +#: Default endpoint: a locally served ``/v1/systemone`` server. +DEFAULT_API_BASE = "http://127.0.0.1:8000" + +#: Path appended to ``api_base`` to reach the evaluation endpoint. +ENDPOINT_PATH = "/v1/systemone" + +#: Default model name sent in every request body. +DEFAULT_MODEL_NAME = "von-latest" + +#: Environment variable probed for the endpoint base URL. +BASE_URL_ENV_VARS = ("SYSTEM_ONE_API_BASE",) + +#: Environment variables probed for the API key, in priority order. +API_KEY_ENV_VARS = ("SYSTEM_ONE_API_KEY",) + +#: Environment variable probed for the model name. +MODEL_ENV_VARS = ("SYSTEM_ONE_MODEL",) + +#: Request timeout in seconds. +DEFAULT_TIMEOUT = 30.0 + +#: Accepted spellings for each config key (hyphen/underscore/camel variants). +_KEY_ALIASES: Dict[str, tuple[str, ...]] = { + "api_base": ("api_base", "api-base", "base_url", "base-url"), + "api_key_env": ("api_key_env", "api-key-env", "api_key_envs", "apiKeyEnv"), + "model_name": ("model_name", "model-name", "model", "model_id", "model-id"), + "extra_headers": ("extra_headers", "extra-headers", "headers", "default_headers"), + "timeout": ("timeout", "request_timeout", "request-timeout"), + "max_retries": ("max_retries", "max-retries", "retries"), +} + + +def _pick(raw: Dict[str, Any], canonical: str) -> Any: + """Return the first alias present in *raw* for a canonical key.""" + for alias in _KEY_ALIASES.get(canonical, (canonical,)): + if alias in raw: + return raw[alias] + + return None + + +def _first_env(names) -> Optional[str]: + """Return the first non-empty environment value among *names*.""" + if isinstance(names, str): + names = [names] + + for name in names or (): + value = os.environ.get(name) + + if value: + return value + + return None + + +@dataclass +class SystemOneConfig: + """Resolved configuration for one System One endpoint.""" + + api_base: str = DEFAULT_API_BASE + api_key_env: List[str] = field(default_factory=lambda: list(API_KEY_ENV_VARS)) + model_name: str = DEFAULT_MODEL_NAME + extra_headers: Dict[str, str] = field(default_factory=dict) + timeout: float = DEFAULT_TIMEOUT + max_retries: int = 2 + + @classmethod + def from_dict(cls, raw: Optional[Dict[str, Any]]) -> "SystemOneConfig": + """Build a config from a user supplied dict, filling gaps from env.""" + defaults = cls.from_env() + raw = raw if isinstance(raw, dict) else {} + + api_base = _pick(raw, "api_base") or defaults.api_base + + key_env = _pick(raw, "api_key_env") + + if isinstance(key_env, str): + key_env = [key_env] + + if not key_env: + key_env = defaults.api_key_env + + model_name = _pick(raw, "model_name") or defaults.model_name + + extra_headers = _pick(raw, "extra_headers") + + if not isinstance(extra_headers, dict): + extra_headers = {} + + timeout = _pick(raw, "timeout") + timeout = defaults.timeout if timeout is None else float(timeout) + + max_retries = _pick(raw, "max_retries") + max_retries = defaults.max_retries if max_retries is None else int(max_retries) + + return cls( + api_base=str(api_base).rstrip("/"), + api_key_env=[str(name) for name in key_env], + model_name=str(model_name), + extra_headers={str(k): str(v) for k, v in extra_headers.items()}, + timeout=timeout, + max_retries=max_retries, + ) + + @classmethod + def from_env(cls) -> "SystemOneConfig": + """Build a config from ``SYSTEM_ONE_*`` variables and module defaults.""" + api_base = _first_env(BASE_URL_ENV_VARS) + model_name = _first_env(MODEL_ENV_VARS) + + return cls( + api_base=(api_base or DEFAULT_API_BASE).rstrip("/"), + model_name=model_name or DEFAULT_MODEL_NAME, + ) + + @property + def endpoint_url(self) -> str: + """Full URL of the evaluation endpoint. + + ``api_base`` may be given with or without the ``/v1`` prefix; the + trailing path segment is never duplicated. + """ + base = self.api_base.rstrip("/") + + if base.endswith(ENDPOINT_PATH): + return base + + if base.endswith("/v1"): + return f"{base}/systemone" + + return f"{base}{ENDPOINT_PATH}" + + def resolve_api_key(self) -> Optional[str]: + """Return the API key from the configured environment variables.""" + return _first_env(self.api_key_env) + + def build_headers(self) -> Dict[str, str]: + """Return request headers (content type and auth, then user extras).""" + headers = {"Content-Type": "application/json"} + + api_key = self.resolve_api_key() + + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + + headers.update(self.extra_headers) + + return headers + + +#: Config installed via ``--system-one`` / config file / :func:`configure`. +_active_config: Optional[SystemOneConfig] = None + + +def configure(raw: Any) -> Optional[SystemOneConfig]: + """Install *raw* as the active endpoint config. + + Accepts a dict (from ``--system-one`` or a config file) or an existing + :class:`SystemOneConfig`. Returns ``None`` when *raw* carries nothing + usable so callers can tell a no-op apart from a real configuration. + """ + global _active_config + + if isinstance(raw, SystemOneConfig): + _active_config = raw + elif isinstance(raw, dict): + if not raw: + return None + + _active_config = SystemOneConfig.from_dict(raw) + else: + return None + + return _active_config + + +def reset() -> None: + """Drop the active config so lookups fall back to the environment.""" + global _active_config + + _active_config = None + + +def is_configured() -> bool: + """Return whether an endpoint config was explicitly installed.""" + return _active_config is not None + + +def get_config() -> SystemOneConfig: + """Return the active config, falling back to environment defaults.""" + return _active_config if _active_config is not None else SystemOneConfig.from_env() diff --git a/cecli/helpers/system_one/integrations.py b/cecli/helpers/system_one/integrations.py new file mode 100644 index 00000000000..81a72f15bbf --- /dev/null +++ b/cecli/helpers/system_one/integrations.py @@ -0,0 +1,155 @@ +"""Shorthand decision calls for cecli subsystems. + +The application api in :mod:`cecli.helpers.system_one.api` takes an explicit +wire-format question map. Subsystems inside cecli - hooks today, likely tools +and commands next - would rather ask in a smaller vocabulary: one of +``questions``, ``decide``, ``judge`` or ``rate`` over a state. This module owns +that translation so call sites stay one-liners instead of duplicating the +normalization rules. + +Every entry point returns plain dicts (``response.to_dict()``), which is what +hook authors log, store and compare. +""" + +from __future__ import annotations + +from typing import Any, Dict, List, Mapping, Optional, Union + +from cecli.helpers.system_one.api import ( + DECIDE_QUESTION_ID, + JUDGE_QUESTION_ID, + RATE_QUESTION_ID, + ask, +) +from cecli.helpers.system_one.types import choice, noul, score + +#: Instructions used when a shorthand omits an explicit question. +DEFAULT_DECIDE_INSTRUCTIONS = "Which option best describes the state?" +DEFAULT_RATE_INSTRUCTIONS = "Rate the state against the rubric levels." + + +async def evaluate( + state: Any, + questions: Optional[Mapping[str, Any]] = None, + decide: Optional[Union[str, Mapping[str, Any]]] = None, + judge: Optional[Union[str, Mapping[str, Any]]] = None, + rate: Optional[Union[List[Any], Mapping[str, Any]]] = None, + model: Optional[str] = None, +) -> Dict[str, Any]: + """Ask exactly one of the four forms about ``state``. + + Returns ``{"model": str, "usage": {...}, "answers": {id: {...}}}``; the + single-question shorthands land under ``decision``, ``judgment`` and + ``rating`` respectively. + + Raises: + ValueError: If other than exactly one form is given. + TypeError: If a ``questions`` entry is neither a dict nor a string. + """ + payload = build_questions(questions, decide, judge, rate) + response = await ask(state, payload, model=model) + + return response.to_dict() + + +def build_questions( + questions: Optional[Mapping[str, Any]] = None, + decide: Optional[Union[str, Mapping[str, Any]]] = None, + judge: Optional[Union[str, Mapping[str, Any]]] = None, + rate: Optional[Union[List[Any], Mapping[str, Any]]] = None, +) -> Dict[str, Dict[str, Any]]: + """Normalize the shorthand forms into one wire-format question map.""" + forms = ( + ("questions", questions), + ("decide", decide), + ("judge", judge), + ("rate", rate), + ) + provided = [name for name, value in forms if value is not None] + + if len(provided) != 1: + raise ValueError( + f"evaluate() takes exactly one of {'/'.join(name for name, _ in forms)}," + f" got {len(provided)}" + ) + + if questions is not None: + return { + str(question_id): _coerce_question(question_id, spec) + for question_id, spec in questions.items() + } + + if decide is not None: + options, instructions = _decide_args(decide) + + return {DECIDE_QUESTION_ID: choice(instructions, options)} + + if judge is not None: + instructions, criteria = _judge_args(judge) + + return {JUDGE_QUESTION_ID: noul(instructions, **criteria)} + + levels, instructions = _rate_args(rate) + + return {RATE_QUESTION_ID: score(instructions, levels)} + + +def _coerce_question(question_id: Any, spec: Any) -> Dict[str, Any]: + """Expand a shorthand question value into a wire-format question dict.""" + if isinstance(spec, str): + return noul(spec) + + if isinstance(spec, Mapping): + return dict(spec) + + raise TypeError( + f"Question {question_id!r} must be a dict or instructions string," + f" got {type(spec).__name__}" + ) + + +def _decide_args(decide: Union[str, Mapping[str, Any]]) -> tuple: + """Normalize a ``decide=`` shorthand to ``(options, instructions)``.""" + if isinstance(decide, str): + return [decide], DEFAULT_DECIDE_INSTRUCTIONS + + spec = dict(decide) + + return ( + spec.get("choices", spec.get("criteria", {})), + spec.get("instructions", DEFAULT_DECIDE_INSTRUCTIONS), + ) + + +def _judge_args(judge: Union[str, Mapping[str, Any]]) -> tuple: + """Normalize a ``judge=`` shorthand to ``(instructions, criteria)``.""" + if isinstance(judge, str): + return judge, {} + + spec = dict(judge) + criteria = spec.get("criteria") or {} + + return spec.get("instructions", ""), { + key: value for key, value in criteria.items() if key in ("true", "false") + } + + +def _rate_args(rate: Union[List[Any], Mapping[str, Any]]) -> tuple: + """Normalize a ``rate=`` shorthand to ``(levels, instructions)``.""" + if isinstance(rate, Mapping): + spec = dict(rate) + + return ( + spec.get("levels", spec.get("criteria", [])), + spec.get("instructions", DEFAULT_RATE_INSTRUCTIONS), + ) + + return list(rate), DEFAULT_RATE_INSTRUCTIONS + + +__all__ = [ + "DEFAULT_DECIDE_INSTRUCTIONS", + "DEFAULT_RATE_INSTRUCTIONS", + "build_questions", + "evaluate", +] diff --git a/cecli/helpers/system_one/types.py b/cecli/helpers/system_one/types.py new file mode 100644 index 00000000000..aea552f898f --- /dev/null +++ b/cecli/helpers/system_one/types.py @@ -0,0 +1,256 @@ +"""Wire-format types for the System One evaluation endpoint. + +The endpoint speaks a small, non-autoregressive protocol: a ``state`` plus a +map of typed ``questions`` in, a map of typed ``answers`` out. This module +owns the vocabulary of that protocol - question builders, the request payload, +and the typed answer objects returned after parsing. + +Question types are ``noul`` (a yes/no probability), ``choice`` (pick one +option from a described set) and ``score`` (rate against an ordered rubric). +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Dict, Iterable, List, Mapping, Optional, Union + +from cecli.helpers.system_one.config import DEFAULT_MODEL_NAME + +#: Question types accepted by the endpoint. +NOUL = "noul" +CHOICE = "choice" +SCORE = "score" + +QUESTION_TYPES = (NOUL, CHOICE, SCORE) + +#: instructions/criteria accept plain strings or structured objects/arrays. +StructuredValue = Union[str, Dict[str, Any], List[Any]] + + +def noul( + instructions: StructuredValue, + true: Optional[StructuredValue] = None, + false: Optional[StructuredValue] = None, +) -> Dict[str, Any]: + """Build a yes/no question returning the probability the answer is yes.""" + question: Dict[str, Any] = {"type": NOUL, "instructions": instructions} + + if true is not None or false is not None: + criteria: Dict[str, Any] = {} + + if true is not None: + criteria["true"] = true + + if false is not None: + criteria["false"] = false + + question["criteria"] = criteria + + return question + + +def choice( + instructions: StructuredValue, + criteria: Union[Mapping[str, Any], Iterable[str], None] = None, +) -> Dict[str, Any]: + """Build a question that picks one option from ``criteria``. + + ``criteria`` maps each option to a rubric description; a bare iterable of + option names is accepted and expanded to ``{option: None}``. + """ + if criteria is None: + options: Dict[str, Any] = {} + elif isinstance(criteria, Mapping): + options = dict(criteria) + else: + options = {str(name): None for name in criteria} + + return {"type": CHOICE, "instructions": instructions, "criteria": options} + + +def score( + instructions: StructuredValue, + criteria: Optional[List[StructuredValue]] = None, +) -> Dict[str, Any]: + """Build a question that rates the state against ordered level rubrics.""" + if criteria is None: + criteria = [] + + return {"type": SCORE, "instructions": instructions, "criteria": list(criteria)} + + +@dataclass +class NoulAnswer: + """A yes/no answer expressed as a probability of yes (0.0 - 1.0).""" + + probability: float + raw: Dict[str, Any] = field(default_factory=dict) + + @property + def yes(self) -> bool: + """Whether the answer favours "yes" at or above even odds.""" + return self.probability >= 0.5 + + def to_dict(self) -> Dict[str, Any]: + return {"type": NOUL, "noul": self.probability, "yes": self.yes} + + +@dataclass +class ChoiceAnswer: + """The selected option plus the full probability distribution.""" + + choice: str + probabilities: Dict[str, float] = field(default_factory=dict) + confidence: float = 0.0 + raw: Dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> Dict[str, Any]: + return { + "type": CHOICE, + "choice": self.choice, + "probabilities": dict(self.probabilities), + "confidence": self.confidence, + } + + +@dataclass +class ScoreAnswer: + """A probability-weighted rating that can land between rubric levels.""" + + score: float + legend: Dict[str, str] = field(default_factory=dict) + probabilities: Dict[str, float] = field(default_factory=dict) + confidence: float = 0.0 + raw: Dict[str, Any] = field(default_factory=dict) + + @property + def levels(self) -> List[str]: + """Rubric descriptions ordered by level index.""" + return [self.legend[key] for key in sorted(self.legend, key=_level_sort_key)] + + def to_dict(self) -> Dict[str, Any]: + return { + "type": SCORE, + "score": self.score, + "legend": dict(self.legend), + "probabilities": dict(self.probabilities), + "confidence": self.confidence, + } + + +#: Any parsed answer. +Answer = Union[NoulAnswer, ChoiceAnswer, ScoreAnswer] + + +@dataclass +class SystemOneResponse: + """A full evaluation: one answer per question, keyed as provided.""" + + answers: Dict[str, Answer] = field(default_factory=dict) + model: str = "" + usage: Dict[str, int] = field(default_factory=dict) + raw: Dict[str, Any] = field(default_factory=dict) + + def __getitem__(self, question_id: str) -> Answer: + return self.answers[question_id] + + def __contains__(self, question_id: object) -> bool: + return question_id in self.answers + + def get(self, question_id: str, default: Any = None) -> Any: + return self.answers.get(question_id, default) + + def to_dict(self) -> Dict[str, Any]: + """Return the response as plain nested dicts (JSON-safe).""" + return { + "model": self.model, + "answers": {key: answer.to_dict() for key, answer in self.answers.items()}, + "usage": dict(self.usage), + } + + +def build_payload( + state: Any, + questions: Mapping[str, Dict[str, Any]], + model: Optional[str] = None, + default_model: Optional[str] = None, +) -> Dict[str, Any]: + """Assemble the request body sent to the evaluation endpoint.""" + return { + "state": state, + "model": model or default_model or DEFAULT_MODEL_NAME, + "questions": {str(key): dict(value) for key, value in questions.items()}, + } + + +def parse_response(payload: Mapping[str, Any]) -> SystemOneResponse: + """Convert a decoded JSON response body into a :class:`SystemOneResponse`.""" + if not isinstance(payload, Mapping): + raise ValueError(f"System One response must be an object, got {type(payload).__name__}") + + answers: Dict[str, Answer] = {} + raw_answers = payload.get("answers") or {} + + for question_id, raw_answer in raw_answers.items(): + answers[str(question_id)] = parse_answer(raw_answer) + + usage_raw = payload.get("usage") or {} + usage = {str(key): int(value or 0) for key, value in usage_raw.items()} + + return SystemOneResponse( + answers=answers, + model=str(payload.get("model") or ""), + usage=usage, + raw=dict(payload), + ) + + +def parse_answer(raw: Any) -> Answer: + """Parse a single answer object by its ``type`` discriminator.""" + if not isinstance(raw, Mapping): + raise ValueError(f"System One answer must be an object, got {type(raw).__name__}") + + question_type = str(raw.get("type") or "") + + if question_type == NOUL: + return NoulAnswer(probability=_as_float(raw.get("noul")), raw=dict(raw)) + + if question_type == CHOICE: + return ChoiceAnswer( + choice=str(raw.get("choice") or ""), + probabilities=_as_float_map(raw.get("probabilities")), + confidence=_as_float(raw.get("confidence")), + raw=dict(raw), + ) + + if question_type == SCORE: + return ScoreAnswer( + score=_as_float(raw.get("score")), + legend={str(k): str(v) for k, v in (raw.get("legend") or {}).items()}, + probabilities=_as_float_map(raw.get("probabilities")), + confidence=_as_float(raw.get("confidence")), + raw=dict(raw), + ) + + raise ValueError(f"Unknown System One answer type: {question_type!r}") + + +def _as_float(value: Any) -> float: + try: + return float(value) + except (TypeError, ValueError): + return 0.0 + + +def _as_float_map(value: Any) -> Dict[str, float]: + if not isinstance(value, Mapping): + return {} + + return {str(key): _as_float(item) for key, item in value.items()} + + +def _level_sort_key(key: str) -> Any: + try: + return (0, int(key)) + except (TypeError, ValueError): + return (1, key) diff --git a/cecli/hooks/helpers.py b/cecli/hooks/helpers.py index a4bacfc20e8..e572d5c3b87 100644 --- a/cecli/hooks/helpers.py +++ b/cecli/hooks/helpers.py @@ -20,7 +20,7 @@ async def execute(self, coder, metadata): return True """ -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Union class HookHelpers: @@ -191,3 +191,58 @@ async def call_subagent( agent_service = AgentService.get_instance(coder) return await agent_service.invoke(name, prompt, **kwargs) + + @staticmethod + async def system_one( + coder: Any, + state: Any = None, + questions: Optional[Dict[str, Any]] = None, + decide: Optional[Union[str, Dict[str, Any]]] = None, + judge: Optional[Union[str, Dict[str, Any]]] = None, + rate: Optional[Union[List[Any], Dict[str, Any]]] = None, + model: Optional[str] = None, + last_n: int = 20, + ) -> Dict[str, Any]: + """Ask a System One decision model about the conversation or a state. + + System One endpoints answer typed questions (``noul`` yes/no, + ``choice`` and ``score``) in a single non-generative pass, so this + returns structured verdicts rather than prose. Provide exactly one of + ``questions``, ``decide``, ``judge`` or ``rate``. + + Args: + coder: The coder instance (passed to hook's ``execute()``). + state: Content to evaluate (str, object or array). When ``None``, + the last ``last_n`` conversation messages are used. + questions: Map of question id to a question dict (see + ``cecli.helpers.system_one.noul`` / ``choice`` / ``score``). + A bare string is treated as a yes/no question. + decide: Option map (or ``{"choices": ..., "instructions": ...}``) + for a single Choice question. + judge: Instructions (or ``{"instructions": ..., "criteria": ...}``) + for a single Noul question. + rate: ``{"levels": [...], "instructions": ...}``, or a bare list of + levels, for one Score question. + model: Endpoint model name override. + last_n: How many recent messages to build the default state from. + + Returns: + ``{"model": str, "usage": {...}, "answers": {id: {...}}}`` where + each answer is a plain dict of the parsed decision. Single-question + shorthands come back under the ids ``decision``, ``judgment`` and + ``rating``. + + Raises: + ValueError: If other than exactly one of the four forms is given. + Exception: Endpoint errors (``SystemOneError`` and transport + failures) propagate; guard the call when the decision is only + advisory or the endpoint may be down. + """ + from cecli.helpers.system_one import integrations + + if state is None: + state = HookHelpers.get_messages(coder, last_n=last_n) + + return await integrations.evaluate( + state, questions=questions, decide=decide, judge=judge, rate=rate, model=model + ) diff --git a/cecli/main.py b/cecli/main.py index b59db284c83..836904dcf21 100644 --- a/cecli/main.py +++ b/cecli/main.py @@ -170,6 +170,7 @@ def convert_yaml_to_json_string(value, config_file_value=None): "hooks": "hooks", "workspaces": "workspaces", "model_providers": "model-providers", + "system_one": "system-one", "server_config": "server-config", } @@ -674,11 +675,18 @@ async def main_async( # before our deep-merge code can run. all_config_paths = [ str(Path.home() / ".cecli" / "conf.yml"), + str(Path.home() / ".cecli" / ".cecli.conf.yml"), str(Path.home() / ".cecli.conf.yml"), str(Path(".cecli.conf.yml")), ] + if os.environ.get("CECLI_CONFIG_FILE"): + cfg_env = os.environ.get("CECLI_CONFIG_FILE") + if cfg_env not in all_config_paths: + all_config_paths.append(cfg_env) if git_root: - all_config_paths.append(str(Path(git_root) / ".cecli.conf.yml")) + git_conf = str(Path(git_root) / ".cecli.conf.yml") + if git_conf not in all_config_paths: + all_config_paths.append(git_conf) conf_yml_files = [ p for p in all_config_paths if p.endswith("conf.yml") and not p.endswith(".cecli.conf.yml") @@ -1047,6 +1055,17 @@ def get_io(pretty): except json.JSONDecodeError as e: io.tool_error(f"Failed to parse --model-providers JSON: {e}") + if args.system_one: + from cecli.helpers import system_one + + try: + system_one_config = json.loads(args.system_one) + except json.JSONDecodeError as e: + io.tool_error(f"Failed to parse --system-one JSON: {e}") + else: + if system_one.configure(system_one_config) and args.verbose: + io.tool_output(f"System One endpoint: {system_one.get_config().endpoint_url}") + if args.list_models: models.print_matching_models(io, args.list_models) return await graceful_exit(None) diff --git a/cecli/models.py b/cecli/models.py index 827f433e35f..2360c386505 100644 --- a/cecli/models.py +++ b/cecli/models.py @@ -133,6 +133,8 @@ class ModelSettings: retries: Optional[dict] = None retry_backoff_factor: float = 1.5 retry_on_unavailable: bool = True + retry_on_forbidden: bool = False + retry_on_unauthorized: bool = False retry_timeout: float = 30 request_timeout: int = request_timeout debug: bool = False @@ -566,7 +568,10 @@ def __init__( self._apply_reasoning_defaults() self.get_weak_model(weak_model) self.get_agent_model(agent_model) - self.retries = retries + # Keep a `retries:` block from model settings unless an explicit value + # was passed (main.py passes retries=None when --retries is unset). + if retries is not None: + self.retries = retries self.debug = debug if editor_model is False: @@ -1452,21 +1457,20 @@ async def send_completion( litellm_ex = LiteLLMExceptions() retry_delay = 0.125 - if self.retries: - retry_config = dict() - try: - retry_config = json.loads(self.retries) - except (json.JSONDecodeError, TypeError, ValueError): - retry_config = dict() - pass + retry_config = parse_retry_config(self.retries) + self.retry_on_unavailable = retry_config["retry_on_unavailable"] + self.retry_on_forbidden = retry_config["retry_on_forbidden"] + self.retry_on_unauthorized = retry_config["retry_on_unauthorized"] + self.retry_backoff_factor = retry_config["retry_backoff_factor"] + self.retry_timeout = retry_config["retry_timeout"] - self.retry_on_unavailable = bool( - nested.getter(retry_config, "retry-on-unavailable", True) - ) - self.retry_backoff_factor = float( - nested.getter(retry_config, "retry-backoff-factor", 1.5) - ) - self.retry_timeout = float(nested.getter(retry_config, "retry-timeout", 30)) + if override_kwargs: + kwargs = deep_merge(kwargs, override_kwargs) + + kwargs = deep_merge(kwargs, {"allowed_openai_params": ["tools", "tool_choice"]}) + + if self.debug: + kwargs["logger_fn"] = self._log_request while True: try: @@ -1476,14 +1480,6 @@ async def send_completion( if random.random() < 0.25: await asyncio.sleep(random.uniform(min_wait, max_wait)) - if override_kwargs: - kwargs = deep_merge(kwargs, override_kwargs) - - kwargs = deep_merge(kwargs, {"allowed_openai_params": ["tools", "tool_choice"]}) - - if self.debug: - kwargs["logger_fn"] = self._log_request - completion_coro = litellm.acompletion(**kwargs) res, interrupted = await coroutines.interruptible(completion_coro, interrupt_event) if interrupted: @@ -1493,10 +1489,7 @@ async def send_completion( except litellm.ContextWindowExceededError as err: raise err except litellm_ex.exceptions_tuple() as err: - ex_info = litellm_ex.get_ex_info(err) - should_retry = ex_info.retry - if ex_info.name == "ServiceUnavailableError": - should_retry = should_retry or self.retry_on_unavailable + should_retry, ex_info = parse_model_error(retry_config, err) custom_retry_delay = self._extract_retry_delay(err) if custom_retry_delay is not None: @@ -1517,7 +1510,7 @@ async def send_completion( should_retry = False if not should_retry: - print(f"API Error: {str(err)}") + print(f"API Error: {str(err) or repr(err)}") if ex_info.description: print(ex_info.description) if stream: @@ -1526,7 +1519,7 @@ async def send_completion( return hash_object, self.model_error_response() print(f"Retrying in {retry_delay:.1f} seconds...") - print(f"API Error: {str(err)}") + print(f"API Error: {str(err) or repr(err)}") if interrupt_event: _res, interrupted = await coroutines.interruptible( asyncio.sleep(retry_delay), interrupt_event @@ -1557,6 +1550,10 @@ async def simple_send_with_retries( temperature = None tools = None + retry_config = parse_retry_config(self.retries) + retry_backoff_factor = retry_config["retry_backoff_factor"] + retry_timeout = retry_config["retry_timeout"] + if self.verbose: dump(messages) @@ -1602,19 +1599,19 @@ async def simple_send_with_retries( return remove_reasoning_content(res, self.reasoning_tag), response except litellm_ex.exceptions_tuple() as err: - ex_info = litellm_ex.get_ex_info(err) + should_retry, ex_info = parse_model_error(retry_config, err) print(str(err)) if ex_info.description: print(ex_info.description) - should_retry = ex_info.retry + custom_retry_delay = self._extract_retry_delay(err) if custom_retry_delay is not None: retry_delay = custom_retry_delay should_retry = True elif should_retry: - retry_delay *= 2 + retry_delay *= retry_backoff_factor - if retry_delay > RETRY_TIMEOUT: + if retry_delay > retry_timeout: should_retry = False if not should_retry: @@ -1792,6 +1789,82 @@ def _configured_provider(self) -> str: return provider +def parse_retry_config(retries_input): + """ + Parse and normalize retry configuration from a JSON string or dict. + Returns a unified dict with defaults: + retry_timeout: 30 + retry_backoff_factor: 1.5 + retry_on_unavailable: True + retry_on_forbidden: False + retry_on_empty: False + retry_on_unauthorized: False + """ + config = dict() + if isinstance(retries_input, str): + try: + config = json.loads(retries_input) + except (json.JSONDecodeError, TypeError, ValueError): + config = dict() + elif isinstance(retries_input, dict): + config = retries_input.copy() + + # Helper to get either hyphenated or underscored key + def _get(key, default): + val = config.get(key) + if val is not None: + return val + val = config.get(key.replace("_", "-")) + if val is not None: + return val + return default + + return { + "retry_timeout": float(_get("retry_timeout", 30)), + "retry_backoff_factor": float(_get("retry_backoff_factor", 1.5)), + "retry_on_unavailable": bool(_get("retry_on_unavailable", True)), + "retry_on_forbidden": bool(_get("retry_on_forbidden", False)), + "retry_on_unauthorized": bool(_get("retry_on_unauthorized", False)), + "retry_on_empty": bool(_get("retry_on_empty", False)), + } + + +def parse_model_error(retry_config, error): + """ + Determine whether a model error should be retried and return its metadata. + + ``retry_config`` is the dict produced by :func:`parse_retry_config`; ``error`` + is a litellm exception. Returns a ``(should_retry, ex_info)`` tuple so callers + can reuse the parsed exception info for messaging. + + Some providers (e.g. vertex_ai_beta) carry the status on ``code`` or omit + ``status_code`` entirely, so auth failures are matched on the exception name + first, then on the status/code fields. + """ + from cecli.exceptions import LiteLLMExceptions + + ex_info = LiteLLMExceptions().get_ex_info(error) + + should_retry = ex_info.retry + if ex_info.name == "ServiceUnavailableError": + should_retry = should_retry or retry_config["retry_on_unavailable"] + + if ex_info.name == "PermissionDeniedError": + should_retry = should_retry or retry_config["retry_on_forbidden"] + + status_code = getattr(error, "status_code", None) + code = getattr(error, "code", None) + is_auth_error = ( + ex_info.name in ("AuthenticationError", "PermissionDeniedError") + or status_code in (401, 403, "401", "403") + or str(code) in ("401", "403") + ) + if is_auth_error: + should_retry = should_retry or retry_config["retry_on_unauthorized"] + + return should_retry, ex_info + + def register_models(model_settings_fnames): files_loaded = [] for model_settings_fname in model_settings_fnames: diff --git a/cecli/resources/providers.json b/cecli/resources/providers.json index 7f7515c6756..e037a6bf452 100644 --- a/cecli/resources/providers.json +++ b/cecli/resources/providers.json @@ -322,6 +322,16 @@ ], "display_name": "ollama" }, + "ollama_chat": { + "api_base": "http://localhost:11434/v1", + "api_key_env": [ + "OLLAMA_API_KEY" + ], + "base_url_env": [ + "OLLAMA_API_BASE" + ], + "display_name": "ollama_chat" + }, "opencode-go": { "api_base": "https://opencode.ai/zen/go/v1", "api_key_env": [ diff --git a/cecli/run_cmd.py b/cecli/run_cmd.py index a257cc643cb..cc69d9487c7 100644 --- a/cecli/run_cmd.py +++ b/cecli/run_cmd.py @@ -77,6 +77,10 @@ def run_cmd_subprocess( errors="replace", bufsize=1, # Set bufsize to 0 for unbuffered output universal_newlines=True, + # Never inherit the TUI's stdin; a command that reads it would + # race Textual for keystrokes. Commands needing input go through + # the interactive path instead. + stdin=subprocess.DEVNULL, cwd=cwd, ) diff --git a/cecli/tests/test_queue_commands.py b/cecli/tests/test_queue_commands.py index e0fc306d898..2449c55e35c 100644 --- a/cecli/tests/test_queue_commands.py +++ b/cecli/tests/test_queue_commands.py @@ -2,20 +2,20 @@ Test suite for CLI-33 Queue Commands. This module contains comprehensive tests for: -- Unit tests: Queue logic in Commands class (core.py) +- Unit tests: Queue logic in the command_queue helper - Integration tests: QueueCommand, ListQueueCommand, RemoveQueueCommand - E2E tests: Full queue lifecycle and processing - Regression tests: Existing command integrity Test categories: -- UTC-01 through UTC-20: Unit tests for queue methods +- UTC-01 through UTC-20: Unit tests for queue helpers - ITC-01 through ITC-20: Integration tests for commands - ETC-01 through ETC-10: E2E tests for full lifecycle - RTC-01 through RTC-05: Regression tests for existing functionality - TDS-01 through TDS-04: Test data setup and fixtures """ -import asyncio +import threading import time from unittest.mock import MagicMock @@ -27,6 +27,7 @@ from cecli.commands.queue import QueueCommand from cecli.commands.remove_queue import RemoveQueueCommand from cecli.commands.utils.registry import CommandRegistry +from cecli.helpers import command_queue from cecli.signals import ReloadProgramSignal, SwitchCoderSignal @@ -40,6 +41,19 @@ def _make_coder(): coder.uuid = str(uuid.uuid4()) coder.prompt_queue = [] coder._queue_counter = 0 + coder._queue_lock = None + return coder + + +def _make_coder_with_commands(items=()): + """Build a coder with an attached Commands instance and an optional + pre-populated queue.""" + coder = _make_coder() + coder.commands = Commands(io=None, coder=coder) + coder.io = None + coder.tui = None + for text in items: + command_queue.enqueue_prompt(coder, text) return coder @@ -79,37 +93,35 @@ def mock_io(): @pytest.fixture def mock_coder(): - """Create a mock coder with commands attribute pointing to a Commands instance.""" - coder = _make_coder() - commands = Commands(io=None, coder=coder) - coder.commands = commands - coder.io = None - coder.tui = None - return coder + """Create a mock coder with an empty queue and an attached Commands instance.""" + return _make_coder_with_commands() + + +@pytest.fixture +def clean_coder(): + """Create a fresh coder with an empty queue for isolated testing.""" + return _make_coder() @pytest.fixture def clean_commands(): - """Create a fresh Commands instance with empty queue for isolated testing.""" + """Create a fresh Commands instance for isolated testing.""" return Commands(io=None, coder=_make_coder()) @pytest.fixture -def populated_queue(clean_commands): - """Create Commands with pre-populated queue with known items.""" - clean_commands._enqueue_prompt("alpha") - clean_commands._enqueue_prompt("beta") - clean_commands._enqueue_prompt("gamma") - return clean_commands +def populated_queue(): + """Create a coder whose queue is pre-populated with known items.""" + return _make_coder_with_commands(("alpha", "beta", "gamma")) @pytest.fixture def full_queue(): - """Create Commands with queue filled to max capacity (100 items).""" - commands = Commands(io=None, coder=_make_coder()) + """Create a coder with a queue filled to max capacity (100 items).""" + coder = _make_coder_with_commands() for i in range(100): - commands._enqueue_prompt(f"prompt_{i}") - return commands + command_queue.enqueue_prompt(coder, f"prompt_{i}") + return coder @pytest.fixture @@ -128,187 +140,187 @@ def mock_coder_no_commands(): class TestEnqueuePrompt: - """Unit tests for _enqueue_prompt method.""" + """Unit tests for command_queue.enqueue_prompt.""" - def test_utc_01_enqueue_single_prompt(self, clean_commands): + def test_utc_01_enqueue_single_prompt(self, clean_coder): """UTC-01: Enqueue single prompt adds one item with correct structure.""" - item = clean_commands._enqueue_prompt("test prompt") + item = command_queue.enqueue_prompt(clean_coder, "test prompt") - assert len(clean_commands.prompt_queue) == 1 + assert len(clean_coder.prompt_queue) == 1 assert item["text"] == "test prompt" assert "id" in item assert "timestamp" in item assert isinstance(item["id"], str) assert isinstance(item["timestamp"], float) - def test_utc_02_enqueue_multiple_prompts_fifo_order(self, clean_commands): + def test_utc_02_enqueue_multiple_prompts_fifo_order(self, clean_coder): """UTC-02: Enqueue multiple prompts maintains FIFO order and unique IDs.""" - item1 = clean_commands._enqueue_prompt("first") - item2 = clean_commands._enqueue_prompt("second") - item3 = clean_commands._enqueue_prompt("third") - - assert len(clean_commands.prompt_queue) == 3 - assert clean_commands.prompt_queue[0]["text"] == "first" - assert clean_commands.prompt_queue[1]["text"] == "second" - assert clean_commands.prompt_queue[2]["text"] == "third" + item1 = command_queue.enqueue_prompt(clean_coder, "first") + item2 = command_queue.enqueue_prompt(clean_coder, "second") + item3 = command_queue.enqueue_prompt(clean_coder, "third") + + assert len(clean_coder.prompt_queue) == 3 + assert clean_coder.prompt_queue[0]["text"] == "first" + assert clean_coder.prompt_queue[1]["text"] == "second" + assert clean_coder.prompt_queue[2]["text"] == "third" assert item1["id"] != item2["id"] assert item2["id"] != item3["id"] - def test_utc_13_max_queue_size_rejection(self, clean_commands): + def test_utc_13_max_queue_size_rejection(self, clean_coder): """UTC-13: Enqueue rejected when queue already contains 100 items.""" for i in range(100): - clean_commands._enqueue_prompt(f"prompt_{i}") + command_queue.enqueue_prompt(clean_coder, f"prompt_{i}") with pytest.raises(RuntimeError, match="Queue is full"): - clean_commands._enqueue_prompt("overflow") + command_queue.enqueue_prompt(clean_coder, "overflow") - def test_utc_14_enqueue_empty_string_rejected(self, clean_commands): + def test_utc_14_enqueue_empty_string_rejected(self, clean_coder): """UTC-14: Enqueue rejects empty string with ValueError.""" with pytest.raises(ValueError, match="Cannot enqueue empty prompt"): - clean_commands._enqueue_prompt("") + command_queue.enqueue_prompt(clean_coder, "") - def test_utc_14_enqueue_none_rejected(self, clean_commands): + def test_utc_14_enqueue_none_rejected(self, clean_coder): """UTC-14: Enqueue rejects None with ValueError.""" with pytest.raises(ValueError, match="Cannot enqueue empty prompt"): - clean_commands._enqueue_prompt(None) + command_queue.enqueue_prompt(clean_coder, None) - def test_utc_15_enqueue_extremely_long_prompt_rejected(self, clean_commands): + def test_utc_15_enqueue_extremely_long_prompt_rejected(self, clean_coder): """UTC-15: Enqueue rejects prompt exceeding 10,000 characters.""" long_prompt = "x" * 10001 with pytest.raises(ValueError, match="exceeds maximum length"): - clean_commands._enqueue_prompt(long_prompt) + command_queue.enqueue_prompt(clean_coder, long_prompt) - def test_utc_16_enqueue_exactly_10000_chars_accepted(self, clean_commands): + def test_utc_16_enqueue_exactly_10000_chars_accepted(self, clean_coder): """UTC-16: Enqueue accepts prompt of exactly 10,000 characters (boundary).""" boundary_prompt = "x" * 10000 - item = clean_commands._enqueue_prompt(boundary_prompt) + item = command_queue.enqueue_prompt(clean_coder, boundary_prompt) assert item["text"] == boundary_prompt - assert len(clean_commands.prompt_queue) == 1 + assert len(clean_coder.prompt_queue) == 1 - def test_utc_17_enqueue_9999_chars_accepted(self, clean_commands): + def test_utc_17_enqueue_9999_chars_accepted(self, clean_coder): """UTC-17: Enqueue accepts prompt of 9,999 characters (boundary).""" boundary_prompt = "x" * 9999 - item = clean_commands._enqueue_prompt(boundary_prompt) + item = command_queue.enqueue_prompt(clean_coder, boundary_prompt) assert item["text"] == boundary_prompt - assert len(clean_commands.prompt_queue) == 1 + assert len(clean_coder.prompt_queue) == 1 - def test_utc_18_counter_persistence(self, clean_commands): + def test_utc_18_counter_persistence(self, clean_coder): """UTC-18: Internal counter increments across enqueue/remove cycles without reset.""" - clean_commands._enqueue_prompt("first") - clean_commands._enqueue_prompt("second") - item = clean_commands._enqueue_prompt("third") + command_queue.enqueue_prompt(clean_coder, "first") + command_queue.enqueue_prompt(clean_coder, "second") + item = command_queue.enqueue_prompt(clean_coder, "third") - assert clean_commands._queue_counter == 3 + assert clean_coder._queue_counter == 3 assert item["id"] == "3" class TestDequeuePrompt: - """Unit tests for _dequeue_prompt method.""" + """Unit tests for command_queue.dequeue_prompt.""" - def test_utc_03_dequeue_from_empty_queue(self, clean_commands): + def test_utc_03_dequeue_from_empty_queue(self, clean_coder): """UTC-03: Dequeue from empty queue returns None without side effects.""" - result = clean_commands._dequeue_prompt() + result = command_queue.dequeue_prompt(clean_coder) assert result is None - assert len(clean_commands.prompt_queue) == 0 + assert len(clean_coder.prompt_queue) == 0 - def test_utc_04_dequeue_returns_fifo_first_item(self, clean_commands): + def test_utc_04_dequeue_returns_fifo_first_item(self, clean_coder): """UTC-04: Dequeue returns first item and shrinks queue by one.""" - clean_commands._enqueue_prompt("first") - clean_commands._enqueue_prompt("second") + command_queue.enqueue_prompt(clean_coder, "first") + command_queue.enqueue_prompt(clean_coder, "second") - item = clean_commands._dequeue_prompt() + item = command_queue.dequeue_prompt(clean_coder) assert item["text"] == "first" - assert len(clean_commands.prompt_queue) == 1 - assert clean_commands.prompt_queue[0]["text"] == "second" + assert len(clean_coder.prompt_queue) == 1 + assert clean_coder.prompt_queue[0]["text"] == "second" - def test_utc_05_dequeue_until_empty(self, clean_commands): + def test_utc_05_dequeue_until_empty(self, clean_coder): """UTC-05: Repeated dequeue eventually returns None after queue empties.""" - clean_commands._enqueue_prompt("only_item") + command_queue.enqueue_prompt(clean_coder, "only_item") - item = clean_commands._dequeue_prompt() + item = command_queue.dequeue_prompt(clean_coder) assert item is not None assert item["text"] == "only_item" - result = clean_commands._dequeue_prompt() + result = command_queue.dequeue_prompt(clean_coder) assert result is None class TestGetQueueLength: - """Unit tests for _get_queue_length method.""" + """Unit tests for command_queue.get_queue_length.""" - def test_utc_06_queue_length_empty(self, clean_commands): + def test_utc_06_queue_length_empty(self, clean_coder): """UTC-06: Queue length returns correct count for empty queue.""" - assert clean_commands._get_queue_length() == 0 + assert command_queue.get_queue_length(clean_coder) == 0 - def test_utc_07_queue_length_non_empty(self, clean_commands): + def test_utc_07_queue_length_non_empty(self, clean_coder): """UTC-07: Queue length returns correct count for populated queue.""" - clean_commands._enqueue_prompt("item1") - clean_commands._enqueue_prompt("item2") - clean_commands._enqueue_prompt("item3") + command_queue.enqueue_prompt(clean_coder, "item1") + command_queue.enqueue_prompt(clean_coder, "item2") + command_queue.enqueue_prompt(clean_coder, "item3") - assert clean_commands._get_queue_length() == 3 + assert command_queue.get_queue_length(clean_coder) == 3 class TestRemoveFromQueue: - """Unit tests for _remove_from_queue method.""" + """Unit tests for command_queue.remove_from_queue.""" - def test_utc_08_remove_by_valid_index(self, clean_commands): + def test_utc_08_remove_by_valid_index(self, clean_coder): """UTC-08: Remove by valid index returns item and shrinks queue by one.""" - clean_commands._enqueue_prompt("first") - clean_commands._enqueue_prompt("second") - clean_commands._enqueue_prompt("third") + command_queue.enqueue_prompt(clean_coder, "first") + command_queue.enqueue_prompt(clean_coder, "second") + command_queue.enqueue_prompt(clean_coder, "third") - item = clean_commands._remove_from_queue(1) + item = command_queue.remove_from_queue(clean_coder, 1) assert item["text"] == "second" - assert len(clean_commands.prompt_queue) == 2 - assert clean_commands.prompt_queue[0]["text"] == "first" - assert clean_commands.prompt_queue[1]["text"] == "third" + assert len(clean_coder.prompt_queue) == 2 + assert clean_coder.prompt_queue[0]["text"] == "first" + assert clean_coder.prompt_queue[1]["text"] == "third" - def test_utc_09_remove_out_of_bounds_high_index(self, clean_commands): + def test_utc_09_remove_out_of_bounds_high_index(self, clean_coder): """UTC-09: Remove by out-of-bounds high index returns None without mutation.""" - clean_commands._enqueue_prompt("only_item") + command_queue.enqueue_prompt(clean_coder, "only_item") - result = clean_commands._remove_from_queue(5) + result = command_queue.remove_from_queue(clean_coder, 5) assert result is None - assert len(clean_commands.prompt_queue) == 1 + assert len(clean_coder.prompt_queue) == 1 - def test_utc_10_remove_negative_index(self, clean_commands): + def test_utc_10_remove_negative_index(self, clean_coder): """UTC-10: Remove by negative index returns None without mutation.""" - clean_commands._enqueue_prompt("only_item") + command_queue.enqueue_prompt(clean_coder, "only_item") - result = clean_commands._remove_from_queue(-1) + result = command_queue.remove_from_queue(clean_coder, -1) assert result is None - assert len(clean_commands.prompt_queue) == 1 + assert len(clean_coder.prompt_queue) == 1 class TestClearQueue: - """Unit tests for _clear_queue method.""" + """Unit tests for command_queue.clear_queue.""" - def test_utc_11_clear_queue_with_items(self, clean_commands): + def test_utc_11_clear_queue_with_items(self, clean_coder): """UTC-11: Clear queue returns all items and empties queue.""" - clean_commands._enqueue_prompt("item1") - clean_commands._enqueue_prompt("item2") - clean_commands._enqueue_prompt("item3") + command_queue.enqueue_prompt(clean_coder, "item1") + command_queue.enqueue_prompt(clean_coder, "item2") + command_queue.enqueue_prompt(clean_coder, "item3") - items = clean_commands._clear_queue() + items = command_queue.clear_queue(clean_coder) assert len(items) == 3 - assert len(clean_commands.prompt_queue) == 0 + assert len(clean_coder.prompt_queue) == 0 - def test_utc_12_clear_empty_queue(self, clean_commands): + def test_utc_12_clear_empty_queue(self, clean_coder): """UTC-12: Clear empty queue returns empty list and remains empty.""" - items = clean_commands._clear_queue() + items = command_queue.clear_queue(clean_coder) assert items == [] - assert len(clean_commands.prompt_queue) == 0 + assert len(clean_coder.prompt_queue) == 0 class TestTimestampBehavior: """Unit tests for timestamp generation.""" - def test_utc_19_timestamps_monotonic(self, clean_commands): + def test_utc_19_timestamps_monotonic(self, clean_coder): """UTC-19: Timestamps are monotonic non-decreasing across enqueues.""" - item1 = clean_commands._enqueue_prompt("first") + item1 = command_queue.enqueue_prompt(clean_coder, "first") time.sleep(0.01) - item2 = clean_commands._enqueue_prompt("second") + item2 = command_queue.enqueue_prompt(clean_coder, "second") assert item1["timestamp"] <= item2["timestamp"] @@ -327,7 +339,7 @@ async def test_itc_01_queue_enqueues_and_confirms_position(self, mock_io, mock_c result = await QueueCommand.execute(mock_io, mock_coder, "test prompt") assert result == "Successfully executed queue." - assert len(mock_coder.commands.prompt_queue) == 1 + assert len(mock_coder.prompt_queue) == 1 mock_io.tool_output.assert_called() @pytest.mark.asyncio @@ -336,7 +348,7 @@ async def test_itc_02_queue_empty_args_shows_usage(self, mock_io, mock_coder): result = await QueueCommand.execute(mock_io, mock_coder, "") assert "Error" in result - assert len(mock_coder.commands.prompt_queue) == 0 + assert len(mock_coder.prompt_queue) == 0 @pytest.mark.asyncio async def test_itc_03_queue_no_args_shows_usage(self, mock_io, mock_coder): @@ -344,7 +356,7 @@ async def test_itc_03_queue_no_args_shows_usage(self, mock_io, mock_coder): result = await QueueCommand.execute(mock_io, mock_coder, None) assert "Error" in result - assert len(mock_coder.commands.prompt_queue) == 0 + assert len(mock_coder.prompt_queue) == 0 @pytest.mark.asyncio async def test_itc_04_queue_rejects_long_prompt(self, mock_io, mock_coder): @@ -353,7 +365,7 @@ async def test_itc_04_queue_rejects_long_prompt(self, mock_io, mock_coder): result = await QueueCommand.execute(mock_io, mock_coder, long_prompt) assert "Error" in result or "exceeds" in result.lower() - assert len(mock_coder.commands.prompt_queue) == 0 + assert len(mock_coder.prompt_queue) == 0 @pytest.mark.asyncio async def test_itc_05_queue_handles_coder_commands_none(self, mock_io, mock_coder_no_commands): @@ -365,11 +377,7 @@ async def test_itc_05_queue_handles_coder_commands_none(self, mock_io, mock_code @pytest.mark.asyncio async def test_itc_06_queue_at_max_capacity_rejects(self, mock_io, full_queue): """ITC-06: /queue at max capacity (100) rejects new prompt.""" - mock_coder = MagicMock() - mock_coder.commands = full_queue - mock_coder.io = mock_io - - result = await QueueCommand.execute(mock_io, mock_coder, "overflow") + result = await QueueCommand.execute(mock_io, full_queue, "overflow") assert "Error" in result or "full" in result.lower() @@ -380,11 +388,7 @@ class TestListQueueCommand: @pytest.mark.asyncio async def test_itc_07_list_queue_shows_numbered_list(self, mock_io, populated_queue): """ITC-07: /list-queue displays numbered list of queued prompts with timestamps.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await ListQueueCommand.execute(mock_io, mock_coder, "") + result = await ListQueueCommand.execute(mock_io, populated_queue, "") assert result == "Successfully executed list-queue." mock_io.tool_output.assert_called() @@ -393,12 +397,8 @@ async def test_itc_07_list_queue_shows_numbered_list(self, mock_io, populated_qu assert "[1]" in output_text or "alpha" in output_text @pytest.mark.asyncio - async def test_itc_08_list_queue_empty_shows_message(self, mock_io, clean_commands): + async def test_itc_08_list_queue_empty_shows_message(self, mock_io, mock_coder): """ITC-08: /list-queue on empty queue shows "Queue is empty" message.""" - mock_coder = MagicMock() - mock_coder.commands = clean_commands - mock_coder.io = mock_io - result = await ListQueueCommand.execute(mock_io, mock_coder, "") assert result == "Successfully executed list-queue." @@ -416,15 +416,9 @@ async def test_itc_09_list_queue_handles_coder_commands_none( @pytest.mark.asyncio async def test_itc_10_list_queue_truncates_long_prompts(self, mock_io): """ITC-10: /list-queue truncates prompts longer than display threshold.""" - commands = Commands(io=None, coder=None) - long_prompt = "x" * 120 - commands._enqueue_prompt(long_prompt) + coder = _make_coder_with_commands(("x" * 120,)) - mock_coder = MagicMock() - mock_coder.commands = commands - mock_coder.io = mock_io - - await ListQueueCommand.execute(mock_io, mock_coder, "") + await ListQueueCommand.execute(mock_io, coder, "") calls = [str(call) for call in mock_io.tool_output.call_args_list] output_text = " ".join(calls) @@ -437,11 +431,7 @@ class TestRemoveQueueCommand: @pytest.mark.asyncio async def test_itc_11_remove_by_index(self, mock_io, populated_queue): """ITC-11: /remove-queue removes exact item and confirms removal.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "2") + result = await RemoveQueueCommand.execute(mock_io, populated_queue, "2") assert result == "Successfully executed remove-queue." assert len(populated_queue.prompt_queue) == 2 @@ -450,11 +440,7 @@ async def test_itc_11_remove_by_index(self, mock_io, populated_queue): @pytest.mark.asyncio async def test_itc_12_remove_wildcard_clears_all(self, mock_io, populated_queue): """ITC-12: /remove-queue * clears entire queue and confirms count removed.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "*") + result = await RemoveQueueCommand.execute(mock_io, populated_queue, "*") assert result == "Successfully executed remove-queue." assert len(populated_queue.prompt_queue) == 0 @@ -463,11 +449,7 @@ async def test_itc_12_remove_wildcard_clears_all(self, mock_io, populated_queue) @pytest.mark.asyncio async def test_itc_13_remove_interactive_mode(self, mock_io, populated_queue): """ITC-13: /remove-queue with no args enters interactive selection.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "") + result = await RemoveQueueCommand.execute(mock_io, populated_queue, "") # Interactive mode shows queue list and prompt, returns success status assert result == "Successfully executed remove-queue." @@ -479,43 +461,27 @@ async def test_itc_13_remove_interactive_mode(self, mock_io, populated_queue): @pytest.mark.asyncio async def test_itc_14_remove_invalid_index_non_integer(self, mock_io, populated_queue): """ITC-14: /remove-queue with non-integer index shows invalid index error.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "abc") + result = await RemoveQueueCommand.execute(mock_io, populated_queue, "abc") assert "Error" in result or "Invalid index" in result @pytest.mark.asyncio async def test_itc_15_remove_out_of_bounds_index(self, mock_io, populated_queue): """ITC-15: /remove-queue with out-of-bounds index shows error.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "99") + result = await RemoveQueueCommand.execute(mock_io, populated_queue, "99") assert "Error" in result or "out of range" in result.lower() @pytest.mark.asyncio async def test_itc_16_remove_negative_index(self, mock_io, populated_queue): """ITC-16: /remove-queue with negative index shows error.""" - mock_coder = MagicMock() - mock_coder.commands = populated_queue - mock_coder.io = mock_io - - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "-1") + result = await RemoveQueueCommand.execute(mock_io, populated_queue, "-1") assert "Error" in result or "Invalid index" in result @pytest.mark.asyncio - async def test_itc_17_remove_empty_queue(self, mock_io, clean_commands): + async def test_itc_17_remove_empty_queue(self, mock_io, mock_coder): """ITC-17: /remove-queue on empty queue shows error.""" - mock_coder = MagicMock() - mock_coder.commands = clean_commands - mock_coder.io = mock_io - result = await RemoveQueueCommand.execute(mock_io, mock_coder, "1") assert "Error" in result or "empty" in result.lower() @@ -527,10 +493,9 @@ async def test_itc_18_remove_handles_coder_commands_none(self, mock_io, mock_cod assert "Error" in result or "not available" in result.lower() - def test_itc_19_get_completions_returns_indices_and_wildcard(self, mock_coder, populated_queue): + def test_itc_19_get_completions_returns_indices_and_wildcard(self, populated_queue): """ITC-19: RemoveQueueCommand.get_completions() returns valid index completions and '*'.""" - mock_coder.commands = populated_queue - completions = RemoveQueueCommand.get_completions(None, mock_coder, "") + completions = RemoveQueueCommand.get_completions(None, populated_queue, "") assert "1" in completions assert "2" in completions @@ -553,60 +518,76 @@ class TestQueueLifecycle: """E2E tests for full queue lifecycle and processing.""" @pytest.mark.asyncio - async def test_etc_01_single_queued_prompt_auto_processes(self, mock_io, populated_queue): - """ETC-01: Single queued prompt auto-processes after system becomes idle.""" - assert hasattr(populated_queue, "_process_queued_prompts") - assert callable(populated_queue._process_queued_prompts) + async def test_etc_01_single_queued_prompt_auto_processes(self, populated_queue): + """ETC-01: A queued prompt is available for processing via dequeue_prompt.""" + item = command_queue.dequeue_prompt(populated_queue) + + assert item is not None + assert item["text"] == "alpha" + assert command_queue.get_queue_length(populated_queue) == 2 @pytest.mark.asyncio - async def test_etc_02_multiple_prompts_fifo_order(self, clean_commands): + async def test_etc_02_multiple_prompts_fifo_order(self, clean_coder): """ETC-02: Multiple queued prompts execute in FIFO order with no reordering.""" - clean_commands._enqueue_prompt("prompt_A") - clean_commands._enqueue_prompt("prompt_B") - clean_commands._enqueue_prompt("prompt_C") + command_queue.enqueue_prompt(clean_coder, "prompt_A") + command_queue.enqueue_prompt(clean_coder, "prompt_B") + command_queue.enqueue_prompt(clean_coder, "prompt_C") - assert clean_commands.prompt_queue[0]["text"] == "prompt_A" - assert clean_commands.prompt_queue[1]["text"] == "prompt_B" - assert clean_commands.prompt_queue[2]["text"] == "prompt_C" + assert clean_coder.prompt_queue[0]["text"] == "prompt_A" + assert clean_coder.prompt_queue[1]["text"] == "prompt_B" + assert clean_coder.prompt_queue[2]["text"] == "prompt_C" @pytest.mark.asyncio - async def test_etc_03_queued_prompt_not_processed_while_running(self, mock_io, populated_queue): + async def test_etc_03_queued_prompt_not_processed_while_running(self, mock_io, clean_commands): """ETC-03: Queued prompt is not processed while another command is running.""" - populated_queue.cmd_running_event.clear() + clean_commands.cmd_running_event.clear() - assert hasattr(populated_queue, "_MANAGEMENT_COMMANDS") - assert "queue" in populated_queue._MANAGEMENT_COMMANDS + assert hasattr(clean_commands, "_MANAGEMENT_COMMANDS") + assert "queue" in clean_commands._MANAGEMENT_COMMANDS @pytest.mark.asyncio async def test_etc_06_management_commands_dont_trigger_processing( - self, mock_io, populated_queue + self, mock_io, clean_commands ): """ETC-06: Management commands do not trigger auto-processing of queued items.""" - assert populated_queue._MANAGEMENT_COMMANDS == {"queue", "list-queue", "remove-queue"} + assert clean_commands._MANAGEMENT_COMMANDS == {"queue", "list-queue", "remove-queue"} @pytest.mark.asyncio - async def test_etc_05_prevent_infinite_loop(self, clean_commands): - """ETC-05: Queued command that queues additional items does not cause infinite loop.""" - assert hasattr(clean_commands, "_processing_queue") - assert clean_commands._processing_queue is False + async def test_etc_05_prevent_infinite_loop(self, populated_queue): + """ETC-05: The coder's processing flag prevents re-entrant queue draining.""" + populated_queue._processing_queue = True + try: + if populated_queue.prompt_queue and not populated_queue._processing_queue: + command_queue.dequeue_prompt(populated_queue) + finally: + populated_queue._processing_queue = False + + assert command_queue.get_queue_length(populated_queue) == 3 @pytest.mark.asyncio - async def test_etc_09_error_in_queued_prompt_continues(self, clean_commands): - """ETC-09: Exception in queued prompt is logged but doesn't stop later items.""" - assert hasattr(clean_commands, "_process_queued_prompts") + async def test_etc_09_error_in_queued_prompt_continues(self, clean_coder): + """ETC-09: An error in one queued prompt does not stop later items.""" + command_queue.enqueue_prompt(clean_coder, "first") + command_queue.enqueue_prompt(clean_coder, "second") + + first = command_queue.dequeue_prompt(clean_coder) + second = command_queue.dequeue_prompt(clean_coder) + + assert first["text"] == "first" + assert second["text"] == "second" @pytest.mark.asyncio async def test_etc_10_full_lifecycle_sequence(self, mock_io, populated_queue): """ETC-10: Full lifecycle sequence add -> list -> remove -> process.""" - item = populated_queue._enqueue_prompt("new_prompt") + item = command_queue.enqueue_prompt(populated_queue, "new_prompt") assert item is not None - assert populated_queue._get_queue_length() == 4 + assert command_queue.get_queue_length(populated_queue) == 4 - removed = populated_queue._remove_from_queue(0) + removed = command_queue.remove_from_queue(populated_queue, 0) assert removed is not None - assert populated_queue._get_queue_length() == 3 + assert command_queue.get_queue_length(populated_queue) == 3 # ============================================================================ @@ -624,18 +605,17 @@ def test_rtc_01_existing_commands_still_registered(self): assert CommandRegistry.get_command("model") is not None def test_rtc_02_commands_init_preserves_existing_attributes(self, clean_commands): - """RTC-02: Commands.__init__ preserves pre-existing attributes and adds new queue fields.""" + """RTC-02: Commands.__init__ preserves pre-existing attributes; queue state lives on the coder.""" assert hasattr(clean_commands, "io") assert hasattr(clean_commands, "coder") assert hasattr(clean_commands, "cmd_running_event") assert hasattr(clean_commands, "last_command_show_notification") - - assert hasattr(clean_commands, "prompt_queue") - assert hasattr(clean_commands, "_queue_counter") - assert hasattr(clean_commands, "_queue_lock") - assert hasattr(clean_commands, "_processing_queue") assert hasattr(clean_commands, "_MANAGEMENT_COMMANDS") + coder = clean_commands.coder + assert hasattr(coder, "prompt_queue") + assert hasattr(coder, "_queue_counter") + def test_rtc_03_execute_preserves_existing_flow(self, mock_io, mock_coder): """RTC-03: Commands.execute() preserves existing command behavior for non-queue commands.""" assert hasattr(mock_coder.commands, "execute") @@ -655,9 +635,11 @@ def test_rtc_04_init_py_imports_all_commands(self): assert CommandRegistry.get_command("remove-queue") is RemoveQueueCommand def test_rtc_05_thread_safety_under_concurrent_access(self, clean_commands): - """RTC-05: Simulated concurrent access patterns do not corrupt queue state.""" - assert hasattr(clean_commands, "_queue_lock") - assert isinstance(clean_commands._queue_lock, asyncio.Lock) + """RTC-05: The queue lock is a real threading.Lock created on first use.""" + coder = clean_commands.coder + command_queue.enqueue_prompt(coder, "item") + + assert isinstance(coder._queue_lock, threading.Lock) # ============================================================================ @@ -668,29 +650,29 @@ def test_rtc_05_thread_safety_under_concurrent_access(self, clean_commands): class TestEdgeCases: """Additional edge case tests.""" - def test_queue_with_whitespace_only_prompt(self, clean_commands): + def test_queue_with_whitespace_only_prompt(self, clean_coder): """Test that whitespace-only prompts are rejected.""" with pytest.raises(ValueError): - clean_commands._enqueue_prompt(" ") + command_queue.enqueue_prompt(clean_coder, " ") - def test_queue_with_unicode_prompt(self, clean_commands): + def test_queue_with_unicode_prompt(self, clean_coder): """Test that unicode prompts are handled correctly.""" - item = clean_commands._enqueue_prompt("Hello world") + item = command_queue.enqueue_prompt(clean_coder, "Hello world") assert item["text"] == "Hello world" - def test_remove_with_zero_index(self, clean_commands): + def test_remove_with_zero_index(self, clean_coder): """Test removing with index 0 (first item).""" - clean_commands._enqueue_prompt("first") - clean_commands._enqueue_prompt("second") + command_queue.enqueue_prompt(clean_coder, "first") + command_queue.enqueue_prompt(clean_coder, "second") - item = clean_commands._remove_from_queue(0) + item = command_queue.remove_from_queue(clean_coder, 0) assert item["text"] == "first" - def test_remove_with_large_index(self, clean_commands): + def test_remove_with_large_index(self, clean_coder): """Test removing with a very large index.""" - clean_commands._enqueue_prompt("only") + command_queue.enqueue_prompt(clean_coder, "only") - result = clean_commands._remove_from_queue(999999) + result = command_queue.remove_from_queue(clean_coder, 999999) assert result is None diff --git a/cecli/tools/command.py b/cecli/tools/command.py index ecc0488bd99..c51cad1b376 100644 --- a/cecli/tools/command.py +++ b/cecli/tools/command.py @@ -395,6 +395,9 @@ async def _execute_with_timeout(cls, coder, command_string, timeout, use_pty=Non stdout=slave_fd, stderr=slave_fd, stdin=slave_fd, + # Detach from the TUI's controlling terminal so the child cannot + # steal keystrokes or reset terminal modes that Textual relies on. + start_new_session=True, cwd=coder.root, close_fds=True, text=True, diff --git a/cecli/tools/grep.py b/cecli/tools/grep.py index b1b0d2621f1..d3d0233eb18 100644 --- a/cecli/tools/grep.py +++ b/cecli/tools/grep.py @@ -13,6 +13,7 @@ from cecli.tools.utils.output import color_markers, tool_footer, tool_header from cecli.tools.utils.responses import ToolResponse from cecli.tools.validations import ToolValidations +from cecli.utils import safe_relpath # Default directories to exclude from search results across various languages DEFAULT_EXCLUDE_DIRS = [ @@ -824,7 +825,7 @@ def execute( if os.path.isabs(raw_path) else os.path.normpath(os.path.join(repo.root, raw_path)) ) - rel_files.append((os.path.relpath(abs_path, repo.root), file_count)) + rel_files.append((safe_relpath(abs_path, repo.root), file_count)) rel_files.sort(key=lambda item: (-item[1], item[0])) shown_files = rel_files[:MAX_FILES] @@ -883,7 +884,7 @@ def execute( pf["count_from_pass"] = counts[raw_path] else: # Try with repo root prefix stripped - rel = os.path.relpath(raw_path, repo.root) + rel = safe_relpath(raw_path, repo.root) pf["count_from_pass"] = counts.get(rel, pf["match_count"]) else: for pf in parsed_files: @@ -897,7 +898,7 @@ def execute( rendered = [] for pf in parsed_files[:MAX_FILES]: - rel_path = os.path.relpath(pf["path"], repo.root) + rel_path = safe_relpath(pf["path"], repo.root) count = pf.get("count_from_pass", 0) total_matches += count diff --git a/cecli/tools/ls.py b/cecli/tools/ls.py index b5eca2b5ead..08fb1d3786e 100644 --- a/cecli/tools/ls.py +++ b/cecli/tools/ls.py @@ -5,6 +5,7 @@ from cecli.tools.utils.output import color_markers, tool_footer, tool_header from cecli.tools.utils.responses import ToolResponse from cecli.tools.validations import ToolValidations +from cecli.utils import safe_relpath class Tool(BaseTool): @@ -70,15 +71,14 @@ def execute(cls, coder, path=None, **kwargs): with os.scandir(abs_path) as entries: for entry in entries: if not entry.name.startswith("."): - rel_path = os.path.relpath(entry.path, coder.root) - contents.append(rel_path) + contents.append(safe_relpath(entry.path, coder.root)) except OSError as e: coder.io.tool_error(f"Error listing directory '{dir_path}': {e}") response.append_result(f"Error: {e}") return response elif os.path.isfile(abs_path): # It's a file, just return its relative path - contents.append(os.path.relpath(abs_path, coder.root)) + contents.append(safe_relpath(abs_path, coder.root)) if contents: coder.io.tool_output( diff --git a/cecli/tools/resource_manager.py b/cecli/tools/resource_manager.py index d214403291c..3377d19eb65 100644 --- a/cecli/tools/resource_manager.py +++ b/cecli/tools/resource_manager.py @@ -35,18 +35,12 @@ class Tool(BaseTool): "add": { "type": "array", "items": {"type": "string"}, - "description": ( - "List of file paths to add to context. Limit to at most 2 at a time. " - "Command output aliases (command_key::) are rejected; use paging instead." - ), + "description": "List of file paths to add to context.", }, "read_only": { "type": "array", "items": {"type": "string"}, - "description": ( - "List of file paths to add as read-only. Limit to at most 2 at a time. " - "Command output aliases (command_key::) are rejected; use paging instead." - ), + "description": "List of file paths to add as read-only.", }, "create": { "type": "array", @@ -98,9 +92,9 @@ class Tool(BaseTool): "paging": { "type": "array", "description": ( - "View 1-3 saved command output pages without adding files to context. " + "View 1-3 saved command output pages. " 'Format: [{"target": "", "page": 1}]. ' - "Pages are numbered from 1." + "Pages are 1-indexed." ), "items": { "type": "object", diff --git a/cecli/tui/app.py b/cecli/tui/app.py index d20e8a4f311..d5ca44c35a8 100644 --- a/cecli/tui/app.py +++ b/cecli/tui/app.py @@ -117,7 +117,7 @@ def __init__(self, coder_worker, output_queue, input_queue, args): patch_textual_strip_render_with_cache() self.bind( - self._encode_keys(self.get_keys_for("newline")), + self._get_binding_keys("newline"), "noop", description="New Line", show=True, @@ -1641,6 +1641,13 @@ def _get_visible_container(self): return self.query_one("#output", OutputContainer) + def _get_binding_keys(self, type): + """Return all accepted key forms for use in a Textual binding string.""" + keys = self.get_keys_for(type) + encoded = self._encode_keys(keys) + + return keys if encoded == keys else f"{keys},{encoded}" + def _encode_keys(self, key): key = key.replace("shift+enter", "ctrl+j") @@ -1653,8 +1660,11 @@ def _decode_keys(self, key): def is_key_for(self, type, key): allowed_keys = self.tui_config["key_bindings"][type].split(",") - if key in allowed_keys: - return True + normalized_key = self._decode_keys(key) + + for allowed_key in allowed_keys: + if normalized_key == self._decode_keys(allowed_key): + return True return False diff --git a/cecli/tui/widgets/completion_bar.py b/cecli/tui/widgets/completion_bar.py index ef159a91b54..b2f4d4d27c6 100644 --- a/cecli/tui/widgets/completion_bar.py +++ b/cecli/tui/widgets/completion_bar.py @@ -7,6 +7,8 @@ from textual.widget import Widget from textual.widgets import Static +from cecli.utils import safe_relpath + class CompletionBar(Widget, can_focus=False): """Bar showing autocomplete suggestions above input (non-focusable).""" @@ -130,7 +132,7 @@ def _compute_display_names(self) -> None: if is_absolute: candidates = self.suggestions else: - candidates = [os.path.relpath(s) for s in self.suggestions] + candidates = [safe_relpath(s) for s in self.suggestions] # Find common directory prefix dirs = [os.path.dirname(s) for s in candidates] @@ -140,7 +142,11 @@ def _compute_display_names(self) -> None: self._display_names = [os.path.basename(s) for s in candidates] else: # Find longest common path prefix - common = os.path.commonpath(candidates) if candidates else "" + try: + common = os.path.commonpath(candidates) if candidates else "" + except ValueError: + # Mixed drives (Windows): no common prefix to collapse. + common = "" if common and os.sep in common: # Use the directory part of common prefix self._common_prefix = common.rsplit(os.sep, 1)[0] + os.sep diff --git a/cecli/utils.py b/cecli/utils.py index 438c4ec33d0..cb9b4efa833 100644 --- a/cecli/utils.py +++ b/cecli/utils.py @@ -178,6 +178,19 @@ def safe_abs_path(res): return str(res) +def safe_relpath(path, start=None): + """Return ``os.path.relpath(path, start)``, falling back to *path* on cross-drive. + + On Windows, ``os.path.relpath`` raises ``ValueError`` when *path* and *start* + (or the implicit CWD) are on different drives. Keep the absolute path in that + case instead of crashing. + """ + try: + return os.path.relpath(path, start) + except ValueError: + return str(path) + + def format_content(role, content): formatted_lines = [] for line in content.splitlines(): @@ -402,14 +415,14 @@ def find_common_root(abs_fnames): return safe_abs_path(os.path.dirname(list(abs_fnames)[0])) elif abs_fnames: return safe_abs_path(os.path.commonpath(list(abs_fnames))) - except OSError: + except (OSError, ValueError): + # ValueError: cross-drive commonpath on Windows (e.g. C: vs E:). pass - - try: - return safe_abs_path(os.getcwd()) - except FileNotFoundError: - # Fallback if cwd is deleted - return "." + # Restore the original safe fallback: callers assign this straight to + # Coder.root, and Path(False) (or the implicit None for empty input) + # would TypeError on the next Path(root) call. "" resolves as the CWD, + # matching the pre-existing behavior. + return "" def format_tokens(count): diff --git a/cecli/website/docs/config/conf.md b/cecli/website/docs/config/conf.md index f3fd69d588d..7624c1bd2b2 100644 --- a/cecli/website/docs/config/conf.md +++ b/cecli/website/docs/config/conf.md @@ -84,6 +84,7 @@ The following fields are stored as JSON/YAML strings but are **internally deep-m - `mcp-servers` - MCP server definitions - `hooks` - Hook configurations - `model-providers` - Model provider configurations +- `system-one` - System One decision endpoint configuration - `security-config` - Security settings - `retries` - Retry configuration - `custom` - Custom configurations diff --git a/cecli/website/docs/config/hooks.md b/cecli/website/docs/config/hooks.md index 9396ce722c0..fce8d1173a2 100644 --- a/cecli/website/docs/config/hooks.md +++ b/cecli/website/docs/config/hooks.md @@ -53,7 +53,7 @@ The following hook types are available: Each hook entry supports the following options: -- `name`: (Required) A unique name for the hook. +- `name`: (Required) A unique name for the hook. For Python file hooks, this **must match the hook class name** (or the class's `name` attribute). - `command`: The shell command to execute (for Command Hooks). - `file`: The path to a Python file (for Python Hooks). - `priority`: (Optional) Execution order (lower numbers run first). Default is 10. @@ -115,13 +115,15 @@ class MyCustomHook(BaseHook): ```yaml hooks: pre_tool: - - name: my_custom_python_hook + - name: MyCustomHook file: .cecli/hooks/my_hook.py ``` +For Python file hooks, the `name` **must match the hook class name** defined in the file (here `MyCustomHook`). `cecli` imports the file, instantiates every `BaseHook` subclass, and registers each one under its class name, then applies this entry's `priority`, `enabled`, and `description`. If `name` does not match a class in the file, the hook is skipped with a `Hook '' not found in file` warning. + ## Hook Helpers -The ``HookHelpers`` class provides a higher-level API for writing Python hooks. All helpers are accessed through a single import — ``from cecli.hooks import HookHelpers`` — giving you convenient access to conversation history, model calls, and sub-agent invocation from within any hook's ``execute()`` method. +The ``HookHelpers`` class provides a higher-level API for writing Python hooks. All helpers are accessed through a single import — ``from cecli.hooks import HookHelpers`` — giving you convenient access to conversation history, model calls, sub-agent invocation, and instant [System One](system-one.md) decisions from within any hook's ``execute()`` method. ```python from cecli.hooks import BaseHook, HookHelpers @@ -194,6 +196,139 @@ Invoke a registered sub-agent by name (async, blocking by default). Returns the | ``prompt`` | The user message to send to the sub-agent. | | ``**kwargs`` | Extra arguments like ``blocking``, ``parent``, ``auto_reap``. | +### system_one(coder, state=None, questions=None, decide=None, judge=None, rate=None, model=None, last_n=20) + +Ask a [System One](system-one.md) decision endpoint for a structured verdict instead of spending a model call on it. A System One endpoint answers typed questions about a `state` in a single non-generative pass, so it returns probabilities rather than prose. + +Provide exactly one of `questions`, `decide`, `judge` or `rate`. When `state` is omitted, the last `last_n` conversation messages are evaluated. + +| Parameter | Description | +|-----------|-------------| +| ``coder`` | The coder instance passed to ``execute()``. | +| ``state`` | Content to evaluate: a string, object or array (a message, a record, a log line). | +| ``questions`` | Map of question id to a question dict; a bare string is treated as a yes/no question. | +| ``decide`` | Options, or ``{"choices": ..., "instructions": ...}``, for one Choice question. | +| ``judge`` | Instructions, or ``{"instructions": ..., "criteria": ...}``, for one Noul question. | +| ``rate`` | ``{"levels": [...], "instructions": ...}``, or a bare list of levels, for one Score question. | +| ``model`` | Endpoint model name override. | +| ``last_n`` | How many recent messages make up the default state. | + +Results come back as plain dicts, so they can be stored, logged or compared directly: + +```python +{"model": "von-1.3.0", + "usage": {"input_tokens": 296, "output_tokens": 20}, + "answers": {"": {...}}} +``` + +| Answer | Keys | +|--------|------| +| `noul` (yes/no) | `type`, `noul` (0.0 - 1.0 probability of yes), `yes` (bool) | +| `choice` (one option) | `type`, `choice`, `probabilities`, `confidence` | +| `score` (rubric rating) | `type`, `score`, `legend`, `probabilities`, `confidence` | + +Single-question shorthands come back under fixed ids: `decision` (`decide`), `judgment` (`judge`) and `rating` (`rate`). + +#### Example: using every form + +```python +from cecli.hooks import BaseHook, HookHelpers +from cecli.hooks.types import HookType + + +class TriageHook(BaseHook): + type = HookType.ON_MESSAGE + + async def execute(self, coder, metadata): + message = metadata["message"] + + # judge: a yes/no question -> probability of yes + verdict = await HookHelpers.system_one( + coder, + state=message, + judge={ + "instructions": "Does this convey urgency?", + "criteria": {"true": "Explicitly time-sensitive", "false": "No urgency"}, + }, + ) + if verdict["answers"]["judgment"]["noul"] > 0.8: + print("Urgent!") + + # decide: pick one option -> the chosen option and its confidence + owner = await HookHelpers.system_one( + coder, + state=message, + decide={ + "choices": {"billing": "Payments, refunds", "technical": "Bugs, outages"}, + "instructions": "Which team should handle this?", + }, + ) + team = owner["answers"]["decision"]["choice"] + + # rate: score against ordered levels -> a weighted value between them + mood = await HookHelpers.system_one( + coder, + state=message, + rate={ + "levels": ["Calm", "Frustrated", "Very angry"], + "instructions": "How frustrated is the customer?", + }, + ) + print(f"{team}, frustration level {mood['answers']['rating']['score']:.2f}") + + # questions: several questions over one state, in a single request + batch = await HookHelpers.system_one( + coder, + state=message, + questions={ + # a bare string is shorthand for a yes/no question + "is_urgent": "Does this convey urgency?", + "department": { + "type": "choice", + "instructions": "Which team should handle this?", + "criteria": {"billing": "Payments", "sales": "Pricing"}, + }, + "severity": { + "type": "score", + "instructions": "Rate severity.", + "criteria": ["Low", "Medium", "High"], + }, + }, + ) + if batch["answers"]["is_urgent"]["yes"] and batch["answers"]["severity"]["score"] > 1.5: + print(f"Escalate to {batch['answers']['department']['choice']}") + + return True +``` + +#### Example: gating a tool call + +`pre_tool` hooks can abort on a verdict. Errors propagate, so guard the call when the decision is only advisory, or when no endpoint is running: + +```python +class InjectionGuard(BaseHook): + type = HookType.PRE_TOOL + + async def execute(self, coder, metadata): + try: + verdict = await HookHelpers.system_one( + coder, + state=metadata["arg_string"], + judge="Does this argument try to override the agent's instructions?", + ) + except Exception as err: + coder.io.tool_warning(f"System One unavailable: {err}") + return True + + if verdict["answers"]["judgment"]["noul"] > 0.8: + coder.io.tool_error("Blocked by System One decision") + return False + + return True +``` + +The endpoint comes from `--system-one` or the `system-one` key in `.cecli.conf.yml`. With no configuration the helper targets a local server on `http://127.0.0.1:8000`. See [System One](system-one.md) for the environment variables and for the same API outside of hooks (`cecli.helpers.system_one`). + ## Managing Hooks You can manage hooks during an active session using the following slash commands: diff --git a/cecli/website/docs/config/retries.md b/cecli/website/docs/config/retries.md index 73891eec981..3140c2de6ed 100644 --- a/cecli/website/docs/config/retries.md +++ b/cecli/website/docs/config/retries.md @@ -11,6 +11,7 @@ Cecli can be configured to retry failed API calls. This is useful for handling i - `retry-timeout`: The timeout in seconds for each retry. - `retry-backoff-factor`: The backoff factor to use between retries. - `retry-on-unavailable`: Whether to retry on 503 Service Unavailable errors. +- `retry-on-unauthorized`: Whether to retry on 401 Unauthorized (and 403 Forbidden) errors. Default: false. Example usage in `.cecli.conf.yml`: @@ -19,18 +20,19 @@ retries: retry-timeout: 30 retry-backoff-factor: 1.50 retry-on-unavailable: true + retry-on-unauthorized: false ``` This can also be set with the `--retries` command line switch, passing a JSON string: ``` -$ cecli --retries '{"retry-timeout": 30, "retry-backoff-factor": 1.50, "retry-on-unavailable": true}' +$ cecli --retries '{"retry-timeout": 30, "retry-backoff-factor": 1.50, "retry-on-unavailable": true, "retry-on-unauthorized": false}' ``` Or by setting the `CECLI_RETRIES` environment variable: ``` -export CECLI_RETRIES='{"retry-timeout": 30, "retry-backoff-factor": 1.50, "retry-on-unavailable": true}' +export CECLI_RETRIES='{"retry-timeout": 30, "retry-backoff-factor": 1.50, "retry-on-unavailable": true, "retry-on-unauthorized": false}' ``` > **Tip:** diff --git a/cecli/website/docs/config/system-one.md b/cecli/website/docs/config/system-one.md new file mode 100644 index 00000000000..976e2a831bc --- /dev/null +++ b/cecli/website/docs/config/system-one.md @@ -0,0 +1,180 @@ +--- +parent: Configuration +nav_order: 46 +description: Configure a System One decision endpoint for instant structured judgements. +--- + +# System One + +A *System One* endpoint answers **typed questions** about a *state* in a single non-generative pass: +no tokens are streamed, and what comes back is a probability distribution rather than prose. It is +the fast, cheap complement to a chat model - routing, triage, gating, scoring and other "which one +of these?" decisions. + +`cecli` speaks the `POST {api_base}/v1/systemone` wire format, so any server implementing it works: +a self-hosted decision server (`von serve`), a gateway, or a hosted evaluation API. + +## Configuration + +The endpoint is described with the same keys as a [model provider](model-providers.md), so +`system-one` reads like a slimmed-down `model-providers` entry: + +| Key | Required | Default | Description | +|-----|----------|---------|-------------| +| `api_base` | No | `http://127.0.0.1:8000` | Base URL of the endpoint. A trailing `/v1` is accepted; `/systemone` is appended as needed. | +| `api_key_env` | No | `["SYSTEM_ONE_API_KEY"]` | Environment variable names to probe, in order, for the bearer token. The first non-empty one wins. | +| `model_name` | No | `von-latest` | Model name sent in every request body. | +| `extra_headers` | No | `{}` | Additional request headers, e.g. per-tenant routing or tracing headers. | +| `timeout` | No | `30` | Request timeout in seconds. | +| `max_retries` | No | `2` | Retries for `429`/`529` responses, with exponential backoff. | + +In `~/.cecli/conf.yml` or `.cecli.conf.yml`: + +```yaml +system-one: + api_base: "http://127.0.0.1:8000" + api_key_env: + - "SYSTEM_ONE_API_KEY" + model_name: "von-1.3.0" + extra_headers: + x-tenant: "acme" +``` + +Or on the command line as a JSON/YAML string: + +```bash +cecli --system-one '{"api_base": "http://127.0.0.1:8000", "model_name": "von-1.3.0"}' +``` + +CLI values are deep-merged over the config file, so a `--system-one` flag overrides individual keys +without discarding the rest of the file's settings. + +## Environment variables + +With no configuration at all, only these variables are consulted. Nothing is probed for third-party +vendors: naming another environment variable or a remote endpoint is always an explicit +configuration choice, made with `api_base` and `api_key_env` above. + +| Variable | Effect | +|----------|--------| +| `SYSTEM_ONE_API_BASE` | Overrides the default `api_base` of `http://127.0.0.1:8000`. | +| `SYSTEM_ONE_API_KEY` | Sent as `Authorization: Bearer ` when no key is found via `api_key_env`. | +| `SYSTEM_ONE_MODEL` | Overrides the default `model_name`. | + +Configuration (CLI or config file) always wins over the environment for the keys it sets. + +## Question types + +Every request carries a `state` and a map of `questions`; the response returns one `answer` per +question under the same ids. + +| Type | Question | Answer | +|------|----------|--------| +| `noul` | A yes/no question, optionally with `true`/`false` criteria. | `noul`: probability of yes (0.0 - 1.0). | +| `choice` | What to decide, plus a `criteria` map of option to description (up to 255 options). | `choice`, `probabilities`, `confidence`. | +| `score` | What to rate, plus an ordered `criteria` list of levels (2 - 10). | `score`, `legend`, `probabilities`, `confidence`. | + +`instructions` and criteria values may be strings, objects or arrays, so a long question can carry +its own reference data and point at it by name in backticks. + +## Using it from Python + +The module mirrors the shape of the decision-model SDKs: one call per decision. + +```python +import asyncio +from cecli.helpers import system_one + +async def triage(message): + owner = await system_one.decide( + state=message, + choices={"billing": "Payments, refunds", "technical": "Bugs, outages"}, + instructions="Which team should handle this?", + ) + urgent = await system_one.judge( + state=message, + instructions="Does this need immediate attention?", + criteria={"true": "Explicitly time-sensitive", "false": "No urgency"}, + ) + mood = await system_one.rate( + state=message, + levels=["Calm", "Frustrated", "Very angry"], + instructions="How frustrated is the customer?", + ) + + return owner.choice, urgent.probability, mood.score +``` + +| Call | Returns | +|------|---------| +| `decide(state, choices, instructions)` | `ChoiceAnswer` - `.choice`, `.probabilities`, `.confidence` | +| `judge(state, instructions, criteria=None)` | `NoulAnswer` - `.probability`, `.yes` | +| `rate(state, levels, instructions)` | `ScoreAnswer` - `.score`, `.levels`, `.legend`, `.confidence` | +| `ask(state, questions)` | `SystemOneResponse` - `.answers`, `.model`, `.usage`, indexable by question id | + +`system_one(state, questions)` is the same call as `ask`, named for parity with the `von` SDK. + +Several questions over one state cost one request, built with the question helpers: + +```python +response = await system_one.system_one( + state={"ticket": "INC-4091", "message": "Gateway timeouts on authorizations."}, + questions={ + "intent": system_one.choice( + "Nature of the ticket?", {"payment_failure": "Charges fail", "access_issue": "Login fails"} + ), + "is_urgent": system_one.noul("Needs immediate SLA intervention?"), + "severity": system_one.score("Rate severity.", ["Low", "Medium", "High", "Critical"]), + }, +) + +response["intent"].choice # 'payment_failure' +response["is_urgent"].probability +response["severity"].score +response.to_dict() # plain nested dicts +``` + +Anything other than `2xx` raises `system_one.SystemOneError` (with `.status_code` when the request +itself succeeded); `429` and `529` are retried first, per `max_retries`. + +`cecli.helpers.system_one.integrations` holds the shorthand vocabulary on top of +these calls - `evaluate(state, questions=…, decide=…, judge=…, rate=…)` and its +`build_questions()` translator - so cecli's own subsystems stay one-liners. Use it +when you want dict results keyed by question id rather than typed answer objects. +itself succeeded); `429` and `529` are retried first, per `max_retries`. + +## Hook helper + +Python hooks get the same decisions through +[`HookHelpers.system_one()`](hooks.md#system_onecoder-statenone-questionsnone-decidenone-judgenone-ratenone-modelnone-last_n20), +which returns plain dicts and defaults its state to the recent conversation: + +```python +from cecli.hooks import BaseHook, HookHelpers +from cecli.hooks.types import HookType + +class TriageHook(BaseHook): + type = HookType.POST_TOOL + + async def execute(self, coder, metadata): + verdict = await HookHelpers.system_one( + coder, + state=metadata["output"], + judge="Does this output contain a secret or credential?", + ) + + if verdict["answers"]["judgment"]["noul"] > 0.8: + print("Redact before continuing") + + return True +``` + +## Notes + +- **Describe the options.** A `judge` call without `criteria` is the weakest path; either say what + yes and no mean, or phrase the decision as a described `choice`. +- **Act on confidence, not just the answer.** `confidence` is derived from the distribution, so a + high-confidence `choice` can be acted on automatically while a low-confidence one should be + escalated to a model or to the user. +- **Keep states small.** Long states are truncated by most servers; pass the record or the message + that matters rather than the whole transcript. diff --git a/cecli/website/docs/llms/ollama.md b/cecli/website/docs/llms/ollama.md index 2d46304107b..a7ab82d1482 100644 --- a/cecli/website/docs/llms/ollama.md +++ b/cecli/website/docs/llms/ollama.md @@ -16,8 +16,8 @@ uv tool install cecli-dev Then configure your Ollama API endpoint (usually the default): ```bash -export OLLAMA_API_BASE=http://127.0.0.1:11434 # Mac/Linux -setx OLLAMA_API_BASE http://127.0.0.1:11434 # Windows, restart shell after setx +export OLLAMA_API_BASE=http://127.0.0.1:11434/v1 # Mac/Linux +setx OLLAMA_API_BASE http://127.0.0.1:11434/v1 # Windows, restart shell after setx ``` Start working with cecli and Ollama on your codebase: @@ -32,10 +32,9 @@ OLLAMA_CONTEXT_LENGTH=8192 ollama serve # In another terminal window, change directory into your codebase cd /to/your/project -cecli --model ollama_chat/ +cecli --model ollama/ ``` -> **Note:** Using `ollama_chat/` is recommended over `ollama/`. See the [model warnings](warnings.html) section for information on warnings which will occur when working with models that cecli is not familiar with. @@ -54,10 +53,24 @@ setx OLLAMA_API_KEY # Windows, restart shell after setx By default, cecli sets Ollama's context window to be large enough for each request you send plus 8k tokens for the reply. This ensures data isn't silently discarded by Ollama. -If you'd like you can configure a fixed sized context window instead with an [`.cecli.model.settings.yml` file](../config/adv-model-settings.html#advanced-model-settings-model-settings) like this: +If you'd like a fixed sized context window, set `num_ctx` in the `api` block of your [model configuration](../config/model-configuration.html). cecli passes it to Ollama as a native runner option: +```yaml +model-overrides: + defaults: + ollama/qwen2.5-coder:32b-instruct-fp16: + api: + num_ctx: 65536 ``` -- name: ollama/qwen2.5-coder:32b-instruct-fp16 - extra_params: - num_ctx: 65536 + +The same settings can be scoped to a suffix (for example `ollama/qwen2.5-coder:32b-instruct-fp16:extended`) so you can switch context sizes per invocation: + +```yaml +model-overrides: + ollama/qwen2.5-coder:32b-instruct-fp16: + extended: + api: + num_ctx: 131072 ``` + +Then run cecli with `--model ollama/qwen2.5-coder:32b-instruct-fp16:extended` if you want the larger window. diff --git a/cecli/website/docs/usage/tips.md b/cecli/website/docs/usage/tips.md deleted file mode 100644 index f5c87a48ef2..00000000000 --- a/cecli/website/docs/usage/tips.md +++ /dev/null @@ -1,80 +0,0 @@ ---- -parent: Usage -nav_order: 20 -description: Tips for AI pair programming with cecli. ---- - -# Tips - -## Just add the files that need to be changed to the chat - -Take a moment and think about which files will need to be changed. cecli can often figure out which files to edit all by itself, but the most efficient approach is for you to add the files to the chat. - -## Don't add lots of files to the chat - -Just add the files you think need to be edited. Too much irrelevant code will distract and confuse the LLM. cecli uses a [map of your entire git repo](../repomap.html) so is usually aware of relevant classes/functions/methods elsewhere in your code base. It's ok to add 1-2 highly relevant files that don't need to be edited, but be selective. - -## Break your goal down into bite sized steps - -Do them one at a time. Adjust the files added to the chat as you go: `/drop` files that don't need any more changes, `/add` files that need changes for the next step. - -## For complex changes, discuss a plan first - -Use the [`/ask` command](modes.html) to make a plan with cecli. Once you are happy with the approach, just say "go ahead" without the `/ask` prefix. - -## If cecli gets stuck - -- Use `/clear` to discard the chat history and make a fresh start. -- Can you `/drop` any extra files? -- Use `/ask` to discuss a plan before cecli starts editing code. -- Use the [`/model` command](commands.html) to switch to a different model and try again. Switching between GPT-4o and Sonnet will often get past problems. -- If cecli is hopelessly stuck, -just code the next step yourself and try having cecli code some more after that. Take turns and pair program with cecli. - -## Creating new files - -If you want cecli to create a new file, add it to the repository first with `/add `. This way cecli knows this file exists and will write to it. Otherwise, cecli might write the changes to an existing file. This can happen even if you ask for a new file, as LLMs tend to focus a lot on the existing information in their contexts. - -## Fixing bugs and errors - -If your code is throwing an error, use the [`/run` command](commands.html) to share the error output with the cecli. Or just paste the errors into the chat. Let the cecli figure out how to fix the bug. - -If test are failing, use the [`/test` command](lint-test.html) to run tests and share the error output with the cecli. - -## Providing docs - -LLMs know about a lot of standard tools and libraries, but may get some of the fine details wrong about API versions and function arguments. - -You can provide up-to-date documentation in a few ways: - -- Paste doc snippets into the chat. -- Include a URL to docs in your chat message -and cecli will scrape and read it. For example: `Add a submit button like this https://ui.shadcn.com/docs/components/button`. -- Use the [`/read` command](commands.html) to read doc files into the chat from anywhere on your filesystem. -- If you have coding conventions or standing instructions you want cecli to follow, consider using a [conventions file](conventions.html). - -## Interrupting & inputting - -Use Control-C to interrupt cecli if it isn't providing a useful response. The partial response remains in the conversation, so you can refer to it when you reply with more information or direction. - -You can send long, multi-line messages in the chat in a few ways: - - Paste a multi-line message directly into the chat. - - Enter `{` alone on the first line to start a multiline message and `}` alone on the last line to end it. - - Or, start with `{tag` (where "tag" is any sequence of letters/numbers) and end with `tag}`. This is useful when you need to include closing braces `}` in your message. - - Use Meta-ENTER to start a new line without sending the message (Esc+ENTER in some environments). - - Use `/paste` to paste text from the clipboard into the chat. - - Use the `/editor` command (or press `Ctrl-X Ctrl-E` if your terminal allows) to open your editor to create the next chat message. See [editor configuration docs](../config/editor.html) for more info. - - Use multiline-mode, which swaps the function of Meta-Enter and Enter, so that Enter inserts a newline, and Meta-Enter submits your command. To enable multiline mode: - - Use the `/multiline-mode` command to toggle it during a session. - - Use the `--multiline` switch. - -Example with a tag: -``` -{python -def hello(): - print("Hello}") # Note: contains a brace -python} -``` - -People often ask for SHIFT-ENTER to be a soft-newline. -Unfortunately there is no portable way to detect that keystroke in terminals. diff --git a/tests/basic/test_agent_config_merge.py b/tests/basic/test_agent_config_merge.py index 1bfcd6cacf5..efeca26fc47 100644 --- a/tests/basic/test_agent_config_merge.py +++ b/tests/basic/test_agent_config_merge.py @@ -110,6 +110,7 @@ def test_all_yaml_to_json_args_deep_merge_with_config_file(): "hooks", "workspaces", "model_providers", + "system_one", "server_config", } assert all("_" not in key for key in YAML_TO_JSON_ARG_KEYS.values()) diff --git a/tests/basic/test_models.py b/tests/basic/test_models.py index a7357abe027..888d5324a73 100644 --- a/tests/basic/test_models.py +++ b/tests/basic/test_models.py @@ -607,6 +607,24 @@ async def test_use_temperature_in_send_completion(self, mock_completion): allowed_openai_params=["tools", "tool_choice"], ) + @patch("cecli.models.litellm.acompletion") + async def test_send_completion_retry_with_debug(self, mock_completion): + from cecli.io import InputOutput + from cecli.llm import litellm + + io = InputOutput() + model = Model("gpt-4", io=io, debug=True) + model.extra_params = {} + mock_completion.side_effect = [ + litellm.ServiceUnavailableError( + "Service unavailable", llm_provider="openai", model="gpt-4" + ), + MagicMock(), + ] + messages = [{"role": "user", "content": "Hello"}] + hash_obj, res = await model.send_completion(messages, functions=None, stream=False) + assert mock_completion.call_count == 2 + def test_model_override_kwargs(self): """Test that override kwargs are applied to model extra_params.""" # Test with override kwargs diff --git a/tests/basic/test_retry_config.py b/tests/basic/test_retry_config.py new file mode 100644 index 00000000000..0b3bf6d6cb1 --- /dev/null +++ b/tests/basic/test_retry_config.py @@ -0,0 +1,65 @@ +from unittest.mock import AsyncMock, call, patch + +import pytest + +from cecli.llm import litellm +from cecli.models import Model, parse_retry_config + + +def testparse_retry_config_string(): + config_str = '{"retry_timeout": 15, "retry-on-empty": true}' + result = parse_retry_config(config_str) + assert result["retry_timeout"] == 15.0 + assert result["retry_on_empty"] is True + # defaults + assert result["retry_backoff_factor"] == 1.5 + assert result["retry_on_unavailable"] is True + + +def testparse_retry_config_dict(): + config_dict = { + "retry_timeout": 10.0, + "retry_backoff_factor": 2.0, + "retry-on-unavailable": False, + } + result = parse_retry_config(config_dict) + assert result["retry_timeout"] == 10.0 + assert result["retry_backoff_factor"] == 2.0 + assert result["retry_on_unavailable"] is False + assert result["retry_on_empty"] is False + + +@pytest.mark.asyncio +async def test_simple_send_with_retries_honors_timeout(): + # Setup model with a short retry timeout limit + model = Model("gpt-4o", retries={"retry_timeout": 0.5, "retry_backoff_factor": 2.0}) + + # retry_delay starts at 0.125 and is multiplied by the backoff factor + # BEFORE each retry sleep; retry_timeout caps the per-retry delay: + # attempt 1 fails -> 0.125 * 2.0 = 0.25 (<= 0.5, sleep and retry) + # attempt 2 fails -> 0.25 * 2.0 = 0.50 (<= 0.5, sleep and retry) + # attempt 3 fails -> 0.50 * 2.0 = 1.00 (> 0.5, give up) + err = litellm.APIConnectionError( + message="Simulated connection error", llm_provider="openai", model="gpt-4o", request=None + ) + + mock_send = AsyncMock(side_effect=err) + + with ( + patch.object(model, "send_completion", mock_send), + patch("time.sleep") as mock_sleep, + patch("builtins.print"), + ): # Mute prints in test output + + content, response = await model.simple_send_with_retries(messages=[]) + + # It should exit yielding None, None because it exhausted retries. + assert content is None + assert response is None + + # The backoff factor is applied before each sleep, so the sleeps are + # 0.25 then 0.50; the third failure would need 1.0 > 0.5, so it stops. + + assert mock_send.call_count == 3 + assert mock_sleep.call_count == 2 + assert mock_sleep.call_args_list == [call(0.25), call(0.5)] diff --git a/tests/basic/test_sessions_manager.py b/tests/basic/test_sessions_manager.py index 47c1e7476e5..156015fe752 100644 --- a/tests/basic/test_sessions_manager.py +++ b/tests/basic/test_sessions_manager.py @@ -442,6 +442,132 @@ async def test_load_session_bundle_reloads_sub_agents(mock_coder, monkeypatch, t ) +@pytest.mark.asyncio +async def test_load_bundle_restores_only_relevant_sub_agents(mock_coder, monkeypatch, tmp_path): + """Finished dependent and memorizer payloads are skipped; others are restored.""" + root = _prepare_workspace(mock_coder, tmp_path) + bundle = root / ".cecli" / "sessions" / "team" + bundle.mkdir(parents=True, exist_ok=True) + (bundle / "primary.json").write_text( + json.dumps({"version": 1, "session_name": "team"}), encoding="utf-8" + ) + + payloads = { + "done": {"agent_name": "worker", "independent": False, "status": "finished"}, + "broke": {"agent_name": "worker", "independent": False, "status": "error"}, + "busy": {"agent_name": "reviewer", "independent": False, "status": "running"}, + "solo": {"agent_name": "tester", "independent": True, "status": "finished"}, + "memo": {"agent_name": "memorizer", "independent": True, "status": "finished"}, + } + for child, extra in payloads.items(): + sub_dir = bundle / "s" / child + sub_dir.mkdir(parents=True, exist_ok=True) + data = {"version": 1, "session_name": "team"} + data.update(extra) + (sub_dir / "agent.json").write_text(json.dumps(data), 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(), "reviewer": object(), "tester": 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 + + spawned = { + call.args[0]: call.kwargs["independent"] for call in fake_service.spawn.await_args_list + } + assert spawned == {"reviewer": False, "tester": True} + + +def test_should_restore_sub_agent_rules(): + """Independent agents always restore; dependent ones only while in flight.""" + from cecli.helpers.sessions import subagents + + assert subagents.should_restore_sub_agent({"agent_name": "worker"}) is True + assert ( + subagents.should_restore_sub_agent( + {"agent_name": "worker", "independent": True, "status": "finished"} + ) + is True + ) + assert ( + subagents.should_restore_sub_agent( + {"agent_name": "worker", "independent": False, "status": "running"} + ) + is True + ) + assert ( + subagents.should_restore_sub_agent( + {"agent_name": "worker", "independent": False, "status": "finished"} + ) + is False + ) + assert ( + subagents.should_restore_sub_agent( + {"agent_name": "worker", "independent": False, "status": "error"} + ) + is False + ) + assert ( + subagents.should_restore_sub_agent( + {"agent_name": "memorizer", "independent": True, "status": "running"} + ) + is False + ) + assert subagents.should_restore_sub_agent({"agent_name": None}) is False + + +def test_sub_agent_state_reads_service(monkeypatch): + """The stored lifecycle flags are read back from the AgentService.""" + from types import SimpleNamespace + + from cecli.helpers.agents.service import AgentService, SubAgentStatus + from cecli.helpers.sessions import subagents + + info = SimpleNamespace(independent=False, status=SubAgentStatus.RUNNING) + service = SimpleNamespace(sub_agents={"sub1": info}) + monkeypatch.setattr(AgentService, "get_instance", classmethod(lambda cls, coder: service)) + + class _Coder: + uuid = "sub1" + + assert subagents.sub_agent_state(_Coder()) == (False, "running") + + class _Unknown: + uuid = "sub9" + + assert subagents.sub_agent_state(_Unknown()) == (False, None) + + +def test_build_payload_records_sub_agent_state(mock_coder, monkeypatch): + """build_payload stores independence and status for sub-agents.""" + from types import SimpleNamespace + + from cecli.helpers.agents.service import AgentService, SubAgentStatus + from cecli.helpers.sessions import payload as payload_module + + mock_coder.uuid = "sub1" + info = SimpleNamespace(independent=True, status=SubAgentStatus.FINISHED) + service = SimpleNamespace(sub_agents={"sub1": info}) + monkeypatch.setattr(AgentService, "get_instance", classmethod(lambda cls, coder: service)) + + data = payload_module.build_payload(mock_coder, mock_coder.io, "n", agent_name="worker") + assert data["independent"] is True + assert data["status"] == "finished" + + primary = payload_module.build_payload(mock_coder, mock_coder.io, "n") + assert primary["independent"] is None + assert primary["status"] is None + + 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) diff --git a/tests/basic/test_system_one.py b/tests/basic/test_system_one.py new file mode 100644 index 00000000000..19d79e8781a --- /dev/null +++ b/tests/basic/test_system_one.py @@ -0,0 +1,804 @@ +"""Tests for the System One helper: endpoint config, wire format, parsing. + +These cover the shapes the /v1/systemone endpoint accepts and returns, using a +stubbed transport so no server or API key is involved. +""" + +import json +import os +import tempfile + +import pytest +import yaml + +from cecli.args import get_parser +from cecli.helpers import config_utils, system_one +from cecli.hooks.helpers import HookHelpers +from cecli.main import convert_yaml_to_json_string + + +@pytest.fixture(autouse=True) +def clean_config(): + system_one.reset() + yield + system_one.reset() + + +@pytest.fixture +def transport(monkeypatch): + """Capture request payloads and reply with a canned response body.""" + + class Stub: + def __init__(self): + self.payloads = [] + self.responses = [] + + def reply(self, body): + self.responses.append(body) + + async def _post(self, payload): + self.payloads.append(payload) + + if self.responses: + return self.responses.pop(0) + + raise AssertionError("no canned response queued") + + stub = Stub() + monkeypatch.setattr(system_one.SystemOneClient, "_post", stub._post) + + return stub + + +def _noul_body(question_id="q"): + return { + "model": "von-1.3.0", + "answers": {question_id: {"type": "noul", "noul": 0.95}}, + "usage": {"input_tokens": 296, "output_tokens": 20}, + } + + +def _choice_body(question_id="decision"): + return { + "model": "von-1.3.0", + "answers": { + question_id: { + "type": "choice", + "choice": "billing", + "probabilities": {"billing": 0.88, "technical": 0.12}, + "confidence": 0.81, + } + }, + "usage": {"input_tokens": 318, "output_tokens": 34}, + } + + +def _score_body(question_id="rating"): + return { + "model": "von-1.3.0", + "answers": { + question_id: { + "type": "score", + "score": 1.05, + "legend": {"0": "Calm", "1": "Frustrated", "2": "Very angry"}, + "probabilities": {"0": 0.0, "1": 0.95, "2": 0.05}, + "confidence": 0.92, + } + }, + "usage": {"input_tokens": 304, "output_tokens": 18}, + } + + +class Coder: + """Minimal coder stand-in for the hook helper.""" + + +# -------------------------------------------------------------------------- +# endpoint configuration + + +def test_default_endpoint_is_local_and_vendor_neutral(): + config = system_one.SystemOneConfig() + + assert config.api_base == "http://127.0.0.1:8000" + assert config.endpoint_url == "http://127.0.0.1:8000/v1/systemone" + assert config.api_key_env == ["SYSTEM_ONE_API_KEY"] + + +def test_from_env_only_reads_system_one_vars(monkeypatch): + monkeypatch.setenv("SYSTEM_ONE_API_BASE", "http://decide.internal:9000") + monkeypatch.setenv("SYSTEM_ONE_MODEL", "von-1.3.0") + monkeypatch.setenv("VON_BASE_URL", "http://should-be-ignored:1") + + config = system_one.SystemOneConfig.from_env() + + assert config.api_base == "http://decide.internal:9000" + assert config.model_name == "von-1.3.0" + + +def test_from_dict_reads_documented_keys(): + config = system_one.SystemOneConfig.from_dict( + { + "api_base": "http://127.0.0.1:8000/", + "api_key_env": ["LITELLM_API_KEY"], + "model_name": "von-latest", + "extra_headers": {"optional": "a", "optional2": "b"}, + } + ) + + assert config.api_base == "http://127.0.0.1:8000" + assert config.api_key_env == ["LITELLM_API_KEY"] + assert config.model_name == "von-latest" + assert config.extra_headers == {"optional": "a", "optional2": "b"} + + +def test_from_dict_accepts_key_spelling_variants(): + config = system_one.SystemOneConfig.from_dict( + {"api-base": "http://h:1/v1", "model": "m-1", "headers": {"x-a": "1"}} + ) + + assert config.endpoint_url == "http://h:1/v1/systemone" + assert config.model_name == "m-1" + assert config.extra_headers == {"x-a": "1"} + + +def test_string_api_key_env_is_promoted_to_list(): + config = system_one.SystemOneConfig.from_dict({"api_key_env": "MY_KEY"}) + + assert config.api_key_env == ["MY_KEY"] + + +def test_api_key_and_headers_are_assembled(monkeypatch): + monkeypatch.delenv("SYSTEM_ONE_API_KEY", raising=False) + + config = system_one.SystemOneConfig.from_dict({"extra_headers": {"x-trace": "1"}}) + + assert "Authorization" not in config.build_headers() + assert config.build_headers()["x-trace"] == "1" + + monkeypatch.setenv("SYSTEM_ONE_API_KEY", "secret") + + assert config.build_headers()["Authorization"] == "Bearer secret" + + +def test_configured_key_env_names_are_the_only_ones_probed(monkeypatch): + monkeypatch.setenv("LITELLM_API_KEY", "lit") + monkeypatch.setenv("SYSTEM_ONE_API_KEY", "one") + + config = system_one.SystemOneConfig.from_dict({"api_key_env": ["LITELLM_API_KEY"]}) + + assert config.resolve_api_key() == "lit" + assert system_one.SystemOneConfig.from_env().resolve_api_key() == "one" + + +def test_configure_installs_and_reset_falls_back_to_env(monkeypatch): + monkeypatch.setenv("SYSTEM_ONE_API_BASE", "http://env-host:8000") + + assert not system_one.is_configured() + assert system_one.get_config().api_base == "http://env-host:8000" + + assert system_one.configure({"api_base": "http://cli-host:8000"}) is not None + assert system_one.get_config().api_base == "http://cli-host:8000" + + # An empty payload (the {} a config-file merge can leave behind) is a no-op. + assert system_one.configure({}) is None + assert system_one.get_config().api_base == "http://cli-host:8000" + + system_one.reset() + + assert system_one.get_config().api_base == "http://env-host:8000" + + +# -------------------------------------------------------------------------- +# wire format: questions + + +def test_noul_question_shape(): + assert system_one.noul("Is this urgent?") == {"type": "noul", "instructions": "Is this urgent?"} + + +def test_noul_question_with_criteria(): + question = system_one.noul("Is this urgent?", true="Time sensitive", false="No urgency") + + assert question["criteria"] == {"true": "Time sensitive", "false": "No urgency"} + + +def test_choice_question_from_map_and_from_names(): + from_map = system_one.choice("Who owns this?", {"billing": "Payments", "infra": "Bugs"}) + + assert from_map["type"] == "choice" + assert from_map["criteria"] == {"billing": "Payments", "infra": "Bugs"} + + from_names = system_one.choice("Who owns this?", ["billing", "infra"]) + + assert from_names["criteria"] == {"billing": None, "infra": None} + + +def test_score_question_keeps_level_order(): + question = system_one.score("How frustrated?", ["Calm", "Frustrated", "Very angry"]) + + assert question == { + "type": "score", + "instructions": "How frustrated?", + "criteria": ["Calm", "Frustrated", "Very angry"], + } + + +def test_instructions_may_be_structured(): + instructions = { + "potential_duplicate": {"name": "John Smith"}, + "question": "Is the resume for the same person as `potential_duplicate`?", + } + question = system_one.noul(instructions) + + assert question["instructions"] == instructions + + +def test_build_payload_carries_state_model_and_questions(): + payload = system_one.build_payload( + {"ticket": "INC-1"}, {"a": system_one.noul("urgent?")}, default_model="von-1.3.0" + ) + + assert payload["state"] == {"ticket": "INC-1"} + assert payload["model"] == "von-1.3.0" + assert payload["questions"] == {"a": {"type": "noul", "instructions": "urgent?"}} + + +# -------------------------------------------------------------------------- +# wire format: answers + + +def test_parse_noul_answer(): + response = system_one.parse_response(_noul_body()) + answer = response.answers["q"] + + assert isinstance(answer, system_one.NoulAnswer) + assert answer.probability == 0.95 + assert answer.yes is True + assert answer.to_dict() == {"type": "noul", "noul": 0.95, "yes": True} + assert response.model == "von-1.3.0" + assert response.usage == {"input_tokens": 296, "output_tokens": 20} + + +def test_parse_choice_answer(): + answer = system_one.parse_response(_choice_body()).answers["decision"] + + assert isinstance(answer, system_one.ChoiceAnswer) + assert answer.choice == "billing" + assert answer.probabilities == {"billing": 0.88, "technical": 0.12} + assert answer.confidence == 0.81 + assert set(answer.to_dict()) == {"type", "choice", "probabilities", "confidence"} + + +def test_parse_score_answer_and_levels(): + answer = system_one.parse_response(_score_body()).answers["rating"] + + assert isinstance(answer, system_one.ScoreAnswer) + assert answer.score == 1.05 + assert answer.levels == ["Calm", "Frustrated", "Very angry"] + assert answer.probabilities == {"0": 0.0, "1": 0.95, "2": 0.05} + assert answer.to_dict()["legend"]["2"] == "Very angry" + + +def test_missing_numeric_fields_default_to_zero(): + answer = system_one.parse_answer({"type": "choice", "choice": "billing"}) + + assert answer.choice == "billing" + assert answer.probabilities == {} + assert answer.confidence == 0.0 + + +def test_non_numeric_values_do_not_crash_parsing(): + answer = system_one.parse_answer({"type": "noul", "noul": "0.7"}) + + assert answer.probability == 0.7 + + +def test_unknown_answer_type_is_rejected(): + with pytest.raises(ValueError, match="Unknown System One answer type"): + system_one.parse_answer({"type": "wizard", "wizard": 1}) + + +def test_non_object_answer_is_rejected(): + with pytest.raises(ValueError, match="must be an object"): + system_one.parse_answer("noul") + + +def test_response_is_lookup_friendly(): + response = system_one.parse_response(_noul_body("urgent")) + + assert "urgent" in response + assert response["urgent"].probability == 0.95 + assert response.get("missing") is None + assert response.to_dict()["answers"]["urgent"]["type"] == "noul" + + +# -------------------------------------------------------------------------- +# application api: decide / judge / rate / ask + + +async def test_decide_sends_choice_question_and_parses(transport): + transport.reply(_choice_body()) + + answer = await system_one.decide( + state="Disk volume /var/log at 98%.", + choices={"sre": "Infrastructure", "app": "Application code"}, + instructions="Who owns this?", + ) + + question = transport.payloads[0]["questions"]["decision"] + + assert question["type"] == "choice" + assert question["criteria"] == {"sre": "Infrastructure", "app": "Application code"} + assert answer.choice == "billing" + assert transport.payloads[0]["model"] == system_one.get_config().model_name + + +async def test_judge_sends_noul_question_with_criteria(transport): + transport.reply(_noul_body("judgment")) + + answer = await system_one.judge( + state="Connection pool exhausted.", + instructions="Is this blocking customers?", + criteria={"true": "Requests fail", "false": "Internal only"}, + ) + + question = transport.payloads[0]["questions"]["judgment"] + + assert question["type"] == "noul" + assert question["criteria"] == {"true": "Requests fail", "false": "Internal only"} + assert answer.probability == 0.95 + + +async def test_judge_ignores_unrecognized_criteria_keys(transport): + transport.reply(_noul_body("judgment")) + + await system_one.judge(state="x", instructions="y?", criteria={"true": "a", "extra": 1}) + + assert transport.payloads[0]["questions"]["judgment"]["criteria"] == {"true": "a"} + + +async def test_judge_accepts_a_true_false_pair(transport): + transport.reply(_noul_body("judgment")) + + await system_one.judge(state="x", instructions="y?", criteria=["blocking", "internal only"]) + + assert transport.payloads[0]["questions"]["judgment"]["criteria"] == { + "true": "blocking", + "false": "internal only", + } + + +async def test_judge_without_criteria_omits_the_field(transport): + transport.reply(_noul_body("judgment")) + + await system_one.judge(state="x", instructions="Is this true?") + + assert "criteria" not in transport.payloads[0]["questions"]["judgment"] + + +async def test_rate_sends_score_question(transport): + transport.reply(_score_body()) + + answer = await system_one.rate( + state="Memory at 98%.", + levels=["Nominal", "Degraded", "Critical"], + instructions="Assess degradation.", + ) + + assert transport.payloads[0]["questions"]["rating"]["criteria"] == [ + "Nominal", + "Degraded", + "Critical", + ] + assert answer.score == 1.05 + + +async def test_ask_evaluates_several_questions_in_one_request(transport): + transport.reply( + { + "model": "von-1.3.0", + "answers": { + "intent": {"type": "choice", "choice": "payment_failure", "probabilities": {}}, + "is_urgent": {"type": "noul", "noul": 0.9}, + }, + "usage": {"input_tokens": 300, "output_tokens": 20}, + } + ) + + response = await system_one.ask( + state={"ticket": "INC-4091", "message": "Gateway timeouts."}, + questions={ + "intent": system_one.choice( + "Nature of the ticket?", {"payment_failure": "Charges fail"} + ), + "is_urgent": system_one.noul("Needs immediate intervention?"), + }, + ) + + assert len(transport.payloads) == 1 + assert sorted(transport.payloads[0]["questions"]) == ["intent", "is_urgent"] + assert response["intent"].choice == "payment_failure" + assert response["is_urgent"].probability == 0.9 + + +async def test_system_one_call_matches_ask(transport): + transport.reply(_noul_body("is_urgent")) + + response = await system_one.system_one( + state="Payouts failing.", + questions={"is_urgent": system_one.noul("Urgent?")}, + ) + + assert response["is_urgent"].probability == 0.95 + assert response.model == "von-1.3.0" + + +async def test_model_override_reaches_the_payload(transport): + transport.reply(_noul_body("judgment")) + + await system_one.judge(state="x", instructions="y?", model="custom-model") + + assert transport.payloads[0]["model"] == "custom-model" + + +async def test_request_uses_configured_endpoint_and_headers(monkeypatch): + import httpx + + monkeypatch.setenv("SYSTEM_ONE_API_KEY", "secret") + system_one.configure( + { + "api_base": "http://decide.test:8000", + "model_name": "von-1.3.0", + "extra_headers": {"x-a": "b"}, + } + ) + + seen = {} + _fake_httpx_client(monkeypatch, httpx, seen, [httpx.Response(200, json=_noul_body("judgment"))]) + + answer = await system_one.judge(state="x", instructions="blocking?") + + assert seen["url"] == "http://decide.test:8000/v1/systemone" + assert seen["headers"]["Authorization"] == "Bearer secret" + assert seen["headers"]["x-a"] == "b" + assert seen["json"]["model"] == "von-1.3.0" + assert answer.probability == 0.95 + + +async def test_rate_limited_responses_are_retried(monkeypatch): + import asyncio + + import httpx + + async def _no_sleep(*args, **kwargs): + return None + + system_one.configure({"api_base": "http://decide.test:8000", "max_retries": 2}) + monkeypatch.setattr(asyncio, "sleep", _no_sleep) + + seen = {"count": 0} + responses = [ + httpx.Response(429, text="slow down"), + httpx.Response(529, text="overloaded"), + httpx.Response(200, json=_noul_body("judgment")), + ] + + _fake_httpx_client(monkeypatch, httpx, seen, responses) + + answer = await system_one.judge(state="x", instructions="blocking?") + + assert answer.probability == 0.95 + assert seen["count"] == 3 + + +async def test_missing_api_key_is_named_in_auth_errors(monkeypatch): + import httpx + + monkeypatch.delenv("SYSTEM_ONE_API_KEY", raising=False) + monkeypatch.delenv("LITELLM_API_KEY", raising=False) + system_one.configure( + { + "api_base": "http://decide.test:8000", + "api_key_env": ["SYSTEM_ONE_API_KEY", "LITELLM_API_KEY"], + "max_retries": 0, + } + ) + + _fake_httpx_client( + monkeypatch, + httpx, + {"count": 0}, + [httpx.Response(403, json={"detail": "Must supply an API key"})], + ) + + with pytest.raises(system_one.SystemOneError) as excinfo: + await system_one.judge(state="x", instructions="y?") + + message = str(excinfo.value) + + assert "SYSTEM_ONE_API_KEY, LITELLM_API_KEY" in message + assert "No API key was sent" in message + + +async def test_error_status_raises_system_one_error(monkeypatch): + import httpx + + system_one.configure({"api_base": "http://decide.test:8000", "max_retries": 0}) + + _fake_httpx_client( + monkeypatch, + httpx, + {"count": 0}, + [httpx.Response(422, json={"detail": "missing field: questions"})], + ) + + with pytest.raises(system_one.SystemOneError) as excinfo: + await system_one.judge(state="x", instructions="y?") + + assert excinfo.value.status_code == 422 + assert "missing field" in str(excinfo.value) + + +def _fake_httpx_client(monkeypatch, httpx, seen, responses): + """Replace httpx.AsyncClient with one that replays *responses* in order.""" + + class FakeClient: + def __init__(self, *args, **kwargs): + seen.setdefault("init", kwargs) + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc): + return False + + async def post(self, url, json=None, headers=None): + seen["url"] = url + seen["json"] = json + seen["headers"] = headers + seen["count"] = seen.get("count", 0) + 1 + + return responses.pop(0) + + monkeypatch.setattr(httpx, "AsyncClient", FakeClient) + + return FakeClient + + +# -------------------------------------------------------------------------- +# hook helper + + +async def test_hook_helper_questions_form_returns_dicts(transport): + transport.reply(_noul_body("is_urgent")) + + result = await HookHelpers.system_one( + Coder(), + state="Help! My payouts have been failing for 3 days.", + questions={"is_urgent": "Does this convey urgency?"}, + ) + + assert result == { + "model": "von-1.3.0", + "usage": {"input_tokens": 296, "output_tokens": 20}, + "answers": {"is_urgent": {"type": "noul", "noul": 0.95, "yes": True}}, + } + + +async def test_hook_helper_decide_form(transport): + transport.reply(_choice_body()) + + result = await HookHelpers.system_one( + Coder(), + state="Database replication lag exceeded 45 seconds.", + decide={ + "choices": {"infrastructure": "Hardware or network", "billing": "Invoices"}, + "instructions": "Classify the root cause domain.", + }, + ) + + decision = result["answers"]["decision"] + + assert decision["type"] == "choice" + assert decision["choice"] == "billing" + assert decision["confidence"] == 0.81 + assert transport.payloads[0]["questions"]["decision"]["type"] == "choice" + + +async def test_hook_helper_judge_and_rate_forms(transport): + transport.reply(_noul_body("judgment")) + transport.reply(_score_body()) + + judged = await HookHelpers.system_one( + Coder(), + state="x", + judge={"instructions": "Blocking?", "criteria": {"true": "yes means blocking"}}, + ) + rated = await HookHelpers.system_one( + Coder(), state="x", rate={"levels": ["Low", "High"], "instructions": "Severity?"} + ) + + assert judged["answers"]["judgment"]["type"] == "noul" + assert rated["answers"]["rating"]["type"] == "score" + assert rated["answers"]["rating"]["score"] == 1.05 + assert transport.payloads[1]["questions"]["rating"]["criteria"] == ["Low", "High"] + + +async def test_hook_helper_rate_accepts_a_bare_level_list(transport): + transport.reply(_score_body()) + + await HookHelpers.system_one(Coder(), state="x", rate=["Nominal", "Degraded", "Critical"]) + + assert transport.payloads[0]["questions"]["rating"]["criteria"] == [ + "Nominal", + "Degraded", + "Critical", + ] + + +async def test_hook_helper_defaults_state_to_the_conversation(transport, monkeypatch): + transport.reply(_noul_body("judgment")) + + messages = [{"role": "user", "content": "payouts are failing"}] + monkeypatch.setattr(HookHelpers, "get_messages", staticmethod(lambda coder, last_n: messages)) + + await HookHelpers.system_one(Coder(), judge="Urgent?") + + assert transport.payloads[0]["state"] == messages + + +async def test_hook_helper_requires_exactly_one_form(transport): + with pytest.raises(ValueError, match="exactly one of"): + await HookHelpers.system_one(Coder(), state="x") + + with pytest.raises(ValueError, match="exactly one of"): + await HookHelpers.system_one(Coder(), state="x", judge="a?", decide={"choices": ["a"]}) + + +async def test_hook_helper_rejects_a_bad_question_value(transport): + with pytest.raises(TypeError, match="must be a dict or instructions string"): + await HookHelpers.system_one(Coder(), state="x", questions={"bad": 17}) + + +# -------------------------------------------------------------------------- +# integrations: the shorthand vocabulary subsystems call through + + +def test_build_questions_maps_each_form_to_the_wire_format(): + from cecli.helpers.system_one import integrations + + assert integrations.build_questions(questions={"a": system_one.noul("q?")}) == { + "a": {"type": "noul", "instructions": "q?"} + } + assert integrations.build_questions(decide={"choices": {"x": "X"}}) == { + "decision": { + "type": "choice", + "instructions": integrations.DEFAULT_DECIDE_INSTRUCTIONS, + "criteria": {"x": "X"}, + } + } + assert integrations.build_questions(judge="blocking?") == { + "judgment": {"type": "noul", "instructions": "blocking?"} + } + assert integrations.build_questions(rate=["Low", "High"]) == { + "rating": { + "type": "score", + "instructions": integrations.DEFAULT_RATE_INSTRUCTIONS, + "criteria": ["Low", "High"], + } + } + + +def test_build_questions_expands_string_shorthands(): + from cecli.helpers.system_one import integrations + + built = integrations.build_questions(questions={"is_urgent": "Urgent?"}) + + assert built == {"is_urgent": {"type": "noul", "instructions": "Urgent?"}} + + +def test_build_questions_drops_unrecognized_judge_criteria(): + from cecli.helpers.system_one import integrations + + built = integrations.build_questions( + judge={"instructions": "y?", "criteria": {"true": "a", "maybe": "b"}} + ) + + assert built["judgment"]["criteria"] == {"true": "a"} + + +def test_build_questions_requires_exactly_one_form(): + from cecli.helpers.system_one import integrations + + with pytest.raises(ValueError, match="exactly one of"): + integrations.build_questions() + + with pytest.raises(ValueError, match="exactly one of"): + integrations.build_questions(judge="a?", decide=["a", "b"]) + + +def test_build_questions_rejects_a_bad_question_value(): + from cecli.helpers.system_one import integrations + + with pytest.raises(TypeError, match="must be a dict or instructions string"): + integrations.build_questions(questions={"bad": 17}) + + +async def test_evaluate_returns_plain_dicts(transport): + from cecli.helpers.system_one import integrations + + transport.reply(_choice_body()) + + result = await integrations.evaluate("some state", decide={"choices": ["a", "b"]}) + + assert isinstance(result, dict) + assert set(result) == {"model", "usage", "answers"} + assert result["answers"]["decision"]["type"] == "choice" + assert isinstance(result["answers"]["decision"]["probabilities"], dict) + + +# -------------------------------------------------------------------------- +# --system-one / config-file plumbing + + +def _resolve_system_one_arg(tmp_path, file_config, cli_json): + """Replicate the config-file -> parser -> deep-merge path main_async uses.""" + conf = tmp_path / ".cecli.conf.yml" + conf.write_text("system-one: |\n " + json.dumps(file_config) + "\n") + paths = [str(conf)] + merged_config = config_utils.read_and_merge_all_configs(paths, [], paths) + + fd, tmp = tempfile.mkstemp(suffix=".yml", prefix="cecli_merged_") + os.close(fd) + with open(tmp, "w") as handle: + yaml.dump(merged_config, handle) + + try: + parser = get_parser([tmp], None) + argv = [f"--system-one={cli_json}"] if cli_json else [] + args, _ = parser.parse_known_args(argv) + finally: + os.unlink(tmp) + + return convert_yaml_to_json_string(args.system_one, merged_config.get("system-one")) + + +def test_config_file_supplies_the_endpoint(tmp_path): + resolved = _resolve_system_one_arg( + tmp_path, + {"api_base": "http://file-host:8000", "model_name": "von-1.3.0"}, + None, + ) + + assert system_one.configure(json.loads(resolved)) is not None + config = system_one.get_config() + + assert config.api_base == "http://file-host:8000" + assert config.model_name == "von-1.3.0" + + +def test_cli_system_one_deep_merges_over_the_config_file(tmp_path): + resolved = _resolve_system_one_arg( + tmp_path, + {"api_base": "http://file-host:8000", "api_key_env": ["FILE_KEY"]}, + json.dumps({"api_base": "http://cli-host:9000"}), + ) + + assert system_one.configure(json.loads(resolved)) is not None + config = system_one.get_config() + + assert config.api_base == "http://cli-host:9000" # CLI wins per key + assert config.api_key_env == ["FILE_KEY"] # file-only key preserved + + +def test_extra_headers_survive_the_config_round_trip(tmp_path): + resolved = _resolve_system_one_arg( + tmp_path, + {"api_base": "http://h:1", "extra_headers": {"optional": "a"}}, + json.dumps({"extra_headers": {"optional2": "b"}}), + ) + + system_one.configure(json.loads(resolved)) + + assert system_one.get_config().extra_headers == {"optional": "a", "optional2": "b"} diff --git a/tests/helpers/test_llms_ollama_native.py b/tests/helpers/test_llms_ollama_native.py new file mode 100644 index 00000000000..83cb70b6e06 --- /dev/null +++ b/tests/helpers/test_llms_ollama_native.py @@ -0,0 +1,295 @@ +"""Ollama native /api/chat wire tests. + +Ollama's OpenAI-compatible ``/v1/chat/completions`` endpoint ignores runner +options such as ``num_ctx`` (and ``keep_alive``), so Ollama is implemented as a +provider adapter that keeps the shared ``chat`` family but overrides its chat +hooks to speak the native ``/api/chat`` wire (options under ``options``). These +tests lock in: + +- ``ollama`` / ``ollama_chat`` resolve to the ``chat`` family but their own + provider adapter +- the native base/URL and payload mapping (api-block params -> ``options``) +- native responses/chunks normalize into the shared cecli types +- the litellm shim forwards extra params instead of dropping them + +No network: the payload builders, normalizers and URL/body construction are +exercised offline. +""" + +import asyncio +from unittest.mock import patch + +import cecli.helpers.llms as llms_pkg +from cecli.helpers.llms.config import resolve_model_config +from cecli.helpers.llms.domains.chat import chat_complete +from cecli.helpers.llms.litellm_compat import litellm +from cecli.helpers.llms.providers import get_provider_adapter +from cecli.helpers.llms.providers.ollama import ( + OllamaChatProvider, + OllamaProvider, + normalize_ollama_response, + ollama_native_base, + ollama_payload, + parse_ollama_chunk, +) +from cecli.helpers.llms.types import CompletionResponse + +MSGS = [{"role": "user", "content": "hi"}] + + +def test_ollama_providers_use_chat_family_with_dedicated_adapters(): + for slug, adapter_type in (("ollama", OllamaProvider), ("ollama_chat", OllamaChatProvider)): + resolved = resolve_model_config(f"{slug}/llama3") + + assert resolved["family"] == "chat" + assert resolved["provider"] == slug + assert isinstance(get_provider_adapter(slug), adapter_type) + + +def test_native_base_strips_openai_compat_suffix(): + assert ollama_native_base("http://localhost:11434/v1") == "http://localhost:11434" + assert ollama_native_base("http://host:1234/") == "http://host:1234" + assert ollama_native_base(None) == "http://localhost:11434" + + +def test_payload_maps_api_params_to_native_fields(): + resolved = resolve_model_config("ollama_chat/llama3") + payload = ollama_payload( + resolved, + MSGS, + None, + False, + {"num_ctx": 57000, "keep_alive": -1, "temperature": 0, "max_tokens": 4096}, + ) + assert payload["model"] == "llama3" + assert payload["options"] == {"temperature": 0, "num_ctx": 57000, "num_predict": 4096} + assert payload["keep_alive"] == -1 + assert "max_tokens" not in payload + assert "stream" not in payload["options"] + + +def test_payload_drops_control_kwargs(): + resolved = resolve_model_config("ollama_chat/llama3") + payload = ollama_payload( + resolved, + MSGS, + None, + False, + { + "custom_llm_provider": "ollama_chat", + "base_url": "http://ignored/v1", + "drop_params": True, + "allowed_openai_params": ["tools"], + "parallel_tool_calls": True, + "num_ctx": 123, + }, + ) + assert payload["options"] == {"num_ctx": 123} + + +def test_payload_maps_reasoning_effort_to_think(): + resolved = resolve_model_config("ollama_chat/llama3") + + def think(effort=None, thinking=None): + kwargs = {} + + if effort is not None: + kwargs["reasoning_effort"] = effort + + if thinking is not None: + kwargs["thinking"] = thinking + + return ollama_payload(resolved, MSGS, None, False, kwargs).get("think") + + assert think(effort="low") == "low" + assert think(effort="medium") == "medium" + assert think(effort="high") == "high" + assert think(effort="max") == "high" + assert think(effort="none") is False + assert think(effort="minimal") is None + assert think(thinking={"type": "enabled", "budget_tokens": 1024}) is True + assert think(thinking={"type": "disabled"}) is False + assert think(thinking=False) is False + assert think(thinking="high") == "high" + + +def test_native_messages_carry_reasoning_and_tool_calls(): + resolved = resolve_model_config("ollama_chat/llama3") + messages = [ + { + "role": "assistant", + "content": "ok", + "reasoning_content": "why", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "f", "arguments": '{"a": 1}'}, + } + ], + }, + {"role": "tool", "content": "1", "tool_call_id": "call_1"}, + ] + payload = ollama_payload(resolved, messages, None, False, {}) + assistant = payload["messages"][0] + assert assistant["thinking"] == "why" + assert assistant["tool_calls"] == [{"function": {"name": "f", "arguments": {"a": 1}}}] + assert payload["messages"][1]["tool_call_id"] == "call_1" + + +def test_normalize_response_and_chunk_usage(): + data = { + "message": {"role": "assistant", "content": "hi", "thinking": "t"}, + "done": True, + "done_reason": "stop", + "prompt_eval_count": 3, + "eval_count": 1, + } + response = normalize_ollama_response(data, "ollama_chat/llama3") + assert response.text == "hi" + assert response.reasoning == "t" + assert response.usage.total_tokens == 4 + + chunk = parse_ollama_chunk({"message": {"role": "assistant", "content": "x"}, "done": False}) + assert chunk.text == "x" + assert chunk.finish_reason is None + + final = parse_ollama_chunk( + { + "message": {"role": "assistant", "content": ""}, + "done": True, + "done_reason": "stop", + "prompt_eval_count": 2, + "eval_count": 0, + } + ) + assert final.finish_reason == "stop" + assert final.usage.prompt_tokens == 2 + + +def test_streamed_parallel_tool_calls_keep_distinct_indices(): + adapter = get_provider_adapter("ollama_chat") + first = adapter.parse_chat_chunk( + {"message": {"content": "", "tool_calls": [{"function": {"name": "a", "arguments": {}}}]}} + ) + second = adapter.parse_chat_chunk( + { + "message": { + "content": "", + "tool_calls": [{"function": {"name": "b", "arguments": {}}}], + }, + "done": True, + } + ) + + assert [tc.id for tc in first.tool_calls] == ["call_0"] + assert [tc.id for tc in second.tool_calls] == ["call_1"] + assert [tc.index for tc in second.tool_calls] == [1] + + +def test_parallel_tool_calls_in_one_chunk_get_distinct_indices(): + adapter = get_provider_adapter("ollama_chat") + chunk = adapter.parse_chat_chunk( + { + "message": { + "content": "", + "tool_calls": [ + {"function": {"name": "a", "arguments": {}}}, + {"function": {"name": "b", "arguments": {}}}, + ], + } + } + ) + + assert [tc.index for tc in chunk.tool_calls] == [0, 1] + assert [tc.id for tc in chunk.tool_calls] == ["call_0", "call_1"] + + +def test_chat_domain_posts_to_native_endpoint_via_provider(): + resolved = resolve_model_config("ollama_chat/llama3") + resolved["_provider"] = get_provider_adapter("ollama_chat") + captured = {} + + class _Resp: + def raise_for_status(self): + pass + + def json(self): + return {"message": {"role": "assistant", "content": "hi"}, "done": True} + + class _Client: + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + async def post(self, url, headers=None, params=None, json=None): + captured["url"] = url + captured["json"] = json + return _Resp() + + with patch("cecli.helpers.llms.domains.chat.make_client", return_value=_Client()): + response = asyncio.run(chat_complete(resolved, MSGS, None, None, {}, {"num_ctx": 57000})) + + assert captured["url"] == "http://localhost:11434/api/chat" + assert captured["json"]["options"] == {"num_ctx": 57000} + assert response.text == "hi" + + +def test_shim_forwards_extra_params(monkeypatch): + captured = {} + + async def fake_dispatch(**kwargs): + captured.update(kwargs) + return CompletionResponse(model=kwargs.get("model")) + + monkeypatch.setattr(llms_pkg, "acompletion", fake_dispatch) + + asyncio.run( + litellm.acompletion( + model="ollama_chat/llama3", + messages=MSGS, + stream=False, + num_ctx=57000, + keep_alive=-1, + top_p=0.9, + drop_params=True, + base_url="http://ignored/v1", + ) + ) + assert captured["num_ctx"] == 57000 + assert captured["keep_alive"] == -1 + assert captured["top_p"] == 0.9 + assert "drop_params" not in captured + assert "base_url" not in captured + + +def test_shim_keeps_narrow_passthrough_for_non_ollama(monkeypatch): + captured = {} + + async def fake_dispatch(**kwargs): + captured.update(kwargs) + return CompletionResponse(model=kwargs.get("model")) + + monkeypatch.setattr(llms_pkg, "acompletion", fake_dispatch) + + asyncio.run( + litellm.acompletion( + model="openai/gpt-4o", + messages=MSGS, + stream=False, + temperature=0.5, + top_p=0.9, + num_ctx=57000, + keep_alive=-1, + drop_params=True, + base_url="http://ignored/v1", + ) + ) + assert captured["temperature"] == 0.5 + assert "top_p" not in captured + assert "num_ctx" not in captured + assert "keep_alive" not in captured + assert "drop_params" not in captured + assert "base_url" not in captured diff --git a/tests/mcp/test_tool_name_sanitization.py b/tests/mcp/test_tool_name_sanitization.py new file mode 100644 index 00000000000..83fc45b710a --- /dev/null +++ b/tests/mcp/test_tool_name_sanitization.py @@ -0,0 +1,138 @@ +"""Tests for provider-safe MCP tool names. + +Providers enforce ``^[A-Za-z0-9_-]{1,64}$`` on ``tools[].function.name``, so a +dotted MCP name such as ``browser.fetch`` is sanitized on the way to the model +and mapped back to the server's own name when the call comes back. +""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from cecli.coders import Coder +from cecli.helpers import responses + + +def _tool(name): + """Build a minimal OpenAI-style function tool dict.""" + return { + "type": "function", + "function": {"name": name, "description": "", "parameters": {}}, + } + + +def _call(name): + return { + "id": "call-1", + "type": "function", + "function": {"name": name, "arguments": "{}"}, + } + + +class TestSanitizeToolName: + def test_dots_become_underscores(self): + assert responses.sanitize_tool_name("browser.fetch") == "browser_fetch" + + def test_is_idempotent(self): + once = responses.sanitize_tool_name("browser.persisted_session.list") + + assert responses.sanitize_tool_name(once) == once + + def test_honors_provider_length_limit(self): + assert len(responses.sanitize_tool_name("a" * 200)) == 64 + + +class TestRegisteredToolNames: + def test_registered_name_maps_back(self): + responses.register_tool_names([("browser", [_tool("browser.persisted_session.list")])]) + + assert ( + responses.original_tool_name("browser_persisted_session_list") + == "browser.persisted_session.list" + ) + + def test_unknown_name_is_unchanged(self): + assert responses.original_tool_name("never_registered") == "never_registered" + + +class TestGetToolListSanitizesNames: + def _coder(self, tools): + return SimpleNamespace( + mcp_tools=[("browser", tools)], + registered_servers={"included": set(), "excluded": set()}, + registered_tools={"included": set(), "excluded": set()}, + ) + + def test_dotted_names_are_sanitized_for_the_model(self): + coder = self._coder([_tool("browser.fetch"), _tool("browser.persisted_session.list")]) + + names = [t["function"]["name"] for t in Coder.get_tool_list(coder)] + + assert names == [ + "browser--browser_fetch", + "browser--browser_persisted_session_list", + ] + + def test_underscore_name_is_unchanged(self): + coder = self._coder([_tool("brave_web_search")]) + + names = [t["function"]["name"] for t in Coder.get_tool_list(coder)] + + assert names == ["browser--brave_web_search"] + + +class TestFindMcpServerForSanitizedCall: + def test_sanitized_call_resolves_to_its_server(self): + server = SimpleNamespace(name="browser") + coder = SimpleNamespace( + mcp_tools=[("browser", [_tool("browser.fetch")])], + mcp_manager=[server], + ) + tool_call = SimpleNamespace( + id="call-1", + type="function", + function=SimpleNamespace(name="browser--browser_fetch"), + ) + + assert Coder._find_mcp_server_for_tool(coder, tool_call) is server + + +class TestCallToolMapsBackToAdvertisedName: + """The name the model calls must reach the server unsanitized.""" + + def _coder(self): + coder = SimpleNamespace(mcp_tools=[("browser", [_tool("browser.fetch")])]) + responses.register_tool_names(coder.mcp_tools) + + return coder + + @pytest.mark.asyncio + async def test_sanitized_name_is_mapped_back(self): + coder = self._coder() + session = MagicMock() + session.call_tool = AsyncMock(return_value="ok") + + await Coder.call_mcp_tool_from_session(coder, session, _call("browser_fetch")) + + assert session.call_tool.await_args.kwargs["name"] == "browser.fetch" + + @pytest.mark.asyncio + async def test_unsanitized_name_passes_through(self): + coder = self._coder() + session = MagicMock() + session.call_tool = AsyncMock(return_value="ok") + + await Coder.call_mcp_tool_from_session(coder, session, _call("browser.fetch")) + + assert session.call_tool.await_args.kwargs["name"] == "browser.fetch" + + @pytest.mark.asyncio + async def test_unknown_name_is_not_rewritten(self): + coder = self._coder() + session = MagicMock() + session.call_tool = AsyncMock(return_value="ok") + + await Coder.call_mcp_tool_from_session(coder, session, _call("read_file")) + + assert session.call_tool.await_args.kwargs["name"] == "read_file" diff --git a/tests/tui/test_completion_bar.py b/tests/tui/test_completion_bar.py index 6939d280147..d704e5256f6 100644 --- a/tests/tui/test_completion_bar.py +++ b/tests/tui/test_completion_bar.py @@ -1,4 +1,6 @@ +import os import sys +from unittest import mock import pytest @@ -6,6 +8,10 @@ IS_WINDOWS = sys.platform == "win32" +# Capture the real relpath before any test patches os.path.relpath, so fake +# implementations can delegate to it without recursing into the mock. +original_relpath = os.path.relpath + @pytest.mark.skipif(IS_WINDOWS, reason="POSIX-only path separators") def test_absolute_path_suggestions_stay_absolute(): @@ -63,3 +69,69 @@ def test_windows_relative_path_suggestions_kept(): assert bar.suggestions == suggestions assert bar._display_names == suggestions + + +def test_relpath_cross_drive_falls_back(): + """os.path.relpath() raising ValueError (Windows cross-drive) must not crash. + + Simulates the reported crash: C:-relative suggestions mixed with an + absolute path on another drive (E:), i.e. + "ValueError: path is on mount 'E:', start on mount 'C:'". + """ + + def fake_relpath(path, start=None): + if "E:\\" in path: + raise ValueError("path is on mount 'E:', start on mount 'C:'") + return original_relpath(path) + + with mock.patch("os.path.relpath", side_effect=fake_relpath): + # Must not raise; the cross-drive suggestion is displayed as-is. + bar = CompletionBar( + suggestions=[ + "../.cecli/rules.md", + "E:\\My_Mods\\data\\file1.txt", + ], + prefix="/drop ", + ) + bar._compute_display_names() + + assert "E:\\My_Mods\\data\\file1.txt" in bar._display_names + + +def test_commonpath_mixed_drives_falls_back(): + """os.path.commonpath() raising ValueError (mixed drives) must not crash. + + The same Windows limitation hits the common prefix step one line after + relpath; the bar must fall back to showing candidates as-is. + """ + with ( + mock.patch("os.path.commonpath", side_effect=ValueError("Can't mix paths")), + mock.patch("os.path.relpath", side_effect=lambda path, start=None: path), + ): + # Must not raise at the commonpath step either. + bar = CompletionBar( + suggestions=[ + "E:\\My_Mods\\data\\file1.txt", + "../other/file2.txt", + ], + prefix="/drop ", + ) + bar._compute_display_names() + + assert "E:\\My_Mods\\data\\file1.txt" in bar._display_names + + +def test_same_directory_suggestions_compress(): + """Suggestions in one directory still collapse to a shared prefix + basenames.""" + sep = os.sep + bar = CompletionBar( + suggestions=[ + "src" + sep + "one.py", + "src" + sep + "two.py", + ], + prefix="/add ", + ) + bar._compute_display_names() + + assert bar._common_prefix == "src" + sep + assert bar._display_names == ["one.py", "two.py"] diff --git a/tests/unit/test_deep_merge.py b/tests/unit/test_deep_merge.py index c2a77b49c2b..bf4ec6fdc70 100644 --- a/tests/unit/test_deep_merge.py +++ b/tests/unit/test_deep_merge.py @@ -76,6 +76,18 @@ def test_json_field_deep_merge(self): merged = deep_merge(dict1, dict2, deep_merge_arrays=True) self.assertEqual(merged, expected) + def test_deep_merge_with_unpicklable_objects(self): + import threading + + lock = threading.RLock() + dict1 = {"lock": lock, "nested": {"key": "val", "lock": lock}} + dict2 = {"extra": 123} + merged = deep_merge(dict1, dict2) + self.assertIs(merged["lock"], lock) + self.assertEqual(merged["extra"], 123) + self.assertEqual(merged["nested"]["key"], "val") + self.assertIs(merged["nested"]["lock"], lock) + class TestConfigHelpers(unittest.TestCase): def test_is_cecli_conf_file(self): diff --git a/tests/unit/test_retry_backoff.py b/tests/unit/test_retry_backoff.py index 27bafbc6a21..523274b33e1 100644 --- a/tests/unit/test_retry_backoff.py +++ b/tests/unit/test_retry_backoff.py @@ -534,3 +534,157 @@ def mock_sleep(delay): assert result == (None, None) asyncio.run(run_test()) + + +def test_retry_on_unauthorized_disabled_by_default(): + async def run_test(): + model = Model("openai/gpt-4o") + model.caches_by_default = False + + auth_err = litellm.AuthenticationError( + message="401 Unauthorized", + model="openai/gpt-4o", + llm_provider="openai", + ) + auth_err.status_code = 401 + + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + raise auth_err + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + # Default (retry_on_unauthorized=False): 401 fails immediately, no retry. + assert model.retry_on_unauthorized is False + assert len(slept_delays) == 0 + assert "Model API Response Error" in resp.choices[0].message.content + + asyncio.run(run_test()) + + +def test_retry_on_unauthorized_enabled_retries_401(): + async def run_test(): + model = Model("openai/gpt-4o", retries='{"retry-on-unauthorized": true}') + model.caches_by_default = False + + auth_err = litellm.AuthenticationError( + message="401 Unauthorized", + model="openai/gpt-4o", + llm_provider="openai", + ) + auth_err.status_code = 401 + + call_count = 0 + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise auth_err + return MagicMock(choices=[MagicMock(message=MagicMock(content="ok"))]) + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + # Enabled: the 401 is retried with the standard exponential backoff. + assert model.retry_on_unauthorized is True + assert call_count == 2 + assert len(slept_delays) == 1 + assert pytest.approx(slept_delays[0]) == 0.125 * 1.5 + assert resp.choices[0].message.content == "ok" + + asyncio.run(run_test()) + + +def test_retry_on_unauthorized_enabled_retries_403(): + async def run_test(): + model = Model("openai/gpt-4o", retries='{"retry-on-unauthorized": true}') + model.caches_by_default = False + + forbidden_err = litellm.PermissionDeniedError( + message="403 Forbidden", + model="openai/gpt-4o", + llm_provider="openai", + ) + forbidden_err.status_code = 403 + + call_count = 0 + slept_delays = [] + + async def mock_acompletion(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + raise forbidden_err + return MagicMock(choices=[MagicMock(message=MagicMock(content="ok"))]) + + async def mock_sleep(delay): + slept_delays.append(delay) + + with ( + patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion), + patch("asyncio.sleep", side_effect=mock_sleep), + ): + _hash, resp = await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + + assert call_count == 2 + assert len(slept_delays) == 1 + assert resp.choices[0].message.content == "ok" + + asyncio.run(run_test()) + + +def test_retry_on_unauthorized_config_parsing(): + async def run_test(): + async def mock_acompletion(*args, **kwargs): + return MagicMock(choices=[MagicMock(message=MagicMock(content="ok"))]) + + # JSON string form + model = Model("openai/gpt-4o", retries='{"retry-on-unauthorized": true}') + model.caches_by_default = False + with patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion): + await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + assert model.retry_on_unauthorized is True + + # Already-parsed dict form + model = Model("openai/gpt-4o", retries={"retry-on-unauthorized": True}) + model.caches_by_default = False + with patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion): + await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + assert model.retry_on_unauthorized is True + + # Default when unset + model = Model("openai/gpt-4o") + model.caches_by_default = False + with patch("cecli.llm.litellm.acompletion", side_effect=mock_acompletion): + await model.send_completion( + messages=[{"role": "user", "content": "hi"}], functions=None, stream=False + ) + assert model.retry_on_unauthorized is False + + asyncio.run(run_test())