Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/client/cli/commands/models.py: 0%
231 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1# stdlib imports
2import re
3from collections import defaultdict
4from collections.abc import Callable
5from dataclasses import dataclass
6from datetime import datetime
7from typing import TYPE_CHECKING, Any, Final, Literal
9# third party imports
10import click
11import rich
12import yaml
13from typing_extensions import NotRequired, ReadOnly, TypedDict
15# local imports
16from ... import Client
17from ._cli_context import cli_context_values
19if TYPE_CHECKING:
20 from rich.console import JustifyMethod
23class _ModelInfoColumnConfig(TypedDict):
24 """Rendering config for one column of the ``models info`` table."""
26 header: ReadOnly[str]
27 style: ReadOnly[str]
28 justify: NotRequired[ReadOnly["JustifyMethod"]]
29 get_value: ReadOnly[Callable[..., str]]
32@dataclass
33class ModelYamlInfo:
34 model_name: str
35 model_params: dict[str, object]
36 model_info: dict[str, object]
37 model_id: str
38 access_groups: list[str]
39 provider: str
41 @property
42 def access_groups_str(self) -> str:
43 return ", ".join(self.access_groups) if self.access_groups else ""
46def _get_model_info_obj_from_yaml(model: dict[str, Any]) -> ModelYamlInfo:
47 """Extract model info from a model dict and return as ModelYamlInfo dataclass."""
48 model_name: Final[str] = model["model_name"]
49 model_params: Final[dict[str, Any]] = model["litellm_params"]
50 model_info: Final[dict[str, Any]] = model.get("model_info", {})
51 model_id: Final[str] = model_params["model"]
52 access_groups: Final = model_info.get("access_groups", [])
53 provider: Final = model_id.split("/", 1)[0] if "/" in model_id else model_id
54 return ModelYamlInfo(
55 model_name=model_name,
56 model_params=model_params,
57 model_info=model_info,
58 model_id=model_id,
59 access_groups=access_groups,
60 provider=provider,
61 )
64def format_iso_datetime_str(iso_datetime_str: str | None) -> str:
65 """Format an ISO format datetime string to human-readable date with minute resolution."""
66 if not iso_datetime_str:
67 return ""
68 try:
69 # Parse ISO format datetime string
70 dt: Final = datetime.fromisoformat(iso_datetime_str.replace("Z", "+00:00"))
71 return dt.strftime("%Y-%m-%d %H:%M")
72 except (TypeError, ValueError):
73 return str(iso_datetime_str)
76def format_timestamp(timestamp: int | None) -> str:
77 """Format a Unix timestamp (integer) to human-readable date with minute resolution."""
78 if timestamp is None:
79 return ""
80 try:
81 dt: Final = datetime.fromtimestamp(timestamp)
82 return dt.strftime("%Y-%m-%d %H:%M")
83 except (TypeError, ValueError):
84 return str(timestamp)
87def format_cost_per_1k_tokens(cost: float | None) -> str:
88 """Format a per-token cost to cost per 1000 tokens."""
89 if cost is None:
90 return ""
91 try:
92 # Convert string to float if needed
93 cost_float: Final = float(cost)
94 # Multiply by 1000 and format to 4 decimal places
95 return f"${cost_float * 1000:.4f}"
96 except (TypeError, ValueError):
97 return str(cost)
100def create_client(ctx: click.Context) -> Client:
101 """Helper function to create a client from context."""
102 context: Final = cli_context_values(ctx)
103 return Client(base_url=context["base_url"], api_key=context["api_key"])
106@click.group()
107def models() -> None:
108 """Manage models on your LiteLLM proxy server"""
111@models.command("list")
112@click.option(
113 "--format",
114 "output_format",
115 type=click.Choice(["table", "json"]),
116 default="table",
117 help="Output format (table or json)",
118)
119@click.pass_context
120def list_models(ctx: click.Context, output_format: Literal["table", "json"]) -> None:
121 """List all available models"""
122 client: Final = create_client(ctx)
123 models_list: Final = client.models.list()
124 assert isinstance(models_list, list)
126 if output_format == "json":
127 rich.print_json(data=models_list)
128 else: # table format
129 table: Final = rich.table.Table(title="Available Models")
131 # Add columns based on the data structure
132 table.add_column("ID", style="cyan")
133 table.add_column("Object", style="green")
134 table.add_column("Created", style="magenta")
135 table.add_column("Owned By", style="yellow")
137 # Add rows
138 for model in models_list:
139 created = model.get("created")
140 # Convert string timestamp to integer if needed
141 if isinstance(created, str) and created.isdigit():
142 created = int(created)
144 table.add_row(
145 str(model.get("id", "")),
146 str(model.get("object", "model")),
147 (format_timestamp(created) if isinstance(created, int) else format_iso_datetime_str(created)),
148 str(model.get("owned_by", "")),
149 )
151 rich.print(table)
154@models.command("add")
155@click.argument("model-name")
156@click.option(
157 "--param",
158 "-p",
159 multiple=True,
160 help="Model parameters in key=value format (can be specified multiple times)",
161)
162@click.option(
163 "--info",
164 "-i",
165 multiple=True,
166 help="Model info in key=value format (can be specified multiple times)",
167)
168@click.pass_context
169def add_model(ctx: click.Context, model_name: str, param: tuple[str, ...], info: tuple[str, ...]) -> None:
170 """Add a new model to the proxy"""
171 # Convert parameters from key=value format to dict
172 model_params: Final = dict(p.split("=", 1) for p in param)
173 model_info: Final = dict(i.split("=", 1) for i in info) if info else None
175 client: Final = create_client(ctx)
176 result: Final = client.models.new(
177 model_name=model_name,
178 model_params=model_params,
179 model_info=model_info,
180 )
181 rich.print_json(data=result)
184@models.command("delete")
185@click.argument("model-id")
186@click.pass_context
187def delete_model(ctx: click.Context, model_id: str) -> None:
188 """Delete a model from the proxy"""
189 client: Final = create_client(ctx)
190 result: Final = client.models.delete(model_id=model_id)
191 rich.print_json(data=result)
194@models.command("get")
195@click.option("--id", "model_id", help="ID of the model to retrieve")
196@click.option("--name", "model_name", help="Name of the model to retrieve")
197@click.pass_context
198def get_model(ctx: click.Context, model_id: str | None, model_name: str | None) -> None:
199 """Get information about a specific model"""
200 if not model_id and not model_name:
201 raise click.UsageError("Either --id or --name must be provided")
203 client: Final = create_client(ctx)
204 result: Final = client.models.get(model_id=model_id, model_name=model_name)
205 rich.print_json(data=result)
208@models.command("info")
209@click.option(
210 "--format",
211 "output_format",
212 type=click.Choice(["table", "json"]),
213 default="table",
214 help="Output format (table or json)",
215)
216@click.option(
217 "--columns",
218 "columns",
219 default="public_model,upstream_model,updated_at",
220 help="Comma-separated list of columns to display. Valid columns: public_model, upstream_model, credential_name, created_at, updated_at, id, input_cost, output_cost. Default: public_model,upstream_model,updated_at",
221)
222@click.pass_context
223def get_models_info(ctx: click.Context, output_format: Literal["table", "json"], columns: str) -> None:
224 """Get detailed information about all models"""
225 client: Final = create_client(ctx)
226 models_info: Final = client.models.info()
227 assert isinstance(models_info, list)
229 if output_format == "json":
230 rich.print_json(data=models_info)
231 else: # table format
232 table: Final = rich.table.Table(title="Models Information")
234 # Define all possible columns with their configurations
235 column_configs: Final[dict[str, _ModelInfoColumnConfig]] = {
236 "public_model": {
237 "header": "Public Model",
238 "style": "cyan",
239 "get_value": lambda m: str(m.get("model_name", "")),
240 },
241 "upstream_model": {
242 "header": "Upstream Model",
243 "style": "green",
244 "get_value": lambda m: str(m.get("litellm_params", {}).get("model", "")),
245 },
246 "credential_name": {
247 "header": "Credential Name",
248 "style": "yellow",
249 "get_value": lambda m: str(m.get("litellm_params", {}).get("litellm_credential_name", "")),
250 },
251 "created_at": {
252 "header": "Created At",
253 "style": "magenta",
254 "get_value": lambda m: format_iso_datetime_str(m.get("model_info", {}).get("created_at")),
255 },
256 "updated_at": {
257 "header": "Updated At",
258 "style": "magenta",
259 "get_value": lambda m: format_iso_datetime_str(m.get("model_info", {}).get("updated_at")),
260 },
261 "id": {
262 "header": "ID",
263 "style": "blue",
264 "get_value": lambda m: str(m.get("model_info", {}).get("id", "")),
265 },
266 "input_cost": {
267 "header": "Input Cost",
268 "style": "green",
269 "justify": "right",
270 "get_value": lambda m: format_cost_per_1k_tokens(m.get("model_info", {}).get("input_cost_per_token")),
271 },
272 "output_cost": {
273 "header": "Output Cost",
274 "style": "green",
275 "justify": "right",
276 "get_value": lambda m: format_cost_per_1k_tokens(m.get("model_info", {}).get("output_cost_per_token")),
277 },
278 }
280 # Add requested columns
281 requested_columns: Final = [col.strip() for col in columns.split(",")]
282 for col_name in requested_columns:
283 if col_name in column_configs:
284 config = column_configs[col_name]
285 table.add_column(
286 config["header"],
287 style=config["style"],
288 justify=config.get("justify", "left"),
289 )
290 else:
291 click.echo(f"Warning: Unknown column '{col_name}'", err=True)
293 # Add rows with only the requested columns
294 for model in models_info:
295 row_values = []
296 for col_name in requested_columns:
297 if col_name in column_configs:
298 row_values.append(column_configs[col_name]["get_value"](model))
299 if row_values:
300 table.add_row(*row_values)
302 rich.print(table)
305@models.command("update")
306@click.argument("model-id")
307@click.option(
308 "--param",
309 "-p",
310 multiple=True,
311 help="Model parameters in key=value format (can be specified multiple times)",
312)
313@click.option(
314 "--info",
315 "-i",
316 multiple=True,
317 help="Model info in key=value format (can be specified multiple times)",
318)
319@click.pass_context
320def update_model(ctx: click.Context, model_id: str, param: tuple[str, ...], info: tuple[str, ...]) -> None:
321 """Update an existing model's configuration"""
322 # Convert parameters from key=value format to dict
323 model_params: Final = dict(p.split("=", 1) for p in param)
324 model_info: Final = dict(i.split("=", 1) for i in info) if info else None
326 client: Final = create_client(ctx)
327 result: Final = client.models.update(
328 model_id=model_id,
329 model_params=model_params,
330 model_info=model_info,
331 )
332 rich.print_json(data=result)
335def _filter_model(model, model_regex, access_group_regex):
336 model_name: Final = model.get("model_name")
337 model_params: Final = model.get("litellm_params")
338 model_info: Final = model.get("model_info", {})
339 if not model_name or not model_params:
340 return False
341 model_id: Final = model_params.get("model")
342 if not model_id or not isinstance(model_id, str):
343 return False
344 if model_regex and not model_regex.search(model_id):
345 return False
346 access_groups: Final = model_info.get("access_groups", [])
347 if access_group_regex:
348 if not isinstance(access_groups, list):
349 return False
350 if not any(isinstance(group, str) and access_group_regex.search(group) for group in access_groups):
351 return False
352 return True
355def _print_models_table(added_models: list[ModelYamlInfo], table_title: str):
356 if not added_models:
357 return
358 table: Final = rich.table.Table(title=table_title)
359 table.add_column("Model Name", style="cyan")
360 table.add_column("Upstream Model", style="green")
361 table.add_column("Access Groups", style="magenta")
362 for m in added_models:
363 table.add_row(m.model_name, m.model_id, m.access_groups_str)
364 rich.print(table)
367def _print_summary_table(provider_counts):
368 summary_table: Final = rich.table.Table(title="Model Import Summary")
369 summary_table.add_column("Provider", style="cyan")
370 summary_table.add_column("Count", style="green")
372 for provider, count in provider_counts.items():
373 summary_table.add_row(str(provider), str(count))
375 total: Final = sum(provider_counts.values())
376 summary_table.add_row("[bold]Total[/bold]", f"[bold]{total}[/bold]")
378 rich.print(summary_table)
381def get_model_list_from_yaml_file(yaml_file: str) -> list[dict[str, Any]]:
382 """Load and validate the model list from a YAML file."""
383 with open(yaml_file, "r") as f:
384 data: Final = yaml.safe_load(f)
385 if not data or "model_list" not in data:
386 raise click.ClickException("YAML file must contain a 'model_list' key with a list of models.")
387 model_list: Final = data["model_list"]
388 if not isinstance(model_list, list):
389 raise click.ClickException("'model_list' must be a list of model definitions.")
390 return model_list
393def _get_filtered_model_list(model_list, only_models_matching_regex, only_access_groups_matching_regex):
394 """Return a list of models that pass the filter criteria."""
395 model_regex: Final = re.compile(only_models_matching_regex) if only_models_matching_regex else None
396 access_group_regex = re.compile(only_access_groups_matching_regex) if only_access_groups_matching_regex else None
397 return [model for model in model_list if _filter_model(model, model_regex, access_group_regex)]
400def _import_models_get_table_title(dry_run: bool) -> str:
401 if dry_run:
402 return "Models that would be imported if [yellow]--dry-run[/yellow] was not provided"
403 else:
404 return "Models Imported"
407@models.command("import")
408@click.argument("yaml_file", type=click.Path(exists=True, dir_okay=False, readable=True))
409@click.option(
410 "--dry-run",
411 is_flag=True,
412 help="Show what would be imported without making any changes.",
413)
414@click.option(
415 "--only-models-matching-regex",
416 default=None,
417 help="Only import models where litellm_params.model matches the given regex.",
418)
419@click.option(
420 "--only-access-groups-matching-regex",
421 default=None,
422 help="Only import models where at least one item in model_info.access_groups matches the given regex.",
423)
424@click.pass_context
425def import_models(
426 ctx: click.Context,
427 yaml_file: str,
428 dry_run: bool,
429 only_models_matching_regex: str | None,
430 only_access_groups_matching_regex: str | None,
431) -> None:
432 """Import models from a YAML file and add them to the proxy."""
433 provider_counts: Final[dict[str, int]] = defaultdict(int)
434 added_models: Final[list[ModelYamlInfo]] = []
435 model_list: Final = get_model_list_from_yaml_file(yaml_file)
436 filtered_model_list: Final = _get_filtered_model_list(
437 model_list, only_models_matching_regex, only_access_groups_matching_regex
438 )
440 if not dry_run:
441 client: Final = create_client(ctx)
443 for model in filtered_model_list:
444 model_info_obj = _get_model_info_obj_from_yaml(model)
445 if not dry_run:
446 try:
447 client.models.new(
448 model_name=model_info_obj.model_name,
449 model_params=model_info_obj.model_params,
450 model_info=model_info_obj.model_info,
451 )
452 except Exception:
453 pass # For summary, ignore errors
454 added_models.append(model_info_obj)
455 provider_counts[model_info_obj.provider] += 1
457 table_title: Final = _import_models_get_table_title(dry_run)
458 _print_models_table(added_models, table_title)
459 _print_summary_table(provider_counts)