fix(auth): reject deactivated JWT users and invalidate cached status

This commit is contained in:
Joshua Valluru 2026-09-19 18:10:49 -07:00
parent d1773d96e9
commit a82f0a0bd2
4 changed files with 115 additions and 9 deletions

View file

@ -1730,6 +1730,16 @@ async def _user_api_key_auth_builder(
)
return JWTAuthManager.user_api_key_auth_from_result(result, parent_otel_span)
if (
user_object is not None
and isinstance(user_object.metadata, dict)
and user_object.metadata.get("scim_active") is False
):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=f"User={user_id} has been deactivated via SCIM. Keys owned by this user cannot be used.",
)
valid_token = JWTAuthManager.user_api_key_auth_from_result(result, parent_otel_span)
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.

View file

@ -1571,7 +1571,7 @@ async def _update_single_user_helper(
await _invalidate_user_spend_counter_if_changed(non_default_values)
if "model_max_budget" in non_default_values:
if "model_max_budget" in non_default_values or "metadata" in data_json:
await evict_and_broadcast(
cache_keys=(non_default_values["user_id"],),
user_api_key_cache=user_api_key_cache,

View file

@ -2091,7 +2091,8 @@ async def test_auto_register_binds_api_key_to_token_hash():
@pytest.mark.asyncio
async def test_auto_register_first_request_propagates_user_email():
@pytest.mark.parametrize("active", [True, False])
async def test_auto_register_first_request_propagates_user_email(active: bool) -> None:
"""
The first auto-registered JWT request must also carry user_email (resolved
from the validated LiteLLM_UserTable), so attribution is consistent with the
@ -2120,6 +2121,7 @@ async def test_auto_register_first_request_propagates_user_email():
user_id="validated-user",
user_email="validated@example.com",
user_role="internal_user",
metadata={"scim_active": active},
)
mock_jwt_result = {
"is_proxy_admin": False,
@ -2150,7 +2152,7 @@ async def test_auto_register_first_request_propagates_user_email():
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None))),
patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler),
patch(
"litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key",
@ -2170,8 +2172,22 @@ async def test_auto_register_first_request_propagates_user_email():
"litellm.proxy.auth.user_api_key_auth._auto_register_jwt_mapping",
new_callable=AsyncMock,
return_value=auto_registered_key,
),
) as auto_register,
):
if not active:
with pytest.raises(ProxyException, match="deactivated via SCIM") as exc:
await _user_api_key_auth_builder(
request=mock_request,
api_key=jwt_token,
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert int(exc.value.code) == 401
auto_register.assert_not_awaited()
return
result = await _user_api_key_auth_builder(
request=mock_request,
api_key=jwt_token,
@ -7315,15 +7331,15 @@ class TestJWTAuthUserEmail:
the Prometheus `user_email` label and `user_api_key_user_email` in
StandardLogging/SpendLogs metadata, which were always None for JWT traffic."""
def _jwt_request(self, jwt_token):
def _jwt_request(self, jwt_token, route="/v1/chat/completions"):
mock_request = MagicMock()
mock_request.url.path = "/v1/chat/completions"
mock_request.method = "POST"
mock_request.url.path = route
mock_request.method = "GET" if route.endswith("/list") else "POST"
mock_request.headers = {"authorization": f"Bearer {jwt_token}"}
mock_request.query_params = {}
return mock_request
async def _run_jwt_auth(self, mock_jwt_result, jwt_token):
async def _run_jwt_auth(self, mock_jwt_result, jwt_token, route="/v1/chat/completions"):
with (
patch(
"litellm.proxy.proxy_server.general_settings",
@ -7344,7 +7360,7 @@ class TestJWTAuthUserEmail:
litellm_jwtauth=LiteLLM_JWTAuth(),
)
return await user_api_key_auth(
request=self._jwt_request(jwt_token),
request=self._jwt_request(jwt_token, route),
api_key=f"Bearer {jwt_token}",
)
@ -7376,6 +7392,41 @@ class TestJWTAuthUserEmail:
assert result.user_id == "jwt-human-user"
assert result.user_email == "resolved@example.com"
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["/mcp-rest/tools/list", "/mcp-rest/tools/call", "/v1/chat/completions"])
@pytest.mark.parametrize("active", [False, True, None, "false", 0])
async def test_jwt_auth_rejects_deactivated_user(self, route: str, active: bool | str | int | None) -> None:
from typing import Final
jwt_token: Final = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"
result: Final = {
"is_proxy_admin": False,
"team_object": None,
"user_object": LiteLLM_UserTable(
user_id="jwt-human-user",
user_role=LitellmUserRoles.INTERNAL_USER.value,
metadata={} if active is None else {"scim_active": active},
),
"end_user_object": None,
"org_object": None,
"token": jwt_token,
"team_id": None,
"user_id": "jwt-human-user",
"user_email": None,
"end_user_id": None,
"org_id": None,
"team_membership": None,
"jwt_claims": {"sub": "user1"},
}
if active is False:
with pytest.raises(ProxyException, match="deactivated via SCIM") as exc:
await self._run_jwt_auth(result, jwt_token, route)
assert int(exc.value.code) == 401
else:
token: Final = await self._run_jwt_auth(result, jwt_token, route)
assert token.user_id == "jwt-human-user"
@pytest.mark.asyncio
async def test_jwt_auth_populates_user_email_on_proxy_admin(self):
jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature"

View file

@ -2228,6 +2228,51 @@ async def test_user_model_budget_update_by_email_refreshes_cached_user(mocker: M
broadcast.assert_awaited_once_with(cache_key=saved_user.user_id)
@pytest.mark.asyncio
@pytest.mark.parametrize("by_email", [False, True])
@pytest.mark.parametrize("active", [False, True, None])
async def test_user_status_update_refreshes_cached_user(
mocker: MockerFixture, by_email: bool, active: bool | None
) -> None:
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.internal_user_endpoints import _update_single_user_helper
saved_user: Final = LiteLLM_UserTable(
user_id="user-spruce",
user_email="spruce@example.test",
metadata={"scim_active": False if active is None else not active, "department": "engineering"},
)
prisma_client: Final = mocker.MagicMock()
prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user)
prisma_client.get_data = mocker.AsyncMock(return_value=[saved_user])
prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": saved_user})
mocker.patch("litellm.proxy.proxy_server.prisma_client", prisma_client) # test-quality-ok: substitute the database dependency
cache: Final = UserApiKeyCache()
await cache.async_set_cache(key=saved_user.user_id, value=saved_user, model_type=LiteLLM_UserTable)
mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) # test-quality-ok: exercise a real isolated cache
broadcast: Final = mocker.patch( # test-quality-ok: observe the Redis publication boundary
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation",
new_callable=mocker.AsyncMock,
)
await _update_single_user_helper(
user_request=UpdateUserRequest(
user_id=None if by_email else saved_user.user_id,
user_email=saved_user.user_email if by_email else None,
metadata={"department": "engineering"} if active is None else {"scim_active": active},
),
user_api_key_dict=UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert prisma_client.update_data.call_args.kwargs["user_id"] == saved_user.user_id
assert prisma_client.update_data.call_args.kwargs["data"]["metadata"] == (
{"department": "engineering"} if active is None else {"scim_active": active}
)
assert await cache.async_get_cache(key=saved_user.user_id, model_type=LiteLLM_UserTable) is None
broadcast.assert_awaited_once_with(cache_key=saved_user.user_id)
@pytest.mark.asyncio
async def test_bulk_user_model_budget_clear_serializes_and_refreshes_cache(mocker: MockerFixture) -> None:
from litellm.proxy._types import LiteLLM_UserTable