mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(mcp): keep team member budget enforcement on the MCP OAuth auth dependency
The MCP OAuth dependency stops at _user_api_key_auth_builder and never reaches common_checks, so removing the builder's inline member budget check would have let over-budget members through there. Enforce it explicitly for that caller.
This commit is contained in:
parent
b9335b6839
commit
ea2b666a68
4 changed files with 152 additions and 7 deletions
|
|
@ -46,6 +46,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
_cache_key_object,
|
||||
_can_object_call_model,
|
||||
_check_end_user_budget,
|
||||
_check_team_member_budget,
|
||||
_delete_cache_key_object,
|
||||
_get_user_role,
|
||||
_is_model_cost_zero,
|
||||
|
|
@ -3097,6 +3098,66 @@ def _resolve_request_principal(request: Request, valid_token: UserAPIKeyAuth) ->
|
|||
)
|
||||
|
||||
|
||||
async def enforce_team_member_budget_without_common_checks(
|
||||
user_api_key_auth_obj: UserAPIKeyAuth,
|
||||
request: Request,
|
||||
request_data: dict,
|
||||
route: str,
|
||||
api_key: str,
|
||||
) -> None:
|
||||
"""Team member budget gate for callers that stop at ``_user_api_key_auth_builder`` and never reach
|
||||
``common_checks`` (the MCP OAuth dependency), so an over-budget member stays blocked there as before.
|
||||
Failures go through the builder's exception handler so the caller still gets the 422.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
team_id: Final = user_api_key_auth_obj.team_id
|
||||
user_id: Final = user_api_key_auth_obj.user_id
|
||||
if prisma_client is None or team_id is None or team_id == UI_TEAM_ID or user_id is None:
|
||||
return
|
||||
parent_otel_span: Final = user_api_key_auth_obj.parent_otel_span
|
||||
try:
|
||||
team_object: Final = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except HTTPException: # no team row means no member budget to enforce
|
||||
return
|
||||
try:
|
||||
user_object = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception: # noqa: BLE001 # the user row only supplies the alert email; enforcement does not need it
|
||||
user_object = None
|
||||
try:
|
||||
await _check_team_member_budget(
|
||||
team_object=team_object,
|
||||
user_object=user_object,
|
||||
valid_token=user_api_key_auth_obj,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except litellm.BudgetExceededError as e:
|
||||
await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
|
||||
e=e,
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
parent_otel_span=parent_otel_span,
|
||||
api_key=api_key,
|
||||
resolved_identity=user_api_key_auth_obj,
|
||||
)
|
||||
|
||||
|
||||
async def _authorize_authenticated_request(
|
||||
user_api_key_auth_obj: UserAPIKeyAuth,
|
||||
request: Request,
|
||||
|
|
|
|||
|
|
@ -205,8 +205,10 @@ if MCP_AVAILABLE:
|
|||
UserMCPManagementMode,
|
||||
is_per_server_oauth_discovery_eligible,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import get_request_route
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_user_api_key_auth_builder,
|
||||
enforce_team_member_budget_without_common_checks,
|
||||
user_api_key_auth,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
|
|
@ -2010,9 +2012,6 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
|
||||
get_request_route,
|
||||
)
|
||||
|
||||
server_id: Final[str] = request.path_params.get("server_id", "")
|
||||
if server_id:
|
||||
|
|
@ -2049,7 +2048,7 @@ if MCP_AVAILABLE:
|
|||
request_data = await _read_request_body(request=request)
|
||||
request_data = populate_request_with_path_params(request_data=request_data, request=request)
|
||||
|
||||
return await _user_api_key_auth_builder(
|
||||
user_api_key_dict: Final = await _user_api_key_auth_builder(
|
||||
request=request,
|
||||
api_key=api_key,
|
||||
azure_api_key_header="",
|
||||
|
|
@ -2058,6 +2057,15 @@ if MCP_AVAILABLE:
|
|||
azure_apim_header=None,
|
||||
request_data=request_data,
|
||||
)
|
||||
# This dependency never reaches common_checks, which is the only other place the member budget is enforced.
|
||||
await enforce_team_member_budget_without_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_dict,
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
route=get_request_route(request),
|
||||
api_key=api_key,
|
||||
)
|
||||
return user_api_key_dict
|
||||
|
||||
async def _get_cached_temporary_mcp_server_or_404(
|
||||
server_id: str,
|
||||
|
|
|
|||
|
|
@ -8194,6 +8194,76 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset
|
|||
assert "Max budget: 2.0" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("team_member_spend, expect_blocked", [(2.4, True), (2.39, False)])
|
||||
async def test_team_member_budget_enforced_for_builder_only_callers(team_member_spend, expect_blocked):
|
||||
"""Callers that stop at the builder and skip common_checks (the MCP OAuth dependency) still reject a member
|
||||
at their team member budget with the same 422 the full auth flow returns."""
|
||||
from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj
|
||||
from litellm.proxy.auth.user_api_key_auth import enforce_team_member_budget_without_common_checks
|
||||
from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key
|
||||
|
||||
team_id = "team-builder-only"
|
||||
user_id = "user-builder-only"
|
||||
user_api_key_cache = DualCache()
|
||||
await user_api_key_cache.async_set_cache(key=f"team_id:{team_id}", value=LiteLLM_TeamTableCachedObj(team_id=team_id))
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=user_id, value=LiteLLM_UserTable(user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
)
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=team_membership_reservation_cache_key(team_id=team_id, user_id=user_id),
|
||||
value=LiteLLM_TeamMembership(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
spend=team_member_spend,
|
||||
budget_id="budget-builder-only",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(max_budget=2.4),
|
||||
),
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/mcp/server/oauth/srv/authorize"
|
||||
mock_request.method = "GET"
|
||||
mock_request.headers = {}
|
||||
mock_request.query_params = {}
|
||||
mock_request.state = SimpleNamespace()
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.budget_alerts = AsyncMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
async def _enforce():
|
||||
await enforce_team_member_budget_without_common_checks(
|
||||
user_api_key_auth_obj=UserAPIKeyAuth(api_key="hashed", team_id=team_id, user_id=user_id),
|
||||
request=mock_request,
|
||||
request_data={},
|
||||
route="/v1/mcp/server/oauth/srv/authorize",
|
||||
api_key="Bearer sk-builder-only",
|
||||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state
|
||||
patch( # test-quality-ok: seed the team, user and membership without a DB
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache
|
||||
),
|
||||
patch( # test-quality-ok: module-global proxy state
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
|
||||
),
|
||||
patch( # test-quality-ok: the live counter needs Redis or a DB; pin the spend the check compares
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(return_value=team_member_spend),
|
||||
),
|
||||
):
|
||||
if not expect_blocked:
|
||||
await _enforce()
|
||||
return
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _enforce()
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
|
||||
assert f"TeamMember={user_id}:{team_id}" in exc_info.value.message
|
||||
|
||||
|
||||
async def _proxy_exception_for_key(
|
||||
api_key: str,
|
||||
general_settings: dict[str, bool],
|
||||
|
|
|
|||
|
|
@ -2552,7 +2552,9 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
expected_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key=api_key_in_cookie
|
||||
)
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=master_key)
|
||||
fake_proxy_server = types.SimpleNamespace(
|
||||
master_key=master_key, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None
|
||||
)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2629,7 +2631,9 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = non_oauth_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None)
|
||||
fake_proxy_server = types.SimpleNamespace(
|
||||
master_key=None, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None
|
||||
)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
@ -2681,7 +2685,9 @@ class TestTemporaryMCPSessionEndpoints:
|
|||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = internal_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
fake_proxy_server = types.SimpleNamespace(master_key=None)
|
||||
fake_proxy_server = types.SimpleNamespace(
|
||||
master_key=None, prisma_client=None, proxy_logging_obj=None, user_api_key_cache=None
|
||||
)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue