mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #14477 from BerriAI/litellm_dev_09_11_2025_p2
`/v1/messages` - don't send content block after message w/ finish reason + usage block + `/key/unblock` - support hashed tokens
This commit is contained in:
commit
663dbc6080
7 changed files with 659 additions and 130 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue