mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
Merge pull request #2650 from BerriAI/litellm_jwt_auth_fixes
feat(handle_jwt.py): enable jwt-project based auth
This commit is contained in:
commit
e2d81722d2
6 changed files with 167 additions and 18 deletions
|
|
@ -602,7 +602,8 @@ general_settings:
|
|||
"completion_model": "string",
|
||||
"disable_spend_logs": "boolean", # turn off writing each transaction to the db
|
||||
"disable_reset_budget": "boolean", # turn off reset budget scheduled task
|
||||
"enable_jwt_auth": "boolean", # allow proxy admin to auth in via jwt tokens with 'litellm_proxy_admin' in claims
|
||||
"enable_jwt_auth": "boolean", # allow proxy admin to auth in via jwt tokens with 'litellm_proxy_admin' in claims
|
||||
"allowed_routes": "list", # list of allowed proxy API routes - a user can access. (currently JWT-Auth only)
|
||||
"key_management_system": "google_kms", # either google_kms or azure_kms
|
||||
"master_key": "string",
|
||||
"database_url": "string",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# [BETA] JWT-based Auth
|
||||
|
||||
Use JWT's to auth admin's into the proxy.
|
||||
Use JWT's to auth admins / projects into the proxy.
|
||||
|
||||
:::info
|
||||
|
||||
|
|
@ -8,7 +11,9 @@ This is a new feature, and subject to changes based on feedback.
|
|||
|
||||
:::
|
||||
|
||||
## Step 1. Set env's
|
||||
## Usage
|
||||
|
||||
### Step 1. Setup Proxy
|
||||
|
||||
- `JWT_PUBLIC_KEY_URL`: This is the public keys endpoint of your OpenID provider. Typically it's `{openid-provider-base-url}/.well-known/openid-configuration/jwks`. For Keycloak it's `{keycloak_base_url}/realms/{your-realm}/protocol/openid-connect/certs`.
|
||||
|
||||
|
|
@ -16,7 +21,26 @@ This is a new feature, and subject to changes based on feedback.
|
|||
export JWT_PUBLIC_KEY_URL="" # "https://demo.duendesoftware.com/.well-known/openid-configuration/jwks"
|
||||
```
|
||||
|
||||
## Step 2. Create JWT with scopes
|
||||
- `enable_jwt_auth` in your config. This will tell the proxy to check if a token is a jwt token.
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
enable_jwt_auth: True
|
||||
|
||||
model_list:
|
||||
- model_name: azure-gpt-3.5
|
||||
litellm_params:
|
||||
model: azure/<your-deployment-name>
|
||||
api_base: os.environ/AZURE_API_BASE
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
api_version: "2023-07-01-preview"
|
||||
```
|
||||
|
||||
### Step 2. Create JWT with scopes
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="admin" label="admin">
|
||||
|
||||
Create a client scope called `litellm_proxy_admin` in your OpenID provider (e.g. Keycloak).
|
||||
|
||||
|
|
@ -32,12 +56,55 @@ curl --location ' 'https://demo.duendesoftware.com/connect/token'' \
|
|||
--data-urlencode 'grant_type=password' \
|
||||
--data-urlencode 'scope=litellm_proxy_admin' # 👈 grant this scope
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="project" label="project">
|
||||
|
||||
## Step 3. Create a proxy key with JWT
|
||||
Create a JWT for your project on your OpenID provider (e.g. Keycloak).
|
||||
|
||||
```bash
|
||||
curl --location ' 'https://demo.duendesoftware.com/connect/token'' \
|
||||
--header 'Content-Type: application/x-www-form-urlencoded' \
|
||||
--data-urlencode 'client_id={CLIENT_ID}' \ # 👈 project id
|
||||
--data-urlencode 'client_secret={CLIENT_SECRET}' \
|
||||
--data-urlencode 'grant_type=client_credential' \
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Step 3. Test your JWT
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="key" label="/key/generate">
|
||||
|
||||
```bash
|
||||
curl --location '{proxy_base_url}/key/generate' \
|
||||
--header 'Authorization: Bearer eyJhbGciOiJSUzI1NiI...' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{}'
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="llm_call" label="/chat/completions">
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/v1/chat/completions' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--header 'Authorization: Bearer eyJhbGciOiJSUzI1...' \
|
||||
--data '{"model": "azure-gpt-3.5", "messages": [ { "role": "user", "content": "What's the weather like in Boston today?" } ]}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Advanced - Allowed Routes
|
||||
|
||||
Configure which routes a non-admin JWT can access via the config.
|
||||
|
||||
By default, a non-admin JWT can call openai + any `/info` endpoints.
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
enable_jwt_auth: True
|
||||
allowed_routes: ["/chat/completions", "/embeddings"]
|
||||
```
|
||||
|
|
@ -531,6 +531,9 @@ class ConfigGeneralSettings(LiteLLMBase):
|
|||
ui_access_mode: Optional[Literal["admin_only", "all"]] = Field(
|
||||
"all", description="Control access to the Proxy UI"
|
||||
)
|
||||
allowed_routes: Optional[List] = Field(
|
||||
None, description="Proxy API Endpoints you want users to be able to access"
|
||||
)
|
||||
|
||||
|
||||
class ConfigYAML(LiteLLMBase):
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ Run checks for:
|
|||
3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_UserTable, LiteLLM_EndUserTable
|
||||
from typing import Optional
|
||||
from typing import Optional, Literal
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.caching import DualCache
|
||||
|
||||
|
|
@ -19,6 +19,13 @@ def common_checks(
|
|||
user_object: LiteLLM_UserTable,
|
||||
end_user_object: Optional[LiteLLM_EndUserTable],
|
||||
) -> bool:
|
||||
"""
|
||||
Common checks across jwt + key-based auth.
|
||||
|
||||
1. If user can call model
|
||||
2. If user is in budget
|
||||
3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
"""
|
||||
_model = request_body.get("model", None)
|
||||
# 1. If user can call model
|
||||
if (
|
||||
|
|
@ -47,6 +54,52 @@ def common_checks(
|
|||
return True
|
||||
|
||||
|
||||
def allowed_routes_check(
|
||||
user_role: Literal["proxy_admin", "app_owner"],
|
||||
route: str,
|
||||
allowed_routes: Optional[list] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if user -> not admin - allowed to access these routes
|
||||
"""
|
||||
openai_routes = [
|
||||
# chat completions
|
||||
"/openai/deployments/{model}/chat/completions",
|
||||
"/chat/completions",
|
||||
"/v1/chat/completions",
|
||||
# completions
|
||||
# embeddings
|
||||
"/openai/deployments/{model}/embeddings",
|
||||
"/embeddings",
|
||||
"/v1/embeddings",
|
||||
# image generation
|
||||
"/images/generations",
|
||||
"/v1/images/generations",
|
||||
# audio transcription
|
||||
"/audio/transcriptions",
|
||||
"/v1/audio/transcriptions",
|
||||
# moderations
|
||||
"/moderations",
|
||||
"/v1/moderations",
|
||||
# models
|
||||
"/models",
|
||||
"/v1/models",
|
||||
]
|
||||
info_routes = ["/key/info", "/team/info", "/user/info", "/model/info"]
|
||||
default_routes = openai_routes + info_routes
|
||||
if user_role == "proxy_admin":
|
||||
return True
|
||||
elif user_role == "app_owner":
|
||||
if allowed_routes is None:
|
||||
if route in default_routes: # check default routes
|
||||
return True
|
||||
elif route in allowed_routes:
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
async def get_end_user_object(
|
||||
end_user_id: Optional[str],
|
||||
prisma_client: Optional[PrismaClient],
|
||||
|
|
|
|||
|
|
@ -8,10 +8,10 @@ JWT token must have 'litellm_proxy_admin' in scope.
|
|||
|
||||
import httpx
|
||||
import jwt
|
||||
from jwt.algorithms import RSAAlgorithm
|
||||
import json
|
||||
import os
|
||||
from litellm.caching import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import LiteLLMProxyRoles, LiteLLM_UserTable
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from typing import Optional
|
||||
|
|
@ -137,6 +137,8 @@ class JWTHandler:
|
|||
return scopes
|
||||
|
||||
async def auth_jwt(self, token: str) -> dict:
|
||||
from jwt.algorithms import RSAAlgorithm
|
||||
|
||||
keys_url = os.getenv("JWT_PUBLIC_KEY_URL")
|
||||
|
||||
if keys_url is None:
|
||||
|
|
@ -147,7 +149,13 @@ class JWTHandler:
|
|||
keys = response.json()["keys"]
|
||||
|
||||
header = jwt.get_unverified_header(token)
|
||||
kid = header["kid"]
|
||||
|
||||
verbose_proxy_logger.debug(f"header: {header}")
|
||||
|
||||
if "kid" in header:
|
||||
kid = header["kid"]
|
||||
else:
|
||||
raise Exception(f"Expected 'kid' in header. header={header}.")
|
||||
|
||||
for key in keys:
|
||||
if key["kid"] == kid:
|
||||
|
|
|
|||
|
|
@ -110,7 +110,11 @@ from litellm.proxy.auth.handle_jwt import JWTHandler
|
|||
from litellm.proxy.hooks.prompt_injection_detection import (
|
||||
_OPTIONAL_PromptInjectionDetection,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import common_checks, get_end_user_object
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
common_checks,
|
||||
get_end_user_object,
|
||||
allowed_routes_check,
|
||||
)
|
||||
|
||||
try:
|
||||
from litellm._version import version
|
||||
|
|
@ -332,7 +336,7 @@ def _get_pydantic_json_dict(pydantic_obj: BaseModel) -> dict:
|
|||
async def user_api_key_auth(
|
||||
request: Request, api_key: str = fastapi.Security(api_key_header)
|
||||
) -> UserAPIKeyAuth:
|
||||
global master_key, prisma_client, llm_model_list, user_custom_auth, custom_db_client
|
||||
global master_key, prisma_client, llm_model_list, user_custom_auth, custom_db_client, general_settings
|
||||
try:
|
||||
if isinstance(api_key, str):
|
||||
passed_in_key = api_key
|
||||
|
|
@ -354,6 +358,7 @@ async def user_api_key_auth(
|
|||
enable_jwt_auth: true
|
||||
```
|
||||
"""
|
||||
route: str = request.url.path
|
||||
if general_settings.get("enable_jwt_auth", False) == True:
|
||||
is_jwt = jwt_handler.is_jwt(token=api_key)
|
||||
verbose_proxy_logger.debug(f"is_jwt: {is_jwt}")
|
||||
|
|
@ -407,15 +412,28 @@ async def user_api_key_auth(
|
|||
user_id=user_id,
|
||||
)
|
||||
else:
|
||||
# return UserAPIKeyAuth object
|
||||
return UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_id=user_object.user_id,
|
||||
tpm_limit=user_object.tpm_limit,
|
||||
rpm_limit=user_object.rpm_limit,
|
||||
models=user_object.models,
|
||||
is_allowed = allowed_routes_check(
|
||||
user_role="app_owner",
|
||||
route=route,
|
||||
allowed_routes=general_settings.get("allowed_routes", None),
|
||||
)
|
||||
if is_allowed:
|
||||
# return UserAPIKeyAuth object
|
||||
return UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_id=user_object.user_id,
|
||||
tpm_limit=user_object.tpm_limit,
|
||||
rpm_limit=user_object.rpm_limit,
|
||||
models=user_object.models,
|
||||
user_role="app_owner",
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": f"User={user_object.user_id} not allowed to access this route={route}."
|
||||
},
|
||||
)
|
||||
#### ELSE ####
|
||||
if master_key is None:
|
||||
if isinstance(api_key, str):
|
||||
|
|
@ -423,7 +441,6 @@ async def user_api_key_auth(
|
|||
else:
|
||||
return UserAPIKeyAuth()
|
||||
|
||||
route: str = request.url.path
|
||||
if route == "/user/auth":
|
||||
if general_settings.get("allow_user_auth", False) == True:
|
||||
return UserAPIKeyAuth()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue