mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(proxy): price pass-through and native routes by the requested model name
Pass-through, MCP, agent, and RAG routes forward the model name as sent and skip the alias steps the unified routes apply, so the zero-cost budget check only follows aliases on the unified inference routes
This commit is contained in:
parent
5b84fc8df9
commit
45d880c4f5
6 changed files with 152 additions and 49 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue