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:
ryan-crabbe-berri 2026-08-27 14:08:51 -07:00
parent 88cb83484b
commit 3ea501430b
9 changed files with 570 additions and 612 deletions

View 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)

View file

@ -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:

View file

@ -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

View file

@ -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:

View file

@ -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",

View 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()
)

View file

@ -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

View file

@ -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