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:
Yucheng He 2026-09-15 04:04:54 -07:00
parent 27bc2a80e2
commit f85cb1dac5
7 changed files with 161 additions and 50 deletions

View file

@ -18,7 +18,7 @@ from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
from fastapi import HTTPException, Request, status from fastapi import HTTPException, Request, status
from pydantic import BaseModel from pydantic import BaseModel, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict from typing_extensions import ReadOnly, TypedDict
import litellm import litellm
@ -133,7 +133,7 @@ from .auth_checks_organization import (
add_team_org_context_to_request_body, add_team_org_context_to_request_body,
organization_role_based_access_check, 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: if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span 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") _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( 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: ) -> 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 That is the case on a public route, on a user-configured pass-through endpoint
auth configured, so the proxy is unauthenticated by configuration) and behind that did not ask for auth, in no-auth dev mode (no master key and no JWT or
a custom auth hook that did not opt in with ``custom_auth_run_common_checks``. 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. 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) 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_opted_out: Final = custom_auth_configured and not general_settings.get(
"custom_auth_run_common_checks", False "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( async def common_checks(

View file

@ -2521,27 +2521,11 @@ async def _run_centralized_common_checks(
user_custom_auth, 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( 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 return

View file

@ -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 from litellm.proxy.proxy_server import master_key, prisma_client, user_api_key_cache, user_custom_auth
if auth_skips_common_checks( 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 () return ()
await tag_max_budget_check_for_tags( await tag_max_budget_check_for_tags(

View file

@ -360,10 +360,9 @@ async def _reserve_counters(
fail_closed_budget_enforcement=fail_closed_budget_enforcement, fail_closed_budget_enforcement=fail_closed_budget_enforcement,
) )
continue continue
except Exception: except BaseException:
await _release_applied_entries_best_effort( await asyncio.shield(
entries=applied_entries, _release_applied_entries_best_effort(entries=applied_entries, default_reserved_cost=reservation_cost)
default_reserved_cost=reservation_cost,
) )
raise raise

View file

@ -2482,27 +2482,60 @@ def test_route_skips_budget_checks_matches_auth_scope(route, expected):
@pytest.mark.parametrize( @pytest.mark.parametrize(
("general_settings", "master_key", "custom_auth_configured", "expected"), ("route", "general_settings", "master_key", "custom_auth_configured", "expected"),
[ [
({}, None, False, True), ("/v1/chat/completions", {}, None, False, True),
({"enable_jwt_auth": True}, None, False, False), ("/v1/chat/completions", {"enable_jwt_auth": True}, None, False, False),
({"enable_oauth2_auth": True}, None, False, False), ("/v1/chat/completions", {"enable_oauth2_auth": True}, None, False, False),
({"enable_oauth2_proxy_auth": True}, None, False, False), ("/v1/chat/completions", {"enable_oauth2_proxy_auth": True}, None, False, False),
({}, "sk-master", False, False), ("/v1/chat/completions", {}, "sk-master", False, False),
({}, "sk-master", True, True), ("/v1/chat/completions", {}, "sk-master", True, True),
({"custom_auth_run_common_checks": True}, "sk-master", True, False), ("/v1/chat/completions", {"custom_auth_run_common_checks": True}, "sk-master", True, False),
({"custom_auth_run_common_checks": False}, "sk-master", True, True), ("/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( def test_auth_skips_common_checks_names_the_requests_that_never_run_them(
general_settings, master_key, custom_auth_configured, expected 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 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 ( assert (
auth_skips_common_checks( 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 is expected
) )

View file

@ -1,7 +1,9 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import json import json
import math import math
from collections.abc import Mapping
from types import MappingProxyType, SimpleNamespace from types import MappingProxyType, SimpleNamespace
from typing import Final from typing import Final
from unittest.mock import AsyncMock, MagicMock 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" HOOK_TAG: Final = "hook-added-tag"
SECOND_HOOK_TAG: Final = "second-hook-added-tag"
BODY_TAG: Final = "body-tag" BODY_TAG: Final = "body-tag"
CHAT_BODY: Final[dict[str, object]] = { CHAT_BODY: Final[dict[str, object]] = {
"model": "gpt-4o", "model": "gpt-4o",
@ -176,11 +179,14 @@ def _budgeted_tag_prisma(tag_names: tuple[str, ...], max_budget: float) -> Magic
async def _reserve_added_tags( 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: ) -> dict[str, object] | None:
return await reserve_budget_for_added_tags( return await reserve_budget_for_added_tags(
tags=tags, tags=tags,
request_body=dict(CHAT_BODY), request_body=dict(request_body),
route=route, route=route,
llm_router=None, llm_router=None,
valid_token=UserAPIKeyAuth(token="hashed-hook-tag-key", max_budget=100.0, spend=0.0), 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 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 @pytest.mark.asyncio
async def test_reserve_budget_for_added_tags_ignores_tags_without_a_budget(spend_counter_cache: DualCache): 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 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" BEDROCK_SONNET: Final = "us.anthropic.claude-sonnet-4-6"
CONVERSE_BODY: Final = { CONVERSE_BODY: Final = {
"messages": [{"role": "user", "content": [{"text": "Reply with one word: pong"}]}], "messages": [{"role": "user", "content": [{"text": "Reply with one word: pong"}]}],

View file

@ -621,13 +621,15 @@ class TestProxyBaseLLMRequestProcessing:
("sk-master", object(), {}, False), ("sk-master", object(), {}, False),
("sk-master", object(), {"custom_auth_run_common_checks": True}, True), ("sk-master", object(), {"custom_auth_run_common_checks": True}, True),
("sk-master", None, {}, 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( 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 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) """A request whose auth wrapper skips common_checks (a public route, no-auth dev mode, custom auth
never budget-checked tags before, so a hook-added tag must not start 429ing it.""" without opt-in) never budget-checked tags before, so a hook-added tag must not start 429ing it."""
processing_obj = ProxyBaseLLMRequestProcessing(data={}) processing_obj = ProxyBaseLLMRequestProcessing(data={})
async def mock_pre_call_hook(user_api_key_dict, data, call_type): 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.master_key", master_key)
monkeypatch.setattr("litellm.proxy.proxy_server.user_custom_auth", user_custom_auth) 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( await processing_obj.common_processing_pre_call_logic(
request=mock_request, request=mock_request,