diff --git a/litellm/proxy/batches_endpoints/batch_guardrail_utils.py b/litellm/proxy/batches_endpoints/batch_guardrail_utils.py index 0c6f63118cf..fc938089634 100644 --- a/litellm/proxy/batches_endpoints/batch_guardrail_utils.py +++ b/litellm/proxy/batches_endpoints/batch_guardrail_utils.py @@ -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 diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index c7254553e1b..4bc74e479d7 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -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: