From 7bfd816d3bb27098bc501706bada7f802549c9d2 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 17 Feb 2025 21:56:00 -0800 Subject: [PATCH] build: merge commit 1b15568af7f8923a014e6e49f5ec09e5abf5f981 Author: Krrish Dholakia Date: Mon Feb 17 21:37:36 2025 -0800 fix(proxy/_types.py): fix linting error commit dc4d5cffa69bbcde5a3de17aa7dc934238d81034 Author: Krrish Dholakia --- docs/my-website/docs/proxy/token_auth.md | 58 +++++++++++++++++ .../proxy/_experimental/out/onboarding.html | 1 - litellm/proxy/_new_secret_config.yaml | 14 +++- litellm/proxy/_types.py | 12 +++- litellm/proxy/auth/handle_jwt.py | 8 +++ litellm/proxy/common_utils/callback_utils.py | 2 +- litellm/proxy/custom_validate.py | 5 ++ .../pass_through_endpoints.py | 2 +- litellm/proxy/proxy_server.py | 2 +- litellm/proxy/types_utils/README.md | 1 + litellm/proxy/types_utils/utils.py | 64 +++++++++++++++++++ litellm/proxy/utils.py | 46 ------------- tests/proxy_unit_tests/test_jwt.py | 39 ++++++++++- .../test_proxy_custom_logger.py | 6 +- 14 files changed, 202 insertions(+), 58 deletions(-) delete mode 100644 litellm/proxy/_experimental/out/onboarding.html create mode 100644 litellm/proxy/custom_validate.py create mode 100644 litellm/proxy/types_utils/README.md create mode 100644 litellm/proxy/types_utils/utils.py diff --git a/docs/my-website/docs/proxy/token_auth.md b/docs/my-website/docs/proxy/token_auth.md index 672d62ec918..c6d280cb822 100644 --- a/docs/my-website/docs/proxy/token_auth.md +++ b/docs/my-website/docs/proxy/token_auth.md @@ -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. diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html deleted file mode 100644 index 1e070dcc422..00000000000 --- a/litellm/proxy/_experimental/out/onboarding.html +++ /dev/null @@ -1 +0,0 @@ -LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index b59a0a57b50..f6c8415021c 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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 \ No newline at end of file + 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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index dbf5270e59d..a1dc0c27c53 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 88a8144b554..29f4b31f9cd 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -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( diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 9859b340573..2280e72e9b8 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -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" diff --git a/litellm/proxy/custom_validate.py b/litellm/proxy/custom_validate.py new file mode 100644 index 00000000000..c80ffa4fe95 --- /dev/null +++ b/litellm/proxy/custom_validate.py @@ -0,0 +1,5 @@ +from typing import Literal + + +def my_custom_validate(token: str) -> Literal[True]: + raise Exception("Custom validate failed") diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 970af05f6db..c3257e47a1b 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 421a588e390..b5db78ef880 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, ) diff --git a/litellm/proxy/types_utils/README.md b/litellm/proxy/types_utils/README.md new file mode 100644 index 00000000000..1d14cf513ba --- /dev/null +++ b/litellm/proxy/types_utils/README.md @@ -0,0 +1 @@ +Utility functions for proxy types.py \ No newline at end of file diff --git a/litellm/proxy/types_utils/utils.py b/litellm/proxy/types_utils/utils.py new file mode 100644 index 00000000000..788849b3d5c --- /dev/null +++ b/litellm/proxy/types_utils/utils.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5e11f61522c..8042a781395 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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): """ diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index a168a91c120..7a9d2f0019b 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -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"}) diff --git a/tests/proxy_unit_tests/test_proxy_custom_logger.py b/tests/proxy_unit_tests/test_proxy_custom_logger.py index eb75c4abf79..ad60335152f 100644 --- a/tests/proxy_unit_tests/test_proxy_custom_logger.py +++ b/tests/proxy_unit_tests/test_proxy_custom_logger.py @@ -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