From b8b78f1fde40fea45125949f4a329fb773c6f44b Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Wed, 7 May 2025 23:39:40 -0700 Subject: [PATCH] Support unified file id (managed files) for batches (#10650) * refactor(managed_files.py): move enterprise feature into enterprise folder prevent unexpected surprises * refactor: safely handle enterprise hooks * fix: fix ruff check errors * fix(files_endpoints.py): cleanup enterprise code from OSS * refactor: complete cleanup * fix(managed_files.py): complete cleanup * fix(managed_files.py): instrument to be able to update deployment values post-router selection and just before making llm call * fix(managed_files.py): instrument to be able to update deployment values post-router selection and just before making llm call * fix: fix linting error * fix: fix linting error --- enterprise/enterprise_hooks/__init__.py | 36 +++++++ .../enterprise_hooks}/managed_files.py | 99 ++++++++----------- litellm/integrations/custom_logger.py | 13 +++ .../prompt_templates/common_utils.py | 8 +- litellm/llms/base_llm/files/transformation.py | 59 ++++++++++- litellm/proxy/_new_secret_config.yaml | 23 +---- litellm/proxy/batches_endpoints/endpoints.py | 54 +++++++--- litellm/proxy/common_request_processing.py | 3 +- litellm/proxy/hooks/__init__.py | 31 ++---- .../openai_files_endpoints/common_utils.py | 45 +++++++++ .../openai_files_endpoints/files_endpoints.py | 64 +++++++----- litellm/router.py | 67 +++++++------ litellm/router_utils/batch_utils.py | 4 +- litellm/types/llms/openai.py | 10 +- litellm/types/utils.py | 4 +- litellm/utils.py | 25 +++++ .../enterprise_hooks}/test_managed_files.py | 4 +- 17 files changed, 370 insertions(+), 179 deletions(-) create mode 100644 enterprise/enterprise_hooks/__init__.py rename {litellm/proxy/hooks => enterprise/enterprise_hooks}/managed_files.py (86%) create mode 100644 litellm/proxy/openai_files_endpoints/common_utils.py rename tests/{litellm/proxy/hooks => enterprise/enterprise_hooks}/test_managed_files.py (97%) diff --git a/enterprise/enterprise_hooks/__init__.py b/enterprise/enterprise_hooks/__init__.py new file mode 100644 index 00000000000..fa51e454625 --- /dev/null +++ b/enterprise/enterprise_hooks/__init__.py @@ -0,0 +1,36 @@ +import os +from typing import Dict, Literal, Type, Union + +from litellm.integrations.custom_logger import CustomLogger + +from .managed_files import _PROXY_LiteLLMManagedFiles +from .parallel_request_limiter_v2 import _PROXY_MaxParallelRequestsHandler + +ENTERPRISE_PROXY_HOOKS: Dict[str, Type[CustomLogger]] = { + "managed_files": _PROXY_LiteLLMManagedFiles, +} + + +## FEATURE FLAG HOOKS ## + +if os.getenv("EXPERIMENTAL_MULTI_INSTANCE_RATE_LIMITING", "false").lower() == "true": + ENTERPRISE_PROXY_HOOKS["max_parallel_requests"] = _PROXY_MaxParallelRequestsHandler + + +def get_enterprise_proxy_hook( + hook_name: Union[ + Literal[ + "managed_files", + "max_parallel_requests", + ], + str, + ] +): + """ + Factory method to get a enterprise hook instance by name + """ + if hook_name not in ENTERPRISE_PROXY_HOOKS: + raise ValueError( + f"Unknown hook: {hook_name}. Available hooks: {list(ENTERPRISE_PROXY_HOOKS.keys())}" + ) + return ENTERPRISE_PROXY_HOOKS[hook_name] diff --git a/litellm/proxy/hooks/managed_files.py b/enterprise/enterprise_hooks/managed_files.py similarity index 86% rename from litellm/proxy/hooks/managed_files.py rename to enterprise/enterprise_hooks/managed_files.py index 9ac6cc580b7..3d5f6f6096a 100644 --- a/litellm/proxy/hooks/managed_files.py +++ b/enterprise/enterprise_hooks/managed_files.py @@ -4,14 +4,18 @@ import base64 import json import uuid -from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast from litellm import Router, verbose_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data +from litellm.llms.base_llm.files.transformation import BaseFileEndpoints from litellm.proxy._types import CallTypes, LiteLLM_ManagedFileTable, UserAPIKeyAuth +from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, + convert_b64_uid_to_unified_uid, +) from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionFileObject, @@ -36,29 +40,7 @@ else: PrismaClient = Any -class BaseFileEndpoints(ABC): - @abstractmethod - async def afile_retrieve( - self, - file_id: str, - litellm_parent_otel_span: Optional[Span], - ) -> OpenAIFileObject: - pass - - @abstractmethod - async def afile_list( - self, custom_llm_provider: str, **data: dict - ) -> List[OpenAIFileObject]: - pass - - @abstractmethod - async def afile_delete( - self, custom_llm_provider: str, file_id: str, **data: dict - ) -> OpenAIFileObject: - pass - - -class _PROXY_LiteLLMManagedFiles(CustomLogger): +class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # Class variables or attributes def __init__( self, internal_usage_cache: InternalUsageCache, prisma_client: PrismaClient @@ -153,12 +135,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): "audio_transcription", "pass_through_endpoint", "rerank", + "acreate_batch", ], ) -> Union[Exception, str, Dict, None]: """ - Detect litellm_proxy/ file_id - add dictionary of mappings of litellm_proxy/ file_id -> provider_file_id => {litellm_proxy/file_id: {"model_id": id, "file_id": provider_file_id}} """ + print("REACHES async_pre_call_hook, call_type:", call_type) if call_type == CallTypes.completion.value: messages = data.get("messages") if messages: @@ -169,9 +153,37 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): ) data["model_file_id_mapping"] = model_file_id_mapping + elif call_type == CallTypes.acreate_batch.value: + input_file_id = cast(Optional[str], data.get("input_file_id")) + if input_file_id: + model_file_id_mapping = await self.get_model_file_id_mapping( + [input_file_id], user_api_key_dict.parent_otel_span + ) + data["model_file_id_mapping"] = model_file_id_mapping return data + async def async_pre_call_deployment_hook( + self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] + ) -> Optional[dict]: + """ + Allow modifying the request just before it's sent to the deployment. + """ + if call_type and call_type == CallTypes.acreate_batch: + input_file_id = cast(Optional[str], kwargs.get("input_file_id")) + model_file_id_mapping = cast( + Optional[Dict[str, Dict[str, str]]], kwargs.get("model_file_id_mapping") + ) + model_id = cast(Optional[str], kwargs.get("model_info", {}).get("id", None)) + mapped_file_id: Optional[str] = None + if input_file_id and model_file_id_mapping and model_id: + mapped_file_id = model_file_id_mapping.get(input_file_id, {}).get( + model_id, None + ) + if mapped_file_id: + kwargs["input_file_id"] = mapped_file_id + return kwargs + def get_file_ids_from_messages(self, messages: List[AllMessageValues]) -> List[str]: """ Gets file ids from messages @@ -192,37 +204,6 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): file_ids.append(file_id) return file_ids - @staticmethod - def _convert_b64_uid_to_unified_uid(b64_uid: str) -> str: - is_base64_unified_file_id = ( - _PROXY_LiteLLMManagedFiles._is_base64_encoded_unified_file_id(b64_uid) - ) - if is_base64_unified_file_id: - return is_base64_unified_file_id - else: - return b64_uid - - @staticmethod - def _is_base64_encoded_unified_file_id(b64_uid: str) -> Union[str, Literal[False]]: - # Add padding back if needed - padded = b64_uid + "=" * (-len(b64_uid) % 4) - # Decode from base64 - try: - decoded = base64.urlsafe_b64decode(padded).decode() - if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value): - return decoded - else: - return False - except Exception: - return False - - def convert_b64_uid_to_unified_uid(self, b64_uid: str) -> str: - is_base64_unified_file_id = self._is_base64_encoded_unified_file_id(b64_uid) - if is_base64_unified_file_id: - return is_base64_unified_file_id - else: - return b64_uid - async def get_model_file_id_mapping( self, file_ids: List[str], litellm_parent_otel_span: Span ) -> dict: @@ -247,7 +228,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): for file_id in file_ids: ## CHECK IF FILE ID IS MANAGED BY LITELM - is_base64_unified_file_id = self._is_base64_encoded_unified_file_id(file_id) + is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: litellm_managed_file_ids.append(file_id) @@ -300,6 +281,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): create_file_request=create_file_request, internal_usage_cache=self.internal_usage_cache, litellm_parent_otel_span=litellm_parent_otel_span, + target_model_names_list=target_model_names_list, ) ## STORE MODEL MAPPINGS IN DB @@ -328,6 +310,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): create_file_request: CreateFileRequest, internal_usage_cache: InternalUsageCache, litellm_parent_otel_span: Span, + target_model_names_list: List[str], ) -> OpenAIFileObject: ## GET THE FILE TYPE FROM THE CREATE FILE REQUEST file_data = extract_file_data(create_file_request["file"]) @@ -335,7 +318,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): file_type = file_data["content_type"] unified_file_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( - file_type, str(uuid.uuid4()) + file_type, str(uuid.uuid4()), ",".join(target_model_names_list) ) # Convert to URL-safe base64 and strip padding @@ -383,7 +366,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger): llm_router: Router, **data: Dict, ) -> OpenAIFileObject: - file_id = self.convert_b64_uid_to_unified_uid(file_id) + file_id = convert_b64_uid_to_unified_uid(file_id) model_file_id_mapping = await self.get_model_file_id_mapping( [file_id], litellm_parent_otel_span ) diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 17441ba1d0a..7b19e8c8b13 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -21,6 +21,7 @@ from litellm.types.integrations.argilla import ArgillaItem from litellm.types.llms.openai import AllMessageValues, ChatCompletionRequest from litellm.types.utils import ( AdapterCompletionStreamWrapper, + CallTypes, LLMResponseTypes, ModelResponse, ModelResponseStream, @@ -127,6 +128,18 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac ) -> List[dict]: return healthy_deployments + async def async_pre_call_deployment_hook( + self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] + ) -> Optional[dict]: + """ + Allow modifying the request just before it's sent to the deployment. + + Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent. + + Used in managed_files.py + """ + pass + async def async_pre_call_check( self, deployment: dict, parent_otel_span: Optional[Span] ) -> Optional[dict]: diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index b6af4a710ad..387c072ffd7 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -346,14 +346,14 @@ def get_format_from_file_id(file_id: Optional[str]) -> Optional[str]: unified_file_id = litellm_proxy:{};unified_id,{} If not a unified file id, returns 'file' as default format """ - from litellm.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + from litellm.proxy.openai_files_endpoints.common_utils import ( + convert_b64_uid_to_unified_uid, + ) if not file_id: return None try: - transformed_file_id = ( - _PROXY_LiteLLMManagedFiles._convert_b64_uid_to_unified_uid(file_id) - ) + transformed_file_id = convert_b64_uid_to_unified_uid(file_id) if transformed_file_id.startswith( SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value ): diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index 9925004c896..4d749af21e1 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -1,5 +1,5 @@ -from abc import abstractmethod -from typing import TYPE_CHECKING, Any, List, Optional, Union +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import httpx @@ -8,6 +8,7 @@ from litellm.types.llms.openai import ( CreateFileRequest, OpenAICreateFileRequestOptionalParams, OpenAIFileObject, + OpenAIFilesPurpose, ) from litellm.types.utils import LlmProviders, ModelResponse @@ -15,10 +16,15 @@ from ..chat.transformation import BaseConfig if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + from litellm.router import Router as _Router LiteLLMLoggingObj = _LiteLLMLoggingObj + Span = Any + Router = _Router else: LiteLLMLoggingObj = Any + Span = Any + Router = Any class BaseFilesConfig(BaseConfig): @@ -99,3 +105,52 @@ class BaseFilesConfig(BaseConfig): raise NotImplementedError( "AudioTranscriptionConfig does not need a response transformation for audio transcription models" ) + + +class BaseFileEndpoints(ABC): + @abstractmethod + async def acreate_file( + self, + create_file_request: CreateFileRequest, + llm_router: Router, + target_model_names_list: List[str], + litellm_parent_otel_span: Span, + ) -> OpenAIFileObject: + pass + + @abstractmethod + async def afile_retrieve( + self, + file_id: str, + litellm_parent_otel_span: Optional[Span], + ) -> OpenAIFileObject: + pass + + @abstractmethod + async def afile_list( + self, + purpose: Optional[OpenAIFilesPurpose], + litellm_parent_otel_span: Optional[Span], + **data: Dict, + ) -> List[OpenAIFileObject]: + pass + + @abstractmethod + async def afile_delete( + self, + file_id: str, + litellm_parent_otel_span: Optional[Span], + llm_router: Router, + **data: Dict, + ) -> OpenAIFileObject: + pass + + @abstractmethod + async def afile_content( + self, + file_id: str, + litellm_parent_otel_span: Optional[Span], + llm_router: Router, + **data: Dict, + ) -> str: + pass diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 11069452097..ff9e0c8f367 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,24 +1,9 @@ model_list: - - model_name: gpt-4o-mini-tts + - model_name: "gemini-2.0-flash" litellm_params: - model: openai/gpt-4o-mini-tts - api_key: os.environ/OPENAI_API_KEY - - model_name: gpt-3.5-turbo - litellm_params: - model: azure/chatgpt-v-3 - api_base: https://openai-gpt-4-test-v-1.openai.azure.com/ - api_version: "2023-05-15" - api_key: os.environ/AZURE_API_KEY - - model_name: "gpt-4o-azure" - litellm_params: - model: azure/gpt-4o - api_key: os.environ/AZURE_API_KEY - api_base: os.environ/AZURE_API_BASE - - model_name: fake-openai-endpoint - litellm_params: - model: openai/fake - api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + model: vertex_ai/gemini-2.0-flash + vertex_project: my-project-id + vertex_location: us-central1 - model_name: "gpt-4o-mini-openai" litellm_params: model: gpt-4o-mini diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 6b7651d48f3..d8f3c58849f 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -11,11 +11,7 @@ from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response import litellm from litellm._logging import verbose_proxy_logger -from litellm.batches.main import ( - CancelBatchRequest, - CreateBatchRequest, - RetrieveBatchRequest, -) +from litellm.batches.main import CancelBatchRequest, RetrieveBatchRequest from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -23,8 +19,13 @@ from litellm.proxy.common_utils.http_parsing_utils import _read_request_body from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, ) +from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, + get_models_from_unified_file_id, +) from litellm.proxy.openai_files_endpoints.files_endpoints import is_known_model from litellm.proxy.utils import handle_exception_on_proxy +from litellm.types.llms.openai import LiteLLMBatchCreateRequest router = APIRouter() @@ -68,7 +69,6 @@ async def create_batch( ``` """ from litellm.proxy.proxy_server import ( - add_litellm_data_to_request, general_settings, llm_router, proxy_config, @@ -82,15 +82,18 @@ async def create_batch( verbose_proxy_logger.debug( "Request received by LiteLLM:\n{}".format(json.dumps(data, indent=4)), ) - - # Include original request and headers in the data - data = await add_litellm_data_to_request( - data=data, + base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) + ( + data, + litellm_logging_obj, + ) = await base_llm_response_processor.common_processing_pre_call_logic( request=request, general_settings=general_settings, user_api_key_dict=user_api_key_dict, version=version, + proxy_logging_obj=proxy_logging_obj, proxy_config=proxy_config, + route_type="acreate_batch", ) ## check if model is a loadbalanced model @@ -103,7 +106,11 @@ async def create_batch( custom_llm_provider = ( provider or data.pop("custom_llm_provider", None) or "openai" ) - _create_batch_data = CreateBatchRequest(**data) + _create_batch_data = LiteLLMBatchCreateRequest(**data) + input_file_id = _create_batch_data.get("input_file_id", None) + unified_file_id: Union[str, Literal[False]] = False + if input_file_id: + unified_file_id = _is_base64_encoded_unified_file_id(input_file_id) if ( litellm.enable_loadbalancing_on_batch_endpoints is True and is_router_model @@ -118,6 +125,31 @@ async def create_batch( ) response = await llm_router.acreate_batch(**_create_batch_data) # type: ignore + elif ( + unified_file_id + ): # litellm_proxy:application/octet-stream;unified_id,c4843482-b176-4901-8292-7523fd0f2c6e;target_model_names,gpt-4o-mini + target_model_names = get_models_from_unified_file_id(unified_file_id) + ## EXPECTS 1 MODEL + if len(target_model_names) != 1: + raise HTTPException( + status_code=400, + detail={ + "error": "Expected 1 model, got {}".format( + len(target_model_names) + ) + }, + ) + model = target_model_names[0] + _create_batch_data["model"] = model + if llm_router is None: + raise HTTPException( + status_code=500, + detail={ + "error": "LLM Router not initialized. Ensure models added to proxy." + }, + ) + + response = await llm_router.acreate_batch(**_create_batch_data) else: response = await litellm.acreate_batch( custom_llm_provider=custom_llm_provider, **_create_batch_data # type: ignore diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 2ea3c18ea80..3c9a3f4a8f3 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -114,6 +114,7 @@ class ProxyBaseLLMRequestProcessing: "_arealtime", "aget_responses", "adelete_responses", + "acreate_batch", ], version: Optional[str] = None, user_model: Optional[str] = None, @@ -163,7 +164,7 @@ class ProxyBaseLLMRequestProcessing: ) ### CALL HOOKS ### - modify/reject incoming data before calling the model self.data = await proxy_logging_obj.pre_call_hook( # type: ignore - user_api_key_dict=user_api_key_dict, data=self.data, call_type="completion" + user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type # type: ignore ) ## LOGGING OBJECT ## - initialize logging object for logging success/failure events for call diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 23bb6c3012b..9c78989cad4 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -1,39 +1,28 @@ -import os -from typing import Literal, Type, Union +from typing import Literal, Union from . import * from .cache_control_check import _PROXY_CacheControlCheck -from .managed_files import _PROXY_LiteLLMManagedFiles from .max_budget_limiter import _PROXY_MaxBudgetLimiter from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler -try: - if ( - os.getenv("EXPERIMENTAL_MULTI_INSTANCE_RATE_LIMITING", "false").lower() - == "true" - ): # FEATURE FLAG as it's still in development - from enterprise.enterprise_hooks.parallel_request_limiter_v2 import ( - _PROXY_MaxParallelRequestsHandler as _PROXY_MaxParallelRequestsHandlerV2, - ) +### CHECK IF ENTERPRISE HOOKS ARE AVAILABLE ### - max_parallel_request_handler: Type[ - Union[ - _PROXY_MaxParallelRequestsHandler, _PROXY_MaxParallelRequestsHandlerV2 - ] - ] = _PROXY_MaxParallelRequestsHandlerV2 - else: - max_parallel_request_handler = _PROXY_MaxParallelRequestsHandler +try: + from enterprise.enterprise_hooks import ENTERPRISE_PROXY_HOOKS except ImportError: - max_parallel_request_handler = _PROXY_MaxParallelRequestsHandler + ENTERPRISE_PROXY_HOOKS = {} # List of all available hooks that can be enabled PROXY_HOOKS = { "max_budget_limiter": _PROXY_MaxBudgetLimiter, - "managed_files": _PROXY_LiteLLMManagedFiles, - "parallel_request_limiter": max_parallel_request_handler, + "parallel_request_limiter": _PROXY_MaxParallelRequestsHandler, "cache_control_check": _PROXY_CacheControlCheck, } +### update PROXY_HOOKS with ENTERPRISE_PROXY_HOOKS ### + +PROXY_HOOKS.update(ENTERPRISE_PROXY_HOOKS) + def get_proxy_hook( hook_name: Union[ diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py new file mode 100644 index 00000000000..c176a383f1f --- /dev/null +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -0,0 +1,45 @@ +import base64 +import re +from typing import List, Literal, Union + +from litellm.types.utils import SpecialEnums + + +def _is_base64_encoded_unified_file_id(b64_uid: str) -> Union[str, Literal[False]]: + # Add padding back if needed + padded = b64_uid + "=" * (-len(b64_uid) % 4) + # Decode from base64 + try: + decoded = base64.urlsafe_b64decode(padded).decode() + if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value): + return decoded + else: + return False + except Exception: + return False + + +def convert_b64_uid_to_unified_uid(b64_uid: str) -> str: + is_base64_unified_file_id = _is_base64_encoded_unified_file_id(b64_uid) + if is_base64_unified_file_id: + return is_base64_unified_file_id + else: + return b64_uid + + +def get_models_from_unified_file_id(unified_file_id: str) -> List[str]: + """ + Extract model names from unified file ID. + + Example: + unified_file_id = "litellm_proxy:application/octet-stream;unified_id,c4843482-b176-4901-8292-7523fd0f2c6e;target_model_names,gpt-4o-mini,gemini-2.0-flash" + returns: ["gpt-4o-mini", "gemini-2.0-flash"] + """ + try: + match = re.search(r"target_model_names,([^;]+)", unified_file_id) + if match: + # Split on comma and strip whitespace from each model name + return [model.strip() for model in match.group(1).split(",")] + return [] + except Exception: + return [] diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index bf29bdf6bd5..568dff7cc19 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -25,13 +25,13 @@ from fastapi import ( import litellm from litellm import CreateFileRequest, get_secret_str from litellm._logging import verbose_proxy_logger +from litellm.llms.base_llm.files.transformation import BaseFileEndpoints from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, ) -from litellm.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles from litellm.proxy.utils import ProxyLogging from litellm.router import Router from litellm.types.llms.openai import ( @@ -40,6 +40,8 @@ from litellm.types.llms.openai import ( OpenAIFilesPurpose, ) +from .common_utils import _is_base64_encoded_unified_file_id + router = APIRouter() files_config = None @@ -150,10 +152,7 @@ async def route_create_file( _create_file_request=_create_file_request, ) elif target_model_names_list: - managed_files_obj = cast( - Optional[_PROXY_LiteLLMManagedFiles], - proxy_logging_obj.get_proxy_hook("managed_files"), - ) + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( message="Managed files hook not found", @@ -168,6 +167,13 @@ async def route_create_file( param="None", code=500, ) + if not isinstance(managed_files_obj, BaseFileEndpoints): + raise ProxyException( + message="Managed files hook is not a BaseFileEndpoints", + type="None", + param="None", + code=500, + ) response = await managed_files_obj.acreate_file( llm_router=llm_router, create_file_request=_create_file_request, @@ -434,14 +440,9 @@ async def get_file_content( ) ## check if file_id is a litellm managed file - is_base64_unified_file_id = ( - _PROXY_LiteLLMManagedFiles._is_base64_encoded_unified_file_id(file_id) - ) + is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: - managed_files_obj = cast( - Optional[_PROXY_LiteLLMManagedFiles], - proxy_logging_obj.get_proxy_hook("managed_files"), - ) + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( message="Managed files hook not found", @@ -456,6 +457,13 @@ async def get_file_content( param="None", code=500, ) + if not isinstance(managed_files_obj, BaseFileEndpoints): + raise ProxyException( + message="Managed files hook is not a BaseFileEndpoints", + type="None", + param="None", + code=500, + ) response = await managed_files_obj.afile_content( file_id=file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, @@ -590,15 +598,10 @@ async def get_file( ) ## check if file_id is a litellm managed file - is_base64_unified_file_id = ( - _PROXY_LiteLLMManagedFiles._is_base64_encoded_unified_file_id(file_id) - ) + is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: - managed_files_obj = cast( - Optional[_PROXY_LiteLLMManagedFiles], - proxy_logging_obj.get_proxy_hook("managed_files"), - ) + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( message="Managed files hook not found", @@ -606,6 +609,13 @@ async def get_file( param="None", code=500, ) + if not isinstance(managed_files_obj, BaseFileEndpoints): + raise ProxyException( + message="Managed files hook is not a BaseFileEndpoints", + type="None", + param="None", + code=500, + ) response = await managed_files_obj.afile_retrieve( file_id=file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, @@ -730,15 +740,10 @@ async def delete_file( ) ## check if file_id is a litellm managed file - is_base64_unified_file_id = ( - _PROXY_LiteLLMManagedFiles._is_base64_encoded_unified_file_id(file_id) - ) + is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id) if is_base64_unified_file_id: - managed_files_obj = cast( - Optional[_PROXY_LiteLLMManagedFiles], - proxy_logging_obj.get_proxy_hook("managed_files"), - ) + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") if managed_files_obj is None: raise ProxyException( message="Managed files hook not found", @@ -753,6 +758,13 @@ async def delete_file( param="None", code=500, ) + if not isinstance(managed_files_obj, BaseFileEndpoints): + raise ProxyException( + message="Managed files hook is not a BaseFileEndpoints", + type="None", + param="None", + code=500, + ) response = await managed_files_obj.afile_delete( file_id=file_id, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, diff --git a/litellm/router.py b/litellm/router.py index e7f98fab483..ffb589a8fdc 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -342,9 +342,9 @@ class Router: ) # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### - cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = ( - "local" # default to an in-memory cache - ) + cache_type: Literal[ + "local", "redis", "redis-semantic", "s3", "disk" + ] = "local" # default to an in-memory cache redis_cache = None cache_config: Dict[str, Any] = {} @@ -565,9 +565,9 @@ class Router: ) ) - self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = ( - model_group_retry_policy - ) + self.model_group_retry_policy: Optional[ + Dict[str, RetryPolicy] + ] = model_group_retry_policy self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None if allowed_fails_policy is not None: @@ -1099,7 +1099,12 @@ class Router: self.fail_calls[model_name] += 1 raise e - def _update_kwargs_before_fallbacks(self, model: str, kwargs: dict) -> None: + def _update_kwargs_before_fallbacks( + self, + model: str, + kwargs: dict, + metadata_variable_name: Optional[str] = "metadata", + ) -> None: """ Adds/updates to kwargs: - num_retries @@ -1108,7 +1113,7 @@ class Router: """ kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries) kwargs.setdefault("litellm_trace_id", str(uuid.uuid4())) - kwargs.setdefault("metadata", {}).update({"model_group": model}) + kwargs.setdefault(metadata_variable_name, {}).update({"model_group": model}) def _update_kwargs_with_default_litellm_params( self, kwargs: dict, metadata_variable_name: Optional[str] = "metadata" @@ -1185,6 +1190,7 @@ class Router: metadata_variable_name = _get_router_metadata_variable_name( function_name=function_name, ) + kwargs.setdefault(metadata_variable_name, {}).update( { "deployment": deployment_model_name, @@ -2847,7 +2853,14 @@ class Router: kwargs["model"] = model kwargs["original_function"] = self._acreate_batch kwargs["num_retries"] = kwargs.get("num_retries", self.num_retries) - self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs) + metadata_variable_name = _get_router_metadata_variable_name( + function_name="_acreate_batch" + ) + self._update_kwargs_before_fallbacks( + model=model, + kwargs=kwargs, + metadata_variable_name=metadata_variable_name, + ) response = await self.async_function_with_fallbacks(**kwargs) return response @@ -2878,21 +2891,13 @@ class Router: specific_deployment=kwargs.pop("specific_deployment", None), request_kwargs=kwargs, ) - metadata_variable_name = _get_router_metadata_variable_name( - function_name="_acreate_batch" - ) - kwargs.setdefault(metadata_variable_name, {}).update( - { - "deployment": deployment["litellm_params"]["model"], - "model_info": deployment.get("model_info", {}), - "api_base": deployment.get("litellm_params", {}).get("api_base"), - } - ) kwargs["model_info"] = deployment.get("model_info", {}) data = deployment["litellm_params"].copy() model_name = data["model"] - self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) + self._update_kwargs_with_deployment( + deployment=deployment, kwargs=kwargs, function_name="_acreate_batch" + ) model_client = self._get_async_openai_model_client( deployment=deployment, @@ -3287,11 +3292,11 @@ class Router: if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: - fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=context_window_fallbacks, - model_group=model_group, - ) + fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=context_window_fallbacks, + model_group=model_group, ) if fallback_model_group is None: raise original_exception @@ -3323,11 +3328,11 @@ class Router: e.message += "\n{}".format(error_message) elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: - fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=content_policy_fallbacks, - model_group=model_group, - ) + fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=content_policy_fallbacks, + model_group=model_group, ) if fallback_model_group is None: raise original_exception @@ -3489,7 +3494,7 @@ class Router: num_retries = kwargs.pop("num_retries") ## ADD MODEL GROUP SIZE TO METADATA - used for model_group_rate_limit_error tracking - _metadata: dict = kwargs.get("metadata") or {} + _metadata: dict = kwargs.get("litellm_metadata", kwargs.get("metadata")) or {} if "model_group" in _metadata and isinstance(_metadata["model_group"], str): model_list = self.get_model_list(model_name=_metadata["model_group"]) if model_list is not None: diff --git a/litellm/router_utils/batch_utils.py b/litellm/router_utils/batch_utils.py index a41bae254c4..50b24ec363b 100644 --- a/litellm/router_utils/batch_utils.py +++ b/litellm/router_utils/batch_utils.py @@ -56,7 +56,9 @@ def _get_router_metadata_variable_name(function_name) -> str: For ALL other endpoints we call this "metadata """ - ROUTER_METHODS_USING_LITELLM_METADATA = set(["batch", "generic_api_call"]) + ROUTER_METHODS_USING_LITELLM_METADATA = set( + ["batch", "generic_api_call", "_acreate_batch"] + ) if function_name in ROUTER_METHODS_USING_LITELLM_METADATA: return "litellm_metadata" else: diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 8e79117f24b..73957a67671 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -378,6 +378,10 @@ class CreateBatchRequest(TypedDict, total=False): timeout: Optional[float] +class LiteLLMBatchCreateRequest(CreateBatchRequest, total=False): + model: str + + class RetrieveBatchRequest(TypedDict, total=False): """ RetrieveBatchRequest @@ -876,7 +880,9 @@ class FineTuningJobCreate(BaseModel): class LiteLLMFineTuningJobCreate(FineTuningJobCreate): custom_llm_provider: Literal["openai", "azure", "vertex_ai"] - model_config = {"extra": "allow"} # This allows the model to accept additional fields + model_config = { + "extra": "allow" + } # This allows the model to accept additional fields AllEmbeddingInputValues = Union[str, List[str], List[int], List[List[int]]] @@ -1331,7 +1337,7 @@ class OpenAIChatCompletionResponse(TypedDict, total=False): system_fingerprint: str service_tier: str + OpenAIChatCompletionFinishReason = Literal[ "stop", "content_filter", "function_call", "tool_calls", "length" ] - diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 50fd5cd9c16..ab3f0d9d2e2 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2307,7 +2307,9 @@ class ExtractedFileData(TypedDict): class SpecialEnums(Enum): LITELM_MANAGED_FILE_ID_PREFIX = "litellm_proxy" - LITELLM_MANAGED_FILE_COMPLETE_STR = "litellm_proxy:{};unified_id,{}" + LITELLM_MANAGED_FILE_COMPLETE_STR = ( + "litellm_proxy:{};unified_id,{};target_model_names,{}" + ) LITELLM_MANAGED_RESPONSE_COMPLETE_STR = ( "litellm:custom_llm_provider:{};model_id:{};response_id:{}" diff --git a/litellm/utils.py b/litellm/utils.py index 8217a1860ad..e28b0ceca15 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -866,6 +866,28 @@ def client(original_function): # noqa: PLR0915 else: return False + async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str): + """ + Allow modifying the request just before it's sent to the deployment. + + Use this instead of 'async_pre_call_hook' when you need to modify the request AFTER a deployment is selected, but BEFORE the request is sent. + """ + try: + typed_call_type = CallTypes(call_type) + except ValueError: + typed_call_type = None # unknown call type + + modified_kwargs = kwargs.copy() + for callback in litellm.callbacks: + if isinstance(callback, CustomLogger): + result = await callback.async_pre_call_deployment_hook( + modified_kwargs, typed_call_type + ) + if result is not None: + modified_kwargs = result + + return modified_kwargs + def post_call_processing(original_response, model, optional_params: Optional[dict]): try: if original_response is None: @@ -1279,6 +1301,9 @@ def client(original_function): # noqa: PLR0915 logging_obj, kwargs = function_setup( original_function.__name__, rules_obj, start_time, *args, **kwargs ) + modified_kwargs = await async_pre_call_deployment_hook(kwargs, call_type) + if modified_kwargs is not None: + kwargs = modified_kwargs kwargs["litellm_logging_obj"] = logging_obj ## LOAD CREDENTIALS diff --git a/tests/litellm/proxy/hooks/test_managed_files.py b/tests/enterprise/enterprise_hooks/test_managed_files.py similarity index 97% rename from tests/litellm/proxy/hooks/test_managed_files.py rename to tests/enterprise/enterprise_hooks/test_managed_files.py index b76e6c76d5e..89b332a506e 100644 --- a/tests/litellm/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/enterprise_hooks/test_managed_files.py @@ -6,13 +6,13 @@ import pytest from fastapi.testclient import TestClient sys.path.insert( - 0, os.path.abspath("../../../..") + 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path from unittest.mock import MagicMock +from enterprise.enterprise_hooks.managed_files import _PROXY_LiteLLMManagedFiles from litellm.caching import DualCache -from litellm.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles from litellm.types.utils import SpecialEnums