From 805069c287376cbc9ff82d08affb461bcbff36de Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 11 Sep 2025 18:52:02 -0700 Subject: [PATCH 1/3] fix(adapters/streaming_iterator.py): Don't send content block after message delta block is sent Fixes https://github.com/BerriAI/litellm/issues/14315 --- litellm/constants.py | 10 +- .../adapters/streaming_iterator.py | 133 +++---- litellm/proxy/_new_secret_config.yaml | 3 + .../hooks/parallel_request_limiter_v3.py | 42 ++- .../test_content_after_stop_reason.py | 343 ++++++++++++++++++ 5 files changed, 448 insertions(+), 83 deletions(-) create mode 100644 tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py diff --git a/litellm/constants.py b/litellm/constants.py index 75c25d9ea9e..c0ce0f265b5 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -15,7 +15,7 @@ DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int( os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10) ) DEFAULT_NUM_WORKERS_LITELLM_PROXY = int( - os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", os.cpu_count() or 4) + os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1) ) DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512)) SQS_SEND_MESSAGE_ACTION = "SendMessage" @@ -60,7 +60,9 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO", 128) ) DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE = int( - os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE", 512) + os.getenv( + "DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE", 512 + ) ) # Generic fallback for unknown models @@ -949,7 +951,9 @@ LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token" DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" PROMETHEUS_EMIT_BUDGET_METRICS_JOB_NAME = "prometheus_emit_budget_metrics" CLOUDZERO_EXPORT_USAGE_DATA_JOB_NAME = "cloudzero_export_usage_data" -CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int(os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000)) +CLOUDZERO_MAX_FETCHED_DATA_RECORDS = int( + os.getenv("CLOUDZERO_MAX_FETCHED_DATA_RECORDS", 50000) +) SPEND_LOG_CLEANUP_JOB_NAME = "spend_log_cleanup" SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500)) SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000)) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index aa95183bb6c..e4191a945f3 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -28,10 +28,6 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): TextBlock, ) - def __init__(self, completion_stream: Any, model: str): - super().__init__(completion_stream) - self.model = model - sent_first_chunk: bool = False sent_content_block_start: bool = False sent_content_block_finish: bool = False @@ -39,6 +35,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): sent_last_message: bool = False holding_chunk: Optional[Any] = None holding_stop_reason_chunk: Optional[Any] = None + queued_usage_chunk: bool = False current_content_block_index: int = 0 current_content_block_start: ContentBlockContentBlockDict = TextBlock( type="text", @@ -47,6 +44,10 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): pending_new_content_block: bool = False chunk_queue: deque = deque() # Queue for buffering multiple chunks + def __init__(self, completion_stream: Any, model: str): + super().__init__(completion_stream) + self.model = model + def __next__(self): from .transformation import LiteLLMAnthropicMessagesAdapter @@ -217,77 +218,83 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # Queue the merged chunk and reset self.chunk_queue.append(merged_chunk) + self.queued_usage_chunk = True self.holding_stop_reason_chunk = None return self.chunk_queue.popleft() # Check if this processed chunk has a stop_reason - hold it for next chunk - if should_start_new_block and not self.sent_content_block_finish: - # Queue the sequence: content_block_stop -> content_block_start -> current_chunk + if not self.queued_usage_chunk: + if should_start_new_block and not self.sent_content_block_finish: + # Queue the sequence: content_block_stop -> content_block_start -> current_chunk - # 1. Stop current content block - self.chunk_queue.append( - { - "type": "content_block_stop", - "index": max(self.current_content_block_index - 1, 0), - } - ) + # 1. Stop current content block + self.chunk_queue.append( + { + "type": "content_block_stop", + "index": max(self.current_content_block_index - 1, 0), + } + ) - # 2. Start new content block - self.chunk_queue.append( - { - "type": "content_block_start", - "index": self.current_content_block_index, - "content_block": self.current_content_block_start, - } - ) + # 2. Start new content block + self.chunk_queue.append( + { + "type": "content_block_start", + "index": self.current_content_block_index, + "content_block": self.current_content_block_start, + } + ) - # 3. Queue the current chunk (don't lose it!) - self.chunk_queue.append(processed_chunk) - - # Reset state for new block - self.sent_content_block_finish = False - - # Return the first queued item - return self.chunk_queue.popleft() - - if ( - processed_chunk["type"] == "message_delta" - and self.sent_content_block_finish is False - ): - # Queue both the content_block_stop and the holding chunk - self.chunk_queue.append( - { - "type": "content_block_stop", - "index": self.current_content_block_index, - } - ) - self.sent_content_block_finish = True - if processed_chunk.get("delta", {}).get("stop_reason") is not None: - - self.holding_stop_reason_chunk = processed_chunk - else: + # 3. Queue the current chunk (don't lose it!) self.chunk_queue.append(processed_chunk) - return self.chunk_queue.popleft() - elif self.holding_chunk is not None: - # Queue both chunks - self.chunk_queue.append(self.holding_chunk) - self.chunk_queue.append(processed_chunk) - self.holding_chunk = None - return self.chunk_queue.popleft() - else: - # Queue the current chunk - self.chunk_queue.append(processed_chunk) - return self.chunk_queue.popleft() + + # Reset state for new block + self.sent_content_block_finish = False + + # Return the first queued item + return self.chunk_queue.popleft() + + if ( + processed_chunk["type"] == "message_delta" + and self.sent_content_block_finish is False + ): + # Queue both the content_block_stop and the holding chunk + self.chunk_queue.append( + { + "type": "content_block_stop", + "index": self.current_content_block_index, + } + ) + self.sent_content_block_finish = True + if ( + processed_chunk.get("delta", {}).get("stop_reason") + is not None + ): + + self.holding_stop_reason_chunk = processed_chunk + else: + self.chunk_queue.append(processed_chunk) + return self.chunk_queue.popleft() + elif self.holding_chunk is not None: + # Queue both chunks + self.chunk_queue.append(self.holding_chunk) + self.chunk_queue.append(processed_chunk) + self.holding_chunk = None + return self.chunk_queue.popleft() + else: + # Queue the current chunk + self.chunk_queue.append(processed_chunk) + return self.chunk_queue.popleft() # Handle any remaining held chunks after stream ends - if self.holding_stop_reason_chunk is not None: - self.chunk_queue.append(self.holding_stop_reason_chunk) - self.holding_stop_reason_chunk = None + if not self.queued_usage_chunk: + if self.holding_stop_reason_chunk is not None: + self.chunk_queue.append(self.holding_stop_reason_chunk) + self.holding_stop_reason_chunk = None - if self.holding_chunk is not None: - self.chunk_queue.append(self.holding_chunk) - self.holding_chunk = None + if self.holding_chunk is not None: + self.chunk_queue.append(self.holding_chunk) + self.holding_chunk = None if not self.sent_last_message: self.sent_last_message = True diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index c785dd05c40..691f67d5778 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -7,3 +7,6 @@ model_list: - model_name: wildcard_models/* litellm_params: model: openai/* + - model_name: xai-grok-3 + litellm_params: + model: xai/grok-3 diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b3840761d2a..b3fa71a1c15 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -6,6 +6,7 @@ This is currently in development and not yet ready for production. import os from datetime import datetime +from math import floor from typing import ( TYPE_CHECKING, Any, @@ -17,7 +18,7 @@ from typing import ( Union, cast, ) -from math import floor + from fastapi import HTTPException from litellm import DualCache @@ -95,6 +96,7 @@ end return results """ + class RateLimitDescriptorRateLimitObject(TypedDict, total=False): requests_per_unit: Optional[int] tokens_per_unit: Optional[int] @@ -480,10 +482,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): }, ) ) - + # Team Member rate limits - if user_api_key_dict.user_id and (user_api_key_dict.team_member_rpm_limit is not None or user_api_key_dict.team_member_tpm_limit is not None): - team_member_value = f"{user_api_key_dict.team_id}:{user_api_key_dict.user_id}" + if user_api_key_dict.user_id and ( + user_api_key_dict.team_member_rpm_limit is not None + or user_api_key_dict.team_member_tpm_limit is not None + ): + team_member_value = ( + f"{user_api_key_dict.team_id}:{user_api_key_dict.user_id}" + ) descriptors.append( RateLimitDescriptor( key="team_member", @@ -557,13 +564,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Find which descriptor hit the limit for i, status in enumerate(response["statuses"]): if status["code"] == "OVER_LIMIT": - descriptor = descriptors[floor(i/2)] + descriptor = descriptors[floor(i / 2)] raise HTTPException( status_code=429, detail=f"Rate limit exceeded for {descriptor['key']}: {descriptor['value']}. Remaining: {status['limit_remaining']}", headers={ "retry-after": str(self.window_size), - "rate_limit_type": str(status["rate_limit_type"]) + "rate_limit_type": str(status["rate_limit_type"]), }, # Retry after 1 minute ) @@ -613,7 +620,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Check if script is available if self.token_increment_script is None: - verbose_proxy_logger.debug("TTL preservation script not available, using regular pipeline") + verbose_proxy_logger.debug( + "TTL preservation script not available, using regular pipeline" + ) await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( increment_list=pipeline_operations, litellm_parent_otel_span=parent_otel_span, @@ -628,7 +637,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): for op in pipeline_operations: # Convert None TTL to 0 for Lua script ttl_value = op["ttl"] if op["ttl"] is not None else 0 - + verbose_proxy_logger.debug( f"Executing TTL-preserving increment for key={op['key']}, " f"increment={op['increment_value']}, ttl={ttl_value}" @@ -693,16 +702,15 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) # Get metadata from kwargs - user_api_key = kwargs["litellm_params"]["metadata"].get("user_api_key") - user_api_key_user_id = kwargs["litellm_params"]["metadata"].get( - "user_api_key_user_id" + litellm_metadata = kwargs["litellm_params"]["metadata"] + if litellm_metadata is None: + return + user_api_key = litellm_metadata.get("user_api_key") + user_api_key_user_id = litellm_metadata.get("user_api_key_user_id") + user_api_key_team_id = litellm_metadata.get("user_api_key_team_id") + user_api_key_end_user_id = kwargs.get("user") or litellm_metadata.get( + "user_api_key_end_user_id" ) - user_api_key_team_id = kwargs["litellm_params"]["metadata"].get( - "user_api_key_team_id" - ) - user_api_key_end_user_id = kwargs.get("user") or kwargs["litellm_params"][ - "metadata" - ].get("user_api_key_end_user_id") model_group = get_model_group_from_litellm_kwargs(kwargs) # Get total tokens from response diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py new file mode 100644 index 00000000000..4a170d666f5 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_content_after_stop_reason.py @@ -0,0 +1,343 @@ +""" +Test for AnthropicStreamWrapper handling content blocks that exist after message_delta with stop_reason and usage. + +This tests the scenario where a streaming response includes: +1. Initial content blocks +2. A message_delta chunk with stop_reason and usage +3. Additional content blocks after the stop_reason + +The wrapper should properly handle this by: +- Holding the stop_reason chunk until usage is available +- Merging usage into the stop_reason chunk +- Properly managing content_block_stop/start events for subsequent content +""" + +import os +import sys +from typing import List + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.experimental_pass_through.adapters.streaming_iterator import ( + AnthropicStreamWrapper, +) +from litellm.types.utils import Delta, ModelResponse, StreamingChoices, Usage + + +class MockCompletionStreamWithContentAfterStopReason: + """Mock stream that simulates content blocks existing after message_delta with stop_reason and usage.""" + + def __init__(self): + self.responses = [ + # Initial text content + ModelResponse( + stream=True, + choices=[ + StreamingChoices( + delta=Delta(content="Hello"), index=0, finish_reason=None + ) + ], + ), + ModelResponse( + stream=True, + choices=[ + StreamingChoices( + delta=Delta(content=" world"), index=0, finish_reason=None + ) + ], + ), + # Message delta with stop_reason AND usage (this is how it actually comes from the API) + ModelResponse( + stream=True, + choices=[ + StreamingChoices( + delta=Delta(content=""), index=0, finish_reason="stop" + ) + ], + usage=Usage(prompt_tokens=230, completion_tokens=65, total_tokens=295), + ), + # Additional content after the stop_reason - this simulates the scenario + # where there might be additional content blocks after the main response + ModelResponse( + stream=True, + choices=[ + StreamingChoices( + delta=Delta(content=" Additional content"), + index=0, + finish_reason=None, + ) + ], + ), + ] + self.index = 0 + + def __iter__(self): + return self + + def __next__(self): + if self.index >= len(self.responses): + raise StopIteration + response = self.responses[self.index] + self.index += 1 + return response + + def __aiter__(self): + return self + + async def __anext__(self): + if self.index >= len(self.responses): + raise StopAsyncIteration + response = self.responses[self.index] + self.index += 1 + return response + + +def test_anthropic_stream_wrapper_content_after_stop_reason(): + """Test that AnthropicStreamWrapper properly handles content blocks after message_delta with stop_reason.""" + + wrapper = AnthropicStreamWrapper( + completion_stream=MockCompletionStreamWithContentAfterStopReason(), + model="claude-3", + ) + + chunks = [] + chunk_types = [] + + # Collect all chunks + for chunk in wrapper: + chunks.append(chunk) + chunk_types.append(chunk.get("type")) + + # Verify the expected sequence of chunk types + expected_types = [ + "message_start", # Initial message start + "content_block_start", # Start of first content block + "content_block_delta", # "Hello" + "content_block_delta", # " world" + "content_block_stop", # End of first content block due to stop_reason + "message_delta", # Stop reason with merged usage + "message_stop", # Final message stop + ] + + print(f"Actual chunk types: {chunk_types}") + print(f"Expected chunk types: {expected_types}") + + # Verify we have the expected number of chunks + assert len(chunk_types) >= len( + expected_types + ), f"Expected at least {len(expected_types)} chunks, got {len(chunk_types)}" + + # Verify key chunk types are present + assert "message_start" in chunk_types + assert "content_block_start" in chunk_types + assert "content_block_delta" in chunk_types + assert "content_block_stop" in chunk_types + assert "message_delta" in chunk_types + assert "message_stop" in chunk_types + + # Find the message_delta chunk with stop_reason + message_delta_chunk = None + for chunk in chunks: + if chunk.get("type") == "message_delta": + message_delta_chunk = chunk + break + + assert message_delta_chunk is not None, "message_delta chunk not found" + + # Verify that the message_delta chunk has both stop_reason and usage + delta = message_delta_chunk.get("delta", {}) + usage = message_delta_chunk.get("usage", {}) + + assert ( + delta.get("stop_reason") == "end_turn" + ), f"Expected stop_reason 'end_turn', got {delta.get('stop_reason')}" + assert ( + usage.get("input_tokens") == 230 + ), f"Expected input_tokens 230, got {usage.get('input_tokens')}" + assert ( + usage.get("output_tokens") == 65 + ), f"Expected output_tokens 65, got {usage.get('output_tokens')}" + + # Verify content_block_stop comes before message_delta + content_block_stop_index = None + message_delta_index = None + + for i, chunk_type in enumerate(chunk_types): + if chunk_type == "content_block_stop" and content_block_stop_index is None: + content_block_stop_index = i + elif chunk_type == "message_delta": + message_delta_index = i + + assert content_block_stop_index is not None, "content_block_stop not found" + assert message_delta_index is not None, "message_delta not found" + assert ( + content_block_stop_index < message_delta_index + ), "content_block_stop should come before message_delta" + + +@pytest.mark.asyncio +async def test_async_anthropic_stream_wrapper_content_after_stop_reason(): + """Test async version of AnthropicStreamWrapper handling content blocks after message_delta with stop_reason.""" + + wrapper = AnthropicStreamWrapper( + completion_stream=MockCompletionStreamWithContentAfterStopReason(), + model="claude-3", + ) + + chunks = [] + chunk_types = [] + + # Collect all chunks asynchronously + async for chunk in wrapper: + chunks.append(chunk) + chunk_types.append(chunk.get("type")) + + print(f"Async - Actual chunk types: {chunk_types}") + + # Verify key chunk types are present + assert "message_start" in chunk_types + assert "content_block_start" in chunk_types + assert "content_block_delta" in chunk_types + assert "content_block_stop" in chunk_types + assert "message_delta" in chunk_types + assert "message_stop" in chunk_types + + # Find the message_delta chunk with stop_reason + message_delta_chunk = None + for chunk in chunks: + if chunk.get("type") == "message_delta": + message_delta_chunk = chunk + break + + assert message_delta_chunk is not None, "message_delta chunk not found" + + # Verify that the message_delta chunk has both stop_reason and usage + delta = message_delta_chunk.get("delta", {}) + usage = message_delta_chunk.get("usage", {}) + + assert ( + delta.get("stop_reason") == "end_turn" + ), f"Expected stop_reason 'end_turn', got {delta.get('stop_reason')}" + assert ( + usage.get("input_tokens") == 230 + ), f"Expected input_tokens 230, got {usage.get('input_tokens')}" + assert ( + usage.get("output_tokens") == 65 + ), f"Expected output_tokens 65, got {usage.get('output_tokens')}" + + +def test_usage_merging_behavior(): + """Test that usage information is properly merged with stop_reason chunk.""" + + wrapper = AnthropicStreamWrapper( + completion_stream=MockCompletionStreamWithContentAfterStopReason(), + model="claude-3", + ) + + # Process chunks and look specifically for the usage merging behavior + chunks = [] + for chunk in wrapper: + chunks.append(chunk) + # If this is a message_delta with stop_reason, verify it has usage + if ( + chunk.get("type") == "message_delta" + and chunk.get("delta", {}).get("stop_reason") is not None + ): + + usage = chunk.get("usage", {}) + assert ( + usage.get("input_tokens") is not None + ), "Usage should be merged with stop_reason chunk" + assert ( + usage.get("output_tokens") is not None + ), "Usage should be merged with stop_reason chunk" + break + + +def test_sse_wrapper_with_content_after_stop_reason(): + """Test SSE wrapper formatting for the content after stop_reason scenario.""" + + wrapper = AnthropicStreamWrapper( + completion_stream=MockCompletionStreamWithContentAfterStopReason(), + model="claude-3", + ) + + # Get SSE formatted chunks + sse_chunks = [] + for chunk in wrapper.anthropic_sse_wrapper(): + sse_chunks.append(chunk) + if len(sse_chunks) >= 10: # Limit to avoid infinite loops in tests + break + + # Verify all chunks are properly formatted as bytes + for chunk in sse_chunks: + assert isinstance(chunk, bytes), "SSE chunks should be bytes" + + # Decode and verify SSE format + chunk_str = chunk.decode("utf-8") + lines = chunk_str.split("\n") + + # Should have event and data lines + assert any( + line.startswith("event: ") for line in lines + ), f"Missing event line in: {chunk_str}" + assert any( + line.startswith("data: ") for line in lines + ), f"Missing data line in: {chunk_str}" + + +@pytest.mark.asyncio +async def test_async_sse_wrapper_with_content_after_stop_reason(): + """Test async SSE wrapper formatting for the content after stop_reason scenario.""" + + wrapper = AnthropicStreamWrapper( + completion_stream=MockCompletionStreamWithContentAfterStopReason(), + model="claude-3", + ) + + # Get SSE formatted chunks asynchronously + sse_chunks = [] + async for chunk in wrapper.async_anthropic_sse_wrapper(): + sse_chunks.append(chunk) + if len(sse_chunks) >= 10: # Limit to avoid infinite loops in tests + break + + # Verify all chunks are properly formatted as bytes + for chunk in sse_chunks: + assert isinstance(chunk, bytes), "Async SSE chunks should be bytes" + + # Decode and verify SSE format + chunk_str = chunk.decode("utf-8") + lines = chunk_str.split("\n") + + # Should have event and data lines + assert any( + line.startswith("event: ") for line in lines + ), f"Missing event line in: {chunk_str}" + assert any( + line.startswith("data: ") for line in lines + ), f"Missing data line in: {chunk_str}" + + +if __name__ == "__main__": + # Run a quick test + test_anthropic_stream_wrapper_content_after_stop_reason() + print("✅ Sync test passed") + + import asyncio + + asyncio.run(test_async_anthropic_stream_wrapper_content_after_stop_reason()) + print("✅ Async test passed") + + test_usage_merging_behavior() + print("✅ Usage merging test passed") + + test_sse_wrapper_with_content_after_stop_reason() + print("✅ SSE wrapper test passed") + + asyncio.run(test_async_sse_wrapper_with_content_after_stop_reason()) + print("✅ Async SSE wrapper test passed") + + print("🎉 All tests passed!") From 0f6898ad0a9d8769bb26038f80107b309c8fb91e Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 11 Sep 2025 19:10:00 -0700 Subject: [PATCH 2/3] fix(key_management_endpoints.py): check if key is a hashed token or sk key before lookup Fixes https://github.com/BerriAI/litellm/issues/13887 --- .../key_management_endpoints.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index bd8faf34be8..a7f15cc6c19 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -346,6 +346,7 @@ def handle_key_type(data: GenerateKeyRequest, data_json: dict) -> dict: data_json["allowed_routes"] = ["info_routes"] return data_json + async def validate_team_id_used_in_service_account_request( team_id: Optional[str], prisma_client: Optional[PrismaClient], @@ -358,13 +359,13 @@ async def validate_team_id_used_in_service_account_request( status_code=400, detail="team_id is required for service account keys. Please specify `team_id` in the request body.", ) - + if prisma_client is None: raise HTTPException( status_code=400, detail="prisma_client is required for service account keys. Please specify `prisma_client` in the request body.", ) - + # check if team_id exists in the database team = await prisma_client.db.litellm_teamtable.find_unique( where={"team_id": team_id}, @@ -376,6 +377,7 @@ async def validate_team_id_used_in_service_account_request( ) return True + async def _common_key_generation_helper( # noqa: PLR0915 data: GenerateKeyRequest, user_api_key_dict: UserAPIKeyAuth, @@ -557,7 +559,7 @@ async def _common_key_generation_helper( # noqa: PLR0915 status_code=400, detail={ "error": f"Invalid key format. LiteLLM Virtual Key must start with 'sk-'. Received: {data.key}" - } + }, ) response = await generate_key_helper_fn( @@ -2885,7 +2887,10 @@ async def unblock_key( param="key", code=status.HTTP_400_BAD_REQUEST, ) - hashed_token = hash_token(token=data.key) + if data.key.startswith("sk-"): + hashed_token = hash_token(token=data.key) + else: + hashed_token = data.key if litellm.store_audit_logs is True: # make an audit log for key update From 0c8b311155f01f5075ca4deb09ac3a5b390ad2c6 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 11 Sep 2025 19:15:15 -0700 Subject: [PATCH 3/3] test: add unit testing for both flows on key unblock --- .../test_key_management_endpoints.py | 245 +++++++++++++++--- 1 file changed, 202 insertions(+), 43 deletions(-) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 3a597adef06..2747fab77fa 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -183,7 +183,9 @@ async def test_budget_reset_and_expires_at_first_of_month(monkeypatch): assert ( response_date.month == expected_month ), f"Expected month {expected_month}, got {response_date.month} for {key}" - assert response_date.day == 1, f"Expected day 1, got {response_date.day} for {key}" + assert ( + response_date.day == 1 + ), f"Expected day 1, got {response_date.day} for {key}" @pytest.mark.asyncio @@ -507,7 +509,6 @@ def test_get_new_token_with_invalid_key(): assert "New key must start with 'sk-'" in str(exc_info.value.detail) - @pytest.mark.asyncio async def test_generate_service_account_requires_team_id(): with pytest.raises(HTTPException): @@ -529,11 +530,12 @@ async def test_generate_service_account_works_with_team_id(): from unittest.mock import patch # Mock the database and router dependencies from proxy_server - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma, \ - patch('litellm.proxy.proxy_server.llm_router') as mock_router, \ - patch('litellm.proxy.proxy_server.premium_user', False), \ - patch('litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn') as mock_generate_key: - + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( + "litellm.proxy.proxy_server.llm_router" + ) as mock_router, patch("litellm.proxy.proxy_server.premium_user", False), patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn" + ) as mock_generate_key: + # Configure mocks mock_prisma.return_value = AsyncMock() mock_router.return_value = None @@ -542,9 +544,9 @@ async def test_generate_service_account_works_with_team_id(): "key": "sk-test-key", "expires": None, "user_id": "test-user", - "team_id": "IJ" + "team_id": "IJ", } - + # This should not raise an exception since team_id is provided await _common_key_generation_helper( data=GenerateKeyRequest( @@ -559,7 +561,6 @@ async def test_generate_service_account_works_with_team_id(): ) - @pytest.mark.asyncio async def test_update_service_account_requires_team_id(): data = UpdateKeyRequest(key="sk-1", metadata={"service_account_id": "sa"}) @@ -571,7 +572,9 @@ async def test_update_service_account_requires_team_id(): @pytest.mark.asyncio async def test_update_service_account_works_with_team_id(): - data = UpdateKeyRequest(key="sk-1", metadata={"service_account_id": "sa"}, team_id="IJ") + data = UpdateKeyRequest( + key="sk-1", metadata={"service_account_id": "sa"}, team_id="IJ" + ) existing_key = LiteLLM_VerificationToken(token="hashed") await prepare_key_update_data(data=data, existing_key_row=existing_key) @@ -580,22 +583,22 @@ async def test_update_service_account_works_with_team_id(): @pytest.mark.asyncio async def test_validate_team_id_used_in_service_account_request_requires_team_id(): """ - Test that validate_team_id_used_in_service_account_request raises HTTPException + Test that validate_team_id_used_in_service_account_request raises HTTPException when team_id is None for service account key generation. """ from litellm.proxy.management_endpoints.key_management_endpoints import ( validate_team_id_used_in_service_account_request, ) - + mock_prisma_client = AsyncMock() - + # Test that HTTPException is raised when team_id is None with pytest.raises(HTTPException) as exc_info: await validate_team_id_used_in_service_account_request( team_id=None, prisma_client=mock_prisma_client, ) - + assert exc_info.value.status_code == 400 assert "team_id is required for service account keys" in str(exc_info.value.detail) @@ -603,7 +606,7 @@ async def test_validate_team_id_used_in_service_account_request_requires_team_id @pytest.mark.asyncio async def test_validate_team_id_used_in_service_account_request_requires_prisma_client(): """ - Test that validate_team_id_used_in_service_account_request raises HTTPException + Test that validate_team_id_used_in_service_account_request raises HTTPException when prisma_client is None for service account key generation. """ from litellm.proxy.management_endpoints.key_management_endpoints import ( @@ -616,78 +619,76 @@ async def test_validate_team_id_used_in_service_account_request_requires_prisma_ team_id="test-team-id", prisma_client=None, ) - + assert exc_info.value.status_code == 400 - assert "prisma_client is required for service account keys" in str(exc_info.value.detail) + assert "prisma_client is required for service account keys" in str( + exc_info.value.detail + ) @pytest.mark.asyncio async def test_validate_team_id_used_in_service_account_request_checks_team_exists(): """ - Test that validate_team_id_used_in_service_account_request validates that + Test that validate_team_id_used_in_service_account_request validates that the team_id exists in the database for service account key generation. """ from litellm.proxy.management_endpoints.key_management_endpoints import ( validate_team_id_used_in_service_account_request, ) - + mock_prisma_client = AsyncMock() - + # Mock the database query to return None (team doesn't exist) mock_find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_teamtable.find_unique = mock_find_unique - + # Test that HTTPException is raised when team doesn't exist in DB with pytest.raises(HTTPException) as exc_info: await validate_team_id_used_in_service_account_request( team_id="non-existent-team-id", prisma_client=mock_prisma_client, ) - + assert exc_info.value.status_code == 400 assert "team_id does not exist in the database" in str(exc_info.value.detail) - + # Verify the database was queried with the correct parameters - mock_find_unique.assert_called_once_with( - where={"team_id": "non-existent-team-id"} - ) + mock_find_unique.assert_called_once_with(where={"team_id": "non-existent-team-id"}) @pytest.mark.asyncio async def test_validate_team_id_used_in_service_account_request_success(): """ - Test that validate_team_id_used_in_service_account_request returns True + Test that validate_team_id_used_in_service_account_request returns True when team_id exists in the database for service account key generation. """ from litellm.proxy.management_endpoints.key_management_endpoints import ( validate_team_id_used_in_service_account_request, ) - + mock_prisma_client = AsyncMock() - + # Mock the database query to return a team object (team exists) mock_team = {"team_id": "existing-team-id", "team_name": "Test Team"} mock_find_unique = AsyncMock(return_value=mock_team) mock_prisma_client.db.litellm_teamtable.find_unique = mock_find_unique - + # Test that function returns True when team exists result = await validate_team_id_used_in_service_account_request( team_id="existing-team-id", prisma_client=mock_prisma_client, ) - + assert result is True - + # Verify the database was queried with the correct parameters - mock_find_unique.assert_called_once_with( - where={"team_id": "existing-team-id"} - ) + mock_find_unique.assert_called_once_with(where={"team_id": "existing-team-id"}) @pytest.mark.asyncio async def test_generate_service_account_key_endpoint_validation(): """ - Test that the /key/service-account/generate endpoint properly validates + Test that the /key/service-account/generate endpoint properly validates team_id requirement and team existence in database. """ from unittest.mock import patch @@ -705,16 +706,16 @@ async def test_generate_service_account_key_endpoint_validation(): ), litellm_changed_by=None, ) - + assert exc_info.value.status_code == 400 assert "team_id is required for service account keys" in str(exc_info.value.detail) - - # Test case 2: Team doesn't exist in database - with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + + # Test case 2: Team doesn't exist in database + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # Mock team not found mock_find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_teamtable.find_unique = mock_find_unique - + with pytest.raises(HTTPException) as exc_info: await generate_service_account_key_fn( data=GenerateKeyRequest(team_id="non-existent-team"), @@ -723,7 +724,165 @@ async def test_generate_service_account_key_endpoint_validation(): ), litellm_changed_by=None, ) - + assert exc_info.value.status_code == 400 assert "team_id does not exist in the database" in str(exc_info.value.detail) + +@pytest.mark.asyncio +async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch): + """ + Test that the unblock_key endpoint correctly handles both sk- prefixed tokens + and hashed tokens by properly converting sk- tokens to hashed format before + database operations. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import BlockKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import unblock_key + + # Mock dependencies + mock_prisma_client = AsyncMock() + mock_user_api_key_cache = MagicMock() + mock_proxy_logging_obj = MagicMock() + + # Use a proper 64-character hex hash for testing + test_hashed_token = ( + "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + ) + + # Mock the key record that will be returned from database + mock_key_record = MagicMock() + mock_key_record.token = test_hashed_token + mock_key_record.blocked = False + mock_key_record.model_dump_json.return_value = ( + f'{{"token": "{test_hashed_token}", "blocked": false}}' + ) + + # Mock database operations + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_key_record + ) + mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock( + return_value=mock_key_record + ) + + # Mock get_key_object and _cache_key_object functions + mock_key_object = MagicMock() + mock_key_object.blocked = True # Initially blocked + + # Mock hash_token function + def mock_hash_token(token): + if token == "sk-test123456789": + return test_hashed_token + return token + + # Apply monkeypatch + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj + ) + monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token) + monkeypatch.setattr( + "litellm.store_audit_logs", False + ) # Disable audit logs for simpler test + + # Mock get_key_object and _cache_key_object + async def mock_get_key_object(**kwargs): + return mock_key_object + + async def mock_cache_key_object(**kwargs): + pass + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_key_object", + mock_get_key_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object", + mock_cache_key_object, + ) + + # Create mock request and user auth + mock_request = MagicMock() + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user" + ) + + # Test Case 1: Using sk- prefixed token + sk_token_request = BlockKeyRequest(key="sk-test123456789") + + result = await unblock_key( + data=sk_token_request, + http_request=mock_request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + # Verify that the database update was called with hashed token + mock_prisma_client.db.litellm_verificationtoken.update.assert_called_with( + where={"token": test_hashed_token}, data={"blocked": False} + ) + + assert result == mock_key_record + assert mock_key_object.blocked == False # Should be updated to unblocked + + # Reset mocks for second test + mock_prisma_client.db.litellm_verificationtoken.update.reset_mock() + mock_key_object.blocked = True # Reset to blocked state + + # Test Case 2: Using already hashed token + hashed_token_request = BlockKeyRequest(key=test_hashed_token) + + result = await unblock_key( + data=hashed_token_request, + http_request=mock_request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + # Verify that the database update was called with the same hashed token + mock_prisma_client.db.litellm_verificationtoken.update.assert_called_with( + where={"token": test_hashed_token}, data={"blocked": False} + ) + + assert result == mock_key_record + assert mock_key_object.blocked == False # Should be updated to unblocked + + +@pytest.mark.asyncio +async def test_unblock_key_invalid_key_format(monkeypatch): + """ + Test that unblock_key properly validates key format and raises appropriate errors + for invalid keys. + """ + from litellm.proxy._types import BlockKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import unblock_key + from litellm.proxy.utils import ProxyException + + # Mock prisma_client to avoid DB connection error + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + # Mock request and user auth + mock_request = MagicMock() + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user" + ) + + # Test with invalid key format + invalid_key_request = BlockKeyRequest(key="invalid-key-format") + + with pytest.raises(ProxyException) as exc_info: + await unblock_key( + data=invalid_key_request, + http_request=mock_request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert exc_info.value.code == "400" + assert "Invalid key format" in str(exc_info.value.message)