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

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 

8 

9# third party imports 

10import click 

11import rich 

12import yaml 

13from typing_extensions import NotRequired, ReadOnly, TypedDict 

14 

15# local imports 

16from ... import Client 

17from ._cli_context import cli_context_values 

18 

19if TYPE_CHECKING: 

20 from rich.console import JustifyMethod 

21 

22 

23class _ModelInfoColumnConfig(TypedDict): 

24 """Rendering config for one column of the ``models info`` table.""" 

25 

26 header: ReadOnly[str] 

27 style: ReadOnly[str] 

28 justify: NotRequired[ReadOnly["JustifyMethod"]] 

29 get_value: ReadOnly[Callable[..., str]] 

30 

31 

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 

40 

41 @property 

42 def access_groups_str(self) -> str: 

43 return ", ".join(self.access_groups) if self.access_groups else "" 

44 

45 

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 ) 

62 

63 

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) 

74 

75 

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) 

85 

86 

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) 

98 

99 

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"]) 

104 

105 

106@click.group() 

107def models() -> None: 

108 """Manage models on your LiteLLM proxy server""" 

109 

110 

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) 

125 

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") 

130 

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") 

136 

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) 

143 

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 ) 

150 

151 rich.print(table) 

152 

153 

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 

174 

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) 

182 

183 

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) 

192 

193 

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") 

202 

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) 

206 

207 

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) 

228 

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") 

233 

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 } 

279 

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) 

292 

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) 

301 

302 rich.print(table) 

303 

304 

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 

325 

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) 

333 

334 

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 

353 

354 

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) 

365 

366 

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") 

371 

372 for provider, count in provider_counts.items(): 

373 summary_table.add_row(str(provider), str(count)) 

374 

375 total: Final = sum(provider_counts.values()) 

376 summary_table.add_row("[bold]Total[/bold]", f"[bold]{total}[/bold]") 

377 

378 rich.print(summary_table) 

379 

380 

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 

391 

392 

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)] 

398 

399 

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" 

405 

406 

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 ) 

439 

440 if not dry_run: 

441 client: Final = create_client(ctx) 

442 

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 

456 

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)