Fix failing tests

This commit is contained in:
Sameer Kankute 2026-05-05 09:18:42 +05:30
parent 13980398fb
commit 32d81ff383
No known key found for this signature in database
3 changed files with 150 additions and 24 deletions

View file

@ -1868,24 +1868,71 @@ async def _noop_none() -> None:
return None
@tracer.wrap()
async def user_api_key_auth(
async def _reserve_budget_after_common_checks(
user_api_key_auth_obj: UserAPIKeyAuth,
request_data: dict,
route: str,
llm_router: Optional[Any],
team_object: Optional[LiteLLM_TeamTableCachedObj],
user_object: Optional[LiteLLM_UserTable],
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
proxy_logging_obj: ProxyLogging,
skip_budget_checks: bool,
end_user_id: Optional[str] = None,
end_user_object: Optional[LiteLLM_EndUserTable] = None,
) -> None:
user_api_key_auth_obj.budget_reservation = None
if skip_budget_checks:
return
from litellm.proxy.spend_tracking.budget_reservation import (
reserve_budget_for_request,
)
user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request(
request_body=request_data,
route=route,
llm_router=llm_router,
valid_token=user_api_key_auth_obj,
team_object=team_object,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
)
def _should_skip_budget_checks(
request_data: dict,
route: str,
request: Optional[Request],
llm_router: Optional[Any],
) -> bool:
model = _get_model_from_request_context(
request_data=request_data,
route=route,
request=request,
)
if model is not None and llm_router is not None:
return _is_model_cost_zero(model=model, llm_router=llm_router)
return False
async def run_user_api_key_auth_pipeline(
request: Request,
api_key: str = fastapi.Security(api_key_header),
azure_api_key_header: str = fastapi.Security(azure_api_key_header),
anthropic_api_key_header: Optional[str] = fastapi.Security(
anthropic_api_key_header
),
google_ai_studio_api_key_header: Optional[str] = fastapi.Security(
google_ai_studio_api_key_header
),
azure_apim_header: Optional[str] = fastapi.Security(azure_apim_header),
custom_litellm_key_header: Optional[str] = fastapi.Security(
custom_litellm_key_header
),
api_key: str,
azure_api_key_header: str,
anthropic_api_key_header: Optional[str],
google_ai_studio_api_key_header: Optional[str],
azure_apim_header: Optional[str],
custom_litellm_key_header: Optional[str],
) -> UserAPIKeyAuth:
"""
Parent function to authenticate user api key / jwt token.
Shared implementation for ``user_api_key_auth`` and for call sites that must
run the same auth pipeline without FastAPI ``Security()`` injection.
"""
request_data = await _read_request_body(request=request)
@ -1941,6 +1988,56 @@ async def user_api_key_auth(
return user_api_key_auth_obj
async def user_api_key_auth_from_request_headers(request: Request) -> UserAPIKeyAuth:
"""
Run the same auth as ``Depends(user_api_key_auth)`` using headers on ``request``.
Used when a route cannot use the FastAPI dependency (e.g. MCP OAuth broker
``/authorize`` / ``/token`` resolving optional ``Authorization``).
"""
h = request.headers
return await run_user_api_key_auth_pipeline(
request=request,
api_key=h.get("authorization") or "",
azure_api_key_header=h.get("api-key") or "",
anthropic_api_key_header=h.get("x-api-key"),
google_ai_studio_api_key_header=h.get("x-goog-api-key"),
azure_apim_header=h.get("ocp-apim-subscription-key"),
custom_litellm_key_header=h.get("x-litellm-api-key"),
)
@tracer.wrap()
async def user_api_key_auth(
request: Request,
api_key: str = fastapi.Security(api_key_header),
azure_api_key_header: str = fastapi.Security(azure_api_key_header),
anthropic_api_key_header: Optional[str] = fastapi.Security(
anthropic_api_key_header
),
google_ai_studio_api_key_header: Optional[str] = fastapi.Security(
google_ai_studio_api_key_header
),
azure_apim_header: Optional[str] = fastapi.Security(azure_apim_header),
custom_litellm_key_header: Optional[str] = fastapi.Security(
custom_litellm_key_header
),
) -> UserAPIKeyAuth:
"""
Parent function to authenticate user api key / jwt token.
"""
return await run_user_api_key_auth_pipeline(
request=request,
api_key=api_key,
azure_api_key_header=azure_api_key_header,
anthropic_api_key_header=anthropic_api_key_header,
google_ai_studio_api_key_header=google_ai_studio_api_key_header,
azure_apim_header=azure_apim_header,
custom_litellm_key_header=custom_litellm_key_header,
)
async def _return_user_api_key_auth_obj(
user_obj: Optional[LiteLLM_UserTable],
api_key: str,

View file

@ -1456,13 +1456,24 @@ if MCP_AVAILABLE:
When absent, unauthenticated access is still allowed for **temp-cache**
servers only (browser OAuth). When present, global-registry access
follows admin / allowlist rules via ``_get_cached_temporary_mcp_server_or_404``.
Only non-empty **string** ``Authorization`` values trigger a full auth
pipeline import (tests and mocks may attach MagicMock headers).
"""
authorization = (
request.headers.get("authorization")
or request.headers.get("Authorization")
or ""
).strip()
if not authorization:
try:
headers = request.headers
except Exception:
return None
raw: object = None
for key in ("authorization", "Authorization"):
try:
candidate = headers.get(key)
except Exception:
candidate = None
if isinstance(candidate, str) and candidate.strip():
raw = candidate
break
if not isinstance(raw, str) or not raw.strip():
return None
from litellm.proxy.auth.user_api_key_auth import (
user_api_key_auth_from_request_headers,

View file

@ -1561,6 +1561,10 @@ class TestTemporaryMCPSessionEndpoints:
server = generate_mock_mcp_server_config_record(server_id="server-1")
authorize_response = MagicMock()
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._try_resolve_mcp_oauth_broker_user",
AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
@ -1583,7 +1587,9 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is authorize_response
get_server.assert_awaited_once_with("server-1", request=request)
get_server.assert_awaited_once_with(
"server-1", user_api_key_dict=None, request=request
)
authorize_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -1607,6 +1613,10 @@ class TestTemporaryMCPSessionEndpoints:
exchange_response = {"access_token": "token"}
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._try_resolve_mcp_oauth_broker_user",
AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
@ -1630,7 +1640,9 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_awaited_once_with("server-1", request=request)
get_server.assert_awaited_once_with(
"server-1", user_api_key_dict=None, request=request
)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -1655,6 +1667,10 @@ class TestTemporaryMCPSessionEndpoints:
exchange_response = {"access_token": "new-token", "refresh_token": "new-rt"}
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._try_resolve_mcp_oauth_broker_user",
AsyncMock(return_value=None),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
@ -1678,7 +1694,9 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_awaited_once_with("server-1", request=request)
get_server.assert_awaited_once_with(
"server-1", user_api_key_dict=None, request=request
)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,