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 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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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"}]}],
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue