Add user auth in standard logging object for bedrock passthrough

This commit is contained in:
Sameer Kankute 2026-01-15 18:36:06 +05:30
parent dca42047b9
commit eb49adb201
2 changed files with 219 additions and 5 deletions

View File

@ -4338,6 +4338,38 @@ class StandardLoggingPayloadSetup:
return messages
@staticmethod
def merge_litellm_metadata(litellm_params: dict) -> dict:
"""
Merge both litellm_metadata and metadata from litellm_params.
litellm_metadata contains model-related fields, metadata contains user API key fields.
We need both for complete standard logging payload.
Args:
litellm_params: Dictionary containing metadata and litellm_metadata
Returns:
dict: Merged metadata with user API key fields taking precedence
"""
merged_metadata: dict = {}
# Start with metadata (user API key fields) - but skip non-serializable objects
if litellm_params.get("metadata") and isinstance(litellm_params.get("metadata"), dict):
for key, value in litellm_params["metadata"].items():
# Skip non-serializable objects like UserAPIKeyAuth
if key == "user_api_key_auth":
continue
merged_metadata[key] = value
# Then merge litellm_metadata (model-related fields) - this will NOT overwrite existing keys
if litellm_params.get("litellm_metadata") and isinstance(litellm_params.get("litellm_metadata"), dict):
for key, value in litellm_params["litellm_metadata"].items():
if key not in merged_metadata: # Don't overwrite existing keys from metadata
merged_metadata[key] = value
return merged_metadata
@staticmethod
def get_standard_logging_metadata(
metadata: Optional[Dict[str, Any]],
@ -5059,11 +5091,8 @@ def get_standard_logging_object_payload(
litellm_params = kwargs.get("litellm_params", {}) or {}
proxy_server_request = litellm_params.get("proxy_server_request") or {}
metadata: dict = (
litellm_params.get("litellm_metadata")
or litellm_params.get("metadata", None)
or {}
)
# Merge both litellm_metadata and metadata to get complete metadata
metadata: dict = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
completion_start_time = kwargs.get("completion_start_time", end_time)
call_type = kwargs.get("call_type")

View File

@ -703,3 +703,188 @@ def test_cost_breakdown_missing_in_standard_logging_payload():
assert payload["response_cost"] == 0.0001
print("✅ Cost breakdown missing test passed!")
def test_merge_litellm_metadata_basic():
"""
Test that merge_litellm_metadata correctly merges metadata and litellm_metadata.
User API key fields (from metadata) should take precedence over model-related fields (from litellm_metadata).
"""
litellm_params = {
"metadata": {
"user_api_key": "test-key-123",
"user_api_key_user_id": "user-456",
"user_api_key_team_id": "team-789",
},
"litellm_metadata": {
"model_group": "gpt-4-group",
"model_info": {"id": "model-123"},
"tags": ["tag1", "tag2"],
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# Check that user API key fields are present
assert result["user_api_key"] == "test-key-123"
assert result["user_api_key_user_id"] == "user-456"
assert result["user_api_key_team_id"] == "team-789"
# Check that model-related fields are present
assert result["model_group"] == "gpt-4-group"
assert result["model_info"] == {"id": "model-123"}
assert result["tags"] == ["tag1", "tag2"]
def test_merge_litellm_metadata_precedence():
"""
Test that metadata fields take precedence over litellm_metadata when there are conflicts.
"""
litellm_params = {
"metadata": {
"tags": ["user-tag1", "user-tag2"],
"custom_field": "from_metadata",
},
"litellm_metadata": {
"tags": ["model-tag1", "model-tag2"], # This should NOT overwrite
"custom_field": "from_litellm_metadata", # This should NOT overwrite
"model_group": "gpt-4-group", # This should be included
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# metadata values should take precedence
assert result["tags"] == ["user-tag1", "user-tag2"]
assert result["custom_field"] == "from_metadata"
# litellm_metadata values should only be included if not in metadata
assert result["model_group"] == "gpt-4-group"
def test_merge_litellm_metadata_skip_non_serializable():
"""
Test that non-serializable objects like UserAPIKeyAuth are skipped.
"""
from litellm.proxy._types import UserAPIKeyAuth
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="test-team",
)
litellm_params = {
"metadata": {
"user_api_key": "test-key-123",
"user_api_key_auth": user_api_key_auth, # This should be skipped
"safe_field": "safe_value",
},
"litellm_metadata": {
"model_group": "gpt-4-group",
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# user_api_key_auth should be skipped
assert "user_api_key_auth" not in result
# Other fields should be present
assert result["user_api_key"] == "test-key-123"
assert result["safe_field"] == "safe_value"
assert result["model_group"] == "gpt-4-group"
def test_merge_litellm_metadata_empty_params():
"""
Test that merge_litellm_metadata handles empty or missing metadata gracefully.
"""
# Test with empty litellm_params
result = StandardLoggingPayloadSetup.merge_litellm_metadata({})
assert result == {}
# Test with only metadata
litellm_params = {
"metadata": {
"user_api_key": "test-key",
}
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
assert result == {"user_api_key": "test-key"}
# Test with only litellm_metadata
litellm_params = {
"litellm_metadata": {
"model_group": "gpt-4-group",
}
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
assert result == {"model_group": "gpt-4-group"}
# Test with None values
litellm_params = {
"metadata": None,
"litellm_metadata": None,
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
assert result == {}
def test_merge_litellm_metadata_bedrock_passthrough_scenario():
"""
Test merge_litellm_metadata in a Bedrock passthrough scenario where both
user API key metadata and model metadata need to be merged.
This is the specific scenario that was fixed - bedrock passthrough requests
should include complete user authentication metadata in logging.
"""
litellm_params = {
"metadata": {
# User API key fields from authentication
"user_api_key": "sk-bedrock-test-key-123",
"user_api_key_hash": "hashed-key-123",
"user_api_key_user_id": "bedrock-user-456",
"user_api_key_team_id": "bedrock-team-789",
"user_api_key_org_id": "bedrock-org-101",
"user_api_key_alias": "bedrock-key-alias",
"user_api_key_team_alias": "bedrock-team-alias",
"user_api_key_end_user_id": "end-user-123",
"user_api_key_request_route": "/bedrock/model/invoke",
},
"litellm_metadata": {
# Model-related fields from Bedrock configuration
"model_group": "bedrock-claude-group",
"model_info": {
"id": "anthropic.claude-3-sonnet",
"mode": "chat",
},
"aws_region_name": "us-east-1",
"tags": ["production", "bedrock"],
},
}
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
# Verify all user API key fields are present
assert result["user_api_key"] == "sk-bedrock-test-key-123"
assert result["user_api_key_hash"] == "hashed-key-123"
assert result["user_api_key_user_id"] == "bedrock-user-456"
assert result["user_api_key_team_id"] == "bedrock-team-789"
assert result["user_api_key_org_id"] == "bedrock-org-101"
assert result["user_api_key_alias"] == "bedrock-key-alias"
assert result["user_api_key_team_alias"] == "bedrock-team-alias"
assert result["user_api_key_end_user_id"] == "end-user-123"
assert result["user_api_key_request_route"] == "/bedrock/model/invoke"
# Verify all model-related fields are present
assert result["model_group"] == "bedrock-claude-group"
assert result["model_info"] == {
"id": "anthropic.claude-3-sonnet",
"mode": "chat",
}
assert result["aws_region_name"] == "us-east-1"
assert result["tags"] == ["production", "bedrock"]
# Verify total number of fields (9 user fields + 4 model fields = 13)
assert len(result) == 13