"""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, 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_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})