From 8e3aa142872a362cc5958a501b5e17323c5aa934 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Wed, 15 May 2024 19:40:34 -0700 Subject: [PATCH] fix revert 3600 --- litellm/proxy/_types.py | 109 ++++++++++++++++++---------------------- 1 file changed, 49 insertions(+), 60 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b1af153e81f..0320b0d0f9e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,37 +1,11 @@ -from pydantic import ConfigDict, BaseModel, Field, root_validator, Json, VERSION +from pydantic import BaseModel, Extra, Field, root_validator, Json, validator +from dataclasses import fields import enum from typing import Optional, List, Union, Dict, Literal, Any from datetime import datetime -import uuid -import json +import uuid, json, sys, os from litellm.types.router import UpdateRouterConfig -try: - from pydantic import model_validator # type: ignore -except ImportError: - from pydantic import root_validator # pydantic v1 - - def model_validator(mode): # type: ignore - pre = mode == "before" - return root_validator(pre=pre) - - -# Function to get Pydantic version -def is_pydantic_v2() -> int: - return int(VERSION.split(".")[0]) - - -def get_model_config(arbitrary_types_allowed: bool = False) -> ConfigDict: - # Version-specific configuration - if is_pydantic_v2() >= 2: - model_config = ConfigDict(extra="allow", arbitrary_types_allowed=arbitrary_types_allowed, protected_namespaces=()) # type: ignore - else: - from pydantic import Extra - - model_config = ConfigDict(extra=Extra.allow, arbitrary_types_allowed=arbitrary_types_allowed) # type: ignore - - return model_config - def hash_token(token: str): import hashlib @@ -61,7 +35,8 @@ class LiteLLMBase(BaseModel): # if using pydantic v1 return self.__fields_set__ - model_config = get_model_config() + class Config: + protected_namespaces = () class LiteLLM_UpperboundKeyGenerateParams(LiteLLMBase): @@ -104,11 +79,6 @@ class LiteLLMRoutes(enum.Enum): "/v1/models", ] - # NOTE: ROUTES ONLY FOR MASTER KEY - only the Master Key should be able to Reset Spend - master_key_only_routes: List = [ - "/global/spend/reset", - ] - info_routes: List = [ "/key/info", "/team/info", @@ -119,6 +89,11 @@ class LiteLLMRoutes(enum.Enum): "/v2/key/info", ] + # NOTE: ROUTES ONLY FOR MASTER KEY - only the Master Key should be able to Reset Spend + master_key_only_routes: List = [ + "/global/spend/reset", + ] + sso_only_routes: List = [ "/key/generate", "/key/update", @@ -259,7 +234,7 @@ class LiteLLMPromptInjectionParams(LiteLLMBase): llm_api_system_prompt: Optional[str] = None llm_api_fail_call_string: Optional[str] = None - @model_validator(mode="before") + @root_validator(pre=True) def check_llm_api_params(cls, values): llm_api_check = values.get("llm_api_check") if llm_api_check is True: @@ -317,7 +292,8 @@ class ProxyChatCompletionRequest(LiteLLMBase): deployment_id: Optional[str] = None request_timeout: Optional[int] = None - model_config = get_model_config() + class Config: + extra = "allow" # allow params not defined here, these fall in litellm.completion(**kwargs) class ModelInfoDelete(LiteLLMBase): @@ -344,9 +320,11 @@ class ModelInfo(LiteLLMBase): ] ] - model_config = get_model_config() + class Config: + extra = Extra.allow # Allow extra fields + protected_namespaces = () - @model_validator(mode="before") + @root_validator(pre=True) def set_model_info(cls, values): if values.get("id") is None: values.update({"id": str(uuid.uuid4())}) @@ -372,9 +350,10 @@ class ModelParams(LiteLLMBase): litellm_params: dict model_info: ModelInfo - model_config = get_model_config() + class Config: + protected_namespaces = () - @model_validator(mode="before") + @root_validator(pre=True) def set_model_info(cls, values): if values.get("model_info") is None: values.update({"model_info": ModelInfo()}) @@ -410,7 +389,8 @@ class GenerateKeyRequest(GenerateRequestBase): {} ) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} - model_config = get_model_config() + class Config: + protected_namespaces = () class GenerateKeyResponse(GenerateKeyRequest): @@ -420,7 +400,7 @@ class GenerateKeyResponse(GenerateKeyRequest): user_id: Optional[str] = None token_id: Optional[str] = None - @model_validator(mode="before") + @root_validator(pre=True) def set_model_info(cls, values): if values.get("token") is not None: values.update({"key": values.get("token")}) @@ -460,7 +440,8 @@ class LiteLLM_ModelTable(LiteLLMBase): created_by: str updated_by: str - model_config = get_model_config() + class Config: + protected_namespaces = () class NewUserRequest(GenerateKeyRequest): @@ -488,7 +469,7 @@ class UpdateUserRequest(GenerateRequestBase): user_role: Optional[str] = None max_budget: Optional[float] = None - @model_validator(mode="before") + @root_validator(pre=True) def check_user_info(cls, values): if values.get("user_id") is None and values.get("user_email") is None: raise ValueError("Either user id or user email must be provided") @@ -508,7 +489,7 @@ class NewEndUserRequest(LiteLLMBase): None # if no equivalent model in allowed region - default all requests to this model ) - @model_validator(mode="before") + @root_validator(pre=True) def check_user_info(cls, values): if values.get("max_budget") is not None and values.get("budget_id") is not None: raise ValueError("Set either 'max_budget' or 'budget_id', not both.") @@ -521,7 +502,7 @@ class Member(LiteLLMBase): user_id: Optional[str] = None user_email: Optional[str] = None - @model_validator(mode="before") + @root_validator(pre=True) def check_user_info(cls, values): if values.get("user_id") is None and values.get("user_email") is None: raise ValueError("Either user id or user email must be provided") @@ -546,7 +527,8 @@ class TeamBase(LiteLLMBase): class NewTeamRequest(TeamBase): model_aliases: Optional[dict] = None - model_config = get_model_config() + class Config: + protected_namespaces = () class GlobalEndUsersSpend(LiteLLMBase): @@ -565,7 +547,7 @@ class TeamMemberDeleteRequest(LiteLLMBase): user_id: Optional[str] = None user_email: Optional[str] = None - @model_validator(mode="before") + @root_validator(pre=True) def check_user_info(cls, values): if values.get("user_id") is None and values.get("user_email") is None: raise ValueError("Either user id or user email must be provided") @@ -599,9 +581,10 @@ class LiteLLM_TeamTable(TeamBase): budget_reset_at: Optional[datetime] = None model_id: Optional[int] = None - model_config = get_model_config() + class Config: + protected_namespaces = () - @model_validator(mode="before") + @root_validator(pre=True) def set_model_info(cls, values): dict_fields = [ "metadata", @@ -637,7 +620,8 @@ class LiteLLM_BudgetTable(LiteLLMBase): model_max_budget: Optional[dict] = None budget_duration: Optional[str] = None - model_config = get_model_config() + class Config: + protected_namespaces = () class NewOrganizationRequest(LiteLLM_BudgetTable): @@ -687,7 +671,8 @@ class KeyManagementSettings(LiteLLMBase): class TeamDefaultSettings(LiteLLMBase): team_id: str - model_config = get_model_config() + class Config: + extra = "allow" # allow params not defined here, these fall in litellm.completion(**kwargs) class DynamoDBArgs(LiteLLMBase): @@ -828,7 +813,8 @@ class ConfigYAML(LiteLLMBase): description="litellm router object settings. See router.py __init__ for all, example router.num_retries=5, router.timeout=5, router.max_retries=5, router.retry_after=5", ) - model_config = get_model_config() + class Config: + protected_namespaces = () class LiteLLM_VerificationToken(LiteLLMBase): @@ -862,7 +848,8 @@ class LiteLLM_VerificationToken(LiteLLMBase): user_id_rate_limits: Optional[dict] = None team_id_rate_limits: Optional[dict] = None - model_config = get_model_config() + class Config: + protected_namespaces = () class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken): @@ -892,7 +879,7 @@ class UserAPIKeyAuth( user_role: Optional[Literal["proxy_admin", "app_owner", "app_user"]] = None allowed_model_region: Optional[Literal["eu"]] = None - @model_validator(mode="before") + @root_validator(pre=True) def check_api_key(cls, values): if values.get("api_key") is not None: values.update({"token": hash_token(values.get("api_key"))}) @@ -919,7 +906,7 @@ class LiteLLM_UserTable(LiteLLMBase): tpm_limit: Optional[int] = None rpm_limit: Optional[int] = None - @model_validator(mode="before") + @root_validator(pre=True) def set_model_info(cls, values): if values.get("spend") is None: values.update({"spend": 0.0}) @@ -927,7 +914,8 @@ class LiteLLM_UserTable(LiteLLMBase): values.update({"models": []}) return values - model_config = get_model_config() + class Config: + protected_namespaces = () class LiteLLM_EndUserTable(LiteLLMBase): @@ -939,13 +927,14 @@ class LiteLLM_EndUserTable(LiteLLMBase): default_model: Optional[str] = None litellm_budget_table: Optional[LiteLLM_BudgetTable] = None - @model_validator(mode="before") + @root_validator(pre=True) def set_model_info(cls, values): if values.get("spend") is None: values.update({"spend": 0.0}) return values - model_config = get_model_config() + class Config: + protected_namespaces = () class LiteLLM_SpendLogs(LiteLLMBase):