From b8dd50dd27c115ba990b101b2d58232981ef9e2f Mon Sep 17 00:00:00 2001 From: Chris <41129709+chrissRo@users.noreply.github.com> Date: Thu, 24 Jul 2025 09:24:21 +0200 Subject: [PATCH] Noticed a difference in the model-list in completer.py and model.py (#230) * fix * reverted parts of the previous changes as they led to a duplication of model-show * now /model shows the predefined models (as before) * the numbering in the cached-models is consistent with model-show and the completer --------- Co-authored-by: christian.pritzl --- src/cai/repl/commands/compact.py | 94 ++--------- src/cai/repl/commands/completer.py | 71 ++++---- src/cai/repl/commands/model.py | 257 +++++++++++++++++------------ 3 files changed, 196 insertions(+), 226 deletions(-) diff --git a/src/cai/repl/commands/compact.py b/src/cai/repl/commands/compact.py index f18d182e..02a26bde 100644 --- a/src/cai/repl/commands/compact.py +++ b/src/cai/repl/commands/compact.py @@ -7,9 +7,17 @@ from typing import List, Optional import os import datetime from rich.console import Console +from rich.table import Table +from rich.panel import Panel from cai.repl.commands.base import Command, register_command from cai.sdk.agents.models.openai_chatcompletions import get_current_active_model +from cai.util import COST_TRACKER +from cai.repl.commands.model import ( + ModelCommand, + get_predefined_model_categories, + get_all_predefined_models +) console = Console() @@ -82,11 +90,6 @@ class CompactCommand(Command): def handle_model(self, args: Optional[List[str]] = None) -> bool: """Set model for compaction.""" - from cai.repl.commands.model import ModelCommand - from rich.table import Table - from rich.panel import Panel - from cai.util import COST_TRACKER - if not args: # Display current model console.print( @@ -97,85 +100,8 @@ class CompactCommand(Command): ) ) - # Create model command instance to reuse its model data - model_cmd = ModelCommand() - - # Define model categories (same as in model.py) - MODEL_CATEGORIES = { - "Alias": [ - { - "name": "alias0", - "description": "Best model for Cybersecurity AI tasks" - } - ], - "Anthropic Claude": [ - { - "name": "claude-sonnet-4-20250514", - "description": "Excellent balance of performance and efficiency" - }, - { - "name": "claude-3-7-sonnet-20250219", - "description": "Excellent model for complex reasoning and creative tasks" - }, - { - "name": "claude-3-5-sonnet-20240620", - "description": "Excellent balance of performance and efficiency" - }, - { - "name": "claude-3-5-haiku-20240307", - "description": "Fast and efficient model" - }, - ], - "OpenAI": [ - { - "name": "o3-mini", - "description": "Latest mini model in the O-series" - }, - { - "name": "gpt-4o", - "description": "Latest GPT-4 model with improved capabilities" - }, - ], - "DeepSeek": [ - { - "name": "deepseek-v3", - "description": "DeepSeek's latest general-purpose model" - }, - { - "name": "deepseek-r1", - "description": "DeepSeek's specialized reasoning model" - } - ] - } - - # Create a flat list of all models - ALL_MODELS = [] - for category, models in MODEL_CATEGORIES.items(): - for model in models: - # Get pricing info - input_cost_per_token, output_cost_per_token = COST_TRACKER.get_model_pricing(model["name"]) - - # Convert to dollars per million tokens - input_cost_per_million = None - output_cost_per_million = None - - if input_cost_per_token is not None and input_cost_per_token > 0: - input_cost_per_million = input_cost_per_token * 1000000 - if output_cost_per_token is not None and output_cost_per_token > 0: - output_cost_per_million = output_cost_per_token * 1000000 - - ALL_MODELS.append({ - "name": model["name"], - "provider": ( - "Anthropic" if "claude" in model["name"] - else "DeepSeek" if "deepseek" in model["name"] - else "OpenAI" - ), - "category": category, - "description": model["description"], - "input_cost": input_cost_per_million, - "output_cost": output_cost_per_million - }) + # Get all predefined models using the shared function + ALL_MODELS = get_all_predefined_models() # Show available models in a table model_table = Table( diff --git a/src/cai/repl/commands/completer.py b/src/cai/repl/commands/completer.py index 3f6d2468..47a347a0 100644 --- a/src/cai/repl/commands/completer.py +++ b/src/cai/repl/commands/completer.py @@ -30,6 +30,7 @@ from cai.repl.commands.base import ( COMMANDS, COMMAND_ALIASES ) +from cai.repl.commands.model import get_predefined_model_names console = Console() @@ -82,12 +83,12 @@ class FuzzyCommandCompleter(Completer): def _background_fetch_models(self): """Fetch models in background to avoid blocking the UI.""" try: - self.fetch_ollama_models() + self.fetch_all_models() except Exception: # pylint: disable=broad-except pass - def fetch_ollama_models(self): # pylint: disable=too-many-branches,too-many-statements,inconsistent-return-statements,line-too-long # noqa: E501 - """Fetch available models from Ollama if it's running.""" + def fetch_all_models(self): # pylint: disable=too-many-branches,too-many-statements,inconsistent-return-statements,line-too-long # noqa: E501 + """Fetch all available models (predefined + LiteLLM + Ollama) to match /model command.""" # Only fetch every 60 seconds to avoid excessive API calls now = datetime.datetime.now() @@ -97,8 +98,28 @@ class FuzzyCommandCompleter(Completer): return self._last_model_fetch = now - ollama_models = [] + + # Start with predefined models from the shared source of truth + all_models = get_predefined_model_names() + # Fetch LiteLLM models (matches the /model command behavior) + try: + litellm_url = ( + "https://raw.githubusercontent.com/BerriAI/litellm/main/" + "model_prices_and_context_window.json" + ) + response = requests.get(litellm_url, timeout=2) + + if response.status_code == 200: + litellm_data = response.json() + # Add LiteLLM models (sorted for consistency) + litellm_models = sorted(litellm_data.keys()) + all_models.extend(litellm_models) + except Exception: # pylint: disable=broad-except + # Silently fail if LiteLLM is not available + pass + + # Fetch Ollama models try: # Get Ollama models with a short timeout to prevent hanging api_base = get_ollama_api_base() @@ -113,47 +134,17 @@ class FuzzyCommandCompleter(Completer): # Fallback for older Ollama versions models = data.get('items', []) - ollama_models = [(model.get('name', ''), []) for model in models] + ollama_models = [model.get('name', '') for model in models] + all_models.extend(ollama_models) except Exception: # pylint: disable=broad-except # Silently fail if Ollama is not available pass - # Standard models always available - standard_models = [ - # Alias models - "alias0", - - # Claude 3.7 models - "claude-3-7-sonnet-20250219", - - # Claude 3.5 models - "claude-3-5-sonnet-20240620", - "claude-3-5-20241122", - - # Claude 3 models - "claude-3-opus-20240229", - "claude-3-sonnet-20240229", - "claude-3-haiku-20240307", - - # OpenAI O-series models - "o1", - "o1-mini", - "o3-mini", - - # OpenAI GPT models - "gpt-4o", - "gpt-4-turbo", - "gpt-3.5-turbo", - - # DeepSeek models - "deepseek-v3", - "deepseek-r1" - ] - - # Combine standard models with Ollama models - self._cached_models = standard_models + ollama_models + # Cache all models that the /model command can handle + self._cached_models = all_models # Create number mappings for models (1-based indexing) + # This matches the /model command numbering exactly self._cached_model_numbers = {} for i, model in enumerate(self._cached_models, 1): self._cached_model_numbers[str(i)] = model @@ -484,7 +475,7 @@ class FuzzyCommandCompleter(Completer): words = text.split() # Refresh Ollama models periodically - self.fetch_ollama_models() + self.fetch_all_models() if not text: # Show all main commands with descriptions diff --git a/src/cai/repl/commands/model.py b/src/cai/repl/commands/model.py index d44485a5..ff23194c 100644 --- a/src/cai/repl/commands/model.py +++ b/src/cai/repl/commands/model.py @@ -5,7 +5,7 @@ This module provides commands for viewing and changing the current LLM model. import os import datetime # Standard library imports -from typing import List, Optional # Dict and Any removed as unused +from typing import List, Optional, Dict, Any # Third-party imports import requests # pylint: disable=import-error @@ -23,6 +23,138 @@ LITELLM_URL = ( ) +def get_predefined_model_categories() -> Dict[str, List[Dict[str, str]]]: + """Get the predefined model categories as the single source of truth. + + This function serves as the authoritative source for all available models + across the CAI system. Other modules should import and use this function + to ensure consistency. + + Returns: + Dictionary mapping category names to lists of model dictionaries + """ + return { + "Alias": [ + { + "name": "alias0", + "description": ( + "Best model for Cybersecurity AI tasks" + ) + }, + { + "name": "alias0-fast", + "description": ( + "Fast version of alias0 for quick tasks" + ) + } + ], + "Anthropic Claude": [ + { + "name": "claude-sonnet-4-20250514", + "description": ( + "Excellent balance of performance and efficiency" + ) + }, + { + "name": "claude-3-7-sonnet-20250219", + "description": ( + "Excellent model for complex reasoning and creative tasks" + ) + }, + { + "name": "claude-3-5-sonnet-20240620", + "description": ( + "Excellent balance of performance and efficiency" + ) + }, + { + "name": "claude-3-5-haiku-20240307", + "description": ( + "Fast and efficient model" + ) + }, + ], + "OpenAI": [ + { + "name": "o3-mini", + "description": "Latest mini model in the O-series" + }, + { + "name": "gpt-4o", + "description": ( + "Latest GPT-4 model with improved capabilities" + ) + }, + ], + "DeepSeek": [ + { + "name": "deepseek-v3", + "description": "DeepSeek's latest general-purpose model" + }, + { + "name": "deepseek-r1", + "description": "DeepSeek's specialized reasoning model" + } + ] + } + + +def get_all_predefined_models() -> List[Dict[str, Any]]: + """Get all predefined models as a flat list with enriched data. + + Returns: + List of model dictionaries with name, provider, category, description, and pricing + """ + model_categories = get_predefined_model_categories() + all_models = [] + + # Simple mapping from category to provider name + category_to_provider = { + "Alias": "OpenAI", # Alias models use OpenAI as base + "Anthropic Claude": "Anthropic", + "OpenAI": "OpenAI", + "DeepSeek": "DeepSeek" + } + + for category, models in model_categories.items(): + provider = category_to_provider.get(category, "Unknown") + + for model in models: + # Get pricing info using COST_TRACKER + input_cost_per_token, output_cost_per_token = COST_TRACKER.get_model_pricing(model["name"]) + + # Convert to dollars per million tokens + input_cost_per_million = None + output_cost_per_million = None + + if input_cost_per_token is not None and input_cost_per_token > 0: + input_cost_per_million = input_cost_per_token * 1000000 + if output_cost_per_token is not None and output_cost_per_token > 0: + output_cost_per_million = output_cost_per_token * 1000000 + + all_models.append({ + "name": model["name"], + "provider": provider, + "category": category, + "description": model["description"], + "input_cost": input_cost_per_million, + "output_cost": output_cost_per_million + }) + + return all_models + + +def get_predefined_model_names() -> List[str]: + """Get a simple list of all predefined model names. + + This is useful for autocompletion and simple model name lists. + + Returns: + List of model name strings + """ + return [model["name"] for model in get_all_predefined_models()] + + class ModelCommand(Command): """Command for viewing and changing the current LLM model.""" @@ -64,108 +196,28 @@ class ModelCommand(Command): Returns: bool: True if the model was changed successfully """ - # Define model categories and their models for easy reference + # Get all predefined models from the shared source of truth # pylint: disable=invalid-name - MODEL_CATEGORIES = { - "Alias": [ - { - "name": "alias0", - "description": ( - "Best model for Cybersecurity AI tasks" - ) - }, - { - "name": "alias0-fast", - "description": ( - "Fast version of alias0 for quick tasks" - ) - } - ], - "Anthropic Claude": [ - { - "name": "claude-sonnet-4-20250514", - "description": ( - "Excellent balance of performance and efficiency" - ) - }, - { - "name": "claude-3-7-sonnet-20250219", - "description": ( - "Excellent model for complex reasoning and creative tasks" - ) - }, - { - "name": "claude-3-5-sonnet-20240620", - "description": ( - "Excellent balance of performance and efficiency" - ) - }, - { - "name": "claude-3-5-haiku-20240307", - "description": ( - "Fast and efficient model" - ) - }, - ], - "OpenAI": [ - { - "name": "o3-mini", - "description": "Latest mini model in the O-series" - }, - { - "name": "gpt-4o", - "description": ( - "Latest GPT-4 model with improved capabilities" - ) - }, - ], - "DeepSeek": [ - { - "name": "deepseek-v3", - "description": "DeepSeek's latest general-purpose model" - }, - { - "name": "deepseek-r1", - "description": "DeepSeek's specialized reasoning model" - } - ] - } + ALL_MODELS = get_all_predefined_models() - # Create a flat list of all models for numeric selection - # pylint: disable=invalid-name - ALL_MODELS = [] - for category, models in MODEL_CATEGORIES.items(): - for model in models: - # Get pricing info using COST_TRACKER - input_cost_per_token, output_cost_per_token = COST_TRACKER.get_model_pricing(model["name"]) + # Also fetch LiteLLM model names to make numbering consistent with /model-show + litellm_model_names = [] + try: + response = requests.get(LITELLM_URL, timeout=5) + if response.status_code == 200: + litellm_data = response.json() + # Add LiteLLM model names (sorted for consistency with /model-show) + litellm_model_names = sorted(litellm_data.keys()) + except Exception: # pylint: disable=broad-except + # Silently fail if LiteLLM is not available + pass - # Convert to dollars per million tokens - input_cost_per_million = None - output_cost_per_million = None - - if input_cost_per_token is not None and input_cost_per_token > 0: - input_cost_per_million = input_cost_per_token * 1000000 - if output_cost_per_token is not None and output_cost_per_token > 0: - output_cost_per_million = output_cost_per_token * 1000000 - - ALL_MODELS.append({ - "name": model["name"], - "provider": ( - "Anthropic" if "claude" in model["name"] - else "DeepSeek" if "deepseek" in model["name"] - else "OpenAI" - ), - "category": category, - "description": model["description"], - "input_cost": input_cost_per_million, - "output_cost": output_cost_per_million - }) - - # Update cached models - self.cached_models = [model["name"] for model in ALL_MODELS] + # Update cached models to include all models for number selection (consistent with /model-show) + predefined_model_names = [model["name"] for model in ALL_MODELS] + self.cached_models = predefined_model_names + litellm_model_names self.cached_model_numbers = { - str(i): model["name"] - for i, model in enumerate(ALL_MODELS, 1) + str(i): model_name + for i, model_name in enumerate(self.cached_models, 1) } if not args: # pylint: disable=too-many-nested-blocks @@ -198,7 +250,7 @@ class ModelCommand(Command): justify="right") model_table.add_column("Description", style="white") - # Add all predefined models with numbers + # Add predefined models with numbers for i, model in enumerate(ALL_MODELS, 1): # Format pricing info as dollars per million tokens input_cost_str = ( @@ -242,7 +294,8 @@ class ModelCommand(Command): ollama_models = data.get('items', []) # Add Ollama models to the table with continuing numbers - start_index = len(ALL_MODELS) + 1 + # (after predefined models + LiteLLM models in numbering) + start_index = len(predefined_model_names) + len(litellm_model_names) + 1 for i, model in enumerate(ollama_models, start_index): model_name = model.get('name', '') model_size = model.get('size', 0) @@ -275,7 +328,7 @@ class ModelCommand(Command): self.cached_model_numbers[str(i)] = model_name except Exception: # pylint: disable=broad-except # Add a note about Ollama if we couldn't fetch models - start_index = len(ALL_MODELS) + 1 + start_index = len(predefined_model_names) + len(litellm_model_names) + 1 model_table.add_row( str(start_index), "llama3",