852 lines
32 KiB
Python
852 lines
32 KiB
Python
"""Provider abstraction for AI CLI tools."""
|
|
|
|
import os
|
|
import re
|
|
import shlex
|
|
import subprocess
|
|
import shutil
|
|
import sys
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
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"
|
|
|
|
|
|
# Fallback chain templates for callers and manual providers.yaml configuration.
|
|
PRESET_CHAINS = {
|
|
"premium": ["claude-opus", "claude-sonnet", "claude-haiku", "codex", "crush", "mock"],
|
|
"free": ["opencode-free", "agy", "codex", "crush", "ollama", "mock"],
|
|
"fast": ["agy", "claude-haiku", "opencode-free", "crush", "mock"],
|
|
"reasoning": ["claude-opus", "opencode-reasoner", "deepseek-api", "crush", "mock"],
|
|
"balanced": ["opencode-pickle", "codex", "agy", "claude-haiku", "crush", "ollama", "mock"],
|
|
}
|
|
|
|
PROVIDERS_CONFIG_VERSION = 2
|
|
V2_DEFAULT_PROVIDER_NAMES = {
|
|
"opencode-free",
|
|
"agy",
|
|
"crush",
|
|
"ollama",
|
|
"openrouter",
|
|
"deepseek-api",
|
|
}
|
|
|
|
|
|
# Known CLIs that CmdForge can auto-discover on the user's PATH.
|
|
# Maps binary name -> default provider config (used during first-run setup).
|
|
KNOWN_PROVIDER_CLIS = {
|
|
"opencode": {
|
|
"name": "opencode-pickle",
|
|
"command": "opencode run --model opencode/big-pickle",
|
|
"description": "OpenCode - Big Pickle (free general model)",
|
|
"tags": ["free", "code", "general"],
|
|
"install_group": "opencode",
|
|
},
|
|
"agy": {
|
|
"name": "agy",
|
|
"command": "agy -p",
|
|
"description": "Antigravity - Google free tier, Gemini models",
|
|
"tags": ["free-tier", "code", "large-context"],
|
|
"install_group": "agy",
|
|
},
|
|
"codex": {
|
|
"name": "codex",
|
|
"command": "codex exec -",
|
|
"description": "Codex CLI - OpenAI free tier available",
|
|
"tags": ["free-tier", "code", "general"],
|
|
"install_group": "codex",
|
|
},
|
|
"claude": {
|
|
"name": "claude",
|
|
"command": "claude -p",
|
|
"description": "Claude Code - auto-routes to best model",
|
|
"tags": ["paid", "subscription", "code"],
|
|
"install_group": "claude",
|
|
},
|
|
"crush": {
|
|
"name": "crush",
|
|
"command": "crush run --quiet",
|
|
"description": "Crush - multi-model via Hyper credits or API keys",
|
|
"tags": ["free-tier", "multi", "code"],
|
|
"install_group": "crush",
|
|
},
|
|
"ollama": {
|
|
"name": "ollama",
|
|
"command": "ollama run llama3.2",
|
|
"description": "Ollama - local, private, free",
|
|
"tags": ["free", "local", "private"],
|
|
"install_group": "ollama",
|
|
},
|
|
}
|
|
|
|
# Known API key environment variables -> provider config
|
|
KNOWN_API_KEYS = {
|
|
"OPENROUTER_API_KEY": {
|
|
"name": "openrouter",
|
|
"command": "https://openrouter.ai/api/v1",
|
|
"model": "openrouter/auto-beta",
|
|
"description": "OpenRouter - 300+ models, one API key",
|
|
"tags": ["api", "per-token", "multi"],
|
|
},
|
|
"DEEPSEEK_API_KEY": {
|
|
"name": "deepseek-api",
|
|
"command": "https://api.deepseek.com/v1",
|
|
"model": "deepseek-chat",
|
|
"description": "DeepSeek API - inexpensive per-token access",
|
|
"tags": ["api", "per-token", "cheap"],
|
|
},
|
|
"OPENAI_API_KEY": {
|
|
"name": "openai-api",
|
|
"command": "https://api.openai.com/v1",
|
|
"model": "gpt-4o",
|
|
"description": "OpenAI API - direct GPT access",
|
|
"tags": ["api", "per-token", "code"],
|
|
},
|
|
}
|
|
|
|
|
|
def discover_installed_providers() -> List[dict]:
|
|
"""Scan the system for installed AI CLIs and configured API keys.
|
|
|
|
Returns a list of discovery results, each with keys:
|
|
- source: "cli" or "api-key"
|
|
- name: provider name
|
|
- command: provider command
|
|
- description, tags, etc.
|
|
"""
|
|
found = []
|
|
|
|
# Check PATH for known CLIs
|
|
for binary, info in KNOWN_PROVIDER_CLIS.items():
|
|
path = shutil.which(binary)
|
|
if path:
|
|
found.append({
|
|
"source": "cli",
|
|
"binary": binary,
|
|
"path": path,
|
|
**{k: v for k, v in info.items() if k != "install_group"},
|
|
})
|
|
|
|
# Check environment for API keys
|
|
for env_var, info in KNOWN_API_KEYS.items():
|
|
if os.environ.get(env_var):
|
|
found.append({
|
|
"source": "api-key",
|
|
"env_var": env_var,
|
|
"type": "api",
|
|
**info,
|
|
})
|
|
|
|
# Check Ollama for available local models (if installed).
|
|
if shutil.which("ollama"):
|
|
try:
|
|
result = subprocess.run(
|
|
["ollama", "list"],
|
|
capture_output=True,
|
|
text=True,
|
|
timeout=5,
|
|
)
|
|
if result.returncode == 0:
|
|
used_names = {item["name"] for item in found}
|
|
# Parse model list (skip header line).
|
|
for line in result.stdout.strip().split("\n")[1:]:
|
|
if line.strip():
|
|
model_name = line.split()[0]
|
|
slug = re.sub(r"[^a-z0-9]+", "-", model_name.lower()).strip("-")
|
|
provider_name = f"ollama-{slug}"
|
|
suffix = 2
|
|
while provider_name in used_names:
|
|
provider_name = f"ollama-{slug}-{suffix}"
|
|
suffix += 1
|
|
used_names.add(provider_name)
|
|
found.append({
|
|
"source": "ollama-model",
|
|
"binary": "ollama",
|
|
"name": provider_name,
|
|
"command": f"ollama run {model_name}",
|
|
"description": f"Ollama local model: {model_name}",
|
|
"tags": ["free", "local", "private"],
|
|
})
|
|
except (subprocess.TimeoutExpired, FileNotFoundError):
|
|
pass
|
|
|
|
return found
|
|
|
|
|
|
@dataclass
|
|
class Provider:
|
|
"""Definition of an AI provider.
|
|
|
|
Types:
|
|
- "subprocess" (default): CLI tool invoked via subprocess.run with stdin.
|
|
`command` is the shell command (e.g. "claude -p", "opencode run --model X").
|
|
- "api": HTTP POST to an OpenAI-compatible endpoint.
|
|
`command` is the endpoint URL, `model` is the model ID, `api_key_env`
|
|
names the environment variable holding the API key.
|
|
- "pty": Interactive CLI wrapped in a pseudo-terminal (experimental).
|
|
`command` is the CLI invocation; pty_config dict holds patterns.
|
|
"""
|
|
|
|
name: str
|
|
command: str
|
|
description: str = ""
|
|
fallback: Optional[str] = None
|
|
type: str = "subprocess" # "subprocess" | "api" | "pty"
|
|
model: Optional[str] = None # Model ID (required for api-type providers)
|
|
tags: List[str] = field(default_factory=list) # e.g. ["free", "code", "reasoning"]
|
|
install: Optional[dict] = None # Structured install metadata
|
|
fallback_chain: Optional[List[str]] = None # Ordered multi-step fallback (new)
|
|
api_key_env: Optional[str] = None # Env var name for api-type providers
|
|
pty_config: Optional[dict] = None # Patterns for pty-type providers
|
|
tools: Optional[List[str]] = None # Allowed CmdForge tools (None = all)
|
|
mcp_servers: Optional[List[str]] = None # MCP servers available to this provider
|
|
|
|
def __post_init__(self) -> None:
|
|
for field_name, values in (
|
|
("tools", self.tools),
|
|
("mcp_servers", self.mcp_servers),
|
|
):
|
|
if values is not None and (
|
|
not isinstance(values, list)
|
|
or not all(isinstance(value, str) and value for value in values)
|
|
):
|
|
raise ValueError(
|
|
f"Provider {field_name} must be a list of non-empty strings or null"
|
|
)
|
|
|
|
def to_dict(self) -> dict:
|
|
d = {
|
|
"name": self.name,
|
|
"command": self.command,
|
|
"description": self.description,
|
|
}
|
|
# Only include type if not default (preserves backward-compat with old YAML)
|
|
if self.type and self.type != "subprocess":
|
|
d["type"] = self.type
|
|
if self.model:
|
|
d["model"] = self.model
|
|
if self.fallback:
|
|
d["fallback"] = self.fallback
|
|
if self.tags:
|
|
d["tags"] = self.tags
|
|
if self.install:
|
|
d["install"] = self.install
|
|
if self.fallback_chain:
|
|
d["fallback_chain"] = self.fallback_chain
|
|
if self.api_key_env:
|
|
d["api_key_env"] = self.api_key_env
|
|
if self.pty_config:
|
|
d["pty_config"] = self.pty_config
|
|
if self.tools is not None:
|
|
d["tools"] = self.tools
|
|
if self.mcp_servers is not None:
|
|
d["mcp_servers"] = self.mcp_servers
|
|
return d
|
|
|
|
@classmethod
|
|
def from_dict(cls, data: dict) -> "Provider":
|
|
return cls(
|
|
name=data["name"],
|
|
command=data["command"],
|
|
description=data.get("description", ""),
|
|
type=data.get("type", "subprocess"),
|
|
model=data.get("model"),
|
|
fallback=data.get("fallback"),
|
|
tags=data.get("tags", []) or [],
|
|
install=data.get("install"),
|
|
fallback_chain=data.get("fallback_chain"),
|
|
api_key_env=data.get("api_key_env"),
|
|
pty_config=data.get("pty_config"),
|
|
tools=data.get("tools"),
|
|
mcp_servers=data.get("mcp_servers"),
|
|
)
|
|
|
|
|
|
@dataclass
|
|
class ProviderResult:
|
|
"""Result from a provider call."""
|
|
text: str
|
|
success: bool
|
|
error: Optional[str] = None
|
|
|
|
|
|
# Default providers that come pre-configured
|
|
# Live-tested July 2026: opencode run, codex exec -, agy -p, ollama run, crush run
|
|
DEFAULT_PROVIDERS = [
|
|
# OPENCODE - best free models (binary auto-updates, 6+ free models)
|
|
Provider("opencode-deepseek", "opencode run --model deepseek/deepseek-chat",
|
|
"DeepSeek V3 - cheap, fast, accurate (paid API key)",
|
|
tags=["paid", "code", "reasoning"]),
|
|
Provider("opencode-pickle", "opencode run --model opencode/big-pickle",
|
|
"Big Pickle - best free general model",
|
|
tags=["free", "code", "general"]),
|
|
Provider("opencode-reasoner", "opencode run --model deepseek/deepseek-reasoner",
|
|
"DeepSeek R1 - complex reasoning, cheap (paid API key)",
|
|
tags=["paid", "reasoning"]),
|
|
Provider("opencode-free", "opencode run --model opencode/deepseek-v4-flash-free",
|
|
"DeepSeek V4 Flash - fast, free",
|
|
tags=["free", "fast"]),
|
|
|
|
# ANTHROPIC CLAUDE - paid, high quality (requires subscription or API key)
|
|
Provider("claude", "claude -p",
|
|
"Claude Code - auto-routes to best model",
|
|
tags=["paid", "subscription", "code", "balanced"]),
|
|
Provider("claude-haiku", "claude -p --model haiku",
|
|
"Claude Haiku - fast, cheapest Claude",
|
|
tags=["paid", "subscription", "fast"]),
|
|
Provider("claude-sonnet", "claude -p --model sonnet",
|
|
"Claude Sonnet - balanced quality/speed",
|
|
tags=["paid", "subscription", "code"]),
|
|
Provider("claude-opus", "claude -p --model opus",
|
|
"Claude Opus - highest quality, expensive",
|
|
tags=["paid", "subscription", "reasoning", "quality"]),
|
|
|
|
# OPENAI CODEX - free tier available
|
|
Provider("codex", "codex exec -",
|
|
"Codex CLI - reliable, auto-routes, free tier available",
|
|
tags=["free-tier", "subscription", "code", "general"]),
|
|
|
|
# GOOGLE ANTIGRAVITY - replaces Gemini CLI (free tier: 1,000 req/day, 60 req/min)
|
|
Provider("agy", "agy -p",
|
|
"Antigravity (Google) - free tier, Gemini models, large context",
|
|
tags=["free-tier", "code", "large-context"]),
|
|
|
|
# CRUSH - multi-provider agent via Hyper credits or API keys
|
|
Provider("crush", "crush run --quiet",
|
|
"Crush - multi-model, requires Hyper credits or API keys",
|
|
tags=["free-tier", "multi", "code"]),
|
|
|
|
# LOCAL MODELS
|
|
Provider("ollama", "ollama run llama3.2",
|
|
"Ollama - local, private, free (GPU recommended)",
|
|
tags=["free", "local", "private"]),
|
|
|
|
# API-TYPE PROVIDERS (pay-per-token, fallback when no CLI covers model)
|
|
Provider("openrouter",
|
|
"https://openrouter.ai/api/v1",
|
|
"OpenRouter - 300+ models, one API key, auto-routing",
|
|
type="api",
|
|
model="openrouter/auto-beta",
|
|
api_key_env="OPENROUTER_API_KEY",
|
|
tags=["api", "per-token", "multi", "fallback"]),
|
|
Provider("deepseek-api",
|
|
"https://api.deepseek.com/v1",
|
|
"DeepSeek API - inexpensive per-token access",
|
|
type="api",
|
|
model="deepseek-chat",
|
|
api_key_env="DEEPSEEK_API_KEY",
|
|
tags=["api", "per-token", "cheap", "reasoning"]),
|
|
|
|
# Mock for testing
|
|
Provider("mock", "mock", "Mock provider for testing",
|
|
tags=["testing"]),
|
|
]
|
|
|
|
|
|
def get_providers_file() -> Path:
|
|
"""Get the providers config file, creating one on first run."""
|
|
if not PROVIDERS_FILE.exists():
|
|
PROVIDERS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
|
discovered = discover_installed_providers()
|
|
selected = []
|
|
if discovered:
|
|
import sys
|
|
print("=" * 60, file=sys.stderr)
|
|
print("CmdForge First-Run Provider Setup", file=sys.stderr)
|
|
print("=" * 60, file=sys.stderr)
|
|
print(file=sys.stderr)
|
|
cli_found = [d for d in discovered if d["source"] == "cli"]
|
|
api_found = [d for d in discovered if d["source"] == "api-key"]
|
|
ollama_found = [d for d in discovered if d["source"] == "ollama-model"]
|
|
|
|
if cli_found:
|
|
print(f"CLIs found on PATH ({len(cli_found)}):", file=sys.stderr)
|
|
for d in cli_found:
|
|
name = d["name"]
|
|
desc = d.get("description", "")
|
|
selected.append(Provider(
|
|
name=name, command=d["command"], description=desc,
|
|
tags=d.get("tags", []),
|
|
))
|
|
print(f" [+] {name:20s} {desc}", file=sys.stderr)
|
|
|
|
if api_found:
|
|
print(f"\nAPI keys detected ({len(api_found)}):", file=sys.stderr)
|
|
for d in api_found:
|
|
name = d["name"]
|
|
desc = d.get("description", "")
|
|
model = d.get("model", "auto")
|
|
env_var = d.get("env_var", "")
|
|
selected.append(Provider(
|
|
name=name, command=d["command"], description=desc,
|
|
type="api", model=model, api_key_env=env_var,
|
|
tags=d.get("tags", []),
|
|
))
|
|
print(f" [+] {name:20s} ({model})", file=sys.stderr)
|
|
|
|
if ollama_found:
|
|
print(f"\nLocal Ollama models ({len(ollama_found)}):", file=sys.stderr)
|
|
for d in ollama_found[:10]:
|
|
name = d["name"]
|
|
desc = d.get("description", "")
|
|
selected.append(Provider(
|
|
name=name, command=d["command"], description=desc,
|
|
tags=d.get("tags", []),
|
|
))
|
|
print(f" [+] {name:30s} {d['command']}", file=sys.stderr)
|
|
if len(ollama_found) > 10:
|
|
print(f" ... and {len(ollama_found) - 10} more", file=sys.stderr)
|
|
|
|
if selected:
|
|
save_providers(selected)
|
|
print(f"\nConfigured {len(selected)} provider(s) from discovery.", file=sys.stderr)
|
|
print(f"Run 'cmdforge providers discover' to re-scan at any time.", file=sys.stderr)
|
|
else:
|
|
save_providers(DEFAULT_PROVIDERS)
|
|
print("\nNo AI providers detected. Default providers written.", file=sys.stderr)
|
|
print("Run 'cmdforge providers install' for an interactive setup guide.", file=sys.stderr)
|
|
|
|
return PROVIDERS_FILE
|
|
|
|
|
|
def load_providers() -> List[Provider]:
|
|
"""Load all defined providers."""
|
|
providers_file = get_providers_file()
|
|
|
|
try:
|
|
data = yaml.safe_load(providers_file.read_text())
|
|
if not data or "providers" not in data:
|
|
return DEFAULT_PROVIDERS.copy()
|
|
providers = [Provider.from_dict(p) for p in data["providers"]]
|
|
except Exception as exc:
|
|
# A present but malformed configuration must not silently fall back to
|
|
# unrestricted defaults, particularly when access policies are invalid.
|
|
print(f"Warning: Failed to load provider configuration: {exc}", file=sys.stderr)
|
|
return []
|
|
|
|
config_version = data.get("version", 1)
|
|
if not isinstance(config_version, int):
|
|
config_version = 1
|
|
if config_version < PROVIDERS_CONFIG_VERSION:
|
|
providers = _merge_missing_defaults(providers)
|
|
try:
|
|
save_providers(providers)
|
|
except OSError:
|
|
# A read-only legacy config should still remain usable.
|
|
pass
|
|
return providers
|
|
|
|
|
|
def _merge_missing_defaults(providers: List[Provider]) -> List[Provider]:
|
|
"""Add defaults introduced since the legacy unversioned config format."""
|
|
merged = list(providers)
|
|
names = {provider.name for provider in merged}
|
|
for default in DEFAULT_PROVIDERS:
|
|
if default.name in V2_DEFAULT_PROVIDER_NAMES and default.name not in names:
|
|
merged.append(Provider.from_dict(default.to_dict()))
|
|
return merged
|
|
|
|
|
|
def save_providers(providers: List[Provider]):
|
|
"""Save providers to config file."""
|
|
PROVIDERS_FILE.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
|
|
data = {
|
|
"version": PROVIDERS_CONFIG_VERSION,
|
|
"providers": [p.to_dict() for p in providers],
|
|
}
|
|
PROVIDERS_FILE.write_text(yaml.safe_dump(data, default_flow_style=False, sort_keys=False))
|
|
try:
|
|
PROVIDERS_FILE.parent.chmod(0o700)
|
|
PROVIDERS_FILE.chmod(0o600)
|
|
except OSError:
|
|
pass
|
|
|
|
|
|
def get_provider(name: str) -> Optional[Provider]:
|
|
"""Get a provider by name."""
|
|
providers = load_providers()
|
|
for p in providers:
|
|
if p.name == name:
|
|
return p
|
|
return None
|
|
|
|
|
|
def add_provider(provider: Provider) -> bool:
|
|
"""Add or update a provider."""
|
|
providers = load_providers()
|
|
|
|
# Update if exists, otherwise add
|
|
for i, p in enumerate(providers):
|
|
if p.name == provider.name:
|
|
providers[i] = provider
|
|
save_providers(providers)
|
|
return True
|
|
|
|
providers.append(provider)
|
|
save_providers(providers)
|
|
return True
|
|
|
|
|
|
def delete_provider(name: str) -> bool:
|
|
"""Delete a provider by name."""
|
|
providers = load_providers()
|
|
original_len = len(providers)
|
|
providers = [p for p in providers if p.name != name]
|
|
|
|
if len(providers) < original_len:
|
|
save_providers(providers)
|
|
return True
|
|
return False
|
|
|
|
|
|
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.
|
|
|
|
Dispatches to call_provider_subprocess, call_provider_api, or
|
|
call_provider_pty based on the provider's `type` field. Falls back
|
|
to the provider's fallback (or fallback_chain) on failure.
|
|
|
|
Args:
|
|
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:
|
|
ProviderResult with the response text or error
|
|
"""
|
|
# Track which providers we've tried to prevent infinite fallback loops
|
|
if _tried is None:
|
|
_tried = set()
|
|
_tried.add(provider_name)
|
|
|
|
# Handle mock provider specially
|
|
if provider_name.lower() == "mock":
|
|
return mock_provider(prompt)
|
|
|
|
# Look up provider
|
|
provider = get_provider(provider_name)
|
|
if not provider:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=f"Provider '{provider_name}' not found. Use 'cmdforge providers' to manage providers."
|
|
)
|
|
|
|
if max_tokens is not None:
|
|
try:
|
|
max_tokens = int(max_tokens)
|
|
except (TypeError, ValueError):
|
|
return ProviderResult(text="", success=False, error="max_tokens must be an integer")
|
|
if max_tokens <= 0 or max_tokens > 1_000_000:
|
|
return ProviderResult(
|
|
text="", success=False, error="max_tokens must be between 1 and 1000000"
|
|
)
|
|
|
|
# Helper to try fallback provider(s) if available
|
|
def try_fallback(error_msg: str) -> ProviderResult:
|
|
import sys
|
|
last_error = error_msg
|
|
# Walk fallback_chain first (ordered, multi-step)
|
|
if provider.fallback_chain:
|
|
for fb in provider.fallback_chain:
|
|
if fb not in _tried:
|
|
print(f"[fallback] {provider_name} failed, trying {fb}...", file=sys.stderr)
|
|
result = call_provider(fb, prompt, timeout, max_tokens, _tried)
|
|
if result.success:
|
|
return result
|
|
last_error = result.error or last_error
|
|
# Fall back to single fallback (backward-compat)
|
|
if provider.fallback and provider.fallback not in _tried:
|
|
print(f"[fallback] {provider_name} failed, trying {provider.fallback}...", file=sys.stderr)
|
|
result = call_provider(provider.fallback, prompt, timeout, max_tokens, _tried)
|
|
if result.success:
|
|
return result
|
|
last_error = result.error or last_error
|
|
return ProviderResult(text="", success=False, error=last_error)
|
|
|
|
# Dispatch by provider type
|
|
ptype = getattr(provider, "type", None) or "subprocess"
|
|
try:
|
|
if ptype == "subprocess":
|
|
result = call_provider_subprocess(provider, prompt, timeout, max_tokens)
|
|
elif ptype == "api":
|
|
result = call_provider_api(provider, prompt, timeout, max_tokens)
|
|
elif ptype == "pty":
|
|
result = call_provider_pty(provider, prompt, timeout, max_tokens)
|
|
else:
|
|
return try_fallback(f"Unknown provider type: {ptype}")
|
|
|
|
if result.success:
|
|
return result
|
|
return try_fallback(result.error or "Provider call failed")
|
|
except Exception as e:
|
|
return try_fallback(f"Provider error: {str(e)}")
|
|
|
|
|
|
def call_provider_subprocess(provider: "Provider", prompt: str, timeout: int, max_tokens: Optional[int]) -> ProviderResult:
|
|
"""Invoke a CLI provider via subprocess.run with stdin piping."""
|
|
cmd = os.path.expandvars(provider.command)
|
|
|
|
# Append max_tokens flags only for CLIs whose interfaces support them.
|
|
if max_tokens is not None:
|
|
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, agy, codex, and crush use their configured/default limits.
|
|
|
|
# Check if base command exists (use shlex for proper quote handling)
|
|
try:
|
|
cmd_parts = shlex.split(cmd)
|
|
base_cmd = cmd_parts[0] if cmd_parts else cmd.split()[0]
|
|
except ValueError:
|
|
base_cmd = cmd.split()[0]
|
|
|
|
base_cmd_expanded = os.path.expanduser(base_cmd)
|
|
if not shutil.which(base_cmd_expanded) and not os.path.isfile(base_cmd_expanded):
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=f"Command '{base_cmd}' not found. Is it installed and in PATH?\n\nTo install AI providers, run: cmdforge providers install"
|
|
)
|
|
|
|
try:
|
|
result = subprocess.run(
|
|
cmd,
|
|
shell=True,
|
|
input=prompt,
|
|
capture_output=True,
|
|
text=True,
|
|
timeout=timeout
|
|
)
|
|
|
|
if result.returncode != 0:
|
|
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 ProviderResult(text="", success=False, error=error_msg)
|
|
|
|
# Warn if output is empty (provider ran but returned nothing)
|
|
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:
|
|
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"
|
|
model_id = model_match.group(1) if model_match else "unknown"
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=(
|
|
f"Model '{model_id}' from provider '{provider_id}' is not available.\n\n"
|
|
f"To fix this, either:\n"
|
|
f" 1. Run 'opencode' to connect the {provider_id} provider\n"
|
|
f" 2. Use --provider to pick a different model (e.g., --provider opencode-pickle)\n"
|
|
f" 3. Run 'cmdforge ui' to edit the tool's default provider"
|
|
)
|
|
)
|
|
|
|
stderr_hint = f" (stderr: {stderr[:200]}...)" if len(stderr) > 200 else (f" (stderr: {stderr})" if stderr else "")
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=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=clean_stdout, success=True)
|
|
|
|
except subprocess.TimeoutExpired:
|
|
return ProviderResult(text="", success=False, error=f"Provider timed out after {timeout} seconds")
|
|
except Exception as e:
|
|
return ProviderResult(text="", success=False, error=f"Provider error: {str(e)}")
|
|
|
|
|
|
def call_provider_api(provider: "Provider", prompt: str, timeout: int, max_tokens: Optional[int]) -> ProviderResult:
|
|
"""Call an OpenAI-compatible HTTP API endpoint.
|
|
|
|
The provider.command is the endpoint URL (e.g. https://api.openrouter.ai/api/v1/chat/completions).
|
|
The provider.model is the model ID (e.g. "deepseek/deepseek-chat").
|
|
The provider.api_key_env names the environment variable holding the API key.
|
|
"""
|
|
if not provider.model:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=f"API provider '{provider.name}' has no model configured."
|
|
)
|
|
|
|
env_var = provider.api_key_env or _infer_api_key_env(provider.name)
|
|
api_key = os.environ.get(env_var) if env_var else None
|
|
if not api_key:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=f"API key not set. Set the {env_var} environment variable to use '{provider.name}'."
|
|
)
|
|
|
|
try:
|
|
import requests
|
|
except ImportError:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error="The 'requests' package is required for API providers. Install with: pip install requests"
|
|
)
|
|
|
|
endpoint = os.path.expandvars(provider.command)
|
|
if not endpoint.endswith("/chat/completions"):
|
|
endpoint = endpoint.rstrip("/") + "/chat/completions"
|
|
|
|
body = {
|
|
"model": provider.model,
|
|
"messages": [{"role": "user", "content": prompt}],
|
|
}
|
|
if max_tokens is not None:
|
|
body["max_tokens"] = max_tokens
|
|
|
|
headers = {
|
|
"Authorization": f"Bearer {api_key}",
|
|
"Content-Type": "application/json",
|
|
}
|
|
|
|
try:
|
|
response = requests.post(endpoint, json=body, headers=headers, timeout=timeout)
|
|
if response.status_code != 200:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=f"API returned HTTP {response.status_code}: {response.text[:500]}"
|
|
)
|
|
|
|
data = response.json()
|
|
text = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
|
if not text:
|
|
return ProviderResult(text="", success=False, error="API returned empty response")
|
|
|
|
return ProviderResult(text=text, success=True)
|
|
except requests.Timeout:
|
|
return ProviderResult(text="", success=False, error=f"API timed out after {timeout} seconds")
|
|
except Exception as e:
|
|
return ProviderResult(text="", success=False, error=f"API error: {str(e)}")
|
|
|
|
|
|
def _infer_api_key_env(provider_name: str) -> str:
|
|
"""Guess the API key env var name from provider name (fallback if api_key_env not set)."""
|
|
name_upper = provider_name.upper().replace("-", "_")
|
|
return f"{name_upper}_API_KEY"
|
|
|
|
|
|
def call_provider_pty(provider: "Provider", prompt: str, timeout: int, max_tokens: Optional[int]) -> ProviderResult:
|
|
"""Wrap an interactive CLI via a pseudo-terminal (experimental).
|
|
|
|
Requires pty_config on the provider with:
|
|
- prompt_pattern: regex/expect pattern for the CLI's ready prompt
|
|
- response_pattern: regex/expect pattern marking end of response
|
|
- exit_command: command to cleanly exit the CLI (e.g. "/exit", "\\q")
|
|
"""
|
|
if not provider.pty_config:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error=f"PTY provider '{provider.name}' has no pty_config set. Cannot wrap interactive CLI."
|
|
)
|
|
|
|
try:
|
|
import pexpect
|
|
except ImportError:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error="The 'pexpect' package is required for PTY providers. Install with: pip install pexpect"
|
|
)
|
|
|
|
cfg = provider.pty_config
|
|
prompt_pattern = cfg.get("prompt_pattern")
|
|
response_pattern = cfg.get("response_pattern")
|
|
exit_command = cfg.get("exit_command", "/exit")
|
|
|
|
if not prompt_pattern or not response_pattern:
|
|
return ProviderResult(
|
|
text="",
|
|
success=False,
|
|
error="pty_config must include prompt_pattern and response_pattern"
|
|
)
|
|
|
|
cmd = os.path.expandvars(provider.command)
|
|
|
|
child = None
|
|
try:
|
|
child = pexpect.spawn(cmd, timeout=timeout, encoding='utf-8')
|
|
# Wait for the CLI's ready prompt
|
|
child.expect(prompt_pattern)
|
|
# Send the user's prompt
|
|
child.sendline(prompt)
|
|
# Wait for the response to complete
|
|
child.expect(response_pattern)
|
|
text = child.before
|
|
# Cleanly exit
|
|
child.sendline(exit_command)
|
|
return ProviderResult(text=strip_ansi(text), success=True)
|
|
except pexpect.TIMEOUT:
|
|
return ProviderResult(text="", success=False, error=f"PTY provider timed out after {timeout} seconds")
|
|
except pexpect.EOF:
|
|
return ProviderResult(text="", success=False, error="PTY provider exited unexpectedly")
|
|
except Exception as e:
|
|
return ProviderResult(text="", success=False, error=f"PTY provider error: {str(e)}")
|
|
finally:
|
|
if child is not None and child.isalive():
|
|
child.close(force=True)
|
|
|
|
|
|
def mock_provider(prompt: str) -> ProviderResult:
|
|
"""
|
|
Return a mock response for testing.
|
|
|
|
Returns structured JSON that matches the default output schema,
|
|
allowing tools to be tested without real AI providers.
|
|
|
|
Args:
|
|
prompt: The prompt (used for generating mock response)
|
|
|
|
Returns:
|
|
ProviderResult with mock JSON response
|
|
"""
|
|
import json
|
|
|
|
lines = prompt.strip().split('\n')
|
|
preview = lines[0][:50] + "..." if len(lines[0]) > 50 else lines[0]
|
|
|
|
# Return structured JSON matching default schema
|
|
mock_response = {
|
|
"output": f"[MOCK] {preview}",
|
|
"reasoning": f"Mock response for prompt with {len(prompt)} chars, {len(lines)} lines."
|
|
}
|
|
|
|
return ProviderResult(
|
|
text=json.dumps(mock_response),
|
|
success=True
|
|
)
|