From e9b7059af4d0aa0ad3da418628f34c1bd02251fa Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Thu, 22 May 2025 23:05:45 -0700 Subject: [PATCH] Litellm add file validation (#11081) * fix: cleanup print statement * feat(managed_files.py): add auth check on managed files Implemented for file retrieve + delete calls * feat(files_endpoints.py): support returning files by model name enables managed file support * feat(managed_files/): filter list of files by the ones created by user prevents user from seeing another file * test: update test * fix(files_endpoints.py): list_files - always default to provider based routing * build: add new table to prisma schema --- enterprise/enterprise_hooks/managed_files.py | 105 +++++++++++++++++- .../migration.sql | 32 ++++++ .../litellm_proxy_extras/schema.prisma | 19 +++- litellm/llms/base_llm/files/transformation.py | 2 + litellm/proxy/_new_secret_config.yaml | 6 +- litellm/proxy/_types.py | 3 + .../openai_files_endpoints/files_endpoints.py | 60 ++++++++-- litellm/proxy/schema.prisma | 3 + litellm/router.py | 5 + schema.prisma | 19 +++- .../enterprise_hooks/test_managed_files.py | 31 +++++- .../llms/azure/test_azure_common_utils.py | 1 + 12 files changed, 265 insertions(+), 21 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20250522223020_managed_object_table/migration.sql diff --git a/enterprise/enterprise_hooks/managed_files.py b/enterprise/enterprise_hooks/managed_files.py index 480ead78386..7410ad793b2 100644 --- a/enterprise/enterprise_hooks/managed_files.py +++ b/enterprise/enterprise_hooks/managed_files.py @@ -5,7 +5,7 @@ import asyncio import base64 import json import uuid -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast from fastapi import HTTPException @@ -26,8 +26,10 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( ) from litellm.types.llms.openai import ( AllMessageValues, + AsyncCursorPage, ChatCompletionFileObject, CreateFileRequest, + FileObject, OpenAIFileObject, OpenAIFilesPurpose, ) @@ -67,6 +69,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): file_object: OpenAIFileObject, litellm_parent_otel_span: Optional[Span], model_mappings: Dict[str, str], + user_api_key_dict: UserAPIKeyAuth, ) -> None: verbose_logger.info( f"Storing LiteLLM Managed File object with id={file_id} in cache" @@ -75,6 +78,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): unified_file_id=file_id, file_object=file_object, model_mappings=model_mappings, + flat_model_file_ids=list(model_mappings.values()), + created_by=user_api_key_dict.user_id, + updated_by=user_api_key_dict.user_id, ) await self.internal_usage_cache.async_set_cache( key=file_id, @@ -87,6 +93,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): "unified_file_id": file_id, "file_object": file_object.model_dump_json(), "model_mappings": json.dumps(model_mappings), + "flat_model_file_ids": list(model_mappings.values()), + "created_by": user_api_key_dict.user_id, + "updated_by": user_api_key_dict.user_id, } ) @@ -169,6 +178,18 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) return initial_value.file_object + async def can_user_call_unified_file_id( + self, unified_file_id: str, user_api_key_dict: UserAPIKeyAuth + ) -> bool: + ## check if the user has access to the unified file id + user_id = user_api_key_dict.user_id + managed_file = await self.prisma_client.db.litellm_managedfiletable.find_first( + where={"unified_file_id": unified_file_id} + ) + if managed_file: + return managed_file.created_by == user_id + return False + async def can_user_call_unified_object_id( self, unified_object_id: str, user_api_key_dict: UserAPIKeyAuth ) -> bool: @@ -184,6 +205,44 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return managed_object.created_by == user_id return False + async def get_user_created_file_ids( + self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str] + ) -> List[OpenAIFileObject]: + """ + Get all file ids created by the user for a list of model object ids + + Returns: + - List of OpenAIFileObject's + """ + file_ids = await self.prisma_client.db.litellm_managedfiletable.find_many( + where={ + "created_by": user_api_key_dict.user_id, + "flat_model_file_ids": {"hasSome": model_object_ids}, + } + ) + return [OpenAIFileObject(**file_object.file_object) for file_object in file_ids] + + async def check_managed_file_id_access( + self, data: Dict, user_api_key_dict: UserAPIKeyAuth + ) -> bool: + retrieve_file_id = cast(Optional[str], data.get("file_id")) + potential_file_id = ( + _is_base64_encoded_unified_file_id(retrieve_file_id) + if retrieve_file_id + else False + ) + if potential_file_id and retrieve_file_id: + if await self.can_user_call_unified_file_id( + retrieve_file_id, user_api_key_dict + ): + return True + else: + raise HTTPException( + status_code=403, + detail=f"User {user_api_key_dict.user_id} does not have access to the file {retrieve_file_id}", + ) + return False + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -200,6 +259,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): "rerank", "acreate_batch", "aretrieve_batch", + "acreate_file", + "afile_list", + "afile_delete", "afile_content", "acreate_fine_tuning_job", "aretrieve_fine_tuning_job", @@ -211,9 +273,14 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): - 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( - "CALLS ASYNC PRE CALL HOOK - DATA={}, CALL_TYPE={}".format(data, call_type) - ) + ### HANDLE FILE ACCESS ### - ensure user has access to the file + if ( + call_type == CallTypes.afile_content.value + or call_type == CallTypes.afile_delete.value + ): + await self.check_managed_file_id_access(data, user_api_key_dict) + + ### HANDLE TRANSFORMATIONS ### if call_type == CallTypes.completion.value: messages = data.get("messages") if messages: @@ -298,7 +365,6 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): [input_file_id], user_api_key_dict.parent_otel_span ) - print("DATA={}".format(data)) return data async def async_pre_call_deployment_hook( @@ -416,6 +482,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): llm_router: Router, target_model_names_list: List[str], litellm_parent_otel_span: Span, + user_api_key_dict: UserAPIKeyAuth, ) -> OpenAIFileObject: responses = await self.create_file_for_each_model( llm_router=llm_router, @@ -448,6 +515,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): file_object=response, litellm_parent_otel_span=litellm_parent_otel_span, model_mappings=model_mappings, + user_api_key_dict=user_api_key_dict, ) return response @@ -560,6 +628,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): async def async_post_call_success_hook( self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes ) -> Any: + print(f"response: {response}, type: {type(response)}") if isinstance(response, LiteLLMBatch): ## Check if unified_file_id is in the response unified_file_id = response._hidden_params.get( @@ -619,6 +688,31 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): user_api_key_dict=user_api_key_dict, ) ) + elif isinstance(response, AsyncCursorPage): + """ + For listing files, filter for the ones created by the user + """ + print("INSIDE ASYNC CURSOR PAGE BLOCK") + ## check if file object + if hasattr(response, "data") and isinstance(response.data, list): + if all( + isinstance(file_object, FileObject) for file_object in response.data + ): + ## Get all file id's + ## Check which file id's were created by the user + ## Filter the response to only include the files created by the user + ## Return the filtered response + file_ids = [ + file_object.id + for file_object in cast(List[FileObject], response.data) # type: ignore + ] + user_created_file_ids = await self.get_user_created_file_ids( + user_api_key_dict, file_ids + ) + ## Filter the response to only include the files created by the user + response.data = user_created_file_ids # type: ignore + return response + return response return response async def afile_retrieve( @@ -638,6 +732,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): litellm_parent_otel_span: Optional[Span], **data: Dict, ) -> List[OpenAIFileObject]: + """Handled in files_endpoints.py""" return [] async def afile_delete( diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20250522223020_managed_object_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20250522223020_managed_object_table/migration.sql new file mode 100644 index 00000000000..95fb8372458 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20250522223020_managed_object_table/migration.sql @@ -0,0 +1,32 @@ +-- AlterTable +ALTER TABLE "LiteLLM_ManagedFileTable" ADD COLUMN "created_by" TEXT, +ADD COLUMN "flat_model_file_ids" TEXT[] DEFAULT ARRAY[]::TEXT[], +ADD COLUMN "updated_by" TEXT; + +-- CreateTable +CREATE TABLE "LiteLLM_ManagedObjectTable" ( + "id" TEXT NOT NULL, + "unified_object_id" TEXT NOT NULL, + "model_object_id" TEXT NOT NULL, + "file_object" JSONB NOT NULL, + "file_purpose" TEXT NOT NULL, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT, + "updated_at" TIMESTAMP(3) NOT NULL, + "updated_by" TEXT, + + CONSTRAINT "LiteLLM_ManagedObjectTable_pkey" PRIMARY KEY ("id") +); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_ManagedObjectTable_unified_object_id_key" ON "LiteLLM_ManagedObjectTable"("unified_object_id"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_ManagedObjectTable_model_object_id_key" ON "LiteLLM_ManagedObjectTable"("model_object_id"); + +-- CreateIndex +CREATE INDEX "LiteLLM_ManagedObjectTable_unified_object_id_idx" ON "LiteLLM_ManagedObjectTable"("unified_object_id"); + +-- CreateIndex +CREATE INDEX "LiteLLM_ManagedObjectTable_model_object_id_idx" ON "LiteLLM_ManagedObjectTable"("model_object_id"); + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 1d6f3b52118..58064abd1dc 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -453,13 +453,30 @@ model LiteLLM_ManagedFileTable { id String @id @default(uuid()) unified_file_id String @unique // The base64 encoded unified file ID file_object Json // Stores the OpenAIFileObject - model_mappings Json // Stores the mapping of model_id -> provider_file_id + model_mappings Json + flat_model_file_ids String[] @default([]) // Flat list of model file id's - for faster querying of model id -> unified file id created_at DateTime @default(now()) + created_by String? updated_at DateTime @updatedAt + updated_by String? @@index([unified_file_id]) } +model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use the + id String @id @default(uuid()) + unified_object_id String @unique // The base64 encoded unified file ID + model_object_id String @unique // the id returned by the backend API provider + file_object Json // Stores the OpenAIFileObject + file_purpose String // either 'batch' or 'fine-tune' + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @updatedAt + updated_by String? + + @@index([unified_object_id]) + @@index([model_object_id]) +} model LiteLLM_ManagedVectorStoresTable { vector_store_id String @id diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py index 4d749af21e1..38a6dc48092 100644 --- a/litellm/llms/base_llm/files/transformation.py +++ b/litellm/llms/base_llm/files/transformation.py @@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import httpx +from litellm.proxy._types import UserAPIKeyAuth from litellm.types.llms.openai import ( AllMessageValues, CreateFileRequest, @@ -115,6 +116,7 @@ class BaseFileEndpoints(ABC): llm_router: Router, target_model_names_list: List[str], litellm_parent_otel_span: Span, + user_api_key_dict: UserAPIKeyAuth, ) -> OpenAIFileObject: pass diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index eac4a69e03c..a255f32f5e3 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -2,10 +2,10 @@ model_list: - model_name: "gemini-2.0-flash-gemini" litellm_params: model: gemini/gemini-2.0-flash - - model_name: "gpt-4o-mini-openai" + - model_name: "gpt-4.1-openai" litellm_params: - model: gpt-4.1-mini-2025-04-14 - api_key: os.environ/OPENAI_API_KEY_2 + model: gpt-4.1 + api_key: os.environ/OPENAI_API_KEY model_info: access_groups: ["default-openai-models"] - model_name: "gpt-4o-realtime-preview" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 1a5af39188b..c1feb3dc330 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2883,6 +2883,9 @@ class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): unified_file_id: str file_object: OpenAIFileObject model_mappings: Dict[str, str] + flat_model_file_ids: List[str] + created_by: Optional[str] + updated_by: Optional[str] class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 3c2c3d80dcb..79345021490 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -179,6 +179,7 @@ async def route_create_file( create_file_request=_create_file_request, target_model_names_list=target_model_names_list, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + user_api_key_dict=user_api_key_dict, ) else: # get configs for custom_llm_provider @@ -869,6 +870,7 @@ async def list_files( fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), provider: Optional[str] = None, + target_model_names: Optional[str] = None, purpose: Optional[str] = None, ): """ @@ -885,8 +887,8 @@ async def list_files( ``` """ from litellm.proxy.proxy_server import ( - add_litellm_data_to_request, general_settings, + llm_router, proxy_config, proxy_logging_obj, version, @@ -894,24 +896,62 @@ async def list_files( data: Dict = {} try: - custom_llm_provider = ( - provider - or await get_custom_llm_provider_from_request_body(request=request) - or "openai" - ) # 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=CallTypes.alist_fine_tuning_jobs.value, ) - response = await litellm.afile_list( - custom_llm_provider=custom_llm_provider, purpose=purpose, **data # type: ignore + response: Optional[Any] = None + if target_model_names and isinstance(target_model_names, str): + target_model_names_list = target_model_names.split(",") + if len(target_model_names_list) != 1: + raise HTTPException( + status_code=400, + detail="target_model_names on list files must be a list of one model name. Example: ['gpt-4o']", + ) + ## Use router to list fine-tuning jobs for that model + if llm_router is None: + raise HTTPException( + status_code=500, + detail="LLM Router not initialized. Ensure models added to proxy.", + ) + data["model"] = target_model_names_list[0] + response = await llm_router.afile_list( + **data, + ) + else: + custom_llm_provider = ( + provider + or await get_custom_llm_provider_from_request_body(request=request) + or "openai" + ) + + response = await litellm.afile_list( + custom_llm_provider=custom_llm_provider, purpose=purpose, **data # type: ignore + ) + + if response is None: + raise HTTPException( + status_code=500, + detail="Either 'provider' or 'target_model_names' must be provided e.g. `?target_model_names=gpt-4o`", + ) + + ## POST CALL HOOKS ### + _response = await proxy_logging_obj.post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response ) + if _response is not None and isinstance(_response, OpenAIFileObject): + response = _response ### ALERTING ### asyncio.create_task( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index e97dc7d2ae1..58064abd1dc 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -454,8 +454,11 @@ model LiteLLM_ManagedFileTable { unified_file_id String @unique // The base64 encoded unified file ID file_object Json // Stores the OpenAIFileObject model_mappings Json + flat_model_file_ids String[] @default([]) // Flat list of model file id's - for faster querying of model id -> unified file id created_at DateTime @default(now()) + created_by String? updated_at DateTime @updatedAt + updated_by String? @@index([unified_file_id]) } diff --git a/litellm/router.py b/litellm/router.py index 6556791a84c..0dd52234000 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -764,6 +764,9 @@ class Router: self.aretrieve_fine_tuning_job = self.factory_function( litellm.aretrieve_fine_tuning_job, call_type="aretrieve_fine_tuning_job" ) + self.afile_list = self.factory_function( + litellm.afile_list, call_type="alist_files" + ) def validate_fallbacks(self, fallback_param: Optional[List]): """ @@ -3185,6 +3188,7 @@ class Router: "acancel_fine_tuning_job", "alist_fine_tuning_jobs", "aretrieve_fine_tuning_job", + "alist_files", ] = "assistants", ): """ @@ -3237,6 +3241,7 @@ class Router: "acancel_fine_tuning_job", "alist_fine_tuning_jobs", "aretrieve_fine_tuning_job", + "alist_files", ): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, diff --git a/schema.prisma b/schema.prisma index 1d6f3b52118..b415a777359 100644 --- a/schema.prisma +++ b/schema.prisma @@ -453,13 +453,30 @@ model LiteLLM_ManagedFileTable { id String @id @default(uuid()) unified_file_id String @unique // The base64 encoded unified file ID file_object Json // Stores the OpenAIFileObject - model_mappings Json // Stores the mapping of model_id -> provider_file_id + model_mappings Json + flat_model_file_ids String[] @default([]) // Flat list of model file id's - for faster querying of model id -> unified file id created_at DateTime @default(now()) + created_by String? updated_at DateTime @updatedAt + updated_by String? @@index([unified_file_id]) } +model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use managed files + id String @id @default(uuid()) + unified_object_id String @unique // The base64 encoded unified object ID + model_object_id String @unique // the id returned by the backend API provider + file_object Json // Stores the OpenAIFileObject + file_purpose String // either 'batch' or 'fine-tune' + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @updatedAt + updated_by String? + + @@index([unified_object_id]) + @@index([model_object_id]) +} model LiteLLM_ManagedVectorStoresTable { vector_store_id String @id diff --git a/tests/enterprise/enterprise_hooks/test_managed_files.py b/tests/enterprise/enterprise_hooks/test_managed_files.py index 04a2717f788..81e27d941fa 100644 --- a/tests/enterprise/enterprise_hooks/test_managed_files.py +++ b/tests/enterprise/enterprise_hooks/test_managed_files.py @@ -3,13 +3,14 @@ import os import sys import pytest +from fastapi import HTTPException from fastapi.testclient import TestClient sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch from enterprise.enterprise_hooks.managed_files import _PROXY_LiteLLMManagedFiles from litellm.caching import DualCache @@ -240,3 +241,31 @@ async def test_async_pre_call_hook_for_unified_finetuning_job(): response = await proxy_managed_files.async_pre_call_hook(**data) assert response["fine_tuning_job_id"] == "ftjob-jTBys7bVsbyZDOwL9GlpYqXR" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("call_type", ["afile_content", "afile_delete"]) +async def test_can_user_call_unified_file_id(call_type): + """ + Test that on file retrieve, delete we check if the user has access to the file + """ + from litellm.proxy._types import UserAPIKeyAuth + + prisma_client = AsyncMock() + return_value = MagicMock() + return_value.created_by = "123" + prisma_client.db.litellm_managedfiletable.find_first.return_value = return_value + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + MagicMock(), prisma_client=prisma_client + ) + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9vY3RldC1zdHJlYW07dW5pZmllZF9pZCxmMTNlNDAzZS01YWM3LTRhZjktOGQzNS0wNDgwZDMxOTgyYTg7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00by1taW5pLW9wZW5haTtsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1Ib3UxZDFXc3c1SDNKcjFMYllpZDJiO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxmODBiNWU2NzQ1NzdkNjkyMjM4YmVhNTIxZDdiMGI5ZGYyY2FmMTEwMTU2YmU5YzBjM2NjMmNkNTBjOTM1ZDI0" + + with pytest.raises(HTTPException) as e: + await proxy_managed_files.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth( + user_id="456", parent_otel_span=MagicMock() + ), + cache=MagicMock(), + data={"file_id": unified_file_id}, + call_type=call_type, + ) diff --git a/tests/litellm/llms/azure/test_azure_common_utils.py b/tests/litellm/llms/azure/test_azure_common_utils.py index 03d9d252198..34f0dc3a973 100644 --- a/tests/litellm/llms/azure/test_azure_common_utils.py +++ b/tests/litellm/llms/azure/test_azure_common_utils.py @@ -395,6 +395,7 @@ def test_select_azure_base_url_called(setup_mocks): "acancel_fine_tuning_job", "alist_fine_tuning_jobs", "aretrieve_fine_tuning_job", + "afile_list", ] ], )