mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
build: merge commit 1b15568af7
Author: Krrish Dholakia <krrishdholakia@gmail.com>
Date: Mon Feb 17 21:37:36 2025 -0800
fix(proxy/_types.py): fix linting error
commit dc4d5cffa6
Author: Krrish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
parent
e8f387200a
commit
7bfd816d3b
14 changed files with 202 additions and 58 deletions
|
|
@ -207,6 +207,64 @@ OIDC Auth for API: [**See Walkthrough**](https://www.loom.com/share/00fe2deab59a
|
|||
- If all checks pass, allow the request
|
||||
|
||||
|
||||
## Advanced - Custom Validate
|
||||
|
||||
Validate a JWT Token using custom logic, if you need an extra way to verify if tokens are valid for LiteLLM Proxy.
|
||||
|
||||
### 1. Setup custom validate function
|
||||
|
||||
```python
|
||||
from typing import Literal
|
||||
|
||||
def my_custom_validate(token: str) -> Literal[True]:
|
||||
"""
|
||||
Only allow tokens with tenant-id == "my-unique-tenant", and claims == ["proxy-admin"]
|
||||
"""
|
||||
allowed_tenants = ["my-unique-tenant"]
|
||||
allowed_claims = ["proxy-admin"]
|
||||
|
||||
if token["tenant_id"] not in allowed_tenants:
|
||||
raise Exception("Invalid JWT token")
|
||||
if token["claims"] not in allowed_claims:
|
||||
raise Exception("Invalid JWT token")
|
||||
return True
|
||||
```
|
||||
|
||||
### 2. Setup config.yaml
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
enable_jwt_auth: True
|
||||
litellm_jwtauth:
|
||||
user_id_jwt_field: "sub"
|
||||
team_id_jwt_field: "tenant_id"
|
||||
user_id_upsert: True
|
||||
custom_validate: custom_validate.my_custom_validate # 👈 custom validate function
|
||||
```
|
||||
|
||||
### 3. Test the flow
|
||||
|
||||
**Expected JWT**
|
||||
|
||||
```
|
||||
{
|
||||
"sub": "my-unique-user",
|
||||
"tenant_id": "INVALID_TENANT",
|
||||
"claims": ["proxy-admin"]
|
||||
}
|
||||
```
|
||||
|
||||
**Expected Response**
|
||||
|
||||
```
|
||||
{
|
||||
"error": "Invalid JWT token"
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Advanced - Allowed Routes
|
||||
|
||||
Configure which routes a JWT can access via the config.
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -1,5 +1,15 @@
|
|||
model_list:
|
||||
- model_name: azure-gpt-35-turbo
|
||||
litellm_params:
|
||||
model: topaz/chatgpt-v-2
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
model: azure/chatgpt-v-2
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
api_base: os.environ/AZURE_API_BASE
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234
|
||||
enable_jwt_auth: True
|
||||
litellm_jwtauth:
|
||||
user_id_jwt_field: "sub"
|
||||
team_id_jwt_field: "client_id"
|
||||
user_id_upsert: True
|
||||
custom_validate: custom_validate.my_custom_validate
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import enum
|
|||
import json
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union
|
||||
|
||||
import httpx
|
||||
from pydantic import (
|
||||
|
|
@ -30,6 +30,8 @@ from litellm.types.utils import (
|
|||
TextCompletionResponse,
|
||||
)
|
||||
|
||||
from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
|
|
@ -2395,6 +2397,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
- public_key_ttl: Default - 600s. TTL for caching public JWT keys.
|
||||
- public_allowed_routes: list of allowed routes for authenticated but unknown litellm role jwt tokens.
|
||||
- enforce_rbac: If true, enforce RBAC for all routes.
|
||||
- custom_validate: A custom function to validates the JWT token.
|
||||
|
||||
See `auth_checks.py` for the specific routes
|
||||
"""
|
||||
|
|
@ -2439,6 +2442,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
scope_mappings: Optional[List[ScopeMapping]] = None
|
||||
enforce_scope_based_access: bool = False
|
||||
enforce_team_based_model_access: bool = False
|
||||
custom_validate: Optional[Callable[..., Literal[True]]] = None
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
# get the attribute names for this Pydantic model
|
||||
|
|
@ -2451,6 +2455,12 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
role_mappings = kwargs.get("role_mappings")
|
||||
scope_mappings = kwargs.get("scope_mappings")
|
||||
enforce_scope_based_access = kwargs.get("enforce_scope_based_access")
|
||||
custom_validate = kwargs.get("custom_validate")
|
||||
|
||||
if custom_validate is not None:
|
||||
fn = get_instance_fn(custom_validate)
|
||||
validate_custom_validate_return_type(fn)
|
||||
kwargs["custom_validate"] = fn
|
||||
|
||||
if invalid_keys:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -862,6 +862,14 @@ class JWTAuthManager:
|
|||
"""Main authentication and authorization builder"""
|
||||
jwt_valid_token: dict = await jwt_handler.auth_jwt(token=api_key)
|
||||
|
||||
# Check custom validate
|
||||
if jwt_handler.litellm_jwtauth.custom_validate:
|
||||
if not jwt_handler.litellm_jwtauth.custom_validate(jwt_valid_token):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Invalid JWT token",
|
||||
)
|
||||
|
||||
# Check RBAC
|
||||
rbac_role = jwt_handler.get_rbac_role(token=jwt_valid_token)
|
||||
await JWTAuthManager.check_rbac_role(
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import litellm
|
|||
from litellm import get_secret
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams
|
||||
from litellm.proxy.utils import get_instance_fn
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
|
||||
blue_color_code = "\033[94m"
|
||||
reset_color_code = "\033[0m"
|
||||
|
|
|
|||
5
litellm/proxy/custom_validate.py
Normal file
5
litellm/proxy/custom_validate.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
from typing import Literal
|
||||
|
||||
|
||||
def my_custom_validate(token: str) -> Literal[True]:
|
||||
raise Exception("Custom validate failed")
|
||||
|
|
@ -609,7 +609,7 @@ def create_pass_through_route(
|
|||
# check if target is an adapter.py or a url
|
||||
import uuid
|
||||
|
||||
from litellm.proxy.utils import get_instance_fn
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
|
||||
try:
|
||||
if isinstance(target, CustomLogger):
|
||||
|
|
|
|||
|
|
@ -238,6 +238,7 @@ from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
|||
router as spend_management_router,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
||||
router as ui_crud_endpoints_router,
|
||||
)
|
||||
|
|
@ -251,7 +252,6 @@ from litellm.proxy.utils import (
|
|||
_is_projected_spend_over_limit,
|
||||
_is_valid_team_configs,
|
||||
get_error_message_str,
|
||||
get_instance_fn,
|
||||
hash_token,
|
||||
update_spend,
|
||||
)
|
||||
|
|
|
|||
1
litellm/proxy/types_utils/README.md
Normal file
1
litellm/proxy/types_utils/README.md
Normal file
|
|
@ -0,0 +1 @@
|
|||
Utility functions for proxy types.py
|
||||
64
litellm/proxy/types_utils/utils.py
Normal file
64
litellm/proxy/types_utils/utils.py
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
import importlib
|
||||
import os
|
||||
from typing import Any, Callable, Literal, Optional, get_type_hints
|
||||
|
||||
|
||||
def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
|
||||
module_name = value
|
||||
instance_name = None
|
||||
try:
|
||||
# Split the path by dots to separate module from instance
|
||||
parts = value.split(".")
|
||||
|
||||
# The module path is all but the last part, and the instance_name is the last part
|
||||
module_name = ".".join(parts[:-1])
|
||||
instance_name = parts[-1]
|
||||
|
||||
# If config_file_path is provided, use it to determine the module spec and load the module
|
||||
if config_file_path is not None:
|
||||
directory = os.path.dirname(config_file_path)
|
||||
module_file_path = os.path.join(directory, *module_name.split("."))
|
||||
module_file_path += ".py"
|
||||
|
||||
spec = importlib.util.spec_from_file_location(module_name, module_file_path) # type: ignore
|
||||
if spec is None:
|
||||
raise ImportError(
|
||||
f"Could not find a module specification for {module_file_path}"
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec) # type: ignore
|
||||
spec.loader.exec_module(module) # type: ignore
|
||||
else:
|
||||
# Dynamically import the module
|
||||
module = importlib.import_module(module_name)
|
||||
|
||||
# Get the instance from the module
|
||||
instance = getattr(module, instance_name)
|
||||
|
||||
return instance
|
||||
except ImportError as e:
|
||||
# Re-raise the exception with a user-friendly message
|
||||
if instance_name and module_name:
|
||||
raise ImportError(
|
||||
f"Could not import {instance_name} from {module_name}"
|
||||
) from e
|
||||
else:
|
||||
raise e
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
||||
def validate_custom_validate_return_type(
|
||||
fn: Optional[Callable[..., Any]]
|
||||
) -> Optional[Callable[..., Literal[True]]]:
|
||||
if fn is None:
|
||||
return None
|
||||
|
||||
hints = get_type_hints(fn)
|
||||
return_type = hints.get("return")
|
||||
|
||||
if return_type != Literal[True]:
|
||||
raise TypeError(
|
||||
f"Custom validator must be annotated to return Literal[True], got {return_type}"
|
||||
)
|
||||
|
||||
return fn
|
||||
|
|
@ -1,7 +1,6 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import hashlib
|
||||
import importlib
|
||||
import json
|
||||
import os
|
||||
import smtplib
|
||||
|
|
@ -2222,51 +2221,6 @@ class PrismaClient:
|
|||
)
|
||||
|
||||
|
||||
### CUSTOM FILE ###
|
||||
def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
|
||||
module_name = value
|
||||
instance_name = None
|
||||
try:
|
||||
# Split the path by dots to separate module from instance
|
||||
parts = value.split(".")
|
||||
|
||||
# The module path is all but the last part, and the instance_name is the last part
|
||||
module_name = ".".join(parts[:-1])
|
||||
instance_name = parts[-1]
|
||||
|
||||
# If config_file_path is provided, use it to determine the module spec and load the module
|
||||
if config_file_path is not None:
|
||||
directory = os.path.dirname(config_file_path)
|
||||
module_file_path = os.path.join(directory, *module_name.split("."))
|
||||
module_file_path += ".py"
|
||||
|
||||
spec = importlib.util.spec_from_file_location(module_name, module_file_path) # type: ignore
|
||||
if spec is None:
|
||||
raise ImportError(
|
||||
f"Could not find a module specification for {module_file_path}"
|
||||
)
|
||||
module = importlib.util.module_from_spec(spec) # type: ignore
|
||||
spec.loader.exec_module(module) # type: ignore
|
||||
else:
|
||||
# Dynamically import the module
|
||||
module = importlib.import_module(module_name)
|
||||
|
||||
# Get the instance from the module
|
||||
instance = getattr(module, instance_name)
|
||||
|
||||
return instance
|
||||
except ImportError as e:
|
||||
# Re-raise the exception with a user-friendly message
|
||||
if instance_name and module_name:
|
||||
raise ImportError(
|
||||
f"Could not import {instance_name} from {module_name}"
|
||||
) from e
|
||||
else:
|
||||
raise e
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
||||
### HELPER FUNCTIONS ###
|
||||
async def _cache_user_row(user_id: str, cache: DualCache, db: PrismaClient):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -26,10 +26,16 @@ from fastapi.routing import APIRoute
|
|||
from fastapi.responses import Response
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_UserTable, LiteLLMRoutes
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_JWTAuth,
|
||||
LiteLLM_UserTable,
|
||||
LiteLLMRoutes,
|
||||
JWTAuthBuilderResult,
|
||||
)
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler, JWTAuthManager
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
from litellm.proxy.proxy_server import chat_completion
|
||||
from typing import Literal
|
||||
|
||||
public_key = {
|
||||
"kty": "RSA",
|
||||
|
|
@ -1247,3 +1253,32 @@ def test_check_scope_based_access(requested_model, should_work):
|
|||
else:
|
||||
with pytest.raises(HTTPException):
|
||||
JWTAuthManager.check_scope_based_access(**args)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_validate_called():
|
||||
# Setup
|
||||
mock_custom_validate = MagicMock(return_value=True)
|
||||
|
||||
jwt_handler = MagicMock()
|
||||
jwt_handler.litellm_jwtauth = MagicMock(
|
||||
custom_validate=mock_custom_validate, allowed_routes=["/chat/completions"]
|
||||
)
|
||||
jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "test_user"})
|
||||
|
||||
try:
|
||||
await JWTAuthManager.auth_builder(
|
||||
api_key="test",
|
||||
jwt_handler=jwt_handler,
|
||||
request_data={},
|
||||
general_settings={},
|
||||
route="/chat/completions",
|
||||
prisma_client=None,
|
||||
user_api_key_cache=MagicMock(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
# Assert custom_validate was called with the jwt token
|
||||
mock_custom_validate.assert_called_once_with({"sub": "test_user"})
|
||||
|
|
|
|||
|
|
@ -51,7 +51,7 @@ print("Testing proxy custom logger")
|
|||
def test_embedding(client):
|
||||
try:
|
||||
litellm.set_verbose = False
|
||||
from litellm.proxy.utils import get_instance_fn
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
|
||||
my_custom_logger = get_instance_fn(
|
||||
value="custom_callbacks.my_custom_logger", config_file_path=python_file_path
|
||||
|
|
@ -122,7 +122,7 @@ def test_chat_completion(client):
|
|||
try:
|
||||
# Your test data
|
||||
litellm.set_verbose = False
|
||||
from litellm.proxy.utils import get_instance_fn
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
|
||||
my_custom_logger = get_instance_fn(
|
||||
value="custom_callbacks.my_custom_logger", config_file_path=python_file_path
|
||||
|
|
@ -217,7 +217,7 @@ def test_chat_completion_stream(client):
|
|||
try:
|
||||
# Your test data
|
||||
litellm.set_verbose = False
|
||||
from litellm.proxy.utils import get_instance_fn
|
||||
from litellm.proxy.types_utils.utils import get_instance_fn
|
||||
|
||||
my_custom_logger = get_instance_fn(
|
||||
value="custom_callbacks.my_custom_logger", config_file_path=python_file_path
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue