From 45c8edbe25b8a1e021c546b01e570f0162e347eb Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 10 Feb 2026 12:25:26 +0530 Subject: [PATCH] Fix: Bug: Batch Rate Limiter Cannot Access User Files --- litellm/proxy/hooks/batch_rate_limiter.py | 6 +- tests/batches_tests/test_batch_rate_limits.py | 230 ++++++++++++++++++ .../proxy/hooks/test_managed_files.py | 2 +- 3 files changed, 236 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index ecad8bc1b1..78a371ad66 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -17,7 +17,7 @@ Quick summary: - async_log_success_event() fires on GET /v1/batches/{id} (batch completion) """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Union +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union from fastapi import HTTPException from pydantic import BaseModel @@ -241,6 +241,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, file_id: str, custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", + user_api_key_dict: Optional[UserAPIKeyAuth] = None, ) -> BatchFileUsage: """ Count number of requests and tokens in a batch input file. @@ -248,6 +249,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): Args: file_id: The file ID to read custom_llm_provider: The custom LLM provider to use for token encoding + user_api_key_dict: User authentication information for file access (required for managed files) Returns: BatchFileUsage with total_tokens and request_count @@ -257,6 +259,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): file_content = await litellm.afile_content( file_id=file_id, custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, ) file_content_as_dict = _get_file_content_as_dictionary( @@ -336,6 +339,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): batch_usage = await self.count_input_file_usage( file_id=input_file_id, custom_llm_provider=custom_llm_provider, + user_api_key_dict=user_api_key_dict, ) verbose_proxy_logger.debug( diff --git a/tests/batches_tests/test_batch_rate_limits.py b/tests/batches_tests/test_batch_rate_limits.py index 776aba438c..13241e94d5 100644 --- a/tests/batches_tests/test_batch_rate_limits.py +++ b/tests/batches_tests/test_batch_rate_limits.py @@ -389,3 +389,233 @@ async def test_batch_rate_limit_multiple_requests(): print(f" Error: {exc_info.value.detail}") finally: os.unlink(file_path_2) + + +@pytest.mark.asyncio() +@pytest.mark.skipif( + os.environ.get("OPENAI_API_KEY") is None, + reason="OPENAI_API_KEY not set - skipping integration test" +) +async def test_batch_rate_limiter_with_managed_files(): + """ + Test for GEN-2166: Verify batch rate limiter can read user files when managed files are enabled. + + This test ensures that: + 1. The batch rate limiter passes user_api_key_dict to afile_content() + 2. The managed files hook can verify file ownership correctly + 3. Rate limiting is enforced (not silently bypassed) + 4. No 403 Permission Denied errors occur for files owned by the user + """ + import tempfile + from unittest.mock import AsyncMock, MagicMock, patch + + CUSTOM_LLM_PROVIDER = "openai" + + # Setup: Create internal usage cache and rate limiter + dual_cache = DualCache() + internal_usage_cache = InternalUsageCache(dual_cache=dual_cache) + rate_limiter = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=internal_usage_cache + ) + + # Setup: Get batch rate limiter + batch_limiter = rate_limiter._get_batch_rate_limiter() + assert batch_limiter is not None, "Batch rate limiter should be available" + + # Setup: Create user API key with TPM = 500, RPM = 10 + test_user_id = "test-user-abc123" + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key-managed-files", + user_id=test_user_id, + tpm_limit=500, + rpm_limit=10, + ) + + print(f"\n=== Testing Batch Rate Limiter with Managed Files ===") + print(f"User ID: {test_user_id}") + + # Create a batch file with ~200 tokens + import json as json_lib + message = "This is a test message for batch rate limiting with managed files. " * 5 + requests = [] + for i in range(1, 4): + request_obj = { + "custom_id": f"request-{i}", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": message}] + } + } + requests.append(json_lib.dumps(request_obj)) + + batch_content = "\n".join(requests) + + with tempfile.NamedTemporaryFile(mode='w', suffix='.jsonl', delete=False) as f: + f.write(batch_content) + file_path = f.name + + try: + # Step 1: Upload file to OpenAI (simulating user upload) + print("\n1. Uploading batch input file...") + file_obj = await litellm.acreate_file( + file=open(file_path, "rb"), + purpose="batch", + custom_llm_provider=CUSTOM_LLM_PROVIDER, + ) + print(f" ✓ File uploaded: {file_obj.id}") + await asyncio.sleep(1) # Give API time to process + + # Step 2: Mock managed files hook to simulate file ownership check + # In a real scenario, the managed files hook would check if the user owns the file + # For this test, we'll verify that user_api_key_dict is passed correctly + print("\n2. Testing rate limiter file access with user context...") + + # Track if user_api_key_dict was passed to afile_content + original_afile_content = litellm.afile_content + user_context_passed = {"value": False} + + async def mock_afile_content(*args, **kwargs): + # Check if user_api_key_dict was passed + if "user_api_key_dict" in kwargs and kwargs["user_api_key_dict"] is not None: + user_context_passed["value"] = True + print(f" ✓ user_api_key_dict passed to afile_content") + print(f" User ID: {kwargs['user_api_key_dict'].user_id}") + else: + print(f" ✗ user_api_key_dict NOT passed to afile_content (BUG!)") + + # Call original function + return await original_afile_content(*args, **kwargs) + + # Patch afile_content to track the call + with patch('litellm.afile_content', side_effect=mock_afile_content): + data = { + "model": "gpt-3.5-turbo", + "input_file_id": file_obj.id, + "custom_llm_provider": CUSTOM_LLM_PROVIDER, + } + + # Step 3: Submit batch and verify rate limiting works + print("\n3. Submitting batch with rate limiting...") + result = await batch_limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=dual_cache, + data=data, + call_type="acreate_batch", + ) + + tokens_used = result.get('_batch_token_count', 0) + requests_count = result.get('_batch_request_count', 0) + print(f" ✓ Batch submitted successfully") + print(f" Tokens counted: {tokens_used}") + print(f" Requests counted: {requests_count}") + print(f" Rate limit usage: {tokens_used}/500 TPM, {requests_count}/10 RPM") + + # Step 4: Verify user context was passed + print("\n4. Verifying fix for GEN-2166...") + assert user_context_passed["value"], ( + "FAILED: user_api_key_dict was not passed to afile_content(). " + "This means the bug GEN-2166 is not fixed!" + ) + print(" ✓ Fix verified: user_api_key_dict is correctly passed") + + # Step 5: Verify rate limiting is actually enforced (not bypassed) + print("\n5. Verifying rate limiting is enforced...") + assert tokens_used > 0, "Token count should be greater than 0" + assert requests_count > 0, "Request count should be greater than 0" + print(" ✓ Rate limiting is active (not silently bypassed)") + + print("\n=== Test Passed: GEN-2166 Fix Verified ===") + print("✓ Batch rate limiter can access user files") + print("✓ User context is correctly passed") + print("✓ Rate limiting is enforced") + print("✓ No silent failures") + + except HTTPException as e: + if e.status_code == 403: + pytest.fail( + f"FAILED: Got 403 Permission Denied error. " + f"This indicates the bug GEN-2166 is not fixed. " + f"Error: {e.detail}" + ) + else: + raise + except Exception as e: + pytest.fail(f"Unexpected error: {str(e)}") + finally: + os.unlink(file_path) + + +@pytest.mark.asyncio() +async def test_batch_rate_limiter_without_user_context(): + """ + Test that verifies the bug scenario from GEN-2166. + + When user_api_key_dict is NOT passed to count_input_file_usage(), + the function should still work for non-managed files, but would fail + for managed files (which is the bug we fixed). + + This test documents the expected behavior with and without user context. + """ + import tempfile + + CUSTOM_LLM_PROVIDER = "openai" + + # Setup + BATCH_LIMITER = _PROXY_BatchRateLimiter( + internal_usage_cache=None, + parallel_request_limiter=None, + ) + + # Create a simple batch file + batch_content = """{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}}""" + + with tempfile.NamedTemporaryFile(mode='w', suffix='.jsonl', delete=False) as f: + f.write(batch_content) + file_path = f.name + + try: + # Upload file + file_obj = await litellm.acreate_file( + file=open(file_path, "rb"), + purpose="batch", + custom_llm_provider=CUSTOM_LLM_PROVIDER, + ) + await asyncio.sleep(1) + + # Test 1: Without user context (old behavior - would fail with managed files) + print("\n=== Test 1: count_input_file_usage WITHOUT user context ===") + try: + usage_without_context = await BATCH_LIMITER.count_input_file_usage( + file_id=file_obj.id, + custom_llm_provider=CUSTOM_LLM_PROVIDER, + user_api_key_dict=None, # Explicitly passing None + ) + print(f"✓ Works for non-managed files (tokens: {usage_without_context.total_tokens})") + print(" Note: Would fail with 403 for managed files (GEN-2166 bug)") + except Exception as e: + print(f"✗ Failed: {str(e)}") + + # Test 2: With user context (new behavior - works with managed files) + print("\n=== Test 2: count_input_file_usage WITH user context ===") + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user-123", + ) + + usage_with_context = await BATCH_LIMITER.count_input_file_usage( + file_id=file_obj.id, + custom_llm_provider=CUSTOM_LLM_PROVIDER, + user_api_key_dict=user_api_key_dict, # Passing user context + ) + print(f"✓ Works with user context (tokens: {usage_with_context.total_tokens})") + print(" Note: This fixes GEN-2166 for managed files") + + # Verify both return the same results + assert usage_with_context.total_tokens == usage_without_context.total_tokens + assert usage_with_context.request_count == usage_without_context.request_count + print("\n✓ Both methods return identical results for non-managed files") + + finally: + os.unlink(file_path) diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 8a22a43c37..ff957ed395 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -934,7 +934,7 @@ async def test_check_file_ids_access_with_unified_file_ids(): Test that check_file_ids_access validates user access to managed file IDs. """ from litellm.proxy._types import UserAPIKeyAuth - + # Create a unified file ID unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" regular_file_id = "file-abc123"