mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(router): authorize config-level fallback targets against the calling key
Router fallbacks configured in router_settings were attempted without re-checking whether the calling key could use the fallback model, so a key limited to one access group was served by any model listed as a fallback for something it could call. Auth only validated the requested model and fallbacks sent in the request body. Add a fallback_access_check predicate to Router, consulted before every cross-model-group fallback attempt; rejected targets are skipped and the primary's own error is raised when none remain. The proxy injects a check that runs the same key, team and project model access checks the requested model goes through.
This commit is contained in:
parent
88cb83484b
commit
3ea501430b
9 changed files with 570 additions and 612 deletions
64
litellm/proxy/auth/fallback_model_access.py
Normal file
64
litellm/proxy/auth/fallback_model_access.py
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
"""
|
||||
Authorize router fallback targets against the caller's key, team and project model access.
|
||||
|
||||
`_enforce_key_and_fallback_model_access` only sees fallbacks the client sends in the request body.
|
||||
Fallbacks configured on the router (`router_settings.fallbacks` and friends) are chosen after auth,
|
||||
inside the router, so this predicate is injected into the router to re-run the same model access
|
||||
checks for each fallback target before it is attempted.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
class _RequestMetadata(BaseModel):
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None
|
||||
|
||||
|
||||
async def is_model_authorized_for_token(*, model: str, valid_token: UserAPIKeyAuth, llm_router: Router) -> bool:
|
||||
try:
|
||||
await can_key_call_resolved_model(
|
||||
model=model,
|
||||
llm_model_list=None,
|
||||
valid_token=valid_token,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyException:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _token_in_metadata(metadata: object) -> UserAPIKeyAuth | None:
|
||||
try:
|
||||
return _RequestMetadata.model_validate(metadata).user_api_key_auth
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _user_api_key_auth_from_request(request_kwargs: Mapping[str, object]) -> UserAPIKeyAuth | None:
|
||||
return next(
|
||||
(
|
||||
token
|
||||
for field in ("metadata", "litellm_metadata")
|
||||
if (token := _token_in_metadata(request_kwargs.get(field))) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
async def router_fallback_access_check(*, model: str, request_kwargs: Mapping[str, object], llm_router: Router) -> bool:
|
||||
"""
|
||||
`FallbackAccessCheck` for the proxy's router: a fallback target is attempted only when the
|
||||
key behind the request could have requested it directly. Requests that carry no key (for
|
||||
example internal health checks) are not restricted.
|
||||
"""
|
||||
valid_token: Final = _user_api_key_auth_from_request(request_kwargs)
|
||||
if valid_token is None:
|
||||
return True
|
||||
return await is_model_authorized_for_token(model=model, valid_token=valid_token, llm_router=llm_router)
|
||||
|
|
@ -290,6 +290,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
is_request_body_safe,
|
||||
warn_once_if_custom_auth_skips_common_checks,
|
||||
)
|
||||
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.litellm_license import LicenseCheck
|
||||
from litellm.proxy.auth.model_checks import (
|
||||
|
|
@ -5580,6 +5581,7 @@ class ProxyConfig:
|
|||
async_only_mode=True # only init async clients
|
||||
),
|
||||
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
)
|
||||
|
||||
if redis_usage_cache is not None and router.cache.redis_cache is None:
|
||||
|
|
@ -6039,6 +6041,7 @@ class ProxyConfig:
|
|||
),
|
||||
search_tools=search_tools,
|
||||
ignore_invalid_deployments=True,
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
)
|
||||
verbose_proxy_logger.debug("updated llm_router: %s", llm_router)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -196,6 +196,7 @@ from litellm.types.router import (
|
|||
CustomRoutingStrategyBase,
|
||||
Deployment,
|
||||
DeploymentTypedDict,
|
||||
FallbackAccessCheck,
|
||||
GuardrailTypedDict,
|
||||
LiteLLM_Params,
|
||||
MockRouterTestingParams,
|
||||
|
|
@ -604,6 +605,7 @@ class Router:
|
|||
health_check_ignore_transient_errors: bool = False,
|
||||
background_health_check_model_groups: Sequence[str] | None = None,
|
||||
enable_weighted_failover: bool = False,
|
||||
fallback_access_check: FallbackAccessCheck | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
|
||||
|
|
@ -640,6 +642,7 @@ class Router:
|
|||
deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600.
|
||||
ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error.
|
||||
enable_weighted_failover (bool): When True and the routing strategy is "simple-shuffle", a retryable failure on one deployment causes the request to re-pick (weighted) across the other deployments in the same model group before any cross-group fallback runs. Bounded by `max_fallbacks`. Async-only: currently honored by `router.acompletion()` and other async entrypoints. The sync `router.completion()` path falls back to the regular fallback flow. Defaults to False.
|
||||
fallback_access_check (Optional[FallbackAccessCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects is skipped. Defaults to None (every configured fallback is attempted).
|
||||
Returns:
|
||||
Router: An instance of the litellm.Router class.
|
||||
|
||||
|
|
@ -679,6 +682,7 @@ class Router:
|
|||
|
||||
self.set_verbose = set_verbose
|
||||
self.ignore_invalid_deployments = ignore_invalid_deployments
|
||||
self.fallback_access_check: Final = fallback_access_check
|
||||
self.debug_level = debug_level
|
||||
self.enable_pre_call_checks = enable_pre_call_checks
|
||||
self.enable_tag_filtering = enable_tag_filtering
|
||||
|
|
|
|||
|
|
@ -263,6 +263,25 @@ def _get_fallback_target_model_group(fallback_entry: str | Mapping[str, object])
|
|||
return target if isinstance(target, str) else None
|
||||
|
||||
|
||||
async def _is_fallback_target_authorized(
|
||||
litellm_router: LitellmRouter,
|
||||
fallback_entry: str | Mapping[str, object],
|
||||
original_model_group: str,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> bool:
|
||||
access_check: Final = litellm_router.fallback_access_check
|
||||
target: Final = _get_fallback_target_model_group(fallback_entry)
|
||||
if access_check is None or target is None or target == original_model_group:
|
||||
return True
|
||||
if await access_check(model=target, request_kwargs=kwargs, llm_router=litellm_router):
|
||||
return True
|
||||
verbose_router_logger.info(
|
||||
"Skipping fallback to model_group = %s: caller is not authorized to call it",
|
||||
mask_sensitive_structure(fallback_entry),
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def references_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
True when the request names a file that only exists under one provider's credentials.
|
||||
|
|
@ -357,6 +376,8 @@ async def run_async_fallback(
|
|||
original_model_group,
|
||||
)
|
||||
continue
|
||||
if not await _is_fallback_target_authorized(litellm_router, mg, original_model_group, kwargs):
|
||||
continue
|
||||
attempt_key = fallback_attempt_key(mg)
|
||||
if attempt_key is not None:
|
||||
if attempt_key in attempted:
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import datetime
|
|||
import enum
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
|
@ -14,6 +14,9 @@ from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_c
|
|||
|
||||
from litellm._uuid import uuid
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
||||
from .completion import CompletionRequest
|
||||
from .embedding import EmbeddingRequest
|
||||
from .llms.openai import OpenAIFileObject
|
||||
|
|
@ -845,6 +848,17 @@ class GenericBudgetWindowDetails(BaseModel):
|
|||
ttl_seconds: int
|
||||
|
||||
|
||||
class FallbackAccessCheck(Protocol):
|
||||
"""
|
||||
Decides whether the caller behind `request_kwargs` may be served by fallback `model`.
|
||||
|
||||
The router runs it before every cross-model-group fallback attempt and skips targets it
|
||||
rejects, so a fallback can never reach a model the caller could not have requested directly.
|
||||
"""
|
||||
|
||||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
|
||||
|
||||
|
||||
OptionalPreCallChecks = list[
|
||||
Literal[
|
||||
"prompt_caching",
|
||||
|
|
|
|||
57
tests/test_litellm/proxy/auth/test_fallback_model_access.py
Normal file
57
tests/test_litellm/proxy/auth/test_fallback_model_access.py
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
import pytest
|
||||
|
||||
from litellm import Router
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.fallback_model_access import (
|
||||
is_model_authorized_for_token,
|
||||
router_fallback_access_check,
|
||||
)
|
||||
|
||||
|
||||
def _router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "open-model",
|
||||
"litellm_params": {"model": "openai/open", "api_key": "k"},
|
||||
"model_info": {"access_groups": ["open-group"]},
|
||||
},
|
||||
{
|
||||
"model_name": "secret-model",
|
||||
"litellm_params": {"model": "openai/secret", "api_key": "k"},
|
||||
"model_info": {"access_groups": ["secret-group"]},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _key_limited_to(access_group: str) -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(api_key="hashed", models=[access_group])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_is_model_authorized_for_token_follows_the_key_access_groups():
|
||||
router = _router()
|
||||
token = _key_limited_to("open-group")
|
||||
|
||||
assert await is_model_authorized_for_token(model="open-model", valid_token=token, llm_router=router) is True
|
||||
assert await is_model_authorized_for_token(model="secret-model", valid_token=token, llm_router=router) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("metadata_field", ["metadata", "litellm_metadata"])
|
||||
async def test_router_fallback_access_check_authorizes_the_key_carried_in_request_metadata(metadata_field: str):
|
||||
router = _router()
|
||||
request_kwargs = {metadata_field: {"user_api_key_auth": _key_limited_to("open-group")}}
|
||||
|
||||
assert await router_fallback_access_check(model="open-model", request_kwargs=request_kwargs, llm_router=router)
|
||||
assert not await router_fallback_access_check(
|
||||
model="secret-model", request_kwargs=request_kwargs, llm_router=router
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_fallback_access_check_does_not_restrict_requests_without_a_key():
|
||||
assert await router_fallback_access_check(
|
||||
model="secret-model", request_kwargs={"metadata": {}}, llm_router=_router()
|
||||
)
|
||||
|
|
@ -11712,3 +11712,18 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the
|
|||
assert real_spend_counter_cache.in_memory_cache.get_cache(key=marker_key) == 0.0, (
|
||||
"the in-flight DB read clobbered the post-reset floor marker with the stale pre-reset value"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_config_router_authorizes_fallback_targets_against_the_calling_key(tmp_path):
|
||||
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
config_file = tmp_path / "config.yaml"
|
||||
config_file.write_text(
|
||||
yaml.dump({"model_list": [{"model_name": "m", "litellm_params": {"model": "openai/m", "api_key": "k"}}]})
|
||||
)
|
||||
|
||||
router, _, _ = await ProxyConfig().load_config(router=None, config_file_path=str(config_file))
|
||||
|
||||
assert router.fallback_access_check is router_fallback_access_check
|
||||
|
|
|
|||
|
|
@ -22,6 +22,8 @@ class StreamingWrapper:
|
|||
|
||||
|
||||
class FakeRouter:
|
||||
fallback_access_check = None
|
||||
|
||||
def log_retry(self, kwargs, e):
|
||||
return kwargs
|
||||
|
||||
|
|
@ -30,6 +32,8 @@ class FakeRouter:
|
|||
|
||||
|
||||
class AlwaysFailRouter:
|
||||
fallback_access_check = None
|
||||
|
||||
def log_retry(self, kwargs, e):
|
||||
return kwargs
|
||||
|
||||
|
|
@ -92,6 +96,8 @@ async def test_run_async_fallback_raises_when_all_fallbacks_fail():
|
|||
|
||||
|
||||
class RecordingRouter:
|
||||
fallback_access_check = None
|
||||
|
||||
def __init__(self):
|
||||
self.received_kwargs = None
|
||||
|
||||
|
|
@ -151,6 +157,8 @@ async def test_run_async_fallback_skips_original_model_group():
|
|||
|
||||
|
||||
class AttemptRecordingRouter:
|
||||
fallback_access_check = None
|
||||
|
||||
def __init__(self):
|
||||
self.attempted_model_groups = []
|
||||
self.received_kwargs = None
|
||||
|
|
@ -339,7 +347,84 @@ async def test_run_async_fallback_records_batch_model_group_outside_provider_met
|
|||
assert router.received_kwargs["litellm_metadata"]["model_group"] == "openai-group"
|
||||
|
||||
|
||||
class AccessCheckedRouter(AttemptRecordingRouter):
|
||||
def __init__(self, allowed_models: frozenset[str]):
|
||||
super().__init__()
|
||||
self.allowed_models = allowed_models
|
||||
self.access_checks = []
|
||||
|
||||
async def fallback_access_check(self, *, model, request_kwargs, llm_router):
|
||||
self.access_checks.append((model, request_kwargs["metadata"]["user_api_key"], llm_router is self))
|
||||
return model in self.allowed_models
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_skips_targets_the_access_check_rejects():
|
||||
router = AccessCheckedRouter(allowed_models=frozenset({"allowed-model"}))
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[
|
||||
{"model": "secret-model", "messages": [{"role": "user", "content": "hi"}]},
|
||||
"allowed-model",
|
||||
],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("primary failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="primary-model",
|
||||
metadata={"user_api_key": "hashed"},
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == ["allowed-model"]
|
||||
assert router.access_checks == [
|
||||
("secret-model", "hashed", True),
|
||||
("allowed-model", "hashed", True),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_raises_original_error_when_no_target_is_authorized():
|
||||
router = AccessCheckedRouter(allowed_models=frozenset())
|
||||
|
||||
with pytest.raises(RuntimeError, match="primary failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=["secret-model", "other-secret-model"],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("primary failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="primary-model",
|
||||
metadata={"user_api_key": "hashed"},
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == []
|
||||
assert [model for model, _, _ in router.access_checks] == ["secret-model", "other-secret-model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_async_fallback_does_not_consult_access_check_for_same_model_group_retries():
|
||||
router = AccessCheckedRouter(allowed_models=frozenset())
|
||||
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[{"model": "primary-model", "_target_order": 2}],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("first order level failed"),
|
||||
max_fallbacks=3,
|
||||
fallback_depth=0,
|
||||
model="primary-model",
|
||||
metadata={"user_api_key": "hashed"},
|
||||
)
|
||||
|
||||
assert router.attempted_model_groups == ["primary-model"]
|
||||
assert router.access_checks == []
|
||||
|
||||
|
||||
class RecordingFailRouter:
|
||||
fallback_access_check = None
|
||||
|
||||
def __init__(self):
|
||||
self.attempted_models = []
|
||||
|
||||
|
|
@ -488,9 +573,7 @@ async def test_run_async_fallback_keeps_a_request_override_distinct_from_the_bar
|
|||
with pytest.raises(RuntimeError, match="fallback model also failed"):
|
||||
await run_async_fallback(
|
||||
litellm_router=router,
|
||||
fallback_model_group=[
|
||||
{"model": "already-attempted", "messages": [{"role": "user", "content": "shorter"}]}
|
||||
],
|
||||
fallback_model_group=[{"model": "already-attempted", "messages": [{"role": "user", "content": "shorter"}]}],
|
||||
original_model_group="primary-model",
|
||||
original_exception=RuntimeError("original failed"),
|
||||
max_fallbacks=3,
|
||||
|
|
@ -774,6 +857,8 @@ class TestTriggerCooldownForFailedDeployment:
|
|||
|
||||
class TestRunAsyncFallbackTriggersCooldown:
|
||||
class RouterWithLoggingKwarg:
|
||||
fallback_access_check = None
|
||||
|
||||
def __init__(self):
|
||||
self.cooldown_time = 60.0
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue