diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 10f95cafff6..7c8cca2ff65 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -564,6 +564,18 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None return True +def _dispatched_model_name(model_name: str, valid_token: UserAPIKeyAuth) -> str: + after_team_alias: Final = alias_map(valid_token.team_model_aliases).get(model_name, model_name) + return alias_map(valid_token.aliases).get(after_team_alias, after_team_alias) + + +def is_dispatched_model_cost_zero( + model: str | list[str] | None, llm_router: Router | None, valid_token: UserAPIKeyAuth +) -> bool: + dispatched_model: Final = _dispatched_model_name(model, valid_token) if isinstance(model, str) else model + return _is_model_cost_zero(model=dispatched_model, llm_router=llm_router) + + _NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({}) _TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True}) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index e3ce9bcd850..fd4509dae66 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -48,7 +48,6 @@ from litellm.proxy.auth.auth_checks import ( _check_end_user_budget, _delete_cache_key_object, _get_user_role, - _is_model_cost_zero, _is_user_proxy_admin, _team_member_max_budget_alert_check, _virtual_key_max_budget_alert_check, @@ -65,6 +64,7 @@ from litellm.proxy.auth.auth_checks import ( get_team_membership, get_team_object, get_user_object, + is_dispatched_model_cost_zero, is_valid_fallback_model, jwt_key_mapping_cache_key, key_model_aliases_for_auth_check, @@ -1795,9 +1795,9 @@ async def _user_api_key_auth_builder( ) skip_budget_checks = False if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero - - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = is_dispatched_model_cost_zero( + model=model, llm_router=llm_router, valid_token=valid_token + ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -2237,9 +2237,9 @@ async def _user_api_key_auth_builder( ) skip_budget_checks = False if model is not None and llm_router is not None: - from litellm.proxy.auth.auth_checks import _is_model_cost_zero - - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = is_dispatched_model_cost_zero( + model=model, llm_router=llm_router, valid_token=valid_token + ) if skip_budget_checks: verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model) @@ -2965,7 +2965,7 @@ async def _run_centralized_common_checks( route=route, request=request, llm_router=llm_router, - team_id=user_api_key_auth_obj.team_id, + valid_token=user_api_key_auth_obj, ) # Pin the metadata variable name (litellm_metadata vs metadata) before @@ -3136,17 +3136,17 @@ def _should_skip_budget_checks( route: str, request: Request | None, llm_router: Any | None, - team_id: str | None = None, + valid_token: UserAPIKeyAuth, ) -> bool: model: Final = _get_model_from_request_context( request_data=request_data, route=route, request=request, llm_router=llm_router, - team_id=team_id, + team_id=valid_token.team_id, ) if model is not None and llm_router is not None: - return _is_model_cost_zero(model=model, llm_router=llm_router) + return is_dispatched_model_cost_zero(model=model, llm_router=llm_router, valid_token=valid_token) return False @@ -3715,7 +3715,7 @@ async def _run_post_custom_auth_checks( # every budget check for these; this path did not, so the same request could # be refused under custom auth and served under the other two. skip_budget_checks: Final = ( - _is_model_cost_zero(model=current_model, llm_router=llm_router) + is_dispatched_model_cost_zero(model=current_model, llm_router=llm_router, valid_token=valid_token) 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 7b92494a6dc..d63c4c88de4 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -25,6 +25,7 @@ from litellm.proxy.utils import PrismaClient from litellm.proxy.auth.auth_checks import ( can_team_access_model, _is_model_cost_zero, + is_dispatched_model_cost_zero, _virtual_key_soft_budget_check, _team_soft_budget_check, ) @@ -1574,3 +1575,27 @@ def test_zero_cost_check_follows_an_alias_repointed_at_runtime( assert _is_model_cost_zero(model=alias, llm_router=router) is not expected_after_repoint router.update_settings(model_group_alias=_aliases_to(repointed_target)) assert _is_model_cost_zero(model=alias, llm_router=router) is expected_after_repoint + + +@pytest.mark.parametrize( + "team_model_aliases, key_aliases, model, expected", + [ + (None, {"visible": "paid-model"}, "visible", False), + (None, {"free-model": "paid-model"}, "free-model", False), + ({"free-model": "paid-model"}, None, "free-model", False), + (None, {"my-free": "free-model"}, "my-free", True), + (None, {"my-free": "visible"}, "my-free", True), + ({"team-name": "key-name"}, {"key-name": "free-model"}, "team-name", True), + ({"team-name": "key-name"}, {"key-name": "paid-model"}, "team-name", False), + (None, {"other": "paid-model"}, "visible", True), + ], +) +def test_zero_cost_check_prices_the_model_the_key_and_team_aliases_dispatch_to( + team_model_aliases: dict[str, str] | None, + key_aliases: dict[str, str] | None, + model: str, + expected: bool, +) -> None: + router: Final = _alias_router(_aliases_to("free-model")) + token: Final = UserAPIKeyAuth(aliases=key_aliases or {}, team_model_aliases=team_model_aliases) + assert is_dispatched_model_cost_zero(model=model, llm_router=router, valid_token=token) is expected 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 9cdac341b1f..2feec498b56 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -1794,3 +1794,47 @@ def test_mapped_key_jwt_falls_through_to_the_shared_user_budget_attach(): "the mapped-key branch returns before the shared virtual-key checks, so the " "user's per-model budget is never attached and never enforced" ) + + +@pytest.mark.parametrize( + "key_aliases, expected", + [ + ({}, True), + ({"free-alias": "paid-model"}, False), + ], +) +def test_budget_skip_judges_the_model_a_key_alias_dispatches_to(key_aliases: dict[str, str], expected: bool) -> None: + 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"}, + ) + skipped = _should_skip_budget_checks( + request_data={"model": "free-alias"}, + route="/chat/completions", + request=None, + llm_router=router, + valid_token=UserAPIKeyAuth(aliases=key_aliases), + ) + assert skipped is expected