diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0c118898f0f..a0a9209475f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -147,6 +147,7 @@ class LitellmTableNames(str, enum.Enum): KEY_TABLE_NAME = "LiteLLM_VerificationToken" PROXY_MODEL_TABLE_NAME = "LiteLLM_ProxyModelTable" MANAGED_FILE_TABLE_NAME = "LiteLLM_ManagedFileTable" + CONFIG_TABLE_NAME = "LiteLLM_Config" class Litellm_EntityType(enum.Enum): diff --git a/litellm/proxy/management_endpoints/router_management_endpoints.py b/litellm/proxy/management_endpoints/router_management_endpoints.py new file mode 100644 index 00000000000..2209f4cbd1f --- /dev/null +++ b/litellm/proxy/management_endpoints/router_management_endpoints.py @@ -0,0 +1,178 @@ +import asyncio +import copy + +from fastapi import APIRouter, Depends, HTTPException, status + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_helpers.audit_logs import create_object_audit_log +from litellm.proxy.proxy_server import ( + LITELLM_PROXY_ADMIN_NAME, + LitellmTableNames, + LitellmUserRoles, + ProxyErrorTypes, + ProxyException, + prisma_client, + store_model_in_db, +) +from litellm.proxy.types import UserAPIKeyAuth +from litellm.types.proxy.management_endpoints.router_management_endpoints import ( + GetRouterSettingsResponse, + PatchRouterSettingsRequest, +) + +router = APIRouter() + + +@router.patch( + "/router_settings", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], +) +async def update_router_settings( + router_settings: PatchRouterSettingsRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update router settings in the database. + Only accessible by proxy admin users. + """ + try: + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={ + "error": "No DB Connected. Here's how to do it - https://docs.litellm.ai/docs/proxy/virtual_keys" + }, + ) + + if store_model_in_db is not True: + raise HTTPException( + status_code=500, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, + ) + + # Only allow proxy admin to update router settings + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail={"error": "Only proxy admin users can update router settings"}, + ) + + existing_router_settings = ( + await prisma_client.db.litellm_config.find_unique( + where={"param_name": "router_settings"} + ) + or {} + ) + + # new router settings + new_router_settings = copy.deepcopy(existing_router_settings) + + # update new router settings with request body + new_router_settings.update(router_settings.model_dump(exclude_none=True)) + + # Update router settings in DB + await prisma_client.db.litellm_config.upsert( + where={"param_name": "router_settings"}, + data={ + "create": { + "param_name": "router_settings", + "param_value": new_router_settings, + "updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, + }, + "update": { + "param_value": new_router_settings, + "updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, + }, + }, + ) + + ## CREATE AUDIT LOG ## + asyncio.create_task( + create_object_audit_log( + object_id="router_settings", + action="updated", + user_api_key_dict=user_api_key_dict, + table_name=LitellmTableNames.CONFIG_TABLE_NAME, + before_value=safe_dumps(existing_router_settings), + after_value=safe_dumps(new_router_settings), + litellm_changed_by=user_api_key_dict.user_id, + litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME, + ) + ) + + return {"message": "Router settings updated successfully"} + + except Exception as e: + verbose_proxy_logger.exception(f"Error updating router settings: {str(e)}") + if isinstance(e, HTTPException): + raise e + raise ProxyException( + message=f"Error updating router settings: {str(e)}", + type=ProxyErrorTypes.internal_server_error, + code=status.HTTP_500_INTERNAL_SERVER_ERROR, + param=None, + ) + + +@router.get( + "/router_settings", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], + response_model=GetRouterSettingsResponse, +) +async def get_router_settings( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get current router settings from the database. + Only accessible by proxy admin users. + """ + try: + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={ + "error": "No DB Connected. Here's how to do it - https://docs.litellm.ai/docs/proxy/virtual_keys" + }, + ) + + if store_model_in_db is not True: + raise HTTPException( + status_code=500, + detail={ + "error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature." + }, + ) + + # Only allow proxy admin to view router settings + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail={"error": "Only proxy admin users can view router settings"}, + ) + + # Get router settings from DB + router_settings = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "router_settings"} + ) + + if router_settings is None: + return {"router_settings": {}} + + return GetRouterSettingsResponse(**router_settings.param_value) + + except Exception as e: + verbose_proxy_logger.exception(f"Error getting router settings: {str(e)}") + if isinstance(e, HTTPException): + raise e + raise ProxyException( + message=f"Error getting router settings: {str(e)}", + type=ProxyErrorTypes.internal_server_error, + code=status.HTTP_500_INTERNAL_SERVER_ERROR, + param=None, + ) diff --git a/litellm/types/proxy/management_endpoints/router_management_endpoints.py b/litellm/types/proxy/management_endpoints/router_management_endpoints.py new file mode 100644 index 00000000000..0d3446fb526 --- /dev/null +++ b/litellm/types/proxy/management_endpoints/router_management_endpoints.py @@ -0,0 +1,23 @@ +from typing import Dict, Optional + +from pydantic import BaseModel + + +class GetRouterSettingsResponse(BaseModel): + """ + Response body for getting router settings + + Add any router params you want to allow UI users to get + """ + + model_group_alias: Optional[Dict[str, str]] = {} + + +class PatchRouterSettingsRequest(BaseModel): + """ + Request body for patching router settings + + Add any router params you want to allow UI users to patch + """ + + model_group_alias: Optional[Dict[str, str]] = {} diff --git a/litellm/types/router.py b/litellm/types/router.py index f5634671350..ff3efa64aab 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -88,6 +88,7 @@ class UpdateRouterConfig(BaseModel): retry_after: Optional[float] = None fallbacks: Optional[List[dict]] = None context_window_fallbacks: Optional[List[dict]] = None + model_group_alias: Optional[Dict[str, str]] = {} model_config = ConfigDict(protected_namespaces=()) @@ -96,18 +97,16 @@ class ModelInfo(BaseModel): id: Optional[ str ] # Allow id to be optional on input, but it will always be present as a str in the model instance - db_model: bool = ( - False # used for proxy - to separate models which are stored in the db vs. config. - ) + db_model: bool = False # used for proxy - to separate models which are stored in the db vs. config. updated_at: Optional[datetime.datetime] = None updated_by: Optional[str] = None created_at: Optional[datetime.datetime] = None created_by: Optional[str] = None - base_model: Optional[str] = ( - None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking - ) + base_model: Optional[ + str + ] = None # specify if the base model is azure/gpt-3.5-turbo etc for accurate cost tracking tier: Optional[Literal["free", "paid"]] = None """ @@ -182,12 +181,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): custom_llm_provider: Optional[str] = None tpm: Optional[int] = None rpm: Optional[int] = None - timeout: Optional[Union[float, str, httpx.Timeout]] = ( - None # if str, pass in as os.environ/ - ) - stream_timeout: Optional[Union[float, str]] = ( - None # timeout when making stream=True calls, if str, pass in as os.environ/ - ) + timeout: Optional[ + Union[float, str, httpx.Timeout] + ] = None # if str, pass in as os.environ/ + stream_timeout: Optional[ + Union[float, str] + ] = None # timeout when making stream=True calls, if str, pass in as os.environ/ max_retries: Optional[int] = None organization: Optional[str] = None # for openai orgs configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None @@ -260,9 +259,9 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): if max_retries is not None and isinstance(max_retries, str): max_retries = int(max_retries) # cast to int # We need to keep max_retries in args since it's a parameter of GenericLiteLLMParams - args["max_retries"] = ( - max_retries # Put max_retries back in args after popping it - ) + args[ + "max_retries" + ] = max_retries # Put max_retries back in args after popping it super().__init__(**args, **params) def __contains__(self, key):