mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
fix(proxy): skip hook-added tag checks where auth skips them and roll back reservations on cancel
The post-hook tag budget check now shares one predicate with the auth wrapper for every request auth never runs common_checks on: public routes, pass-through endpoints without auth, no-auth mode and custom auth without the opt-in. A cancellation while the hook-added tags are being reserved releases the counters already taken instead of leaving them charged.
This commit is contained in:
parent
27bc2a80e2
commit
f85cb1dac5
7 changed files with 161 additions and 50 deletions
|
|
@ -18,7 +18,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
|
|
@ -133,7 +133,7 @@ from .auth_checks_organization import (
|
|||
add_team_org_context_to_request_body,
|
||||
organization_role_based_access_check,
|
||||
)
|
||||
from .auth_utils import get_model_from_request, get_request_route_template
|
||||
from .auth_utils import get_model_from_request, get_request_route_template, route_in_additonal_public_routes
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -863,22 +863,52 @@ def route_skips_budget_checks(route: str) -> bool:
|
|||
_AUTHN_FLAGS: Final = ("enable_jwt_auth", "enable_oauth2_auth", "enable_oauth2_proxy_auth")
|
||||
|
||||
|
||||
class _PassThroughEndpointAuth(BaseModel):
|
||||
"""The two fields of a ``pass_through_endpoints`` entry that decide whether auth runs on it."""
|
||||
|
||||
path: str = ""
|
||||
auth: bool | str | None = None
|
||||
|
||||
|
||||
_PASS_THROUGH_ENDPOINTS_ADAPTER: Final = TypeAdapter(tuple[_PassThroughEndpointAuth, ...])
|
||||
|
||||
|
||||
def _is_unauthenticated_pass_through(route: str, general_settings: Mapping[str, object]) -> bool:
|
||||
configured: Final = general_settings.get("pass_through_endpoints")
|
||||
if configured is None:
|
||||
return False
|
||||
try:
|
||||
endpoints: Final = _PASS_THROUGH_ENDPOINTS_ADAPTER.validate_python(configured)
|
||||
except ValidationError:
|
||||
return False
|
||||
return any(endpoint.path == route and endpoint.auth is not True for endpoint in endpoints)
|
||||
|
||||
|
||||
def auth_skips_common_checks(
|
||||
general_settings: Mapping[str, object], master_key: str | None, custom_auth_configured: bool
|
||||
route: str, general_settings: Mapping[str, object], master_key: str | None, custom_auth_configured: bool
|
||||
) -> bool:
|
||||
"""
|
||||
Whether ``user_api_key_auth`` runs no ``common_checks`` at all for this deployment.
|
||||
Whether ``user_api_key_auth`` runs no ``common_checks`` at all for this request.
|
||||
|
||||
That is the case in no-auth dev mode (no master key and no JWT or OAuth2
|
||||
auth configured, so the proxy is unauthenticated by configuration) and behind
|
||||
a custom auth hook that did not opt in with ``custom_auth_run_common_checks``.
|
||||
That is the case on a public route, on a user-configured pass-through endpoint
|
||||
that did not ask for auth, in no-auth dev mode (no master key and no JWT or
|
||||
OAuth2 auth configured, so the proxy is unauthenticated by configuration) and
|
||||
behind a custom auth hook that did not opt in with ``custom_auth_run_common_checks``.
|
||||
Post-auth checks that mirror ``common_checks`` skip themselves on the same terms.
|
||||
"""
|
||||
public_route: Final = route in LiteLLMRoutes.public_routes.value or route_in_additonal_public_routes(
|
||||
current_route=route
|
||||
)
|
||||
no_auth_mode: Final = master_key is None and not any(general_settings.get(flag, False) for flag in _AUTHN_FLAGS)
|
||||
custom_auth_opted_out: Final = custom_auth_configured and not general_settings.get(
|
||||
"custom_auth_run_common_checks", False
|
||||
)
|
||||
return no_auth_mode or custom_auth_opted_out
|
||||
return (
|
||||
public_route
|
||||
or _is_unauthenticated_pass_through(route=route, general_settings=general_settings)
|
||||
or no_auth_mode
|
||||
or custom_auth_opted_out
|
||||
)
|
||||
|
||||
|
||||
async def common_checks(
|
||||
|
|
|
|||
|
|
@ -2521,27 +2521,11 @@ async def _run_centralized_common_checks(
|
|||
user_custom_auth,
|
||||
)
|
||||
|
||||
# Public routes (e.g. /health/liveness) are exempt from
|
||||
# auth in the builder — the wrapper must not retroactively apply
|
||||
# authz on top, or k8s readiness probes and other unauthenticated
|
||||
# callers get 401.
|
||||
if route in LiteLLMRoutes.public_routes.value or route_in_additonal_public_routes(current_route=route):
|
||||
return
|
||||
|
||||
# User-configured pass-through endpoints with ``auth: false`` are
|
||||
# explicitly unauthenticated — the builder returns an empty
|
||||
# UserAPIKeyAuth() and the request is forwarded as-is. Running
|
||||
# common_checks on the empty token would reject the request as
|
||||
# admin-only. The "auth" flag on the endpoint config is the
|
||||
# contract; honor it.
|
||||
pass_through_endpoints: Final = general_settings.get("pass_through_endpoints", None)
|
||||
if pass_through_endpoints is not None:
|
||||
for endpoint in pass_through_endpoints:
|
||||
if isinstance(endpoint, dict) and endpoint.get("path", "") == route and endpoint.get("auth") is not True:
|
||||
return
|
||||
|
||||
if auth_skips_common_checks(
|
||||
general_settings=general_settings, master_key=master_key, custom_auth_configured=user_custom_auth is not None
|
||||
route=route,
|
||||
general_settings=general_settings,
|
||||
master_key=master_key,
|
||||
custom_auth_configured=user_custom_auth is not None,
|
||||
):
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -695,7 +695,10 @@ async def _enforce_tag_budgets_for_added_tags(
|
|||
from litellm.proxy.proxy_server import master_key, prisma_client, user_api_key_cache, user_custom_auth
|
||||
|
||||
if auth_skips_common_checks(
|
||||
general_settings=general_settings, master_key=master_key, custom_auth_configured=user_custom_auth is not None
|
||||
route=route,
|
||||
general_settings=general_settings,
|
||||
master_key=master_key,
|
||||
custom_auth_configured=user_custom_auth is not None,
|
||||
):
|
||||
return ()
|
||||
await tag_max_budget_check_for_tags(
|
||||
|
|
|
|||
|
|
@ -360,10 +360,9 @@ async def _reserve_counters(
|
|||
fail_closed_budget_enforcement=fail_closed_budget_enforcement,
|
||||
)
|
||||
continue
|
||||
except Exception:
|
||||
await _release_applied_entries_best_effort(
|
||||
entries=applied_entries,
|
||||
default_reserved_cost=reservation_cost,
|
||||
except BaseException:
|
||||
await asyncio.shield(
|
||||
_release_applied_entries_best_effort(entries=applied_entries, default_reserved_cost=reservation_cost)
|
||||
)
|
||||
raise
|
||||
|
||||
|
|
|
|||
|
|
@ -2482,27 +2482,60 @@ def test_route_skips_budget_checks_matches_auth_scope(route, expected):
|
|||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("general_settings", "master_key", "custom_auth_configured", "expected"),
|
||||
("route", "general_settings", "master_key", "custom_auth_configured", "expected"),
|
||||
[
|
||||
({}, None, False, True),
|
||||
({"enable_jwt_auth": True}, None, False, False),
|
||||
({"enable_oauth2_auth": True}, None, False, False),
|
||||
({"enable_oauth2_proxy_auth": True}, None, False, False),
|
||||
({}, "sk-master", False, False),
|
||||
({}, "sk-master", True, True),
|
||||
({"custom_auth_run_common_checks": True}, "sk-master", True, False),
|
||||
({"custom_auth_run_common_checks": False}, "sk-master", True, True),
|
||||
("/v1/chat/completions", {}, None, False, True),
|
||||
("/v1/chat/completions", {"enable_jwt_auth": True}, None, False, False),
|
||||
("/v1/chat/completions", {"enable_oauth2_auth": True}, None, False, False),
|
||||
("/v1/chat/completions", {"enable_oauth2_proxy_auth": True}, None, False, False),
|
||||
("/v1/chat/completions", {}, "sk-master", False, False),
|
||||
("/v1/chat/completions", {}, "sk-master", True, True),
|
||||
("/v1/chat/completions", {"custom_auth_run_common_checks": True}, "sk-master", True, False),
|
||||
("/v1/chat/completions", {"custom_auth_run_common_checks": False}, "sk-master", True, True),
|
||||
("/health/liveliness", {}, "sk-master", False, True),
|
||||
("/v1/chat/completions", {"public_routes": ["/v1/chat/completions"]}, "sk-master", False, True),
|
||||
("/v1/chat/completions", {"public_routes": ["/v1/embeddings"]}, "sk-master", False, False),
|
||||
("/bria", {"pass_through_endpoints": [{"path": "/bria", "target": "https://x"}]}, "sk-master", False, True),
|
||||
(
|
||||
"/bria",
|
||||
{"pass_through_endpoints": [{"path": "/bria", "target": "https://x", "auth": False}]},
|
||||
"sk-master",
|
||||
False,
|
||||
True,
|
||||
),
|
||||
(
|
||||
"/bria",
|
||||
{"pass_through_endpoints": [{"path": "/bria", "target": "https://x", "auth": True}]},
|
||||
"sk-master",
|
||||
False,
|
||||
False,
|
||||
),
|
||||
(
|
||||
"/other",
|
||||
{"pass_through_endpoints": [{"path": "/bria", "target": "https://x", "auth": False}]},
|
||||
"sk-master",
|
||||
False,
|
||||
False,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_auth_skips_common_checks_names_the_deployments_that_never_run_them(
|
||||
general_settings, master_key, custom_auth_configured, expected
|
||||
def test_auth_skips_common_checks_names_the_requests_that_never_run_them(
|
||||
monkeypatch, route, general_settings, master_key, custom_auth_configured, expected
|
||||
):
|
||||
"""No-auth dev mode and a custom auth hook without the opt-in run no common_checks, so no budget checks."""
|
||||
"""Public routes, pass-through endpoints without auth, no-auth dev mode and a custom auth hook without
|
||||
the opt-in run no common_checks, so no budget checks."""
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.auth_checks import auth_skips_common_checks
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", general_settings)
|
||||
monkeypatch.setattr(proxy_server, "premium_user", True)
|
||||
|
||||
assert (
|
||||
auth_skips_common_checks(
|
||||
general_settings=general_settings, master_key=master_key, custom_auth_configured=custom_auth_configured
|
||||
route=route,
|
||||
general_settings=general_settings,
|
||||
master_key=master_key,
|
||||
custom_auth_configured=custom_auth_configured,
|
||||
)
|
||||
is expected
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
|
@ -139,6 +141,7 @@ async def test_repeated_token_counting_never_touches_a_tiny_budget(
|
|||
|
||||
|
||||
HOOK_TAG: Final = "hook-added-tag"
|
||||
SECOND_HOOK_TAG: Final = "second-hook-added-tag"
|
||||
BODY_TAG: Final = "body-tag"
|
||||
CHAT_BODY: Final[dict[str, object]] = {
|
||||
"model": "gpt-4o",
|
||||
|
|
@ -176,11 +179,14 @@ def _budgeted_tag_prisma(tag_names: tuple[str, ...], max_budget: float) -> Magic
|
|||
|
||||
|
||||
async def _reserve_added_tags(
|
||||
route: str, prisma: MagicMock, tags: tuple[str, ...] = (HOOK_TAG,)
|
||||
route: str,
|
||||
prisma: MagicMock,
|
||||
tags: tuple[str, ...] = (HOOK_TAG,),
|
||||
request_body: Mapping[str, object] = CHAT_BODY,
|
||||
) -> dict[str, object] | None:
|
||||
return await reserve_budget_for_added_tags(
|
||||
tags=tags,
|
||||
request_body=dict(CHAT_BODY),
|
||||
request_body=dict(request_body),
|
||||
route=route,
|
||||
llm_router=None,
|
||||
valid_token=UserAPIKeyAuth(token="hashed-hook-tag-key", max_budget=100.0, spend=0.0),
|
||||
|
|
@ -231,11 +237,63 @@ async def test_reserve_budget_for_added_tags_skips_routes_auth_never_reserves(sp
|
|||
assert spend_counter_cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") is None
|
||||
|
||||
|
||||
class _ParkingIncrementCache(DualCache):
|
||||
"""A spend-counter cache whose increment of ``parked_key`` never returns, so a test can cancel mid-reservation."""
|
||||
|
||||
def __init__(self, parked_key: str) -> None:
|
||||
super().__init__()
|
||||
self.parked_key: Final = parked_key
|
||||
self.parked: Final = asyncio.Event()
|
||||
|
||||
async def async_increment_cache(self, key: str, value: float, **kwargs: object) -> float | None:
|
||||
if key == self.parked_key:
|
||||
self.parked.set()
|
||||
await asyncio.Event().wait()
|
||||
return await super().async_increment_cache(key=key, value=value, **kwargs)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_cancelled_mid_acquisition(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""A client disconnect under SSE keepalives cancels the request while the second tag is being reserved;
|
||||
the first tag's counter must not stay charged for a request that never reached the provider."""
|
||||
cache: Final = _ParkingIncrementCache(parked_key=f"spend:tag:{SECOND_HOOK_TAG}")
|
||||
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
prisma: Final = _budgeted_tag_prisma((HOOK_TAG, SECOND_HOOK_TAG), max_budget=1.0)
|
||||
|
||||
reserving: Final = asyncio.create_task(
|
||||
_reserve_added_tags("/v1/chat/completions", prisma, tags=(HOOK_TAG, SECOND_HOOK_TAG))
|
||||
)
|
||||
await cache.parked.wait()
|
||||
assert cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") > 0
|
||||
reserving.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await reserving
|
||||
|
||||
assert cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") == pytest.approx(0.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_budget_for_added_tags_ignores_tags_without_a_budget(spend_counter_cache: DualCache):
|
||||
assert await _reserve_added_tags("/v1/chat/completions", _budgeted_tag_prisma((), max_budget=1.0)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_budget_for_added_tags_skips_a_request_with_no_model_to_price(spend_counter_cache: DualCache):
|
||||
"""Without a model there is no estimate to reserve, same as the auth-time reservation."""
|
||||
body: Final = MappingProxyType({key: value for key, value in CHAT_BODY.items() if key != "model"})
|
||||
|
||||
assert (
|
||||
await _reserve_added_tags(
|
||||
"/v1/chat/completions", _budgeted_tag_prisma((HOOK_TAG,), max_budget=1.0), request_body=body
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert spend_counter_cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") is None
|
||||
|
||||
|
||||
BEDROCK_SONNET: Final = "us.anthropic.claude-sonnet-4-6"
|
||||
CONVERSE_BODY: Final = {
|
||||
"messages": [{"role": "user", "content": [{"text": "Reply with one word: pong"}]}],
|
||||
|
|
|
|||
|
|
@ -621,13 +621,15 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
("sk-master", object(), {}, False),
|
||||
("sk-master", object(), {"custom_auth_run_common_checks": True}, True),
|
||||
("sk-master", None, {}, True),
|
||||
("sk-master", None, {"public_routes": ["/v1/chat/completions"]}, False),
|
||||
("sk-master", None, {"public_routes": ["/v1/embeddings"]}, True),
|
||||
],
|
||||
)
|
||||
async def test_common_processing_pre_call_logic_enforces_hook_added_tags_only_where_auth_runs_common_checks(
|
||||
self, monkeypatch, master_key, user_custom_auth, general_settings, checked
|
||||
):
|
||||
"""A deployment whose auth wrapper skips common_checks (no-auth dev mode, custom auth without opt-in)
|
||||
never budget-checked tags before, so a hook-added tag must not start 429ing it."""
|
||||
"""A request whose auth wrapper skips common_checks (a public route, no-auth dev mode, custom auth
|
||||
without opt-in) never budget-checked tags before, so a hook-added tag must not start 429ing it."""
|
||||
processing_obj = ProxyBaseLLMRequestProcessing(data={})
|
||||
|
||||
async def mock_pre_call_hook(user_api_key_dict, data, call_type):
|
||||
|
|
@ -639,6 +641,8 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", master_key)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_auth", user_custom_auth)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
|
||||
await processing_obj.common_processing_pre_call_logic(
|
||||
request=mock_request,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue