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 <christian.pritzl@iu.org>
This commit is contained in:
Chris 2025-07-24 09:24:21 +02:00 committed by GitHub
parent 83ab74f585
commit b8dd50dd27
3 changed files with 196 additions and 226 deletions

View File

@ -7,9 +7,17 @@ from typing import List, Optional
import os import os
import datetime import datetime
from rich.console import Console 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.repl.commands.base import Command, register_command
from cai.sdk.agents.models.openai_chatcompletions import get_current_active_model 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() console = Console()
@ -82,11 +90,6 @@ class CompactCommand(Command):
def handle_model(self, args: Optional[List[str]] = None) -> bool: def handle_model(self, args: Optional[List[str]] = None) -> bool:
"""Set model for compaction.""" """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: if not args:
# Display current model # Display current model
console.print( console.print(
@ -97,85 +100,8 @@ class CompactCommand(Command):
) )
) )
# Create model command instance to reuse its model data # Get all predefined models using the shared function
model_cmd = ModelCommand() ALL_MODELS = get_all_predefined_models()
# 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
})
# Show available models in a table # Show available models in a table
model_table = Table( model_table = Table(

View File

@ -30,6 +30,7 @@ from cai.repl.commands.base import (
COMMANDS, COMMANDS,
COMMAND_ALIASES COMMAND_ALIASES
) )
from cai.repl.commands.model import get_predefined_model_names
console = Console() console = Console()
@ -82,12 +83,12 @@ class FuzzyCommandCompleter(Completer):
def _background_fetch_models(self): def _background_fetch_models(self):
"""Fetch models in background to avoid blocking the UI.""" """Fetch models in background to avoid blocking the UI."""
try: try:
self.fetch_ollama_models() self.fetch_all_models()
except Exception: # pylint: disable=broad-except except Exception: # pylint: disable=broad-except
pass pass
def fetch_ollama_models(self): # pylint: disable=too-many-branches,too-many-statements,inconsistent-return-statements,line-too-long # noqa: E501 def fetch_all_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.""" """Fetch all available models (predefined + LiteLLM + Ollama) to match /model command."""
# Only fetch every 60 seconds to avoid excessive API calls # Only fetch every 60 seconds to avoid excessive API calls
now = datetime.datetime.now() now = datetime.datetime.now()
@ -97,8 +98,28 @@ class FuzzyCommandCompleter(Completer):
return return
self._last_model_fetch = now 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: try:
# Get Ollama models with a short timeout to prevent hanging # Get Ollama models with a short timeout to prevent hanging
api_base = get_ollama_api_base() api_base = get_ollama_api_base()
@ -113,47 +134,17 @@ class FuzzyCommandCompleter(Completer):
# Fallback for older Ollama versions # Fallback for older Ollama versions
models = data.get('items', []) 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 except Exception: # pylint: disable=broad-except
# Silently fail if Ollama is not available # Silently fail if Ollama is not available
pass pass
# Standard models always available # Cache all models that the /model command can handle
standard_models = [ self._cached_models = all_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
# Create number mappings for models (1-based indexing) # Create number mappings for models (1-based indexing)
# This matches the /model command numbering exactly
self._cached_model_numbers = {} self._cached_model_numbers = {}
for i, model in enumerate(self._cached_models, 1): for i, model in enumerate(self._cached_models, 1):
self._cached_model_numbers[str(i)] = model self._cached_model_numbers[str(i)] = model
@ -484,7 +475,7 @@ class FuzzyCommandCompleter(Completer):
words = text.split() words = text.split()
# Refresh Ollama models periodically # Refresh Ollama models periodically
self.fetch_ollama_models() self.fetch_all_models()
if not text: if not text:
# Show all main commands with descriptions # Show all main commands with descriptions

View File

@ -5,7 +5,7 @@ This module provides commands for viewing and changing the current LLM model.
import os import os
import datetime import datetime
# Standard library imports # Standard library imports
from typing import List, Optional # Dict and Any removed as unused from typing import List, Optional, Dict, Any
# Third-party imports # Third-party imports
import requests # pylint: disable=import-error 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): class ModelCommand(Command):
"""Command for viewing and changing the current LLM model.""" """Command for viewing and changing the current LLM model."""
@ -64,108 +196,28 @@ class ModelCommand(Command):
Returns: Returns:
bool: True if the model was changed successfully 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 # pylint: disable=invalid-name
MODEL_CATEGORIES = { ALL_MODELS = get_all_predefined_models()
"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"
}
]
}
# Create a flat list of all models for numeric selection # Also fetch LiteLLM model names to make numbering consistent with /model-show
# pylint: disable=invalid-name litellm_model_names = []
ALL_MODELS = [] try:
for category, models in MODEL_CATEGORIES.items(): response = requests.get(LITELLM_URL, timeout=5)
for model in models: if response.status_code == 200:
# Get pricing info using COST_TRACKER litellm_data = response.json()
input_cost_per_token, output_cost_per_token = COST_TRACKER.get_model_pricing(model["name"]) # 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 # Update cached models to include all models for number selection (consistent with /model-show)
input_cost_per_million = None predefined_model_names = [model["name"] for model in ALL_MODELS]
output_cost_per_million = None self.cached_models = predefined_model_names + litellm_model_names
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]
self.cached_model_numbers = { self.cached_model_numbers = {
str(i): model["name"] str(i): model_name
for i, model in enumerate(ALL_MODELS, 1) for i, model_name in enumerate(self.cached_models, 1)
} }
if not args: # pylint: disable=too-many-nested-blocks if not args: # pylint: disable=too-many-nested-blocks
@ -198,7 +250,7 @@ class ModelCommand(Command):
justify="right") justify="right")
model_table.add_column("Description", style="white") 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): for i, model in enumerate(ALL_MODELS, 1):
# Format pricing info as dollars per million tokens # Format pricing info as dollars per million tokens
input_cost_str = ( input_cost_str = (
@ -242,7 +294,8 @@ class ModelCommand(Command):
ollama_models = data.get('items', []) ollama_models = data.get('items', [])
# Add Ollama models to the table with continuing numbers # 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): for i, model in enumerate(ollama_models, start_index):
model_name = model.get('name', '') model_name = model.get('name', '')
model_size = model.get('size', 0) model_size = model.get('size', 0)
@ -275,7 +328,7 @@ class ModelCommand(Command):
self.cached_model_numbers[str(i)] = model_name self.cached_model_numbers[str(i)] = model_name
except Exception: # pylint: disable=broad-except except Exception: # pylint: disable=broad-except
# Add a note about Ollama if we couldn't fetch models # 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( model_table.add_row(
str(start_index), str(start_index),
"llama3", "llama3",