Add provider for-tools command and max_tokens support

- Add 'cmdforge providers for-tools <tool>...' to list providers used by tools
- Add --warm flag to pre-load local models (Ollama) for faster first inference
- Warm larger models first to claim contiguous GPU memory
- Skip cloud providers (claude, opencode, gemini) during warm-up
- Add max_tokens field to PromptStep for controlling output length
- Pass max_tokens to call_provider with provider-specific flags
- Add comprehensive tests for provider functionality

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
rob 2026-04-04 17:50:26 -03:00
parent 468a999b80
commit b673ea019c
6 changed files with 280 additions and 14 deletions

View File

@ -62,12 +62,14 @@ def main():
p_test.set_defaults(func=cmd_test) p_test.set_defaults(func=cmd_test)
# 'run' command # '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("name", help="Tool name")
p_run.add_argument("-i", "--input", help="Input file (reads from stdin if piped)") 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("-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("--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("--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("--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") 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.add_argument("name", help="Provider name")
p_prov_test.set_defaults(func=cmd_providers) 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 # 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))) p_providers.set_defaults(func=lambda args: cmd_providers(args) if args.providers_cmd else (setattr(args, 'providers_cmd', 'list') or cmd_providers(args)))

View File

@ -67,6 +67,8 @@ def cmd_providers(args):
return _cmd_providers_test(args) return _cmd_providers_test(args)
elif args.providers_cmd == "check": elif args.providers_cmd == "check":
return _cmd_providers_check(args) return _cmd_providers_check(args)
elif args.providers_cmd == "for-tools":
return _cmd_providers_for_tools(args)
return 0 return 0
@ -269,6 +271,117 @@ def _cmd_providers_test(args):
return 0 if result.success else 1 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): def _cmd_providers_check(args):
"""Check which providers are available.""" """Check which providers are available."""
providers = load_providers() providers = load_providers()

View File

@ -1,6 +1,7 @@
"""Provider abstraction for AI CLI tools.""" """Provider abstraction for AI CLI tools."""
import os import os
import re
import shlex import shlex
import subprocess import subprocess
import shutil import shutil
@ -11,6 +12,16 @@ from typing import Optional, List
import yaml 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 # Default providers config location
PROVIDERS_FILE = Path.home() / ".cmdforge" / "providers.yaml" PROVIDERS_FILE = Path.home() / ".cmdforge" / "providers.yaml"
@ -144,7 +155,7 @@ def delete_provider(name: str) -> bool:
return False 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. 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 provider_name: Name of the provider to use
prompt: The prompt to send prompt: The prompt to send
timeout: Maximum execution time in seconds 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) _tried: Internal set of already-tried providers (prevents infinite loops)
Returns: Returns:
@ -178,6 +190,15 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: O
# Parse command (expand environment variables) # Parse command (expand environment variables)
cmd = os.path.expandvars(provider.command) 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) # Check if base command exists (use shlex for proper quote handling)
try: try:
cmd_parts = shlex.split(cmd) 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: if provider.fallback and provider.fallback not in _tried:
import sys import sys
print(f"[fallback] {provider_name} failed, trying {provider.fallback}...", file=sys.stderr) 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) return ProviderResult(text="", success=False, error=error_msg)
# Expand ~ for the which check # 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: if result.returncode != 0:
error_msg = f"Provider exited with code {result.returncode}: {result.stderr}" stderr_clean = strip_ansi(result.stderr)
if "not found" in result.stderr.lower() or "not installed" in result.stderr.lower(): 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" error_msg += "\n\nTo install AI providers, run: cmdforge providers install"
return try_fallback(error_msg) return try_fallback(error_msg)
# Warn if output is empty (provider ran but returned nothing) # Warn if output is empty (provider ran but returned nothing)
if not result.stdout.strip(): clean_stdout = strip_ansi(result.stdout)
stderr = result.stderr.strip() if not clean_stdout.strip():
stderr = strip_ansi(result.stderr).strip()
# Check for OpenCode's ProviderModelNotFoundError # Check for OpenCode's ProviderModelNotFoundError
if "ProviderModelNotFoundError" in stderr or "ModelNotFoundError" in stderr: if "ProviderModelNotFoundError" in stderr or "ModelNotFoundError" in stderr:
# Extract provider and model info if possible # Extract provider and model info if possible
import re
provider_match = re.search(r'providerID:\s*"([^"]+)"', stderr) provider_match = re.search(r'providerID:\s*"([^"]+)"', stderr)
model_match = re.search(r'modelID:\s*"([^"]+)"', stderr) model_match = re.search(r'modelID:\s*"([^"]+)"', stderr)
provider_id = provider_match.group(1) if provider_match else "unknown" 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" 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: except subprocess.TimeoutExpired:
return try_fallback(f"Provider timed out after {timeout} seconds") return try_fallback(f"Provider timed out after {timeout} seconds")

View File

@ -546,7 +546,7 @@ def execute_prompt_step(
if provider.lower() == "mock": if provider.lower() == "mock":
result = mock_provider(prompt) result = mock_provider(prompt)
else: else:
result = call_provider(provider, prompt) result = call_provider(provider, prompt, max_tokens=step.max_tokens)
if not result.success: if not result.success:
print(f"Error in prompt step: {result.error}", file=sys.stderr) 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": if provider.lower() == "mock":
result = mock_provider(current_prompt) result = mock_provider(current_prompt)
else: else:
result = call_provider(provider, current_prompt) result = call_provider(provider, current_prompt, max_tokens=step.max_tokens)
if not result.success: if not result.success:
print(f"Error in prompt step: {result.error}", file=sys.stderr) print(f"Error in prompt step: {result.error}", file=sys.stderr)

View File

@ -105,6 +105,7 @@ class PromptStep:
output_schema: Optional[dict] = None # JSON schema for validation (None = use default) output_schema: Optional[dict] = None # JSON schema for validation (None = use default)
max_retries: int = 1 # Retry count on validation failure max_retries: int = 1 # Retry count on validation failure
plain_text: bool = False # Bypass structured output enforcement plain_text: bool = False # Bypass structured output enforcement
max_tokens: Optional[int] = None # Max output tokens (provider-dependent)
def to_dict(self) -> dict: def to_dict(self) -> dict:
d = { d = {
@ -127,6 +128,8 @@ class PromptStep:
d["max_retries"] = self.max_retries d["max_retries"] = self.max_retries
if self.plain_text: if self.plain_text:
d["plain_text"] = self.plain_text d["plain_text"] = self.plain_text
if self.max_tokens:
d["max_tokens"] = self.max_tokens
return d return d
@classmethod @classmethod
@ -141,7 +144,8 @@ class PromptStep:
strip_fences=data.get("strip_fences", False), strip_fences=data.get("strip_fences", False),
output_schema=data.get("output_schema"), output_schema=data.get("output_schema"),
max_retries=data.get("max_retries", 1), 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")
) )

View File

@ -12,7 +12,7 @@ from cmdforge.providers import (
load_providers, save_providers, get_provider, load_providers, save_providers, get_provider,
add_provider, delete_provider, add_provider, delete_provider,
call_provider, mock_provider, call_provider, mock_provider,
DEFAULT_PROVIDERS DEFAULT_PROVIDERS, strip_ansi
) )
@ -507,3 +507,122 @@ class TestProviderFallback:
assert result.success is True assert result.success is True
assert "[MOCK]" in result.text 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]