diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 7673c3983b4..1cfda5cd9bf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -450,10 +450,6 @@ def _get_router_zero_cost_cache(llm_router: Router) -> dict[str, bool] | None: return cache if isinstance(cache, dict) else None -def _resolve_cost_model_group(model_name: str, llm_router: Router) -> str: - return resolve_model_group_alias(llm_router.model_group_alias, model_name) or model_name - - def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: """ Check if a model has zero cost (no configured pricing). @@ -467,7 +463,25 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None Returns: bool: True if all costs for the model are zero, False otherwise """ - if model is None or llm_router is None: + if llm_router is None: + return False + return _is_target_group_cost_zero( + model=model, + llm_router=llm_router, + target_group_of=lambda name: resolve_model_group_alias(llm_router.model_group_alias, name) or name, + ) + + +def is_requested_model_cost_zero(model: str | list[str] | None, llm_router: Router | None) -> bool: + if llm_router is None: + return False + return _is_target_group_cost_zero(model=model, llm_router=llm_router, target_group_of=lambda name: name) + + +def _is_target_group_cost_zero( + model: str | list[str] | None, llm_router: Router, target_group_of: Callable[[str], str] +) -> bool: + if model is None: return False # Handle list of models @@ -477,7 +491,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None for model_name in model_list: try: - target_group = _resolve_cost_model_group(model_name, llm_router) + target_group = target_group_of(model_name) if zero_cost_cache is not None: cached = zero_cost_cache.get(target_group) if cached is not None: diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 2ab76a7a101..b19b8c9232f 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -395,13 +395,7 @@ class RouteChecks: if not isinstance(route, str): return False - if route in LiteLLMRoutes.openai_routes.value: - return True - - if route in LiteLLMRoutes.anthropic_routes.value: - return True - - if route in LiteLLMRoutes.google_routes.value: + if RouteChecks.is_unified_llm_api_route(route): return True if RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.mcp_inference_routes.value): @@ -413,6 +407,22 @@ class RouteChecks: if route in LiteLLMRoutes.litellm_native_routes.value: return True + for _llm_passthrough_route in LiteLLMRoutes.mapped_pass_through_routes.value: + if route == _llm_passthrough_route or route.startswith(_llm_passthrough_route + "/"): + return True + return False + + @staticmethod + def is_unified_llm_api_route(route: str) -> bool: + if route in LiteLLMRoutes.openai_routes.value: + return True + + if route in LiteLLMRoutes.anthropic_routes.value: + return True + + if route in LiteLLMRoutes.google_routes.value: + return True + # fuzzy match routes like "/v1/threads/thread_49EIN5QF32s4mH20M7GFKdlZ" # Check for routes with placeholders or wildcard patterns for openai_route in LiteLLMRoutes.openai_routes.value: @@ -438,13 +448,7 @@ class RouteChecks: if RouteChecks._route_matches_pattern(route=route, pattern=anthropic_route): return True - if RouteChecks._is_azure_openai_route(route=route): - return True - - for _llm_passthrough_route in LiteLLMRoutes.mapped_pass_through_routes.value: - if route == _llm_passthrough_route or route.startswith(_llm_passthrough_route + "/"): - return True - return False + return RouteChecks._is_azure_openai_route(route=route) @staticmethod def _is_get_mcp_server_discovery_route(route: str, request: Request | None) -> bool: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 8ced8eed481..0bc891722d5 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -65,6 +65,7 @@ from litellm.proxy.auth.auth_checks import ( get_team_object, get_user_object, is_dispatched_model_cost_zero, + is_requested_model_cost_zero, is_valid_fallback_model, jwt_key_mapping_cache_key, key_model_aliases_for_auth_check, @@ -1796,7 +1797,7 @@ async def _user_api_key_auth_builder( skip_budget_checks = False if model is not None and llm_router is not None: skip_budget_checks = await _is_dispatched_model_cost_zero( - model=model, llm_router=llm_router, valid_token=valid_token, request=request + model=model, llm_router=llm_router, valid_token=valid_token, request=request, route=route ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -2238,7 +2239,7 @@ async def _user_api_key_auth_builder( skip_budget_checks = False if model is not None and llm_router is not None: skip_budget_checks = await _is_dispatched_model_cost_zero( - model=model, llm_router=llm_router, valid_token=valid_token, request=request + model=model, llm_router=llm_router, valid_token=valid_token, request=request, route=route ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -3147,14 +3148,21 @@ async def _should_skip_budget_checks( ) if model is not None and llm_router is not None: return await _is_dispatched_model_cost_zero( - model=model, llm_router=llm_router, valid_token=valid_token, request=request + model=model, llm_router=llm_router, valid_token=valid_token, request=request, route=route ) return False async def _is_dispatched_model_cost_zero( - model: str | list[str], llm_router: litellm.Router, valid_token: UserAPIKeyAuth, request: Request | None + model: str | list[str], + llm_router: litellm.Router, + valid_token: UserAPIKeyAuth, + request: Request | None, + route: str, ) -> bool: + if not RouteChecks.is_unified_llm_api_route(route): + return is_requested_model_cost_zero(model=model, llm_router=llm_router) + from litellm.proxy.proxy_server import prisma_client, proxy_config, proxy_logging_obj settings: Final = await proxy_config.get_hierarchical_router_settings( @@ -3736,7 +3744,7 @@ async def _run_post_custom_auth_checks( # be refused under custom auth and served under the other two. skip_budget_checks: Final = ( await _is_dispatched_model_cost_zero( - model=current_model, llm_router=llm_router, valid_token=valid_token, request=request + model=current_model, llm_router=llm_router, valid_token=valid_token, request=request, route=route ) if current_model is not None and llm_router is not None else False diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index 3f1af71c06a..3d2a1aaebf9 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -26,6 +26,7 @@ from litellm.proxy.auth.auth_checks import ( can_team_access_model, _is_model_cost_zero, is_dispatched_model_cost_zero, + is_requested_model_cost_zero, _virtual_key_soft_budget_check, _team_soft_budget_check, ) @@ -1560,6 +1561,23 @@ def test_zero_cost_check_prices_an_alias_as_its_target_group( assert _is_model_cost_zero(model=model, llm_router=router) is expected +@pytest.mark.parametrize( + "target, model, expected", + [ + ("free-model", "visible", False), + ("free-model", "hidden", False), + ("free-model", "free-model", True), + ("paid-model", "paid-model", False), + ], +) +def test_requested_name_zero_cost_check_prices_the_name_without_the_global_alias( + target: str, model: str, expected: bool +) -> None: + router: Final = _alias_router(_aliases_to(target)) + assert _is_model_cost_zero(model=model, llm_router=router) is (target == "free-model") + assert is_requested_model_cost_zero(model=model, llm_router=router) is expected + + @pytest.mark.parametrize( "first_target, repointed_target, expected_after_repoint", [ diff --git a/tests/unit/proxy/auth/test_proxy_routes.py b/tests/unit/proxy/auth/test_proxy_routes.py index 129a93ea08d..fc4e5ff46a3 100644 --- a/tests/unit/proxy/auth/test_proxy_routes.py +++ b/tests/unit/proxy/auth/test_proxy_routes.py @@ -136,6 +136,28 @@ def test_anthropic_api_routes(): assert RouteChecks.is_llm_api_route(route="/v1/messages") is True +@pytest.mark.parametrize( + "route, expected", + [ + ("/v1/chat/completions", True), + ("/v1/messages", True), + ("/v1/responses", True), + ("/v1beta/models/gemini-pro:generateContent", True), + ("/engines/gpt-4/chat/completions", True), + ("/openai/deployments/gpt-4o/chat/completions", True), + ("/openai/deployments/vertex_ai/gemini-1.5-flash/chat/completions", True), + ("/openai/v1/chat/completions", False), + ("/anthropic/v1/messages", False), + ("/bedrock/model/cohere.command-r-v1:0/converse", False), + ("/rag/query", False), + ("/mcp/tools/call", False), + ], +) +def test_unified_llm_api_route_excludes_routes_that_forward_the_model_name_as_sent(route: str, expected: bool): + assert RouteChecks.is_unified_llm_api_route(route) is expected + assert RouteChecks.is_llm_api_route(route) is True + + def create_request(path: str, base_url: str = "http://testserver") -> Request: return Request( { diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index aa09f6ae45f..61af2e270a0 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -1796,6 +1796,34 @@ def test_mapped_key_jwt_falls_through_to_the_shared_user_budget_attach(): ) +def _free_and_paid_router() -> litellm.Router: + from litellm.router import Router + + return Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + }, + { + "model_name": "paid-model", + "litellm_params": { + "model": "openai/paid-model", + "api_key": "sk-fake", + "input_cost_per_token": 1e-06, + "output_cost_per_token": 2e-06, + }, + }, + ], + model_group_alias={"free-alias": "free-model"}, + ) + + @pytest.mark.asyncio @pytest.mark.parametrize( "model, key_aliases, key_router_settings, router_settings_rewritten, expected", @@ -1819,31 +1847,8 @@ async def test_budget_skip_judges_the_model_the_key_aliases_dispatch_to( from litellm.constants import MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY from litellm.proxy.auth.user_api_key_auth import _should_skip_budget_checks - from litellm.router import Router - router = Router( - model_list=[ - { - "model_name": "free-model", - "litellm_params": { - "model": "ollama/llama2", - "api_base": "http://localhost:11434", - "input_cost_per_token": 0.0, - "output_cost_per_token": 0.0, - }, - }, - { - "model_name": "paid-model", - "litellm_params": { - "model": "openai/paid-model", - "api_key": "sk-fake", - "input_cost_per_token": 1e-06, - "output_cost_per_token": 2e-06, - }, - }, - ], - model_group_alias={"free-alias": "free-model"}, - ) + router = _free_and_paid_router() skipped = await _should_skip_budget_checks( request_data={"model": model}, route="/chat/completions", @@ -1865,3 +1870,35 @@ async def test_budget_skip_judges_the_model_the_key_aliases_dispatch_to( valid_token=UserAPIKeyAuth(aliases=key_aliases, router_settings=key_router_settings), ) assert skipped is expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "route, model, key_aliases, key_router_settings, expected", + [ + ("/chat/completions", "free-alias", {}, None, True), + ("/anthropic/v1/messages", "free-alias", {}, None, False), + ("/anthropic/v1/messages", "my-alias", {"my-alias": "free-model"}, None, False), + ("/anthropic/v1/messages", "rs-alias", {}, {"model_group_alias": {"rs-alias": "free-model"}}, False), + ("/anthropic/v1/messages", "free-model", {"free-model": "paid-model"}, None, True), + ("/rag/query", "free-alias", {}, None, False), + ], +) +async def test_budget_skip_prices_the_requested_name_on_routes_that_forward_it_unaliased( + route: str, + model: str, + key_aliases: dict[str, str], + key_router_settings: dict[str, dict[str, str]] | None, + expected: bool, +) -> None: + from litellm.proxy.auth.user_api_key_auth import _should_skip_budget_checks + + router = _free_and_paid_router() + skipped = await _should_skip_budget_checks( + request_data={"model": model}, + route=route, + request=None, + llm_router=router, + valid_token=UserAPIKeyAuth(aliases=key_aliases, router_settings=key_router_settings), + ) + assert skipped is expected