#!/usr/bin/env python3 """ Test suite for the model command functionality. Tests all handle methods and input possibilities for the model command. """ import os import sys import pytest import datetime from unittest.mock import patch, Mock, MagicMock # Add src to path sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "src")) from cai.repl.commands.model import ModelCommand from cai.repl.commands.base import Command class TestModelCommand: """Test cases for ModelCommand.""" @pytest.fixture(autouse=True) def setup_and_cleanup(self): """Setup and cleanup for each test.""" # Set up test environment os.environ["CAI_TELEMETRY"] = "false" os.environ["CAI_TRACING"] = "false" # Store original CAI_MODEL if it exists self.original_model = os.environ.get("CAI_MODEL") yield # Restore original CAI_MODEL or remove if it didn't exist if self.original_model is not None: os.environ["CAI_MODEL"] = self.original_model elif "CAI_MODEL" in os.environ: del os.environ["CAI_MODEL"] @pytest.fixture def model_command(self): """Create a ModelCommand instance for testing.""" return ModelCommand() @pytest.fixture def mock_litellm_response(self): """Create a mock response for LiteLLM model data.""" return { "gpt-4": { "input_cost_per_token": 0.00003, "output_cost_per_token": 0.00006, "max_tokens": 8192, "supports_function_calling": True, "supports_vision": False, "litellm_provider": "openai", }, "claude-3-sonnet-20240229": { "input_cost_per_token": 0.000015, "output_cost_per_token": 0.000075, "max_tokens": 200000, "supports_function_calling": True, "supports_vision": True, "litellm_provider": "anthropic", }, "deepseek/deepseek-v3": { "input_cost_per_token": 0.000001, "output_cost_per_token": 0.000002, "max_tokens": 128000, "supports_function_calling": True, "supports_vision": False, "litellm_provider": "deepseek", }, } @pytest.fixture def mock_ollama_response(self): """Create a mock response for Ollama models.""" return { "models": [ { "name": "llama3", "size": 4661211648, # ~4.3 GB }, { "name": "mistral:7b", "size": 7365960192, # ~6.9 GB }, ] } def test_command_initialization(self, model_command): """Test that ModelCommand initializes correctly.""" assert model_command.name == "/model" assert model_command.description == "View or change the current LLM model" assert model_command.aliases == ["/mod"] # Check that cached models and numbers are initialized assert hasattr(model_command, "cached_models") assert hasattr(model_command, "cached_model_numbers") assert hasattr(model_command, "last_model_fetch") @patch("requests.get") def test_handle_no_args_with_mock_data( self, mock_get, model_command, mock_litellm_response, mock_ollama_response ): """Test showing current model and available models with no arguments.""" # Mock LiteLLM response mock_litellm = Mock() mock_litellm.status_code = 200 mock_litellm.json.return_value = mock_litellm_response # Mock Ollama response mock_ollama = Mock() mock_ollama.status_code = 200 mock_ollama.json.return_value = mock_ollama_response # Configure the mock to return different responses based on URL def side_effect(url, timeout=None): if "litellm" in url: return mock_litellm elif "ollama" in url: return mock_ollama else: return Mock(status_code=404) mock_get.side_effect = side_effect # Set a model first os.environ["CAI_MODEL"] = "gpt-4" result = model_command.handle([]) assert result is True @patch("requests.get") def test_handle_select_model_by_name(self, mock_get, model_command, mock_litellm_response): """Test selecting a model by name.""" # Mock LiteLLM response mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response mock_get.return_value = mock_response result = model_command.handle(["gpt-4"]) assert result is True assert os.environ.get("CAI_MODEL") == "gpt-4" @patch("requests.get") def test_handle_select_model_by_number(self, mock_get, model_command, mock_litellm_response): """Test selecting a model by number.""" # Mock LiteLLM response mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response mock_get.return_value = mock_response # First call to populate cache model_command.handle([]) # Then select by number result = model_command.handle(["1"]) assert result is True assert "CAI_MODEL" in os.environ @patch("requests.get") def test_rejects_unknown_model_name(self, mock_get, model_command, mock_litellm_response): """Unknown names are rejected and do not change CAI_MODEL.""" mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response mock_get.return_value = mock_response before = os.environ.get("CAI_MODEL") result = model_command.handle(["custom-model-name-not-in-catalog"]) assert result is True assert os.environ.get("CAI_MODEL") == before @patch("requests.get") def test_handle_with_network_error(self, mock_get, model_command): """Test handling when network requests fail.""" # Mock network failure mock_get.side_effect = Exception("Network error") result = model_command.handle([]) assert result is True # Should still work, just without external data @patch("requests.get") def test_handle_model_pricing_data_error(self, mock_get, model_command): """Test handling when LiteLLM API returns error.""" # Mock HTTP error mock_response = Mock() mock_response.status_code = 404 mock_get.return_value = mock_response result = model_command.handle([]) assert result is True # Should still work with built-in models def test_command_base_functionality(self, model_command): """Test that the command inherits from base Command properly.""" assert isinstance(model_command, Command) assert model_command.name == "/model" assert "/mod" in model_command.aliases class TestModelShowSubcommand: """Tests for ``/model show`` subcommand.""" @pytest.fixture(autouse=True) def setup_and_cleanup(self): """Setup and cleanup for each test.""" # Set up test environment os.environ["CAI_TELEMETRY"] = "false" os.environ["CAI_TRACING"] = "false" yield @pytest.fixture def model_show_command(self): """Fresh ``ModelCommand`` (invoke ``handle([\"show\", ...])``).""" return ModelCommand() @pytest.fixture def mock_litellm_response(self): """Create a mock response for LiteLLM model data.""" return { "gpt-4": { "input_cost_per_token": 0.00003, "output_cost_per_token": 0.00006, "max_tokens": 8192, "supports_function_calling": True, "supports_vision": False, "litellm_provider": "openai", }, "claude-3-sonnet-20240229": { "input_cost_per_token": 0.000015, "output_cost_per_token": 0.000075, "max_tokens": 200000, "supports_function_calling": True, "supports_vision": True, "litellm_provider": "anthropic", }, "gpt-3.5-turbo": { "input_cost_per_token": 0.000001, "output_cost_per_token": 0.000002, "max_tokens": 4096, "supports_function_calling": False, "supports_vision": False, "litellm_provider": "openai", }, } @pytest.fixture def mock_ollama_response(self): """Create a mock response for Ollama models.""" return { "models": [ {"name": "llama3", "size": 4661211648}, {"name": "mistral:7b", "size": 7365960192}, ] } @patch("requests.get") def test_handle_no_args( self, mock_get, model_show_command, mock_litellm_response, mock_ollama_response ): """Test showing all models with no arguments.""" # Mock LiteLLM response mock_litellm = Mock() mock_litellm.status_code = 200 mock_litellm.json.return_value = mock_litellm_response # Mock Ollama response mock_ollama = Mock() mock_ollama.status_code = 200 mock_ollama.json.return_value = mock_ollama_response # Configure the mock to return different responses based on URL def side_effect(url, timeout=None): if "litellm" in url: return mock_litellm elif "ollama" in url: return mock_ollama else: return Mock(status_code=404) mock_get.side_effect = side_effect result = model_show_command.handle(["show"]) assert result is True @patch("requests.get") def test_handle_supported_filter(self, mock_get, model_show_command, mock_litellm_response): """Test showing only supported models (with function calling).""" # Mock LiteLLM response mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response mock_get.return_value = mock_response result = model_show_command.handle(["show", "supported"]) assert result is True @patch("requests.get") def test_handle_search_filter(self, mock_get, model_show_command, mock_litellm_response): """Test filtering models by search term.""" # Mock LiteLLM response mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response mock_get.return_value = mock_response result = model_show_command.handle(["show", "gpt"]) assert result is True @patch("requests.get") def test_handle_supported_and_search(self, mock_get, model_show_command, mock_litellm_response): """Test combining supported filter with search term.""" # Mock LiteLLM response mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response mock_get.return_value = mock_response result = model_show_command.handle(["show", "supported", "claude"]) assert result is True @patch("requests.get") def test_handle_network_error(self, mock_get, model_show_command): """Test handling when network request fails.""" # Mock network failure mock_get.side_effect = Exception("Network error") result = model_show_command.handle(["show"]) assert result is True # Should handle gracefully @patch("requests.get") def test_handle_http_error(self, mock_get, model_show_command): """Test handling when API returns HTTP error.""" # Mock HTTP error mock_response = Mock() mock_response.status_code = 500 mock_get.return_value = mock_response result = model_show_command.handle(["show"]) assert result is True # Should handle gracefully @patch("requests.get") def test_handle_with_ollama_error(self, mock_get, model_show_command, mock_litellm_response): """Test handling when Ollama is not available but LiteLLM works.""" # Mock LiteLLM success but Ollama failure def side_effect(url, timeout=None): if "litellm" in url: mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response return mock_response elif "ollama" in url: raise Exception("Ollama not available") else: return Mock(status_code=404) mock_get.side_effect = side_effect result = model_show_command.handle(["show"]) assert result is True @pytest.mark.integration class TestModelCommandIntegration: """Integration tests for model command functionality.""" @pytest.fixture(autouse=True) def setup_integration(self): """Setup for integration tests.""" # Store original CAI_MODEL if it exists self.original_model = os.environ.get("CAI_MODEL") yield # Restore original CAI_MODEL or remove if it didn't exist if self.original_model is not None: os.environ["CAI_MODEL"] = self.original_model elif "CAI_MODEL" in os.environ: del os.environ["CAI_MODEL"] @patch("requests.get") def test_full_model_workflow(self, mock_get): """Test a complete workflow of listing and selecting models.""" # Mock responses mock_litellm_response = { "gpt-4": { "input_cost_per_token": 0.00003, "output_cost_per_token": 0.00006, "max_tokens": 8192, "supports_function_calling": True, }, "claude-3-sonnet-20240229": { "input_cost_per_token": 0.000015, "output_cost_per_token": 0.000075, "max_tokens": 200000, "supports_function_calling": True, }, } mock_ollama_response = {"models": [{"name": "llama3", "size": 4661211648}]} # Configure mock responses def side_effect(url, timeout=None): if "litellm" in url: mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_litellm_response return mock_response elif "ollama" in url: mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = mock_ollama_response return mock_response else: return Mock(status_code=404) mock_get.side_effect = side_effect model_cmd = ModelCommand() # List all models result1 = model_cmd.handle([]) assert result1 is True # Show detailed model info result2 = model_cmd.handle(["show"]) assert result2 is True # Select a model by name result3 = model_cmd.handle(["gpt-4"]) assert result3 is True assert os.environ.get("CAI_MODEL") == "gpt-4" # Show current model again result4 = model_cmd.handle([]) assert result4 is True # Select by number (after cache is populated) result5 = model_cmd.handle(["1"]) assert result5 is True @patch("requests.get") def test_model_selection_edge_cases(self, mock_get): """Test edge cases in model selection.""" # Mock minimal response to avoid network dependency mock_response = Mock() mock_response.status_code = 200 mock_response.json.return_value = {"gpt-4": {}} mock_get.return_value = mock_response cmd = ModelCommand() before = os.environ.get("CAI_MODEL") # Out-of-range index: error, env unchanged result1 = cmd.handle(["999"]) assert result1 is True assert os.environ.get("CAI_MODEL") == before # Valid id from mocked LiteLLM catalog result2 = cmd.handle(["gpt-4"]) assert result2 is True assert os.environ.get("CAI_MODEL") == "gpt-4" # Empty / whitespace name: error result3 = cmd.handle([""]) assert result3 is True assert os.environ.get("CAI_MODEL") == "gpt-4" result4 = cmd.handle([" "]) assert result4 is True assert os.environ.get("CAI_MODEL") == "gpt-4" @patch("requests.get") def test_model_show_filters_combination(self, mock_get): """Test various combinations of filters for ``/model show``.""" mock_response = { "gpt-4": { "supports_function_calling": True, "input_cost_per_token": 0.00003, "output_cost_per_token": 0.00006, }, "gpt-3.5-turbo": { "supports_function_calling": False, "input_cost_per_token": 0.000001, "output_cost_per_token": 0.000002, }, "claude-3-sonnet": { "supports_function_calling": True, "input_cost_per_token": 0.000015, "output_cost_per_token": 0.000075, }, } mock_http_response = Mock() mock_http_response.status_code = 200 mock_http_response.json.return_value = mock_response mock_get.return_value = mock_http_response cmd = ModelCommand() # Test supported only result1 = cmd.handle(["show", "supported"]) assert result1 is True # Test search only result2 = cmd.handle(["show", "gpt"]) assert result2 is True # Test supported + search result3 = cmd.handle(["show", "supported", "claude"]) assert result3 is True # Test search + supported (different order) result4 = cmd.handle(["show", "claude", "supported"]) assert result4 is True if __name__ == "__main__": pytest.main([__file__, "-v"])