diff --git a/src/cmdforge/cli/__init__.py b/src/cmdforge/cli/__init__.py index 67a0ff2..5cda11c 100644 --- a/src/cmdforge/cli/__init__.py +++ b/src/cmdforge/cli/__init__.py @@ -62,12 +62,14 @@ def main(): p_test.set_defaults(func=cmd_test) # 'run' command - p_run = subparsers.add_parser("run", help="Run a tool") + # NOTE: Options like --provider must come BEFORE the tool name due to argparse.REMAINDER + # Example: cmdforge run --provider mock my-tool (not: cmdforge run my-tool --provider mock) + p_run = subparsers.add_parser("run", help="Run a tool (options must come before tool name)") p_run.add_argument("name", help="Tool name") p_run.add_argument("-i", "--input", help="Input file (reads from stdin if piped)") p_run.add_argument("-o", "--output", help="Output file (writes to stdout if omitted)") p_run.add_argument("--stdin", action="store_true", help="Read input interactively (type then Ctrl+D)") - p_run.add_argument("-p", "--provider", help="Override provider") + p_run.add_argument("-p", "--provider", help="Override provider (must come before tool name)") p_run.add_argument("--dry-run", action="store_true", help="Show what would happen without executing") p_run.add_argument("--show-prompt", action="store_true", help="Show prompts in addition to output") p_run.add_argument("-v", "--verbose", action="store_true", help="Show debug information") @@ -126,6 +128,12 @@ def main(): p_prov_test.add_argument("name", help="Provider name") p_prov_test.set_defaults(func=cmd_providers) + # providers for-tools + p_prov_for_tools = providers_sub.add_parser("for-tools", help="List providers used by specified tools") + p_prov_for_tools.add_argument("tools", nargs="+", help="Tool names to check") + p_prov_for_tools.add_argument("--warm", action="store_true", help="Warm up the providers (send minimal prompt)") + p_prov_for_tools.set_defaults(func=cmd_providers) + # Default for providers with no subcommand p_providers.set_defaults(func=lambda args: cmd_providers(args) if args.providers_cmd else (setattr(args, 'providers_cmd', 'list') or cmd_providers(args))) diff --git a/src/cmdforge/cli/provider_commands.py b/src/cmdforge/cli/provider_commands.py index 588da33..e40be32 100644 --- a/src/cmdforge/cli/provider_commands.py +++ b/src/cmdforge/cli/provider_commands.py @@ -67,6 +67,8 @@ def cmd_providers(args): return _cmd_providers_test(args) elif args.providers_cmd == "check": return _cmd_providers_check(args) + elif args.providers_cmd == "for-tools": + return _cmd_providers_for_tools(args) return 0 @@ -269,6 +271,117 @@ def _cmd_providers_test(args): return 0 if result.success else 1 +def _cmd_providers_for_tools(args): + """List (and optionally warm) providers used by specified tools.""" + import sys + import re + from ..tool import load_tool + from ..resolver import resolve_tool, ToolNotFoundError + + # Cloud providers don't need warming - they're always "warm" + CLOUD_PREFIXES = ('claude', 'opencode', 'gemini', 'codex', 'gpt', 'openai') + + def is_local_provider(name): + """Check if provider is local (needs warming) vs cloud.""" + name_lower = name.lower() + return not any(name_lower.startswith(prefix) for prefix in CLOUD_PREFIXES) + + def estimate_model_size(provider_name): + """Estimate model size in billions from provider name/command for sorting. + + Larger models should be warmed first to claim contiguous GPU memory. + Returns estimated size in billions (e.g., 70 for 70b model). + """ + # Check provider command for model size hints + provider = get_provider(provider_name) + text_to_check = f"{provider_name} {provider.command if provider else ''}" + + # Look for patterns like "70b", "14b", "7b", "3b", "1.5b" + match = re.search(r'(\d+(?:\.\d+)?)\s*[bB]', text_to_check) + if match: + return float(match.group(1)) + + # Known large model keywords + if any(kw in text_to_check.lower() for kw in ['opus', 'large', 'big', '32b', '70b']): + return 70.0 + if any(kw in text_to_check.lower() for kw in ['medium', '14b', '13b']): + return 14.0 + if any(kw in text_to_check.lower() for kw in ['small', 'fast', '3b', '1b', 'tiny', 'mini']): + return 3.0 + + # Default to medium size + return 7.0 + + def get_provider(name): + """Get provider by name.""" + providers = load_providers() + for p in providers: + if p.name == name: + return p + return None + + tool_names = args.tools + providers_used = set() + + for tool_name in tool_names: + try: + resolved = resolve_tool(tool_name) + tool = resolved.tool + except ToolNotFoundError: + print(f"Warning: Tool '{tool_name}' not found, skipping", file=sys.stderr) + continue + + # Extract providers from prompt steps + for step in tool.steps: + if hasattr(step, 'provider') and step.provider: + providers_used.add(step.provider) + + # Also check tool steps for nested tools (recursive would be complex, just note dependency) + for step in tool.steps: + if hasattr(step, 'tool') and step.tool: + # Could recurse here, but for now just note it + pass + + # Sort for consistent output + providers_list = sorted(providers_used) + + if args.warm: + # Only warm local providers (Ollama, etc.) + local_providers = [p for p in providers_list if is_local_provider(p)] + cloud_providers = [p for p in providers_list if not is_local_provider(p)] + + if cloud_providers: + print(f"Skipping {len(cloud_providers)} cloud provider(s): {', '.join(cloud_providers)}", file=sys.stderr) + + if local_providers: + # Sort by model size DESCENDING - warm largest models first + # This lets large models claim contiguous GPU memory before small models fragment it + local_providers_sorted = sorted( + local_providers, + key=lambda p: estimate_model_size(p), + reverse=True + ) + + print(f"Warming {len(local_providers_sorted)} local provider(s) (largest first)...", file=sys.stderr) + for provider_name in local_providers_sorted: + size_est = estimate_model_size(provider_name) + try: + result = call_provider(provider_name, "Hi", timeout=120) + if result.success: + print(f" [+] {provider_name} (~{size_est:.0f}B): warm", file=sys.stderr) + else: + print(f" [-] {provider_name} (~{size_est:.0f}B): failed ({result.error[:50]}...)", file=sys.stderr) + except Exception as e: + print(f" [-] {provider_name} (~{size_est:.0f}B): error ({e})", file=sys.stderr) + print(file=sys.stderr) + + # Output provider names (one per line for easy parsing) + for p in providers_list: + print(p) + + return 0 + + def _cmd_providers_check(args): """Check which providers are available.""" providers = load_providers() diff --git a/src/cmdforge/providers.py b/src/cmdforge/providers.py index e251ba1..f3c7a7b 100644 --- a/src/cmdforge/providers.py +++ b/src/cmdforge/providers.py @@ -1,6 +1,7 @@ """Provider abstraction for AI CLI tools.""" import os +import re import shlex import subprocess import shutil @@ -11,6 +12,16 @@ from typing import Optional, List import yaml +# Regex to strip ANSI escape codes (spinners, colors, cursor movements) +# Handles: CSI sequences (\x1b[...X), including private modes (\x1b[?...X), OSC sequences, and \r +ANSI_ESCAPE_RE = re.compile(r'\x1b\[[\x20-\x3f]*[\x40-\x7e]|\x1b\].*?\x07|\r') + + +def strip_ansi(text: str) -> str: + """Strip ANSI escape codes from text (e.g., Ollama spinner output).""" + return ANSI_ESCAPE_RE.sub('', text) + + # Default providers config location PROVIDERS_FILE = Path.home() / ".cmdforge" / "providers.yaml" @@ -144,7 +155,7 @@ def delete_provider(name: str) -> bool: return False -def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: Optional[set] = None) -> ProviderResult: +def call_provider(provider_name: str, prompt: str, timeout: int = 300, max_tokens: Optional[int] = None, _tried: Optional[set] = None) -> ProviderResult: """ Call an AI provider with the given prompt. @@ -152,6 +163,7 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: O provider_name: Name of the provider to use prompt: The prompt to send timeout: Maximum execution time in seconds + max_tokens: Optional max output tokens (appends provider-specific flag) _tried: Internal set of already-tried providers (prevents infinite loops) Returns: @@ -178,6 +190,15 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: O # Parse command (expand environment variables) cmd = os.path.expandvars(provider.command) + # Append max_tokens flag for known providers + if max_tokens: + name_lower = provider_name.lower() + if name_lower.startswith("claude") or "claude" in cmd.lower(): + cmd = f"{cmd} --max-tokens {max_tokens}" + elif name_lower.startswith("gemini") or "gemini" in cmd.lower(): + cmd = f"{cmd} --max-output-tokens {max_tokens}" + # opencode and codex use default model limits, no flag available + # Check if base command exists (use shlex for proper quote handling) try: cmd_parts = shlex.split(cmd) @@ -191,7 +212,7 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: O if provider.fallback and provider.fallback not in _tried: import sys print(f"[fallback] {provider_name} failed, trying {provider.fallback}...", file=sys.stderr) - return call_provider(provider.fallback, prompt, timeout, _tried) + return call_provider(provider.fallback, prompt, timeout, max_tokens, _tried) return ProviderResult(text="", success=False, error=error_msg) # Expand ~ for the which check @@ -212,19 +233,20 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: O ) if result.returncode != 0: - error_msg = f"Provider exited with code {result.returncode}: {result.stderr}" - if "not found" in result.stderr.lower() or "not installed" in result.stderr.lower(): + stderr_clean = strip_ansi(result.stderr) + error_msg = f"Provider exited with code {result.returncode}: {stderr_clean}" + if "not found" in stderr_clean.lower() or "not installed" in stderr_clean.lower(): error_msg += "\n\nTo install AI providers, run: cmdforge providers install" return try_fallback(error_msg) # Warn if output is empty (provider ran but returned nothing) - if not result.stdout.strip(): - stderr = result.stderr.strip() + clean_stdout = strip_ansi(result.stdout) + if not clean_stdout.strip(): + stderr = strip_ansi(result.stderr).strip() # Check for OpenCode's ProviderModelNotFoundError if "ProviderModelNotFoundError" in stderr or "ModelNotFoundError" in stderr: # Extract provider and model info if possible - import re provider_match = re.search(r'providerID:\s*"([^"]+)"', stderr) model_match = re.search(r'modelID:\s*"([^"]+)"', stderr) provider_id = provider_match.group(1) if provider_match else "unknown" @@ -243,7 +265,7 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: O f"Provider returned empty output{stderr_hint}.\n\nThis may mean the model is not available. Try a different provider or run: cmdforge providers install" ) - return ProviderResult(text=result.stdout, success=True) + return ProviderResult(text=clean_stdout, success=True) except subprocess.TimeoutExpired: return try_fallback(f"Provider timed out after {timeout} seconds") diff --git a/src/cmdforge/runner.py b/src/cmdforge/runner.py index 8b70109..b7a9ea0 100644 --- a/src/cmdforge/runner.py +++ b/src/cmdforge/runner.py @@ -546,7 +546,7 @@ def execute_prompt_step( if provider.lower() == "mock": result = mock_provider(prompt) else: - result = call_provider(provider, prompt) + result = call_provider(provider, prompt, max_tokens=step.max_tokens) if not result.success: print(f"Error in prompt step: {result.error}", file=sys.stderr) @@ -590,7 +590,7 @@ Please try again with valid JSON matching the schema exactly.""" if provider.lower() == "mock": result = mock_provider(current_prompt) else: - result = call_provider(provider, current_prompt) + result = call_provider(provider, current_prompt, max_tokens=step.max_tokens) if not result.success: print(f"Error in prompt step: {result.error}", file=sys.stderr) diff --git a/src/cmdforge/tool.py b/src/cmdforge/tool.py index 8f3d4ec..3abfd61 100644 --- a/src/cmdforge/tool.py +++ b/src/cmdforge/tool.py @@ -105,6 +105,7 @@ class PromptStep: output_schema: Optional[dict] = None # JSON schema for validation (None = use default) max_retries: int = 1 # Retry count on validation failure plain_text: bool = False # Bypass structured output enforcement + max_tokens: Optional[int] = None # Max output tokens (provider-dependent) def to_dict(self) -> dict: d = { @@ -127,6 +128,8 @@ class PromptStep: d["max_retries"] = self.max_retries if self.plain_text: d["plain_text"] = self.plain_text + if self.max_tokens: + d["max_tokens"] = self.max_tokens return d @classmethod @@ -141,7 +144,8 @@ class PromptStep: strip_fences=data.get("strip_fences", False), output_schema=data.get("output_schema"), max_retries=data.get("max_retries", 1), - plain_text=data.get("plain_text", False) + plain_text=data.get("plain_text", False), + max_tokens=data.get("max_tokens") ) diff --git a/tests/test_providers.py b/tests/test_providers.py index 3db3a5a..5882d3b 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -12,7 +12,7 @@ from cmdforge.providers import ( load_providers, save_providers, get_provider, add_provider, delete_provider, call_provider, mock_provider, - DEFAULT_PROVIDERS + DEFAULT_PROVIDERS, strip_ansi ) @@ -507,3 +507,122 @@ class TestProviderFallback: assert result.success is True assert "[MOCK]" in result.text + + +class TestStripAnsi: + """Tests for ANSI escape code stripping.""" + + def test_strip_basic_color_codes(self): + """Strip basic color codes like \\x1b[0m.""" + text = "\x1b[31mred text\x1b[0m" + assert strip_ansi(text) == "red text" + + def test_strip_cursor_movements(self): + """Strip cursor movement codes.""" + text = "\x1b[2Jhello\x1b[H" + assert strip_ansi(text) == "hello" + + def test_strip_private_mode_sequences(self): + """Strip private mode sequences like \\x1b[?25l (hide cursor).""" + text = "\x1b[?2026h\x1b[?25lspinner\x1b[?25h\x1b[?2026l" + assert strip_ansi(text) == "spinner" + + def test_strip_ollama_spinner(self): + """Strip Ollama spinner output.""" + text = "\x1b[?2026h\x1b[?25l⠋ \x1b[?25h\x1b[?2026l\x1b[?2026h\x1b[?25l⠙ \x1b[?25h" + result = strip_ansi(text) + # Should only have spinner characters left, not escape codes + assert "\x1b" not in result + assert "⠋" in result or "⠙" in result + + def test_strip_carriage_return(self): + """Strip carriage returns.""" + text = "line1\rline2" + assert strip_ansi(text) == "line1line2" + + def test_preserve_normal_text(self): + """Normal text should be preserved.""" + text = "Hello, World! 123 [test] {data}" + assert strip_ansi(text) == text + + def test_strip_osc_sequences(self): + """Strip OSC (Operating System Command) sequences.""" + text = "\x1b]0;Window Title\x07content" + assert strip_ansi(text) == "content" + + +class TestMaxTokens: + """Tests for max_tokens parameter.""" + + @pytest.fixture + def temp_providers_file(self, tmp_path): + providers_file = tmp_path / ".cmdforge" / "providers.yaml" + with patch('cmdforge.providers.PROVIDERS_FILE', providers_file): + yield providers_file + + @patch('subprocess.run') + @patch('shutil.which') + def test_max_tokens_appended_for_claude(self, mock_which, mock_run, temp_providers_file): + """max_tokens should append --max-tokens for claude providers.""" + mock_which.return_value = "/usr/bin/claude" + mock_run.return_value = MagicMock(returncode=0, stdout="response", stderr="") + save_providers([Provider("claude-haiku", "claude -p --model haiku")]) + + call_provider("claude-haiku", "Test prompt", max_tokens=4096) + + # Verify the command included --max-tokens + call_args = mock_run.call_args + cmd = call_args[0][0] + assert "--max-tokens 4096" in cmd + + @patch('subprocess.run') + @patch('shutil.which') + def test_max_tokens_appended_for_gemini(self, mock_which, mock_run, temp_providers_file): + """max_tokens should append --max-output-tokens for gemini providers.""" + mock_which.return_value = "/usr/bin/gemini" + mock_run.return_value = MagicMock(returncode=0, stdout="response", stderr="") + save_providers([Provider("gemini", "gemini --model gemini-2.5-pro")]) + + call_provider("gemini", "Test prompt", max_tokens=8192) + + call_args = mock_run.call_args + cmd = call_args[0][0] + assert "--max-output-tokens 8192" in cmd + + @patch('subprocess.run') + @patch('shutil.which') + def test_max_tokens_not_appended_when_none(self, mock_which, mock_run, temp_providers_file): + """No flag should be added when max_tokens is None.""" + mock_which.return_value = "/usr/bin/claude" + mock_run.return_value = MagicMock(returncode=0, stdout="response", stderr="") + save_providers([Provider("claude", "claude -p")]) + + call_provider("claude", "Test prompt", max_tokens=None) + + call_args = mock_run.call_args + cmd = call_args[0][0] + assert "--max-tokens" not in cmd + + @patch('subprocess.run') + @patch('shutil.which') + def test_max_tokens_passed_through_fallback(self, mock_which, mock_run, temp_providers_file): + """max_tokens should be passed to fallback provider.""" + mock_which.return_value = "/usr/bin/claude" + # First call fails, second succeeds + mock_run.side_effect = [ + MagicMock(returncode=1, stdout="", stderr="Error"), + MagicMock(returncode=0, stdout="fallback response", stderr="") + ] + save_providers([ + Provider("primary", "claude -p --model opus", fallback="fallback"), + Provider("fallback", "claude -p --model haiku") + ]) + + result = call_provider("primary", "Test prompt", max_tokens=4096) + + assert result.success is True + # Both calls should have had max_tokens + calls = mock_run.call_args_list + assert "--max-tokens 4096" in calls[0][0][0] + assert "--max-tokens 4096" in calls[1][0][0] +