CmdForge/tests/test_providers.py

1078 lines
37 KiB
Python

"""Tests for providers.py - AI provider abstraction."""
import tempfile
from pathlib import Path
from unittest.mock import patch, MagicMock
import pytest
import yaml
from cmdforge.providers import (
Provider, ProviderResult, ProviderExecutionPolicy,
load_providers, save_providers, get_provider, get_providers_file,
add_provider, delete_provider,
call_provider, discover_installed_providers, mock_provider,
DEFAULT_PROVIDERS, strip_ansi
)
class TestProvider:
"""Tests for Provider dataclass."""
def test_create_basic(self):
provider = Provider(name="test", command="echo test")
assert provider.name == "test"
assert provider.command == "echo test"
assert provider.description == ""
def test_create_with_description(self):
provider = Provider(
name="claude",
command="claude -p",
description="Anthropic Claude"
)
assert provider.description == "Anthropic Claude"
def test_to_dict(self):
provider = Provider(
name="gpt4",
command="openai chat",
description="OpenAI GPT-4"
)
d = provider.to_dict()
assert d["name"] == "gpt4"
assert d["command"] == "openai chat"
assert d["description"] == "OpenAI GPT-4"
def test_from_dict(self):
data = {
"name": "local",
"command": "ollama run llama2",
"description": "Local Ollama"
}
provider = Provider.from_dict(data)
assert provider.name == "local"
assert provider.command == "ollama run llama2"
assert provider.description == "Local Ollama"
def test_from_dict_missing_description(self):
data = {"name": "test", "command": "test cmd"}
provider = Provider.from_dict(data)
assert provider.description == ""
def test_legacy_policy_field_aliases_are_canonicalized(self):
provider = Provider.from_dict({
"name": "legacy-policy",
"command": "provider",
"model_id": "model-a",
"max_input_tokens": 8192,
})
assert provider.model == "model-a"
assert provider.max_context_tokens == 8192
assert provider.to_dict()["model"] == "model-a"
assert provider.to_dict()["max_context_tokens"] == 8192
def test_roundtrip(self):
original = Provider(
name="custom",
command="my-ai --prompt",
description="My custom provider"
)
restored = Provider.from_dict(original.to_dict())
assert restored.name == original.name
assert restored.command == original.command
assert restored.description == original.description
def test_legacy_positional_fallback_is_preserved(self):
provider = Provider("primary", "primary-cmd", "Primary", "backup")
assert provider.fallback == "backup"
assert provider.type == "subprocess"
def test_extended_fields_roundtrip(self):
original = Provider(
name="api-test",
command="https://example.test/v1",
description="Test API",
fallback="mock",
type="api",
model="test/model",
tags=["api", "test"],
install={"cost": "test"},
fallback_chain=["backup", "mock"],
api_key_env="TEST_API_KEY",
pty_config={"prompt_pattern": ">"},
locality="remote",
capabilities=["text", "structured-json"],
model_digest="sha256:abc",
cost_class="low",
latency_class="fast",
max_context_tokens=131072,
data_policy="internal",
)
assert Provider.from_dict(original.to_dict()) == original
class TestProviderResult:
"""Tests for ProviderResult dataclass."""
def test_success_result(self):
result = ProviderResult(text="Hello!", success=True)
assert result.text == "Hello!"
assert result.success is True
assert result.error is None
def test_error_result(self):
result = ProviderResult(text="", success=False, error="API timeout")
assert result.text == ""
assert result.success is False
assert result.error == "API timeout"
class TestMockProvider:
"""Tests for the mock provider."""
def test_mock_returns_success(self):
result = mock_provider("Test prompt")
assert result.success is True
# Mock returns structured JSON with output and reasoning
import json
parsed = json.loads(result.text)
assert "[MOCK]" in parsed["output"]
def test_mock_includes_prompt_info(self):
result = mock_provider("This is a test prompt")
import json
parsed = json.loads(result.text)
# Prompt info is in reasoning field
assert "chars" in parsed["reasoning"]
assert "lines" in parsed["reasoning"]
def test_mock_shows_first_line_preview(self):
result = mock_provider("First line here\nSecond line")
assert "First line" in result.text
def test_mock_truncates_long_first_line(self):
long_line = "x" * 100
result = mock_provider(long_line)
assert "..." in result.text
def test_mock_counts_lines(self):
prompt = "line1\nline2\nline3"
result = mock_provider(prompt)
assert "3 lines" in result.text
class TestProviderPersistence:
"""Tests for provider save/load operations."""
@pytest.fixture
def temp_providers_file(self, tmp_path):
"""Create a temporary providers file."""
providers_file = tmp_path / ".cmdforge" / "providers.yaml"
with patch('cmdforge.providers.PROVIDERS_FILE', providers_file):
yield providers_file
def test_save_and_load_providers(self, temp_providers_file):
providers = [
Provider("test1", "cmd1", "Description 1"),
Provider("test2", "cmd2", "Description 2")
]
save_providers(providers)
loaded = load_providers()
assert len(loaded) == 2
assert loaded[0].name == "test1"
assert loaded[1].name == "test2"
def test_save_providers_uses_private_file_permissions(self, temp_providers_file):
save_providers([Provider("test", "cmd")])
assert oct(temp_providers_file.stat().st_mode & 0o777) == "0o600"
def test_first_run_without_discovered_providers_writes_defaults(
self, temp_providers_file, capsys
):
with patch(
"cmdforge.providers.discover_installed_providers", return_value=[]
):
result = get_providers_file()
assert result == temp_providers_file
assert temp_providers_file.is_file()
assert "No AI providers detected" in capsys.readouterr().err
def test_malformed_access_policy_fails_closed(self, temp_providers_file, capsys):
temp_providers_file.parent.mkdir(parents=True)
temp_providers_file.write_text(yaml.safe_dump({
"version": 2,
"providers": [{
"name": "claude",
"command": "claude -p",
"tools": "should-have-been-a-list",
}],
}))
assert load_providers() == []
assert "Failed to load provider configuration" in capsys.readouterr().err
def test_legacy_config_gains_missing_defaults_without_losing_custom_provider(self, temp_providers_file):
temp_providers_file.parent.mkdir(parents=True)
temp_providers_file.write_text(yaml.safe_dump({
"providers": [
{
"name": "custom",
"command": "custom-command",
"description": "Keep me",
"fallback": "mock",
}
]
}))
loaded = load_providers()
custom = next(provider for provider in loaded if provider.name == "custom")
assert custom.command == "custom-command"
assert custom.fallback == "mock"
assert any(provider.name == "opencode-free" for provider in loaded)
assert any(provider.name == "openrouter" for provider in loaded)
saved = yaml.safe_load(temp_providers_file.read_text())
assert saved["version"] == 2
def test_get_provider_exists(self, temp_providers_file):
providers = [
Provider("target", "target-cmd", "Target provider")
]
save_providers(providers)
result = get_provider("target")
assert result is not None
assert result.name == "target"
assert result.command == "target-cmd"
def test_get_provider_not_exists(self, temp_providers_file):
save_providers([])
result = get_provider("nonexistent")
assert result is None
def test_add_new_provider(self, temp_providers_file):
save_providers([])
new_provider = Provider("new", "new-cmd", "New provider")
add_provider(new_provider)
loaded = load_providers()
assert any(p.name == "new" for p in loaded)
def test_add_provider_updates_existing(self, temp_providers_file):
save_providers([
Provider("existing", "old-cmd", "Old description")
])
updated = Provider("existing", "new-cmd", "New description")
add_provider(updated)
loaded = load_providers()
existing = next(p for p in loaded if p.name == "existing")
assert existing.command == "new-cmd"
assert existing.description == "New description"
def test_delete_provider(self, temp_providers_file):
save_providers([
Provider("keep", "keep-cmd"),
Provider("delete", "delete-cmd")
])
result = delete_provider("delete")
assert result is True
loaded = load_providers()
assert not any(p.name == "delete" for p in loaded)
assert any(p.name == "keep" for p in loaded)
def test_delete_nonexistent_provider(self, temp_providers_file):
save_providers([])
result = delete_provider("nonexistent")
assert result is False
class TestDefaultProviders:
"""Tests for default providers."""
def test_default_providers_exist(self):
assert len(DEFAULT_PROVIDERS) > 0
def test_mock_in_defaults(self):
assert any(p.name == "mock" for p in DEFAULT_PROVIDERS)
def test_claude_in_defaults(self):
assert any(p.name == "claude" for p in DEFAULT_PROVIDERS)
def test_all_defaults_have_commands(self):
for provider in DEFAULT_PROVIDERS:
assert provider.command, f"Provider {provider.name} has no command"
class TestCallProvider:
"""Tests for call_provider function."""
@pytest.fixture
def temp_providers_file(self, tmp_path):
"""Create a temporary providers file."""
providers_file = tmp_path / ".cmdforge" / "providers.yaml"
with patch('cmdforge.providers.PROVIDERS_FILE', providers_file):
yield providers_file
def test_call_mock_provider(self, temp_providers_file):
"""Mock provider should work without subprocess."""
result = call_provider("mock", "Test prompt")
assert result.success is True
# Mock returns structured JSON
import json
parsed = json.loads(result.text)
assert "[MOCK]" in parsed["output"]
def test_call_nonexistent_provider(self, temp_providers_file):
save_providers([])
result = call_provider("nonexistent", "Test")
assert result.success is False
assert "not found" in result.error.lower()
@patch('subprocess.run')
def test_call_provider_rejects_non_integer_max_tokens(self, mock_run, temp_providers_file):
save_providers([Provider("claude", "claude -p")])
result = call_provider("claude", "Test", max_tokens="1; touch /tmp/pwned")
assert result.success is False
assert "max_tokens" in result.error
mock_run.assert_not_called()
@patch('subprocess.run')
@patch('shutil.which')
def test_call_real_provider_success(self, mock_which, mock_run, temp_providers_file):
# Setup
mock_which.return_value = "/usr/bin/echo"
mock_run.return_value = MagicMock(
returncode=0,
stdout="AI response here",
stderr=""
)
save_providers([Provider("echo-test", "echo test")])
result = call_provider("echo-test", "Prompt")
assert result.success is True
assert result.text == "AI response here"
@patch('subprocess.run')
@patch('shutil.which')
def test_call_provider_nonzero_exit(self, mock_which, mock_run, temp_providers_file):
mock_which.return_value = "/usr/bin/cmd"
mock_run.return_value = MagicMock(
returncode=1,
stdout="",
stderr="Error occurred"
)
save_providers([Provider("failing", "failing-cmd")])
result = call_provider("failing", "Prompt")
assert result.success is False
assert "exited with code 1" in result.error
@patch('subprocess.run')
@patch('shutil.which')
def test_call_provider_empty_output(self, mock_which, mock_run, temp_providers_file):
mock_which.return_value = "/usr/bin/cmd"
mock_run.return_value = MagicMock(
returncode=0,
stdout=" ", # Only whitespace
stderr=""
)
save_providers([Provider("empty", "empty-cmd")])
result = call_provider("empty", "Prompt")
assert result.success is False
assert "empty output" in result.error.lower()
@patch('subprocess.run')
@patch('shutil.which')
def test_call_provider_timeout(self, mock_which, mock_run, temp_providers_file):
import subprocess
mock_which.return_value = "/usr/bin/slow"
mock_run.side_effect = subprocess.TimeoutExpired(cmd="slow", timeout=10)
save_providers([Provider("slow", "slow-cmd")])
result = call_provider("slow", "Prompt", timeout=10)
assert result.success is False
assert "timed out" in result.error.lower()
@patch('shutil.which')
def test_call_provider_command_not_found(self, mock_which, temp_providers_file):
mock_which.return_value = None
save_providers([Provider("missing", "nonexistent-binary")])
result = call_provider("missing", "Prompt")
assert result.success is False
assert "not found" in result.error.lower()
@patch('subprocess.run')
@patch('shutil.which')
def test_provider_receives_prompt_as_stdin(self, mock_which, mock_run, temp_providers_file):
mock_which.return_value = "/usr/bin/cat"
mock_run.return_value = MagicMock(returncode=0, stdout="output", stderr="")
save_providers([Provider("cat", "cat")])
call_provider("cat", "My prompt text")
# Verify prompt was passed as input
call_kwargs = mock_run.call_args[1]
assert call_kwargs["input"] == "My prompt text"
def test_environment_variable_expansion(self, temp_providers_file):
"""Provider commands should expand $HOME etc."""
save_providers([
Provider("home-test", "$HOME/bin/my-ai")
])
# This will fail because the command doesn't exist,
# but we can check the error message to verify expansion happened
result = call_provider("home-test", "Test")
# The error should mention the expanded path, not $HOME
assert "$HOME" not in result.error
class TestProviderCommandParsing:
"""Tests for command parsing with shlex."""
@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('shutil.which')
def test_command_with_quotes(self, mock_which, temp_providers_file):
"""Commands with quotes should be parsed correctly."""
mock_which.return_value = None # Will fail, but we test parsing
save_providers([
Provider("quoted", 'my-cmd --arg "value with spaces"')
])
result = call_provider("quoted", "Test")
# Should fail at command-not-found, not at parsing
assert "not found" in result.error.lower()
assert "my-cmd" in result.error
@patch('shutil.which')
def test_command_with_env_vars(self, mock_which, temp_providers_file):
"""Environment variables in commands should be expanded."""
import os
mock_which.return_value = None
save_providers([
Provider("env-test", "$HOME/.local/bin/my-ai")
])
result = call_provider("env-test", "Test")
# Error should show expanded path
home = os.environ.get("HOME", "")
assert "$HOME" not in result.error or home in result.error
class TestProviderFallback:
"""Tests for provider fallback functionality."""
@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
def test_provider_with_fallback_to_dict(self):
"""Provider with fallback should include it in dict."""
provider = Provider(
name="cloud",
command="cloud-ai",
description="Cloud AI",
fallback="local"
)
d = provider.to_dict()
assert d["fallback"] == "local"
def test_provider_without_fallback_to_dict(self):
"""Provider without fallback should not include fallback key."""
provider = Provider(name="simple", command="simple-ai")
d = provider.to_dict()
assert "fallback" not in d
def test_provider_from_dict_with_fallback(self):
"""Provider.from_dict should handle fallback field."""
data = {
"name": "cloud",
"command": "cloud-ai",
"fallback": "local"
}
provider = Provider.from_dict(data)
assert provider.fallback == "local"
def test_provider_from_dict_without_fallback(self):
"""Provider.from_dict should handle missing fallback."""
data = {"name": "simple", "command": "simple-ai"}
provider = Provider.from_dict(data)
assert provider.fallback is None
def test_fallback_triggers_on_failure(self, temp_providers_file):
"""When primary provider fails, fallback should be tried."""
save_providers([
Provider("primary", "nonexistent-cmd", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("primary", "Test prompt")
# Fallback to mock should succeed
assert result.success is True
assert "[MOCK]" in result.text
assert result.requested_provider == "primary"
assert result.actual_provider == "mock"
assert result.attempted_providers == ["primary", "mock"]
assert result.fallback_used is True
def test_no_fallback_fails_closed(self, temp_providers_file):
save_providers([
Provider("primary", "nonexistent-cmd", fallback="mock"),
Provider("mock", "mock"),
])
result = call_provider(
"primary", "Test prompt",
execution_policy=ProviderExecutionPolicy(allow_fallback=False),
)
assert result.success is False
assert result.requested_provider == "primary"
assert result.actual_provider is None
assert result.attempted_providers == ["primary"]
assert result.fallback_used is False
def test_private_data_requires_explicit_compatible_policy(
self, temp_providers_file
):
save_providers([
Provider(
"local", "unused", locality="local",
capabilities=["structured-json"],
),
])
result = call_provider(
"local", "private packet",
execution_policy=ProviderExecutionPolicy(
allow_fallback=False,
required_locality="local",
required_capabilities=("structured-json",),
data_classification="private",
),
)
assert result.success is False
assert "data_policy is unspecified" in result.error
def test_require_local_rejects_ollama_redirected_to_remote_host(
self, temp_providers_file, monkeypatch
):
save_providers([
Provider(
"ollama-private", "ollama run safe", locality="local",
data_policy="private",
),
])
monkeypatch.setenv("OLLAMA_HOST", "gpu.example.test:11434")
result = call_provider(
"ollama-private", "private packet",
execution_policy=ProviderExecutionPolicy(
allow_fallback=False,
required_locality="local",
data_classification="private",
),
)
assert result.success is False
assert "locality is remote" in result.error
def test_ollama_digest_is_observed_from_local_api(
self, temp_providers_file, monkeypatch
):
save_providers([
Provider(
"ollama-safe", "ollama run safe:latest", model="safe:latest",
locality="local", data_policy="private",
),
])
monkeypatch.delenv("OLLAMA_HOST", raising=False)
class Response:
def __enter__(self):
return self
def __exit__(self, *args):
return False
def read(self):
return (
b'{"models":[{"name":"safe:latest","model":"safe:latest",'
b'"digest":"abcdef"}]}'
)
with (
patch("urllib.request.urlopen", return_value=Response()),
patch(
"cmdforge.providers.call_provider_subprocess",
return_value=ProviderResult(text="safe", success=True),
),
):
result = call_provider(
"ollama-safe", "private",
execution_policy=ProviderExecutionPolicy(
allow_fallback=False,
required_locality="local",
data_classification="private",
require_model_identity=True,
require_model_digest=True,
),
)
assert result.model == "safe:latest"
assert result.model_digest == "sha256:abcdef"
assert result.model_identity_source == "ollama-api"
def test_policy_eligible_provider_reports_verified_metadata(
self, temp_providers_file
):
provider = Provider(
"local", "local-ai", model="model-a",
model_digest="sha256:123", locality="local",
capabilities=["structured-json"], data_policy="private",
)
save_providers([provider])
low_level = ProviderResult(text='{"ok": true}', success=True)
with patch(
"cmdforge.providers.call_provider_subprocess",
return_value=low_level,
):
result = call_provider(
"local", "private packet",
execution_policy=ProviderExecutionPolicy(
allow_fallback=False,
required_locality="local",
required_capabilities=("structured-json",),
data_classification="private",
require_model_identity=True,
require_model_digest=True,
),
)
assert result.success is True
assert result.provenance() == {
"success": True,
"requested_provider": "local",
"actual_provider": "local",
"attempted_providers": ["local"],
"fallback_used": False,
"model": "model-a",
"model_digest": "sha256:123",
"locality": "local",
"model_identity_source": "provider-config",
}
def test_fallback_chain_skips_policy_ineligible_provider(
self, temp_providers_file
):
save_providers([
Provider(
"primary", "primary", fallback_chain=["remote", "local"],
locality="local", data_policy="private",
),
Provider(
"remote", "remote", locality="remote", data_policy="private",
),
Provider(
"local", "local", locality="local", data_policy="private",
model="safe-model",
),
])
def invoke(provider, prompt, timeout, max_tokens):
if provider.name == "primary":
return ProviderResult(text="", success=False, error="offline")
assert provider.name == "local"
return ProviderResult(text="safe", success=True)
with patch(
"cmdforge.providers.call_provider_subprocess", side_effect=invoke
) as low_level:
result = call_provider(
"primary", "private packet",
execution_policy=ProviderExecutionPolicy(
required_locality="local",
data_classification="private",
),
)
assert result.success is True
assert result.actual_provider == "local"
assert result.attempted_providers == ["primary", "remote", "local"]
assert [call.args[0].name for call in low_level.call_args_list] == [
"primary", "local"
]
def test_fallback_chain(self, temp_providers_file):
"""Fallback can chain to another provider with fallback."""
save_providers([
Provider("first", "nonexistent1", fallback="second"),
Provider("second", "nonexistent2", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("first", "Test prompt")
# Should chain: first -> second -> mock
assert result.success is True
assert "[MOCK]" in result.text
def test_fallback_chain_continues_after_failed_candidate(self, temp_providers_file):
save_providers([
Provider("primary", "missing-primary", fallback_chain=["secondary", "mock"]),
Provider("secondary", "missing-secondary"),
Provider("mock", "mock"),
])
result = call_provider("primary", "Test prompt")
assert result.success is True
assert "[MOCK]" in result.text
def test_fallback_prevents_infinite_loop(self, temp_providers_file):
"""Circular fallback references should not cause infinite loop."""
save_providers([
Provider("a", "nonexistent-a", fallback="b"),
Provider("b", "nonexistent-b", fallback="a"),
])
result = call_provider("a", "Test prompt")
# Should fail gracefully, not loop forever
assert result.success is False
assert "not found" in result.error.lower()
@patch('subprocess.run')
@patch('shutil.which')
def test_fallback_on_timeout(self, mock_which, mock_run, temp_providers_file):
"""Fallback should trigger on timeout."""
import subprocess
mock_which.return_value = "/usr/bin/slow-ai"
mock_run.side_effect = subprocess.TimeoutExpired("slow-ai", 5)
save_providers([
Provider("slow", "slow-ai", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("slow", "Test prompt", timeout=5)
assert result.success is True
assert "[MOCK]" in result.text
@patch('subprocess.run')
@patch('shutil.which')
def test_fallback_on_nonzero_exit(self, mock_which, mock_run, temp_providers_file):
"""Fallback should trigger on non-zero exit code."""
mock_which.return_value = "/usr/bin/failing-ai"
mock_run.return_value = MagicMock(
returncode=1,
stdout="",
stderr="API error"
)
save_providers([
Provider("failing", "failing-ai", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("failing", "Test prompt")
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]
class TestApiProviders:
@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('requests.post')
def test_openai_compatible_api_dispatch(self, mock_post, temp_providers_file):
response = MagicMock(status_code=200)
response.json.return_value = {
"choices": [{"message": {"content": "API response"}}]
}
mock_post.return_value = response
save_providers([
Provider(
"test-api",
"https://example.test/v1",
type="api",
model="example/model",
api_key_env="TEST_API_KEY",
)
])
with patch.dict('os.environ', {"TEST_API_KEY": "secret"}):
result = call_provider("test-api", "Prompt", max_tokens=128)
assert result.success is True
assert result.text == "API response"
mock_post.assert_called_once_with(
"https://example.test/v1/chat/completions",
json={
"model": "example/model",
"messages": [{"role": "user", "content": "Prompt"}],
"max_tokens": 128,
},
headers={
"Authorization": "Bearer secret",
"Content-Type": "application/json",
},
timeout=300,
)
def test_api_provider_requires_configured_key(self, temp_providers_file):
save_providers([
Provider(
"test-api",
"https://example.test/v1",
type="api",
model="example/model",
api_key_env="MISSING_TEST_API_KEY",
)
])
with patch.dict('os.environ', {"MISSING_TEST_API_KEY": ""}):
result = call_provider("test-api", "Prompt")
assert result.success is False
assert "MISSING_TEST_API_KEY" in result.error
class TestProviderDiscovery:
def test_ollama_model_names_include_tags_and_are_unique(self):
ollama_list = MagicMock(
returncode=0,
stdout=(
"NAME ID SIZE MODIFIED\n"
"hermes4.3:latest abc 1 GB now\n"
"hermes4.3:q4_k_m def 1 GB now\n"
),
)
def which(binary):
return "/usr/bin/ollama" if binary == "ollama" else None
with patch('cmdforge.providers.shutil.which', side_effect=which):
with patch('cmdforge.providers.subprocess.run', return_value=ollama_list):
found = discover_installed_providers()
models = [item for item in found if item["source"] == "ollama-model"]
assert len(models) == 2
assert len({item["name"] for item in models}) == 2
assert models[0]["name"] == "ollama-hermes4-3-latest"
assert models[1]["name"] == "ollama-hermes4-3-q4-k-m"
class TestProviderAccessControl:
def test_tools_allowlist_defaults_to_none(self):
provider = Provider("test", "cmd")
assert provider.tools is None
assert provider.mcp_servers is None
def test_tools_allowlist_roundtrip(self):
provider = Provider("test", "cmd", tools=["tool-a", "tool-b"])
d = provider.to_dict()
assert d["tools"] == ["tool-a", "tool-b"]
restored = Provider.from_dict(d)
assert restored.tools == ["tool-a", "tool-b"]
def test_mcp_servers_roundtrip(self):
provider = Provider("test", "cmd", mcp_servers=["filesystem"])
d = provider.to_dict()
assert d["mcp_servers"] == ["filesystem"]
restored = Provider.from_dict(d)
assert restored.mcp_servers == ["filesystem"]
def test_tools_not_in_dict_when_none(self):
provider = Provider("test", "cmd")
assert "tools" not in provider.to_dict()
assert "mcp_servers" not in provider.to_dict()
def test_empty_allowlists_survive_roundtrip(self):
provider = Provider("locked", "cmd", tools=[], mcp_servers=[])
serialized = provider.to_dict()
assert serialized["tools"] == []
assert serialized["mcp_servers"] == []
restored = Provider.from_dict(serialized)
assert restored.tools == []
assert restored.mcp_servers == []
@pytest.mark.parametrize(
("field", "value"),
[
("tools", "tool-a"),
("tools", [""]),
("tools", [1]),
("mcp_servers", "filesystem"),
("mcp_servers", [""]),
("mcp_servers", [1]),
],
)
def test_rejects_invalid_access_policies(self, field, value):
with pytest.raises(ValueError, match=field):
Provider("test", "cmd", **{field: value})