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)
|
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)))
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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]
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue