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
This commit is contained in:
Krish Dholakia 2025-05-07 23:39:40 -07:00 • committed by GitHub
parent fcaa4a9f30
commit b8b78f1fde
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 370 additions and 179 deletions

View file

@ -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]

View file

@ -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
)

View file

@ -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]:

View file

@ -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
):

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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[

View file

@ -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 []

View file

@ -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,

View file

@ -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:

View file

@ -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:

View file

@ -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"
]

View file

@ -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:{}"

View file

@ -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

View file

@ -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