mirror of https://github.com/aliasrobotics/cai.git
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:
parent
83ab74f585
commit
b8dd50dd27
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Reference in New Issue