fix: req chagnes

This commit is contained in:
Harshit Jain 2026-02-24 14:08:47 +05:30
parent 4560c4fa30
commit 208ac5a5a5
No known key found for this signature in database
GPG key ID: 36C392CD4415B4CF
2 changed files with 137 additions and 73 deletions

View file

@ -1,8 +1,14 @@
import json
from typing import Optional
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypesLiteral
from litellm.proxy.utils import ProxyLogging
def _get_call_type_from_endpoint(endpoint: str) -> CallTypesLiteral:
"""Map batch JSONL endpoint to CallTypesLiteral."""
@ -13,14 +19,19 @@ def _get_call_type_from_endpoint(endpoint: str) -> CallTypesLiteral:
else:
return "acompletion" # default fallback
async def run_pre_call_guardrails_on_batch_file(
file_content: bytes,
proxy_logging_obj: ProxyLogging,
user_api_key_cache: DualCache,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""
Parse a batch JSONL file and run pre-call guardrails on each request.
Directly invokes guardrail callbacks instead of going through
pre_call_hook() which requires internal proxy state (logging objects,
pipelines, etc.) that batch file items don't have.
Raises an exception if any guardrail rejects a request.
"""
lines = file_content.decode("utf-8").strip().splitlines()
@ -28,29 +39,54 @@ async def run_pre_call_guardrails_on_batch_file(
for line_num, line in enumerate(lines, start=1):
if not line.strip():
continue
try:
json_obj = json.loads(line)
except json.JSONDecodeError:
continue
body = json_obj.get("body", {})
if not body or "messages" not in body:
continue
# Determine call_type from the endpoint in the JSONL line
endpoint = json_obj.get("url", "")
call_type = _get_call_type_from_endpoint(endpoint)
custom_id = json_obj.get("custom_id", f"line_{line_num}")
try:
# Run pre_call_hook (which triggers all guardrails)
await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_dict,
data=body,
call_type=call_type,
)
for callback in litellm.callbacks:
if isinstance(callback, CustomGuardrail):
if (
callback.should_run_guardrail(
data=body, event_type=GuardrailEventHooks.pre_call
)
is not True
):
continue
response = await callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=user_api_key_cache,
data=body,
call_type=call_type,
)
if response is not None and isinstance(response, dict):
body = response
elif (
isinstance(callback, CustomLogger)
and "async_pre_call_hook" in vars(callback.__class__)
and callback.__class__.async_pre_call_hook
!= CustomLogger.async_pre_call_hook
):
await callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=user_api_key_cache,
data=body,
call_type=call_type,
)
except Exception as e:
custom_id = json_obj.get("custom_id", f"line_{line_num}")
# Reraise exception to abort the upload, appending item information
raise Exception(f"Guardrail rejected batch item {custom_id} (line {line_num}): {str(e)}") from e
raise Exception(
f"Guardrail rejected batch item '{custom_id}' (line {line_num}): {str(e)}"
) from e

View file

@ -145,7 +145,7 @@ async def route_create_file(
) -> OpenAIFileObject:
"""
Route file creation request to the appropriate provider.
Priority:
1. If target_storage is specified and not "default" -> use storage backend
2. If model parameter provided -> use model credentials and encode ID
@ -153,7 +153,7 @@ async def route_create_file(
4. If enable_loadbalancing_on_batch_endpoints -> deprecated loadbalancing
5. Else -> use custom_llm_provider with files_settings
"""
# Handle custom storage backend
if target_storage and target_storage != "default":
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -162,7 +162,7 @@ async def route_create_file(
# Extract file data
file_data = extract_file_data(cast(Any, _create_file_request.get("file")))
# Use storage backend service to handle upload
file_object = await StorageBackendFileService.upload_file_to_storage_backend(
file_data=file_data,
@ -172,9 +172,9 @@ async def route_create_file(
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
)
return file_object
# NEW: Handle model-based routing (no DB required)
if model is not None:
# Get credentials from model_list via router
@ -183,19 +183,19 @@ async def route_create_file(
model_id=model,
operation_context="file upload",
)
# Merge credentials into the request
prepare_data_with_credentials(
data=_create_file_request, # type: ignore
credentials=credentials,
)
# Create the file with model credentials
response = await litellm.acreate_file(
**_create_file_request,
custom_llm_provider=credentials["custom_llm_provider"]
**_create_file_request,
custom_llm_provider=credentials["custom_llm_provider"],
) # type: ignore
# Encode the file ID with model information
if response and hasattr(response, "id") and response.id:
original_id = response.id
@ -204,9 +204,9 @@ async def route_create_file(
verbose_proxy_logger.debug(
f"Encoded file ID: {original_id} -> {encoded_id} (model: {model})"
)
return response
# Handle managed files (supports loadbalancing via llm_router.acreate_file)
# Priority: Check for managed files BEFORE deprecated loadbalancing
if target_model_names_list:
@ -339,7 +339,7 @@ async def create_file( # noqa: PLR0915
target_model_names_form=target_model_names,
target_storage_form=target_storage,
)
target_storage = file_params.target_storage
target_model_names_list = file_params.target_model_names
model_param = file_params.model
@ -358,18 +358,19 @@ async def create_file( # noqa: PLR0915
purpose = cast(OpenAIFilesPurpose, purpose)
data = {}
# Parse expires_after if provided
expires_after: Optional[FileExpiresAfter] = None
form_data_raw = await request.form()
form_data_dict: Dict[str, Any] = dict(form_data_raw)
extracted_litellm_metadata: Optional[Dict[str, Any]] = extract_nested_form_metadata(
form_data=form_data_dict,
prefix="litellm_metadata["
extracted_litellm_metadata: Optional[
Dict[str, Any]
] = extract_nested_form_metadata(
form_data=form_data_dict, prefix="litellm_metadata["
)
expires_after_anchor = form_data_raw.get("expires_after[anchor]")
expires_after_seconds_str = form_data_raw.get("expires_after[seconds]")
# Add litellm_metadata to data if provided (from form field)
if extracted_litellm_metadata is not None:
data["litellm_metadata"] = extracted_litellm_metadata
@ -382,7 +383,7 @@ async def create_file( # noqa: PLR0915
"error": "Both expires_after[anchor] and expires_after[seconds] must be provided if expires_after is specified",
},
)
# Validate expires_after[anchor] is a string (not UploadFile)
if isinstance(expires_after_anchor, UploadFile):
raise HTTPException(
@ -391,7 +392,7 @@ async def create_file( # noqa: PLR0915
"error": "expires_after[anchor] must be a string, not a file upload",
},
)
# Validate expires_after[seconds] is a string (not UploadFile)
# Use positive isinstance check for proper type narrowing (matches codebase pattern)
if not isinstance(expires_after_seconds_str, str):
@ -403,7 +404,7 @@ async def create_file( # noqa: PLR0915
)
# After this check, mypy knows expires_after_seconds_str is str
expires_after_seconds_str_validated: str = expires_after_seconds_str
# Validate anchor is "created_at"
if expires_after_anchor != "created_at":
raise HTTPException(
@ -412,7 +413,7 @@ async def create_file( # noqa: PLR0915
"error": f"expires_after[anchor] must be 'created_at', got '{expires_after_anchor}'",
},
)
# Convert seconds to int
try:
expires_after_seconds = int(expires_after_seconds_str_validated)
@ -423,7 +424,7 @@ async def create_file( # noqa: PLR0915
"error": f"expires_after[seconds] must be a valid integer, got '{expires_after_seconds_str}': {e}",
},
)
# Use literal "created_at" (not variable) for TypedDict to satisfy Literal type
expires_after = FileExpiresAfter(
anchor="created_at", # Literal, not expires_after_anchor variable
@ -455,10 +456,10 @@ async def create_file( # noqa: PLR0915
)
_create_file_request = CreateFileRequest(
file=file_data,
file=file_data,
purpose=cast(CREATE_FILE_REQUESTS_PURPOSE, purpose),
expires_after=expires_after,
**data
**data,
)
# Run pre-call guardrails on batch file content
@ -466,9 +467,10 @@ async def create_file( # noqa: PLR0915
from litellm.proxy.batches_endpoints.batch_guardrail_utils import (
run_pre_call_guardrails_on_batch_file,
)
await run_pre_call_guardrails_on_batch_file(
file_content=file_content,
proxy_logging_obj=proxy_logging_obj,
user_api_key_cache=proxy_logging_obj.call_details["user_api_key_cache"],
user_api_key_dict=user_api_key_dict,
)
@ -641,9 +643,11 @@ async def get_file_content( # noqa: PLR0915
param="None",
code=500,
)
# Check if file is stored in a storage backend (check DB)
if hasattr(managed_files_obj, "prisma_client") and getattr(managed_files_obj, "prisma_client", None):
if hasattr(managed_files_obj, "prisma_client") and getattr(
managed_files_obj, "prisma_client", None
):
prisma_client = getattr(managed_files_obj, "prisma_client")
db_file = await prisma_client.db.litellm_managedfiletable.find_first(
where={"unified_file_id": file_id}
@ -653,17 +657,18 @@ async def get_file_content( # noqa: PLR0915
from litellm.llms.base_llm.files.storage_backend_factory import (
get_storage_backend,
)
storage_backend_name = db_file.storage_backend
storage_url = db_file.storage_url
try:
# Get storage backend (uses same env vars as callback)
storage_backend = get_storage_backend(storage_backend_name)
file_content = await storage_backend.download_file(storage_url)
# Return file content
from fastapi.responses import Response as FastAPIResponse
return FastAPIResponse(
content=file_content,
media_type="application/octet-stream",
@ -675,7 +680,7 @@ async def get_file_content( # noqa: PLR0915
param="file_id",
code=400,
)
model = cast(Optional[str], data.get("model"))
if model:
response = await llm_router.afile_content(
@ -697,14 +702,19 @@ async def get_file_content( # noqa: PLR0915
)
else:
# Check for model-based credential routing
should_route, model_used, original_file_id, credentials = handle_model_based_routing(
(
should_route,
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
check_file_id_encoding=True,
)
if should_route:
# Use model-based routing with credentials from config
prepare_data_with_credentials(
@ -712,15 +722,19 @@ async def get_file_content( # noqa: PLR0915
credentials=credentials, # type: ignore
file_id=original_file_id, # Use decoded file ID if from encoded ID
)
response = await litellm.afile_content(
custom_llm_provider=credentials["custom_llm_provider"], # type: ignore
**data
**data,
) # type: ignore
verbose_proxy_logger.debug(
f"Retrieved file content using model: {model_used}"
+ (f", file_id: {file_id} -> {original_file_id}" if original_file_id else "")
+ (
f", file_id: {file_id} -> {original_file_id}"
if original_file_id
else ""
)
)
else:
# Fallback to default behavior (uses env variables or provider-based routing)
@ -838,7 +852,6 @@ async def get_file(
data: Dict = {"file_id": file_id}
try:
custom_llm_provider = (
provider
or get_custom_llm_provider_from_request_headers(request=request)
@ -864,15 +877,20 @@ async def get_file(
## Check for model-based credential routing
from litellm.proxy.proxy_server import llm_router
should_route, model_used, original_file_id, credentials = handle_model_based_routing(
(
should_route,
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
check_file_id_encoding=True,
)
if should_route:
# Use model-based routing with credentials from config
prepare_data_with_credentials(
@ -882,16 +900,21 @@ async def get_file(
)
response = await litellm.afile_retrieve(**data) # type: ignore
# Keep the encoded ID in response if it was originally encoded
if original_file_id and response and hasattr(response, "id") and response.id:
if (
original_file_id
and response
and hasattr(response, "id")
and response.id
):
response.id = file_id
verbose_proxy_logger.debug(
f"Retrieved file using model: {model_used}"
+ (f", original_id: {original_file_id}" if original_file_id else "")
)
## EXISTING: check if file_id is a litellm managed file
elif _is_base64_encoded_unified_file_id(file_id):
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
@ -1028,7 +1051,7 @@ async def delete_file(
or await get_custom_llm_provider_from_request_body(request=request)
or "openai"
)
# Call common_processing_pre_call_logic to trigger permission checks
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
@ -1043,7 +1066,7 @@ async def delete_file(
proxy_config=proxy_config,
route_type="afile_delete",
)
# Include original request and headers in the data
data = await add_litellm_data_to_request(
data=data,
@ -1055,14 +1078,19 @@ async def delete_file(
)
# Check for model-based credential routing
should_route, model_used, original_file_id, credentials = handle_model_based_routing(
(
should_route,
model_used,
original_file_id,
credentials,
) = handle_model_based_routing(
file_id=file_id,
request=request,
llm_router=llm_router,
data=data,
check_file_id_encoding=True,
)
if should_route:
# Use model-based routing with credentials from config
prepare_data_with_credentials(
@ -1070,14 +1098,14 @@ async def delete_file(
credentials=credentials, # type: ignore
file_id=original_file_id,
)
response = await litellm.afile_delete(**data) # type: ignore
verbose_proxy_logger.debug(
f"Deleted file using model: {model_used}"
+ (f", original_id: {original_file_id}" if original_file_id else "")
)
## EXISTING: check if file_id is a litellm managed file
elif _is_base64_encoded_unified_file_id(file_id):
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
@ -1230,7 +1258,7 @@ async def list_files(
)
response: Optional[Any] = None
# Check for model-based credential routing (no file_id encoding check for list)
should_route, model_used, _, credentials = handle_model_based_routing(
file_id="", # No file_id for list endpoint
@ -1239,18 +1267,18 @@ async def list_files(
data=data,
check_file_id_encoding=False,
)
if should_route:
# Use model-based routing with credentials from config
data.update(credentials) # type: ignore
response = await litellm.afile_list(
custom_llm_provider=credentials["custom_llm_provider"], # type: ignore
purpose=purpose,
**data # type: ignore
**data, # type: ignore
)
verbose_proxy_logger.debug(f"Listed files using model: {model_used}")
elif target_model_names and isinstance(target_model_names, str):
target_model_names_list = target_model_names.split(",")
if len(target_model_names_list) != 1: