Remove duplicates and re-invokation to obtain the pricing, leverage COST_TRACKER

Signed-off-by: Víctor Mayoral Vilches <v.mayoralv@gmail.com>
This commit is contained in:
Víctor Mayoral Vilches 2025-06-01 13:30:14 +00:00
parent b906fff4dd
commit cfa839dce6
1 changed files with 8 additions and 35 deletions

View File

@ -12,7 +12,7 @@ import requests # pylint: disable=import-error
from rich.console import Console # pylint: disable=import-error
from rich.table import Table # pylint: disable=import-error
from rich.panel import Panel # pylint: disable=import-error
from cai.util import get_ollama_api_base
from cai.util import get_ollama_api_base, COST_TRACKER
from cai.repl.commands.base import Command, register_command
console = Console()
@ -125,49 +125,22 @@ class ModelCommand(Command):
]
}
# Fetch model pricing data from LiteLLM GitHub repository
model_pricing_data = {}
try:
response = requests.get(LITELLM_URL, timeout=2)
if response.status_code == 200:
model_pricing_data = response.json()
# Add DeepSeek models with their pricing if not in data
deepseek_v3_path = "deepseek/deepseek-v3"
deepseek_r1_path = "deepseek/deepseek-r1"
if (deepseek_v3_path in model_pricing_data and
"deepseek-v3" not in model_pricing_data):
model_pricing_data["deepseek-v3"] = (
model_pricing_data[deepseek_v3_path]
)
if (deepseek_r1_path in model_pricing_data and
"deepseek-r1" not in model_pricing_data):
model_pricing_data["deepseek-r1"] = (
model_pricing_data[deepseek_r1_path]
)
except Exception: # pylint: disable=broad-except
console.print(
"[yellow]Warning: Could not fetch model pricing data[/yellow]"
)
# 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 if available
pricing_info = model_pricing_data.get(model["name"], {})
input_cost = pricing_info.get("input_cost_per_token", None)
output_cost = pricing_info.get("output_cost_per_token", None)
# 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 if values exist
# Convert to dollars per million tokens
input_cost_per_million = None
output_cost_per_million = None
if input_cost is not None:
input_cost_per_million = input_cost * 1000000
if output_cost is not None:
output_cost_per_million = output_cost * 1000000
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"],