diff --git a/tests/test_litellm/test_muse_spark_1_2_model_metadata.py b/tests/test_litellm/test_muse_spark_1_2_model_metadata.py index 81ef96f9b7a..dcce92b0c17 100644 --- a/tests/test_litellm/test_muse_spark_1_2_model_metadata.py +++ b/tests/test_litellm/test_muse_spark_1_2_model_metadata.py @@ -3,6 +3,7 @@ from pathlib import Path import pytest +import litellm from litellm.cost_calculator import cost_per_token from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider @@ -20,6 +21,21 @@ def _load_cost_map(filename: str = "model_prices_and_context_window.json") -> di return json.load(f) +@pytest.fixture +def local_model_cost_map(monkeypatch): + """Force the bundled backup cost map so assertions don't depend on the + network-fetched ``main`` copy (which lags this branch until merge).""" + original_model_cost = litellm.model_cost + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm.get_model_info.cache_clear() + try: + yield + finally: + litellm.model_cost = original_model_cost + litellm.get_model_info.cache_clear() + + @pytest.mark.parametrize("model, input_cost, cached_cost, output_cost", PRICING) def test_muse_spark_1_2_model_info(model: str, input_cost: float, cached_cost: float, output_cost: float): info = _load_cost_map().get(model) @@ -54,7 +70,9 @@ def test_muse_spark_1_2_model_info(model: str, input_cost: float, cached_cost: f @pytest.mark.parametrize("model, input_cost, cached_cost, output_cost", PRICING) -def test_muse_spark_1_2_cost_per_token(model: str, input_cost: float, cached_cost: float, output_cost: float): +def test_muse_spark_1_2_cost_per_token( + local_model_cost_map, model: str, input_cost: float, cached_cost: float, output_cost: float +): prompt_cost, completion_cost = cost_per_token(model=model, prompt_tokens=1000, completion_tokens=500) assert prompt_cost == pytest.approx(1000 * input_cost)