Merge pull request #26197 from BerriAI/litellm_yj_apr20

[Infra] Merge dev branch
This commit is contained in:
yuneng-jiang
2026-04-21 16:55:00 -07:00
committed by GitHub
7 changed files with 92 additions and 29 deletions
@@ -323,6 +323,14 @@ async def authorize_with_server(
)
parsed = urlparse(redirect_uri)
if parsed.scheme not in ("http", "https"):
raise HTTPException(
status_code=400,
detail={
"error": "invalid_redirect_uri",
"message": "redirect_uri must use http or https scheme",
},
)
base_url = urlunparse(parsed._replace(query=""))
request_base_url = get_request_base_url(request)
encoded_state = encode_state_with_base_url(
+21 -7
View File
@@ -626,11 +626,17 @@ async def common_checks( # noqa: PLR0915
and user_object.max_budget is not None
):
user_budget = user_object.max_budget
if user_budget < user_object.spend:
from litellm.proxy.proxy_server import get_current_spend
user_spend = await get_current_spend(
counter_key=f"spend:user:{user_object.user_id}",
fallback_spend=user_object.spend or 0.0,
)
if user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_object.spend,
current_cost=user_spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_object.spend}, Budget={user_budget}",
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
)
## 4.2 check team member budget, if team key
@@ -3665,12 +3671,20 @@ async def _organization_max_budget_check(
if org_max_budget is None or org_max_budget <= 0:
return
# Read spend from cross-pod counter (Redis-first) or cached object (fallback)
from litellm.proxy.proxy_server import get_current_spend
org_spend = await get_current_spend(
counter_key=f"spend:org:{org_id}",
fallback_spend=org_table.spend or 0.0,
)
# Check if organization spend exceeds max budget
if org_table.spend >= org_max_budget:
if org_spend >= org_max_budget:
# Trigger budget alert
call_info = CallInfo(
token=valid_token.token,
spend=org_table.spend,
spend=org_spend,
max_budget=org_max_budget,
user_id=valid_token.user_id,
team_id=valid_token.team_id,
@@ -3686,9 +3700,9 @@ async def _organization_max_budget_check(
)
raise litellm.BudgetExceededError(
current_cost=org_table.spend,
current_cost=org_spend,
max_budget=org_max_budget,
message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_table.spend}, Max budget: {org_max_budget}",
message=f"Budget has been exceeded! Organization={org_id} Current cost: {org_spend}, Max budget: {org_max_budget}",
)
+22 -12
View File
@@ -21,20 +21,30 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
):
try:
verbose_proxy_logger.debug("Inside Max Budget Limiter Pre-Call Hook")
cache_key = f"{user_api_key_dict.user_id}_user_api_key_user_id"
user_row = await cache.async_get_cache(
cache_key, parent_otel_span=user_api_key_dict.parent_otel_span
max_budget = user_api_key_dict.user_max_budget
user_id = user_api_key_dict.user_id
if max_budget is None or user_id is None:
return
# Personal budget applies only to non-team requests, matching
# the explicit team-key exemption in common_checks section 4.1.
if user_api_key_dict.team_id is not None:
return
from litellm.proxy.proxy_server import get_current_spend
curr_spend = await get_current_spend(
counter_key=f"spend:user:{user_id}",
fallback_spend=user_api_key_dict.user_spend or 0.0,
)
if user_row is None: # value not yet cached
return
max_budget = user_row["max_budget"]
curr_spend = user_row["spend"]
if max_budget is None:
return
if curr_spend is None:
return
verbose_proxy_logger.debug(
"MaxBudgetLimiter: user_id=%s, spend=%.6f, max=%.6f",
user_id,
curr_spend,
max_budget,
)
# CHECK IF REQUEST ALLOWED
if curr_spend >= max_budget:
@@ -213,6 +213,7 @@ class _ProxyDBLogger(CustomLogger):
team_id=team_id,
user_id=user_id,
response_cost=response_cost,
org_id=org_id,
)
# update cache (fire-and-forget for backward compat:
@@ -1336,7 +1336,9 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials(temp_record)
def _get_cached_temporary_mcp_server_or_404(server_id: str) -> MCPServer:
def _get_cached_temporary_mcp_server_or_404(
server_id: str, request: Optional[Request] = None
) -> MCPServer:
server = get_cached_temporary_mcp_server(server_id)
if server is None:
# Fall back to real DB/config server (e.g. for the user-side OAuth flow
@@ -1344,10 +1346,14 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
client_ip = IPAddressUtils.get_mcp_client_ip(request) if request else None
server = global_mcp_server_manager.get_mcp_server_by_id(
server_id
) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
) or global_mcp_server_manager.get_mcp_server_by_name(
server_id, client_ip=client_ip
)
if server is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
@@ -1358,10 +1364,12 @@ if MCP_AVAILABLE:
@router.get(
"/server/oauth/{server_id}/authorize",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def mcp_authorize(
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
client_id: Optional[str] = None,
redirect_uri: str = Query(...),
state: str = "",
@@ -1370,7 +1378,7 @@ if MCP_AVAILABLE:
response_type: Optional[str] = None,
scope: Optional[str] = None,
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
# Use the server's stored client_id when the caller doesn't supply one
resolved_client_id = mcp_server.client_id or client_id or ""
if not resolved_client_id:
@@ -1399,10 +1407,12 @@ if MCP_AVAILABLE:
@router.post(
"/server/oauth/{server_id}/token",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def mcp_token(
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
grant_type: str = Form(...),
code: Optional[str] = Form(None),
redirect_uri: Optional[str] = Form(None),
@@ -1412,7 +1422,7 @@ if MCP_AVAILABLE:
refresh_token: Optional[str] = Form(None),
scope: Optional[str] = Form(None),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
resolved_client_id = mcp_server.client_id or client_id or ""
if not resolved_client_id:
raise HTTPException(
@@ -1441,9 +1451,14 @@ if MCP_AVAILABLE:
@router.post(
"/server/oauth/{server_id}/register",
include_in_schema=False,
dependencies=[Depends(user_api_key_auth)],
)
async def mcp_register(request: Request, server_id: str):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
async def mcp_register(
request: Request,
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
request_data = await _read_request_body(request=request)
data: dict = {**request_data}
+15
View File
@@ -1795,6 +1795,7 @@ async def increment_spend_counters(
team_id: Optional[str],
user_id: Optional[str],
response_cost: Optional[float],
org_id: Optional[str] = None,
):
"""
Atomically increment spend counters for budget enforcement.
@@ -1881,6 +1882,20 @@ async def increment_spend_counters(
increment=response_cost,
)
if user_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:user:{user_id}",
source_cache_key=user_id,
increment=response_cost,
)
if org_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=f"org_id:{org_id}",
increment=response_cost,
)
async def _init_and_increment_spend_counter(
counter_key: str,
@@ -1486,7 +1486,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is authorize_response
get_server.assert_called_once_with("server-1")
get_server.assert_called_once_with("server-1", request=request)
authorize_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@@ -1533,7 +1533,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_called_once_with("server-1")
get_server.assert_called_once_with("server-1", request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@@ -1581,7 +1581,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_called_once_with("server-1")
get_server.assert_called_once_with("server-1", request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@@ -1628,7 +1628,7 @@ class TestTemporaryMCPSessionEndpoints:
result = await mcp_register(request=request, server_id="server-1")
assert result is register_response
get_server.assert_called_once_with("server-1")
get_server.assert_called_once_with("server-1", request=request)
read_body.assert_awaited_once_with(request=request)
register_mock.assert_awaited_once_with(
request=request,