mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-05 18:24:49 +00:00
Merge pull request #26197 from BerriAI/litellm_yj_apr20
[Infra] Merge dev branch
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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}",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user