From 1db1af1154a03acc7f6caab601ce02f54277e0a2 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 14 May 2024 17:09:36 -0700 Subject: [PATCH] fix(types): fix typing --- litellm/proxy/_types.py | 4 +--- litellm/types/embedding.py | 21 +++++++++++++++++++-- 2 files changed, 20 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index fd166a1757d..b1af153e81f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -687,9 +687,7 @@ class KeyManagementSettings(LiteLLMBase): class TeamDefaultSettings(LiteLLMBase): team_id: str - model_config = ConfigDict( - extra="allow", # allow params not defined here, these fall in litellm.completion(**kwargs) - ) + model_config = get_model_config() class DynamoDBArgs(LiteLLMBase): diff --git a/litellm/types/embedding.py b/litellm/types/embedding.py index 4690b133246..831c4266c32 100644 --- a/litellm/types/embedding.py +++ b/litellm/types/embedding.py @@ -1,6 +1,23 @@ from typing import List, Optional, Union -from pydantic import ConfigDict, BaseModel, validator +from pydantic import ConfigDict, BaseModel, validator, VERSION + + +# 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 class EmbeddingRequest(BaseModel): @@ -17,4 +34,4 @@ class EmbeddingRequest(BaseModel): litellm_call_id: Optional[str] = None litellm_logging_obj: Optional[dict] = None logger_fn: Optional[str] = None - model_config = ConfigDict(extra="allow") + model_config = get_model_config()