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:
parent
468a999b80
commit
b673ea019c
|
|
@ -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)))
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue