mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
fcaa4a9f30
commit
b8b78f1fde
17 changed files with 370 additions and 179 deletions
36
enterprise/enterprise_hooks/__init__.py
Normal file
36
enterprise/enterprise_hooks/__init__.py
Normal 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]
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
45
litellm/proxy/openai_files_endpoints/common_utils.py
Normal file
45
litellm/proxy/openai_files_endpoints/common_utils.py
Normal 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 []
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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:{}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
Loading…
Add table
Reference in a new issue