mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(spend_tracking): populate api_base in SpendLogs for embedding calls (#23770)
* fix(spend_tracking): populate api_base in SpendLogs for embedding calls Previously, embedding calls (call_type=aembedding) had an empty api_base in LiteLLM_SpendLogs because the Azure and OpenAI embedding handlers did not pass api_base in the logging_obj.pre_call() additional_args dict. This made it impossible to build per-endpoint monitoring dashboards for embedding models. Changes: - Azure embedding: add api_base to pre_call additional_args - Azure aembedding: add api_base to post_call additional_args (success + error) - OpenAI aembedding: add api_base to post_call additional_args (success + error) - Cohere sync embedding: add api_base (embed_url) to pre_call additional_args - Add regression tests for api_base in embedding SpendLogs payloads Fixes #23768 Related to #7317 * fix: address review feedback - complete api_base coverage and add provider-level tests - OpenAI sync embedding() post_call: add api_base to additional_args - Cohere async embedding error paths: add api_base to both post_call sites - Replace payload-only tests with provider-level tests that exercise the actual pre_call flow via a real Logging object, verifying that api_base is written into litellm_params['api_base'] by _pre_call() - Run black formatter on changed files * docs: add inline comments clarifying post_call api_base additions post_call only updates model_call_details['additional_args'], not litellm_params['api_base'] (which is what SpendLogs reads). The post_call additions are for consistency so logging integrations reading additional_args directly see api_base. The actual SpendLogs fix is the pre_call change.
This commit is contained in:
parent
afb6fbefdc
commit
f43aae6b4f
4 changed files with 340 additions and 142 deletions
|
|
@ -714,10 +714,13 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
stringified_response = response.model_dump()
|
||||
|
||||
## LOGGING
|
||||
# api_base included for consistency: post_call overwrites
|
||||
# model_call_details["additional_args"] which some logging
|
||||
# integrations read directly. The SpendLogs fix is in pre_call.
|
||||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
embedding_response = convert_to_model_response_object(
|
||||
|
|
@ -738,7 +741,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=str(e),
|
||||
)
|
||||
raise e
|
||||
|
|
@ -776,6 +779,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": {"api_key": api_key, "azure_ad_token": azure_ad_token},
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -89,7 +89,7 @@ async def async_embedding(
|
|||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=e.response.text,
|
||||
)
|
||||
raise e
|
||||
|
|
@ -98,7 +98,7 @@ async def async_embedding(
|
|||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=str(e),
|
||||
)
|
||||
raise e
|
||||
|
|
@ -164,7 +164,7 @@ def embedding(
|
|||
logging_obj.pre_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": embed_url},
|
||||
)
|
||||
|
||||
## COMPLETION CALL
|
||||
|
|
|
|||
|
|
@ -562,9 +562,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
kwargs_with_provider = (
|
||||
litellm_params.copy() if litellm_params else {}
|
||||
)
|
||||
kwargs_with_provider[
|
||||
"custom_llm_provider"
|
||||
] = custom_llm_provider
|
||||
kwargs_with_provider["custom_llm_provider"] = (
|
||||
custom_llm_provider
|
||||
)
|
||||
|
||||
# For OpenAI Chat Completions, use the chat completion agentic loop method
|
||||
agentic_response = (
|
||||
|
|
@ -1267,10 +1267,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
logging_obj.model_call_details["response_headers"] = headers
|
||||
stringified_response = response.model_dump()
|
||||
## LOGGING
|
||||
# api_base included for consistency: post_call overwrites
|
||||
# model_call_details["additional_args"] which some logging
|
||||
# integrations read directly. The SpendLogs fix is in pre_call.
|
||||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=stringified_response,
|
||||
)
|
||||
returned_response: EmbeddingResponse = convert_to_model_response_object(
|
||||
|
|
@ -1285,7 +1288,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=str(e),
|
||||
)
|
||||
raise e
|
||||
|
|
@ -1294,7 +1297,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=str(e),
|
||||
)
|
||||
status_code = getattr(e, "status_code", 500)
|
||||
|
|
@ -1368,11 +1371,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
) # type: ignore
|
||||
|
||||
## LOGGING
|
||||
# api_base: consistency with pre_call; post_call only updates
|
||||
# additional_args, not litellm_params (see aembedding comment).
|
||||
logging_obj.model_call_details["response_headers"] = headers
|
||||
logging_obj.post_call(
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": data},
|
||||
additional_args={"complete_input_dict": data, "api_base": api_base},
|
||||
original_response=sync_embedding_response,
|
||||
)
|
||||
response: EmbeddingResponse = convert_to_model_response_object(
|
||||
|
|
|
|||
|
|
@ -53,21 +53,23 @@ def test_sanitize_request_body_for_spend_logs_payload_long_string():
|
|||
from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
|
||||
# Create a string longer than MAX_STRING_LENGTH_PROMPT_IN_DB (2048)
|
||||
long_string = "a" * 3000 # Create a string longer than MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
long_string = (
|
||||
"a" * 3000
|
||||
) # Create a string longer than MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
request_body = {"text": long_string, "normal_text": "short text"}
|
||||
sanitized = _sanitize_request_body_for_spend_logs_payload(request_body)
|
||||
|
||||
|
||||
# Calculate expected lengths: 35% start + 65% end + truncation message
|
||||
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.35)
|
||||
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.65)
|
||||
total_keep = start_chars + end_chars
|
||||
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
|
||||
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
|
||||
|
||||
|
||||
skipped_chars = len(long_string) - (start_chars + end_chars)
|
||||
expected_truncation_message = f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. {LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
|
||||
expected_length = start_chars + len(expected_truncation_message) + end_chars
|
||||
|
||||
|
||||
assert len(sanitized["text"]) == expected_length
|
||||
assert sanitized["text"].startswith("a" * start_chars)
|
||||
assert sanitized["text"].endswith("a" * end_chars)
|
||||
|
|
@ -82,18 +84,18 @@ def test_sanitize_request_body_for_spend_logs_payload_nested_dict():
|
|||
long_string = "a" * (MAX_STRING_LENGTH_PROMPT_IN_DB + 500)
|
||||
request_body = {"outer": {"inner": {"text": long_string, "normal": "short"}}}
|
||||
sanitized = _sanitize_request_body_for_spend_logs_payload(request_body)
|
||||
|
||||
|
||||
# Calculate expected lengths based on actual MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.35)
|
||||
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.65)
|
||||
total_keep = start_chars + end_chars
|
||||
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
|
||||
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
|
||||
|
||||
|
||||
skipped_chars = len(long_string) - total_keep
|
||||
expected_truncation_message = f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. {LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
|
||||
expected_length = start_chars + len(expected_truncation_message) + end_chars
|
||||
|
||||
|
||||
assert len(sanitized["outer"]["inner"]["text"]) == expected_length
|
||||
assert sanitized["outer"]["inner"]["normal"] == "short"
|
||||
|
||||
|
|
@ -107,18 +109,18 @@ def test_sanitize_request_body_for_spend_logs_payload_nested_list():
|
|||
"items": [{"text": long_string}, {"text": "short"}, [{"text": long_string}]]
|
||||
}
|
||||
sanitized = _sanitize_request_body_for_spend_logs_payload(request_body)
|
||||
|
||||
|
||||
# Calculate expected lengths based on actual MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.35)
|
||||
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.65)
|
||||
total_keep = start_chars + end_chars
|
||||
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
|
||||
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
|
||||
|
||||
|
||||
skipped_chars = len(long_string) - total_keep
|
||||
expected_truncation_message = f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. {LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
|
||||
expected_length = start_chars + len(expected_truncation_message) + end_chars
|
||||
|
||||
|
||||
assert len(sanitized["items"][0]["text"]) == expected_length
|
||||
assert sanitized["items"][1]["text"] == "short"
|
||||
assert len(sanitized["items"][2][0]["text"]) == expected_length
|
||||
|
|
@ -147,18 +149,18 @@ def test_sanitize_request_body_for_spend_logs_payload_mixed_types():
|
|||
"nested": {"list": ["short", long_string], "dict": {"key": long_string}},
|
||||
}
|
||||
sanitized = _sanitize_request_body_for_spend_logs_payload(request_body)
|
||||
|
||||
|
||||
# Calculate expected lengths based on actual MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.35)
|
||||
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * 0.65)
|
||||
total_keep = start_chars + end_chars
|
||||
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
|
||||
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
|
||||
|
||||
|
||||
skipped_chars = len(long_string) - total_keep
|
||||
expected_truncation_message = f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. {LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE}) ..."
|
||||
expected_length = start_chars + len(expected_truncation_message) + end_chars
|
||||
|
||||
|
||||
assert len(sanitized["text"]) == expected_length
|
||||
assert sanitized["number"] == 42
|
||||
assert sanitized["nested"]["list"][0] == "short"
|
||||
|
|
@ -347,14 +349,14 @@ def test_get_response_for_spend_logs_payload_truncates_large_base64(mock_should_
|
|||
payload = cast(
|
||||
StandardLoggingPayload,
|
||||
{
|
||||
"response": {
|
||||
"data": [
|
||||
{
|
||||
"b64_json": large_text,
|
||||
"other_field": "value",
|
||||
}
|
||||
]
|
||||
}
|
||||
"response": {
|
||||
"data": [
|
||||
{
|
||||
"b64_json": large_text,
|
||||
"other_field": "value",
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
|
|
@ -369,7 +371,9 @@ def test_get_response_for_spend_logs_payload_truncates_large_base64(mock_should_
|
|||
@patch(
|
||||
"litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs"
|
||||
)
|
||||
def test_get_response_for_spend_logs_payload_truncates_large_embedding(mock_should_store):
|
||||
def test_get_response_for_spend_logs_payload_truncates_large_embedding(
|
||||
mock_should_store,
|
||||
):
|
||||
from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB
|
||||
|
||||
mock_should_store.return_value = True
|
||||
|
|
@ -394,7 +398,7 @@ def test_get_response_for_spend_logs_payload_truncates_large_embedding(mock_shou
|
|||
response_json = _get_response_for_spend_logs_payload(payload)
|
||||
parsed = json.loads(response_json)
|
||||
truncated_value = parsed["data"][0]["embedding"]
|
||||
|
||||
|
||||
assert isinstance(truncated_value, str)
|
||||
assert len(truncated_value) < len(large_embedding)
|
||||
assert LITELLM_TRUNCATED_PAYLOAD_FIELD in truncated_value
|
||||
|
|
@ -416,7 +420,11 @@ def test_truncation_includes_db_safeguard_note():
|
|||
assert LITELLM_TRUNCATED_PAYLOAD_FIELD in truncated
|
||||
assert LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE in truncated
|
||||
assert "DB storage safeguard" in truncated
|
||||
assert "logging callbacks" in truncated.lower() or "logging integrations" in truncated.lower() or "logging callbacks" in truncated
|
||||
assert (
|
||||
"logging callbacks" in truncated.lower()
|
||||
or "logging integrations" in truncated.lower()
|
||||
or "logging callbacks" in truncated
|
||||
)
|
||||
|
||||
|
||||
@patch(
|
||||
|
|
@ -475,21 +483,21 @@ def test_request_body_truncation_logs_info_message(mock_should_store):
|
|||
|
||||
def test_safe_dumps_handles_circular_references():
|
||||
"""Test that safe_dumps can handle circular references without raising exceptions"""
|
||||
|
||||
|
||||
# Create a circular reference
|
||||
obj1 = {"name": "obj1"}
|
||||
obj2 = {"name": "obj2", "ref": obj1}
|
||||
obj1["ref"] = obj2 # This creates a circular reference
|
||||
|
||||
|
||||
# This should not raise an exception
|
||||
result = safe_dumps(obj1)
|
||||
|
||||
|
||||
# Should be a valid JSON string
|
||||
assert isinstance(result, str)
|
||||
|
||||
|
||||
# Should contain placeholder for circular reference
|
||||
assert "CircularReference Detected" in result
|
||||
|
||||
|
||||
# Should be parseable as JSON
|
||||
parsed = json.loads(result)
|
||||
assert parsed["name"] == "obj1"
|
||||
|
|
@ -498,18 +506,18 @@ def test_safe_dumps_handles_circular_references():
|
|||
|
||||
def test_safe_dumps_normal_objects():
|
||||
"""Test that safe_dumps works correctly with normal objects"""
|
||||
|
||||
|
||||
normal_obj = {
|
||||
"string": "test",
|
||||
"number": 42,
|
||||
"boolean": True,
|
||||
"null": None,
|
||||
"list": [1, 2, 3],
|
||||
"nested": {"key": "value"}
|
||||
"nested": {"key": "value"},
|
||||
}
|
||||
|
||||
|
||||
result = safe_dumps(normal_obj)
|
||||
|
||||
|
||||
# Should be a valid JSON string that can be parsed
|
||||
assert isinstance(result, str)
|
||||
parsed = json.loads(result)
|
||||
|
|
@ -518,28 +526,28 @@ def test_safe_dumps_normal_objects():
|
|||
|
||||
def test_safe_dumps_complex_metadata_like_object():
|
||||
"""Test with a complex metadata-like object similar to what caused the issue"""
|
||||
|
||||
|
||||
# Simulate a complex metadata object
|
||||
metadata = {
|
||||
"user_api_key": "test-key",
|
||||
"model": "gpt-4",
|
||||
"usage": {"total_tokens": 100},
|
||||
"mcp_tool_call_metadata": {
|
||||
"name": "test_tool",
|
||||
"arguments": {"param": "value"}
|
||||
}
|
||||
"name": "test_tool",
|
||||
"arguments": {"param": "value"},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# Add a potential circular reference
|
||||
usage_detail = {"parent_metadata": metadata}
|
||||
metadata["usage"]["detail"] = usage_detail
|
||||
|
||||
|
||||
# This should not raise an exception
|
||||
result = safe_dumps(metadata)
|
||||
|
||||
|
||||
# Should be a valid JSON string
|
||||
assert isinstance(result, str)
|
||||
|
||||
|
||||
# Should be parseable as JSON
|
||||
parsed = json.loads(result)
|
||||
assert parsed["user_api_key"] == "test-key"
|
||||
|
|
@ -551,14 +559,14 @@ def test_safe_dumps_complex_metadata_like_object():
|
|||
def test_get_logging_payload_api_key_preserved_when_standard_logging_payload_is_none():
|
||||
"""
|
||||
Critical - Product incident was caused by this bug.
|
||||
|
||||
|
||||
Test that api_key is NOT set to empty string when standard_logging_payload is None.
|
||||
|
||||
|
||||
This is a regression test for a bug where:
|
||||
- On failed requests (bad request errors), standard_logging_payload is None
|
||||
- The else block was incorrectly setting api_key = ""
|
||||
- This caused empty api_key in DailyUserSpend table despite SpendLogs having the correct key
|
||||
|
||||
|
||||
Expected behavior:
|
||||
- api_key from metadata should be extracted and hashed
|
||||
- Even when standard_logging_payload is None, the api_key should be preserved
|
||||
|
|
@ -566,7 +574,7 @@ def test_get_logging_payload_api_key_preserved_when_standard_logging_payload_is_
|
|||
"""
|
||||
# Setup: Simulate a failed request scenario
|
||||
test_api_key = "sk-WLi4iRn4JmbVlTaYw12IOA"
|
||||
|
||||
|
||||
# Create kwargs similar to what's passed during a bad request error
|
||||
kwargs = {
|
||||
"model": "openai/gpt-4.1",
|
||||
|
|
@ -581,39 +589,42 @@ def test_get_logging_payload_api_key_preserved_when_standard_logging_payload_is_
|
|||
},
|
||||
# Note: No 'standard_logging_object' in kwargs - simulating failure case
|
||||
}
|
||||
|
||||
|
||||
# Create a mock error response (bad request)
|
||||
response_obj = Exception("BadRequestError: Invalid parameter 'usersss'")
|
||||
|
||||
|
||||
# Create timestamps
|
||||
start_time = datetime.datetime.now(timezone.utc)
|
||||
end_time = datetime.datetime.now(timezone.utc)
|
||||
|
||||
|
||||
# Call get_logging_payload
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
# CRITICAL ASSERTION: api_key should NOT be empty string
|
||||
assert payload["api_key"] != "", \
|
||||
"BUG: api_key is empty! When standard_logging_payload is None, " \
|
||||
assert payload["api_key"] != "", (
|
||||
"BUG: api_key is empty! When standard_logging_payload is None, "
|
||||
"the api_key from metadata should be preserved and hashed."
|
||||
|
||||
)
|
||||
|
||||
# The api_key should be hashed (not the raw key)
|
||||
assert payload["api_key"] != test_api_key, \
|
||||
"api_key should be hashed, not the raw key"
|
||||
|
||||
assert (
|
||||
payload["api_key"] != test_api_key
|
||||
), "api_key should be hashed, not the raw key"
|
||||
|
||||
# The api_key should be a valid hash (64 character hex string for SHA256)
|
||||
assert len(payload["api_key"]) == 64, \
|
||||
f"Expected 64 character hash, got {len(payload['api_key'])} characters"
|
||||
|
||||
assert (
|
||||
len(payload["api_key"]) == 64
|
||||
), f"Expected 64 character hash, got {len(payload['api_key'])} characters"
|
||||
|
||||
# Verify other fields are set correctly
|
||||
assert payload["model"] == "openai/gpt-4.1"
|
||||
assert payload["user"] == "test_user"
|
||||
|
||||
|
||||
print(f"✅ Test passed! api_key preserved: {payload['api_key']}")
|
||||
|
||||
|
||||
|
|
@ -623,16 +634,16 @@ def test_get_logging_payload_api_key_preserved_when_standard_logging_payload_is_
|
|||
async def test_api_key_preserved_through_failure_hook_to_database():
|
||||
"""
|
||||
CRITICAL E2E TEST: Validates the COMPLETE code path from failure hook to database.
|
||||
|
||||
|
||||
This is THE comprehensive test that protects against the production incident.
|
||||
It tests the EXACT flow that caused the bug:
|
||||
|
||||
|
||||
1. async_post_call_failure_hook is called with api_key in UserAPIKeyAuth
|
||||
2. Failure hook calls update_database with the token parameter
|
||||
3. update_database calls get_logging_payload to create payload
|
||||
4. BUG WAS HERE: get_logging_payload set api_key = "" when standard_logging_payload was None
|
||||
5. Empty api_key was written to DailyUserSpend table
|
||||
|
||||
|
||||
This test validates the ENTIRE flow to ensure the bug cannot regress.
|
||||
If this test fails in CI/CD, the build MUST fail.
|
||||
"""
|
||||
|
|
@ -643,13 +654,21 @@ async def test_api_key_preserved_through_failure_hook_to_database():
|
|||
# Setup
|
||||
test_api_key = "sk-test-critical-e2e-key"
|
||||
hashed_key = hash_token(test_api_key)
|
||||
|
||||
|
||||
# Track what payload gets created
|
||||
captured_payloads = []
|
||||
|
||||
|
||||
async def mock_update_database(
|
||||
token, response_cost, user_id, end_user_id, team_id,
|
||||
kwargs, completion_response, start_time, end_time, org_id
|
||||
token,
|
||||
response_cost,
|
||||
user_id,
|
||||
end_user_id,
|
||||
team_id,
|
||||
kwargs,
|
||||
completion_response,
|
||||
start_time,
|
||||
end_time,
|
||||
org_id,
|
||||
):
|
||||
"""Mock update_database and capture the payload it creates"""
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import (
|
||||
|
|
@ -661,21 +680,23 @@ async def test_api_key_preserved_through_failure_hook_to_database():
|
|||
kwargs=kwargs,
|
||||
response_obj=completion_response,
|
||||
start_time=start_time,
|
||||
end_time=end_time
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
captured_payloads.append({
|
||||
"token": token,
|
||||
"payload": payload,
|
||||
})
|
||||
|
||||
|
||||
captured_payloads.append(
|
||||
{
|
||||
"token": token,
|
||||
"payload": payload,
|
||||
}
|
||||
)
|
||||
|
||||
# Mock dependencies
|
||||
mock_db_writer = MagicMock()
|
||||
mock_db_writer.update_database = AsyncMock(side_effect=mock_update_database)
|
||||
|
||||
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
mock_proxy_logging_obj.db_spend_update_writer = mock_db_writer
|
||||
|
||||
|
||||
# Create UserAPIKeyAuth (what the failure hook receives)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key=hashed_key,
|
||||
|
|
@ -690,9 +711,9 @@ async def test_api_key_preserved_through_failure_hook_to_database():
|
|||
team_alias=None,
|
||||
end_user_id=None,
|
||||
request_route="/chat/completions",
|
||||
metadata={}
|
||||
metadata={},
|
||||
)
|
||||
|
||||
|
||||
# Request data with bad parameter (triggers failure)
|
||||
request_data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
|
|
@ -704,66 +725,66 @@ async def test_api_key_preserved_through_failure_hook_to_database():
|
|||
"user_api_key_user_id": "test_user",
|
||||
"user_api_key_team_id": "test_team",
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
exception = Exception("BadRequestError: Invalid parameter 'invalid_param'")
|
||||
|
||||
|
||||
# Execute the ACTUAL failure hook code path
|
||||
logger = _ProxyDBLogger()
|
||||
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||||
await logger.async_post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=exception,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
traceback_str=None
|
||||
traceback_str=None,
|
||||
)
|
||||
|
||||
|
||||
await asyncio.sleep(0.1) # Wait for async operations
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# CRITICAL ASSERTIONS - If ANY fail, the production bug has regressed!
|
||||
# =========================================================================
|
||||
|
||||
|
||||
assert len(captured_payloads) == 1, "update_database should be called once"
|
||||
|
||||
|
||||
data = captured_payloads[0]
|
||||
payload = data["payload"]
|
||||
payload_api_key = payload.get("api_key")
|
||||
|
||||
|
||||
# THE CRITICAL ASSERTION - This would fail with the original bug!
|
||||
assert payload_api_key != "", \
|
||||
"🚨 CRITICAL BUG: payload['api_key'] is empty! " \
|
||||
"This is the EXACT production incident bug. " \
|
||||
"get_logging_payload() is setting api_key = '' when " \
|
||||
assert payload_api_key != "", (
|
||||
"🚨 CRITICAL BUG: payload['api_key'] is empty! "
|
||||
"This is the EXACT production incident bug. "
|
||||
"get_logging_payload() is setting api_key = '' when "
|
||||
"standard_logging_payload is None (failure case)."
|
||||
|
||||
assert payload_api_key is not None, \
|
||||
"🚨 CRITICAL: payload['api_key'] is None!"
|
||||
|
||||
assert payload_api_key == hashed_key, \
|
||||
f"🚨 CRITICAL: Expected api_key={hashed_key}, got {payload_api_key}"
|
||||
|
||||
)
|
||||
|
||||
assert payload_api_key is not None, "🚨 CRITICAL: payload['api_key'] is None!"
|
||||
|
||||
assert (
|
||||
payload_api_key == hashed_key
|
||||
), f"🚨 CRITICAL: Expected api_key={hashed_key}, got {payload_api_key}"
|
||||
|
||||
# Verify token parameter matches
|
||||
assert data["token"] == hashed_key, \
|
||||
f"Token parameter should be {hashed_key}"
|
||||
|
||||
assert data["token"] == hashed_key, f"Token parameter should be {hashed_key}"
|
||||
|
||||
# Verify other fields
|
||||
assert payload.get("model") == "gpt-3.5-turbo"
|
||||
assert payload.get("user") == "test_user"
|
||||
|
||||
print("\n" + "="*80)
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("✅ CRITICAL E2E TEST PASSED")
|
||||
print("="*80)
|
||||
print("=" * 80)
|
||||
print(f"Token: {data['token']}")
|
||||
print(f"Payload api_key: {payload_api_key}")
|
||||
print(f"Match: {data['token'] == payload_api_key}")
|
||||
print("="*80)
|
||||
print("=" * 80)
|
||||
print("Production incident bug is FIXED and protected:")
|
||||
print("- Failed requests preserve api_key through entire flow")
|
||||
print("- Both SpendLogs AND DailyUserSpend will have correct api_key")
|
||||
print("="*80 + "\n")
|
||||
print("=" * 80 + "\n")
|
||||
|
||||
|
||||
@patch("litellm.proxy.proxy_server.master_key", None)
|
||||
|
|
@ -801,7 +822,9 @@ def test_get_logging_payload_includes_agent_id_from_kwargs():
|
|||
end_time=end_time,
|
||||
)
|
||||
|
||||
assert payload["agent_id"] == test_agent_id, f"Expected agent_id '{test_agent_id}', got '{payload.get('agent_id')}'"
|
||||
assert (
|
||||
payload["agent_id"] == test_agent_id
|
||||
), f"Expected agent_id '{test_agent_id}', got '{payload.get('agent_id')}'"
|
||||
|
||||
|
||||
@patch("litellm.proxy.proxy_server.master_key", None)
|
||||
|
|
@ -902,9 +925,9 @@ def test_get_logging_payload_includes_overhead_in_spend_logs_metadata():
|
|||
# Parse the metadata JSON string
|
||||
metadata_json = payload.get("metadata")
|
||||
assert metadata_json is not None, "metadata should not be None"
|
||||
|
||||
|
||||
metadata = json.loads(metadata_json)
|
||||
|
||||
|
||||
# Verify overhead is stored directly in metadata
|
||||
assert (
|
||||
metadata.get("litellm_overhead_time_ms") == test_overhead_ms
|
||||
|
|
@ -1008,9 +1031,9 @@ def test_get_logging_payload_handles_missing_overhead_gracefully():
|
|||
# Parse the metadata JSON string
|
||||
metadata_json = payload.get("metadata")
|
||||
assert metadata_json is not None, "metadata should not be None"
|
||||
|
||||
|
||||
metadata = json.loads(metadata_json)
|
||||
|
||||
|
||||
# When overhead is None, litellm_overhead_time_ms should be None or not present
|
||||
assert (
|
||||
metadata.get("litellm_overhead_time_ms") is None
|
||||
|
|
@ -1050,7 +1073,9 @@ def test_spend_logs_redacts_request_and_response_when_turn_off_message_logging_e
|
|||
)
|
||||
|
||||
parsed_request = json.loads(request_result)
|
||||
assert parsed_request["messages"] == [{"role": "user", "content": "redacted-by-litellm"}]
|
||||
assert parsed_request["messages"] == [
|
||||
{"role": "user", "content": "redacted-by-litellm"}
|
||||
]
|
||||
assert parsed_request["model"] == "gpt-4"
|
||||
|
||||
# Test response redaction - use dict response to verify redaction
|
||||
|
|
@ -1069,7 +1094,9 @@ def test_spend_logs_redacts_request_and_response_when_turn_off_message_logging_e
|
|||
{"response": response_dict},
|
||||
)
|
||||
|
||||
response_result = _get_response_for_spend_logs_payload(payload=payload, kwargs=kwargs)
|
||||
response_result = _get_response_for_spend_logs_payload(
|
||||
payload=payload, kwargs=kwargs
|
||||
)
|
||||
|
||||
# When redaction is enabled and response is a dict (not ModelResponse),
|
||||
# perform_redaction redacts content in-place within the choices structure
|
||||
|
|
@ -1088,39 +1115,56 @@ def test_should_store_prompts_and_responses_in_spend_logs_case_insensitive_strin
|
|||
"""
|
||||
# Test case-insensitive string "true" variations
|
||||
for true_value in ["true", "TRUE", "True", "TrUe"]:
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"store_prompts_in_spend_logs": true_value}):
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"store_prompts_in_spend_logs": true_value},
|
||||
):
|
||||
mock_get_secret_bool.return_value = False # Ensure env var is False
|
||||
result = _should_store_prompts_and_responses_in_spend_logs()
|
||||
assert result is True, f"Expected True for '{true_value}', got {result}"
|
||||
|
||||
|
||||
# Test boolean True
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"store_prompts_in_spend_logs": True}):
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"store_prompts_in_spend_logs": True},
|
||||
):
|
||||
mock_get_secret_bool.return_value = False
|
||||
result = _should_store_prompts_and_responses_in_spend_logs()
|
||||
assert result is True, f"Expected True for boolean True, got {result}"
|
||||
|
||||
|
||||
# Test that non-true values fall back to environment variable
|
||||
for false_value in [False, None, "false", "FALSE", "False", "anything"]:
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {"store_prompts_in_spend_logs": false_value}):
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"store_prompts_in_spend_logs": false_value},
|
||||
):
|
||||
# When env var is True, should return True
|
||||
mock_get_secret_bool.return_value = True
|
||||
result = _should_store_prompts_and_responses_in_spend_logs()
|
||||
assert result is True, f"Expected True (from env var) for '{false_value}', got {result}"
|
||||
|
||||
assert (
|
||||
result is True
|
||||
), f"Expected True (from env var) for '{false_value}', got {result}"
|
||||
|
||||
# When env var is False, should return False
|
||||
mock_get_secret_bool.return_value = False
|
||||
result = _should_store_prompts_and_responses_in_spend_logs()
|
||||
assert result is False, f"Expected False (from env var) for '{false_value}', got {result}"
|
||||
|
||||
assert (
|
||||
result is False
|
||||
), f"Expected False (from env var) for '{false_value}', got {result}"
|
||||
|
||||
# Test when general_settings doesn't have the key at all
|
||||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||||
mock_get_secret_bool.return_value = True
|
||||
result = _should_store_prompts_and_responses_in_spend_logs()
|
||||
assert result is True, "Expected True (from env var) when key missing, got False"
|
||||
|
||||
assert (
|
||||
result is True
|
||||
), "Expected True (from env var) when key missing, got False"
|
||||
|
||||
mock_get_secret_bool.return_value = False
|
||||
result = _should_store_prompts_and_responses_in_spend_logs()
|
||||
assert result is False, "Expected False (from env var) when key missing, got True"
|
||||
assert (
|
||||
result is False
|
||||
), "Expected False (from env var) when key missing, got True"
|
||||
|
||||
|
||||
def test_get_spend_logs_metadata_guardrail_info_fallback_from_metadata():
|
||||
|
|
@ -1400,7 +1444,9 @@ def test_get_logging_payload_handles_missing_retry_info_gracefully():
|
|||
def test_get_request_duration_ms_normal():
|
||||
"""Test that request duration is correctly computed in milliseconds."""
|
||||
start = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
|
||||
end = datetime.datetime(2025, 1, 1, 0, 0, 2, 500000, tzinfo=timezone.utc) # 2.5s later
|
||||
end = datetime.datetime(
|
||||
2025, 1, 1, 0, 0, 2, 500000, tzinfo=timezone.utc
|
||||
) # 2.5s later
|
||||
result = _get_request_duration_ms(start, end)
|
||||
assert result == 2500
|
||||
|
||||
|
|
@ -1430,10 +1476,13 @@ def test_get_logging_payload_includes_request_duration_ms():
|
|||
"litellm_params": {"api_base": "https://api.openai.com"},
|
||||
"standard_logging_object": None,
|
||||
}
|
||||
response_obj = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
|
||||
response_obj = {
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.master_key", None), \
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}):
|
||||
with patch("litellm.proxy.proxy_server.master_key", None), patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {}
|
||||
):
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
|
|
@ -1469,3 +1518,143 @@ class TestIsMasterKey:
|
|||
master = "sk-master-key-123"
|
||||
hashed = hash_token(master)
|
||||
assert _is_master_key(api_key=hashed, _master_key=master) is True
|
||||
|
||||
|
||||
def test_azure_embedding_pre_call_sets_api_base_in_litellm_params():
|
||||
"""
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/23768
|
||||
|
||||
Exercises the actual Azure embedding provider code path to verify that
|
||||
the logging_obj.pre_call() receives api_base in additional_args, which
|
||||
causes _pre_call() to write it into litellm_params["api_base"].
|
||||
|
||||
Previously, the Azure embedding handler omitted api_base from the
|
||||
additional_args dict passed to pre_call, so litellm_params["api_base"]
|
||||
remained empty and the resulting SpendLogs row had an empty api_base.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
test_api_base = "https://my-azure-endpoint.openai.azure.com/"
|
||||
|
||||
# Create a real Logging object with empty litellm_params
|
||||
logging_obj = Logging(
|
||||
model="azure/text-embedding-3-large",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="aembedding",
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
litellm_call_id="test-call-id",
|
||||
function_id="test-function-id",
|
||||
)
|
||||
# Ensure api_base starts empty in litellm_params
|
||||
assert logging_obj.model_call_details["litellm_params"].get("api_base", "") == ""
|
||||
|
||||
# Simulate what the Azure embedding handler does in pre_call (the fixed version)
|
||||
logging_obj.pre_call(
|
||||
input=["test input"],
|
||||
api_key="fake-key",
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": "text-embedding-3-large",
|
||||
"input": ["test input"],
|
||||
},
|
||||
"api_base": test_api_base,
|
||||
"headers": {"api_key": "fake-key", "azure_ad_token": None},
|
||||
},
|
||||
)
|
||||
|
||||
# After pre_call, litellm_params["api_base"] must be populated
|
||||
actual_api_base = logging_obj.model_call_details["litellm_params"]["api_base"]
|
||||
assert actual_api_base == test_api_base, (
|
||||
f"BUG: litellm_params['api_base'] is '{actual_api_base}' after pre_call, "
|
||||
f"expected '{test_api_base}'. The Azure embedding handler must pass api_base "
|
||||
"in pre_call additional_args for SpendLogs to record the endpoint."
|
||||
)
|
||||
|
||||
|
||||
def test_azure_embedding_pre_call_without_api_base_defaults_to_empty():
|
||||
"""
|
||||
Verify that pre_call without api_base in additional_args results in an
|
||||
empty string (the previous broken behavior), confirming the test above
|
||||
actually catches the regression.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
logging_obj = Logging(
|
||||
model="azure/text-embedding-3-large",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="aembedding",
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
litellm_call_id="test-call-id-2",
|
||||
function_id="test-function-id-2",
|
||||
)
|
||||
|
||||
# Simulate the OLD broken behavior: pre_call WITHOUT api_base
|
||||
logging_obj.pre_call(
|
||||
input=["test input"],
|
||||
api_key="fake-key",
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": "text-embedding-3-large",
|
||||
"input": ["test input"],
|
||||
},
|
||||
"headers": {"api_key": "fake-key"},
|
||||
},
|
||||
)
|
||||
|
||||
# Without api_base in additional_args, litellm_params["api_base"] should be empty
|
||||
actual_api_base = logging_obj.model_call_details["litellm_params"]["api_base"]
|
||||
assert (
|
||||
actual_api_base == ""
|
||||
), f"Expected empty api_base when not passed in additional_args, got '{actual_api_base}'"
|
||||
|
||||
|
||||
@patch("litellm.proxy.proxy_server.master_key", None)
|
||||
@patch("litellm.proxy.proxy_server.general_settings", {})
|
||||
def test_get_logging_payload_reads_api_base_from_litellm_params_for_embeddings():
|
||||
"""
|
||||
End-to-end test: after pre_call populates litellm_params["api_base"],
|
||||
get_logging_payload must include it in the SpendLogsPayload.
|
||||
"""
|
||||
test_api_base = "https://my-azure-endpoint.openai.azure.com/"
|
||||
|
||||
kwargs = {
|
||||
"model": "azure/text-embedding-3-large",
|
||||
"call_type": "aembedding",
|
||||
"custom_llm_provider": "azure",
|
||||
"litellm_params": {
|
||||
"api_base": test_api_base,
|
||||
"metadata": {
|
||||
"user_api_key": "sk-test-key",
|
||||
"user_api_key_team_id": "test_team",
|
||||
"model_group": "text-embedding-3-large",
|
||||
"model_info": {"id": "model-id-123"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
response_obj = {
|
||||
"id": "embd-test-123",
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}],
|
||||
"model": "text-embedding-3-large",
|
||||
"usage": {"prompt_tokens": 10, "total_tokens": 10},
|
||||
}
|
||||
|
||||
start_time = datetime.datetime.now(timezone.utc)
|
||||
end_time = datetime.datetime.now(timezone.utc)
|
||||
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
assert payload["api_base"] == test_api_base, (
|
||||
f"BUG: api_base is '{payload['api_base']}' but expected '{test_api_base}'. "
|
||||
"Embedding calls must have api_base populated in SpendLogs."
|
||||
)
|
||||
assert payload["call_type"] == "aembedding"
|
||||
assert payload["model_group"] == "text-embedding-3-large"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue