diff --git a/docs/my-website/docs/proxy/token_auth.md b/docs/my-website/docs/proxy/token_auth.md index 82d2266dd0d..4e6ff30a188 100644 --- a/docs/my-website/docs/proxy/token_auth.md +++ b/docs/my-website/docs/proxy/token_auth.md @@ -130,28 +130,57 @@ general_settings: Set the field in the jwt token, which corresponds to a litellm user / team / org. +**Note:** All JWT fields support dot notation to access nested claims (e.g., `"user.sub"`, `"resource_access.client.roles"`). + ```yaml general_settings: master_key: sk-1234 enable_jwt_auth: True litellm_jwtauth: admin_jwt_scope: "litellm-proxy-admin" - team_id_jwt_field: "client_id" # 👈 CAN BE ANY FIELD - user_id_jwt_field: "sub" # 👈 CAN BE ANY FIELD - org_id_jwt_field: "org_id" # 👈 CAN BE ANY FIELD - end_user_id_jwt_field: "customer_id" # 👈 CAN BE ANY FIELD + team_id_jwt_field: "client_id" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims) + user_id_jwt_field: "sub" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims) + org_id_jwt_field: "org_id" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims) + end_user_id_jwt_field: "customer_id" # 👈 CAN BE ANY FIELD (supports dot notation for nested claims) ``` -Expected JWT: +Expected JWT (flat structure): -``` +```json { "client_id": "my-unique-team", "sub": "my-unique-user", - "org_id": "my-unique-org", + "org_id": "my-unique-org" } ``` +**Or with nested structure using dot notation:** + +```json +{ + "user": { + "sub": "my-unique-user", + "email": "user@example.com" + }, + "tenant": { + "team_id": "my-unique-team" + }, + "organization": { + "id": "my-unique-org" + } +} +``` + +**Configuration for nested example:** + +```yaml +litellm_jwtauth: + user_id_jwt_field: "user.sub" + user_email_jwt_field: "user.email" + team_id_jwt_field: "tenant.team_id" + org_id_jwt_field: "organization.id" +``` + Now litellm will automatically update the spend for the user/team/org in the db for each call. ### JWT Scopes @@ -407,9 +436,15 @@ environment_variables: JWT_AUDIENCE: "api://LiteLLM_Proxy" # ensures audience is validated ``` -- `object_id_jwt_field`: The field in the JWT token that contains the object id. This id can be either a user id or a team id. Use this instead of `user_id_jwt_field` and `team_id_jwt_field`. If the same field could be both. +- `object_id_jwt_field`: The field in the JWT token that contains the object id. This id can be either a user id or a team id. Use this instead of `user_id_jwt_field` and `team_id_jwt_field`. If the same field could be both. **Supports dot notation** for nested claims (e.g., `"profile.object_id"`). -- `roles_jwt_field`: The field in the JWT token that contains the roles. This field is a list of roles that the user has. To index into a nested field, use dot notation - eg. `resource_access.litellm-test-client-id.roles`. +- `roles_jwt_field`: The field in the JWT token that contains the roles. This field is a list of roles that the user has. **Supports dot notation** for nested fields - e.g., `resource_access.litellm-test-client-id.roles`. + +**Additional JWT Field Configuration Options:** + +- `team_ids_jwt_field`: Field containing team IDs (as a list). **Supports dot notation** (e.g., `"groups"`, `"teams.ids"`). +- `user_email_jwt_field`: Field containing user email. **Supports dot notation** (e.g., `"email"`, `"user.email"`). +- `end_user_id_jwt_field`: Field containing end-user ID for cost tracking. **Supports dot notation** (e.g., `"customer_id"`, `"customer.id"`). - `role_mappings`: A list of role mappings. Map the received role in the JWT token to an internal role on LiteLLM. diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 07033801d51..12146f0e866 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -162,12 +162,13 @@ class JWTHandler: return False def get_team_ids_from_jwt(self, token: dict) -> List[str]: - if ( - self.litellm_jwtauth.team_ids_jwt_field is not None - and token.get(self.litellm_jwtauth.team_ids_jwt_field) is not None - ): - - return token[self.litellm_jwtauth.team_ids_jwt_field] + if self.litellm_jwtauth.team_ids_jwt_field is not None: + team_ids: Optional[List[str]] = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.team_ids_jwt_field, + default=[], + ) + return team_ids or [] return [] @@ -176,7 +177,11 @@ class JWTHandler: ) -> Optional[str]: try: if self.litellm_jwtauth.end_user_id_jwt_field is not None: - user_id = token[self.litellm_jwtauth.end_user_id_jwt_field] + user_id = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.end_user_id_jwt_field, + default=default_value, + ) else: user_id = None except KeyError: @@ -210,7 +215,21 @@ class JWTHandler: def get_team_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: try: if self.litellm_jwtauth.team_id_jwt_field is not None: - team_id = token[self.litellm_jwtauth.team_id_jwt_field] + # Use a sentinel value to detect if the path actually exists + sentinel = object() + team_id = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.team_id_jwt_field, + default=sentinel, + ) + if team_id is sentinel: + # Path doesn't exist, use team_id_default if available + if self.litellm_jwtauth.team_id_default is not None: + return self.litellm_jwtauth.team_id_default + else: + return default_value + # At this point, team_id is not the sentinel, so it should be a string + return team_id # type: ignore[return-value] elif self.litellm_jwtauth.team_id_default is not None: team_id = self.litellm_jwtauth.team_id_default else: @@ -232,7 +251,11 @@ class JWTHandler: def get_user_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: try: if self.litellm_jwtauth.user_id_jwt_field is not None: - user_id = token[self.litellm_jwtauth.user_id_jwt_field] + user_id = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.user_id_jwt_field, + default=default_value, + ) else: user_id = default_value except KeyError: @@ -319,7 +342,11 @@ class JWTHandler: ) -> Optional[str]: try: if self.litellm_jwtauth.user_email_jwt_field is not None: - user_email = token[self.litellm_jwtauth.user_email_jwt_field] + user_email = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.user_email_jwt_field, + default=default_value, + ) else: user_email = None except KeyError: @@ -329,7 +356,11 @@ class JWTHandler: def get_object_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: try: if self.litellm_jwtauth.object_id_jwt_field is not None: - object_id = token[self.litellm_jwtauth.object_id_jwt_field] + object_id = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.object_id_jwt_field, + default=default_value, + ) else: object_id = default_value except KeyError: @@ -339,7 +370,11 @@ class JWTHandler: def get_org_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: try: if self.litellm_jwtauth.org_id_jwt_field is not None: - org_id = token[self.litellm_jwtauth.org_id_jwt_field] + org_id = get_nested_value( + data=token, + key_path=self.litellm_jwtauth.org_id_jwt_field, + default=default_value, + ) else: org_id = None except KeyError: diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 8fb30e46b54..5627b9aaf90 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -447,3 +447,241 @@ async def test_map_jwt_role_to_litellm_role(): token = {"roles": ["team_"]} # No character after underscore result = jwt_handler.map_jwt_role_to_litellm_role(token) assert result is None + + +@pytest.mark.asyncio +async def test_nested_jwt_field_access(): + """ + Test that all JWT fields support dot notation for nested access + + This test verifies that: + 1. All JWT field methods can access nested values using dot notation + 2. Backward compatibility is maintained for flat field names + 3. Missing nested paths return appropriate defaults + """ + from litellm.proxy.auth.handle_jwt import JWTHandler + from litellm.proxy._types import LiteLLM_JWTAuth + + # Create JWT handler + jwt_handler = JWTHandler() + + # Test token with nested claims + nested_token = { + "user": { + "sub": "u123", + "email": "user@example.com" + }, + "resource_access": { + "my-client": { + "roles": ["admin", "user"] + } + }, + "groups": ["team1", "team2"], + "organization": { + "id": "org456" + }, + "profile": { + "object_id": "obj789" + }, + "customer": { + "end_user_id": "customer123" + }, + "tenant": { + "team_id": "team456" + } + } + + # Test flat token for backward compatibility + flat_token = { + "sub": "u123", + "email": "user@example.com", + "roles": ["admin", "user"], + "groups": ["team1", "team2"], + "org_id": "org456", + "object_id": "obj789", + "end_user_id": "customer123", + "team_id": "team456" + } + + # Test 1: user_id_jwt_field with nested access + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="user.sub") + assert jwt_handler.get_user_id(nested_token, None) == "u123" + + # Test 1b: user_id_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="sub") + assert jwt_handler.get_user_id(flat_token, None) == "u123" + + # Test 2: user_email_jwt_field with nested access + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="user.email") + assert jwt_handler.get_user_email(nested_token, None) == "user@example.com" + + # Test 2b: user_email_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="email") + assert jwt_handler.get_user_email(flat_token, None) == "user@example.com" + + # Test 3: team_ids_jwt_field with nested access + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") + assert jwt_handler.get_team_ids_from_jwt(nested_token) == ["team1", "team2"] + + # Test 3b: team_ids_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") + assert jwt_handler.get_team_ids_from_jwt(flat_token) == ["team1", "team2"] + + # Test 4: org_id_jwt_field with nested access + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_id_jwt_field="organization.id") + assert jwt_handler.get_org_id(nested_token, None) == "org456" + + # Test 4b: org_id_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_id_jwt_field="org_id") + assert jwt_handler.get_org_id(flat_token, None) == "org456" + + # Test 5: object_id_jwt_field with nested access (requires role_mappings) + from litellm.proxy._types import RoleMapping, LitellmUserRoles + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + object_id_jwt_field="profile.object_id", + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)] + ) + assert jwt_handler.get_object_id(nested_token, None) == "obj789" + + # Test 5b: object_id_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + object_id_jwt_field="object_id", + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)] + ) + assert jwt_handler.get_object_id(flat_token, None) == "obj789" + + # Test 6: end_user_id_jwt_field with nested access + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") + assert jwt_handler.get_end_user_id(nested_token, None) == "customer123" + + # Test 6b: end_user_id_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="end_user_id") + assert jwt_handler.get_end_user_id(flat_token, None) == "customer123" + + # Test 7: team_id_jwt_field with nested access + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="tenant.team_id") + assert jwt_handler.get_team_id(nested_token, None) == "team456" + + # Test 7b: team_id_jwt_field with flat access (backward compatibility) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id") + assert jwt_handler.get_team_id(flat_token, None) == "team456" + + # Test 8: roles_jwt_field with deeply nested access (already supported, but testing) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") + assert jwt_handler.get_jwt_role(nested_token, []) == ["admin", "user"] + + # Test 9: user_roles_jwt_field with nested access (already supported, but testing) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_roles_jwt_field="resource_access.my-client.roles", + user_allowed_roles=["admin", "user"] + ) + assert jwt_handler.get_user_roles(nested_token, []) == ["admin", "user"] + + +@pytest.mark.asyncio +async def test_nested_jwt_field_missing_paths(): + """ + Test handling of missing nested paths in JWT tokens + + This test verifies that: + 1. Missing nested paths return appropriate defaults + 2. Partial paths that exist but don't have the final key return defaults + 3. team_id_default fallback works with nested fields + """ + from litellm.proxy.auth.handle_jwt import JWTHandler + from litellm.proxy._types import LiteLLM_JWTAuth + + # Create JWT handler + jwt_handler = JWTHandler() + + # Test token with missing nested paths + incomplete_token = { + "user": { + "name": "test user" + # missing "sub" and "email" + }, + "resource_access": { + "other-client": { + "roles": ["viewer"] + } + # missing "my-client" + } + # missing "organization", "profile", "customer", "tenant", "groups" + } + + # Test 1: Missing user.sub should return default + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="user.sub") + assert jwt_handler.get_user_id(incomplete_token, "default_user") == "default_user" + + # Test 2: Missing user.email should return default + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="user.email") + assert jwt_handler.get_user_email(incomplete_token, "default@example.com") == "default@example.com" + + # Test 3: Missing groups should return empty list + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") + assert jwt_handler.get_team_ids_from_jwt(incomplete_token) == [] + + # Test 4: Missing organization.id should return default + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_id_jwt_field="organization.id") + assert jwt_handler.get_org_id(incomplete_token, "default_org") == "default_org" + + # Test 5: Missing profile.object_id should return default (requires role_mappings) + from litellm.proxy._types import RoleMapping, LitellmUserRoles + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + object_id_jwt_field="profile.object_id", + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)] + ) + assert jwt_handler.get_object_id(incomplete_token, "default_obj") == "default_obj" + + # Test 6: Missing customer.end_user_id should return default + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") + assert jwt_handler.get_end_user_id(incomplete_token, "default_customer") == "default_customer" + + # Test 7: Missing tenant.team_id should use team_id_default fallback + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_id_jwt_field="tenant.team_id", + team_id_default="fallback_team" + ) + assert jwt_handler.get_team_id(incomplete_token, "default_team") == "fallback_team" + + # Test 8: Missing resource_access.my-client.roles should return default + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") + assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == ["default_role"] + + # Test 9: Missing nested user roles should return default + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_roles_jwt_field="resource_access.my-client.roles", + user_allowed_roles=["admin", "user"] + ) + assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == ["default_user_role"] + + +@pytest.mark.asyncio +async def test_metadata_prefix_handling_in_nested_fields(): + """ + Test that metadata. prefix is properly handled in nested JWT field access + + The get_nested_value function should remove metadata. prefix before traversing + """ + from litellm.proxy.auth.handle_jwt import JWTHandler + from litellm.proxy._types import LiteLLM_JWTAuth + + # Create JWT handler + jwt_handler = JWTHandler() + + # Test token with proper structure for metadata prefix removal + token = { + "user": { + "email": "user@example.com" # This will be accessed when metadata.user.email is used + }, + "sub": "u123" + } + + # Test 1: metadata.user.email should access user.email after prefix removal + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="metadata.user.email") + # The get_nested_value function removes "metadata." prefix, so "metadata.user.email" becomes "user.email" + assert jwt_handler.get_user_email(token, None) == "user@example.com" + + # Test 2: user.sub should work normally without metadata prefix + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="sub") + assert jwt_handler.get_user_id(token, None) == "u123"