mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
Merge pull request #40372 from BerriAI/litellm_cli_skip_cost_map_fetch
fix(cli): skip remote model cost map fetch in lite CLI processes
This commit is contained in:
commit
b23995ee29
3 changed files with 43 additions and 5 deletions
|
|
@ -2,6 +2,7 @@
|
|||
Pulls the cost + context window + provider route for known models from https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json
|
||||
|
||||
This can be disabled by setting the LITELLM_LOCAL_MODEL_COST_MAP environment variable to True.
|
||||
The ``lite`` and ``litellm-proxy`` CLI entry points also use the bundled map without fetching.
|
||||
|
||||
```
|
||||
export LITELLM_LOCAL_MODEL_COST_MAP=True
|
||||
|
|
@ -13,12 +14,14 @@ import hashlib
|
|||
import json
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime, timezone
|
||||
from importlib.resources import files
|
||||
from pathlib import Path
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
|
@ -34,6 +37,12 @@ from litellm.litellm_core_utils.fallback_generalizations import (
|
|||
)
|
||||
|
||||
FALLBACK_GENERALIZATIONS_KEY: Final = "fallback_generalizations"
|
||||
_CLI_ENTRYPOINT_NAMES: Final = frozenset({"lite", "litellm-proxy"})
|
||||
|
||||
|
||||
def _is_cli_process() -> bool:
|
||||
return Path(sys.argv[0]).stem in _CLI_ENTRYPOINT_NAMES
|
||||
|
||||
|
||||
# Reserved top-level keys that are not model entries. They must be excluded
|
||||
# from the model-count integrity check so a real upstream shrink can't be masked.
|
||||
|
|
@ -600,8 +609,12 @@ def get_model_cost_map(
|
|||
"""
|
||||
Public entry point — returns the model cost map dict.
|
||||
|
||||
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
|
||||
2. Otherwise fetches from ``url``, retrying transient errors in a background thread.
|
||||
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set or this is a ``lite`` /
|
||||
``litellm-proxy`` CLI process, uses the local backup only.
|
||||
2. Otherwise fetches from ``url``, retrying transient HTTP errors
|
||||
(429/5xx/transport) with Retry-After-aware backoff in a background
|
||||
thread, validates integrity, and falls back to the local backup on any
|
||||
failure.
|
||||
|
||||
Only the backup model count is cached (a single int) for validation.
|
||||
The full backup dict is only parsed when it must be *returned* as a
|
||||
|
|
@ -610,7 +623,7 @@ def get_model_cost_map(
|
|||
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
|
||||
# Note: can't use get_secret_bool here — this runs during litellm.__init__
|
||||
# before litellm._key_management_settings is set.
|
||||
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true":
|
||||
if os.getenv("LITELLM_LOCAL_MODEL_COST_MAP", "").lower() == "true" or _is_cli_process():
|
||||
_cost_map_source_info.source = "local"
|
||||
_cost_map_source_info.url = None
|
||||
_cost_map_source_info.is_env_forced = True
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ count actual model entries, not reserved meta keys) and the extraction of the
|
|||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
|
@ -749,3 +750,29 @@ def test_boot_load_that_fails_the_integrity_check_reports_the_backup_not_the_rej
|
|||
assert source["etag"] is None
|
||||
assert source["source_revision"] == _bundled_blob_id()
|
||||
assert source["source_revision"] != git_blob_id(shrunk_body)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("argv0", "request_count"),
|
||||
[
|
||||
("/some/venv/bin/lite", 0),
|
||||
("/some/venv/bin/lite.exe", 0),
|
||||
("/some/venv/bin/python", 1),
|
||||
],
|
||||
)
|
||||
def test_boot_load_skips_remote_fetch_for_cli_processes(
|
||||
monkeypatch: pytest.MonkeyPatch, argv0: str, request_count: int
|
||||
) -> None:
|
||||
monkeypatch.setattr(sys, "argv", [argv0, "--version"])
|
||||
monkeypatch.delenv("LITELLM_LOCAL_MODEL_COST_MAP", raising=False)
|
||||
client, calls = _mock_client([httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client)
|
||||
|
||||
cost_map = get_model_cost_map(url=_URL, client=client)
|
||||
|
||||
assert calls["count"] == request_count
|
||||
assert cost_map
|
||||
source = get_model_cost_map_source_info()
|
||||
if request_count == 0:
|
||||
assert source["source"] == "local"
|
||||
else:
|
||||
assert source["source"] == "remote"
|
||||
|
|
|
|||
|
|
@ -7,8 +7,6 @@ from unittest.mock import Mock, patch
|
|||
import pytest
|
||||
from click.testing import CliRunner
|
||||
|
||||
|
||||
|
||||
import litellm.proxy.client.cli
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.proxy.client.cli import cli
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue