mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(logging): preserve pass-through endpoint response_cost (#21844)
* fix(logging): preserve pass-through endpoint response_cost in async_success_handler Two places in the logging pipeline were overwriting response_cost that pass-through handlers (Gemini/Vertex) had already calculated: 1. _process_hidden_params_and_response_cost fell through to _response_cost_calculator which returns None for pass-through calls 2. async_success_handler pass-through branch unconditionally set response_cost = None (introduced in PR #19887) Now both places check if response_cost is already set before overwriting. * test: add regression test for pass-through endpoint response_cost preservation
This commit is contained in:
parent
3278fee714
commit
9f459c5c57
2 changed files with 79 additions and 16 deletions
|
|
@ -1636,6 +1636,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details["response_cost"] = 0.0
|
||||
elif "response_cost" in hidden_params:
|
||||
self.model_call_details["response_cost"] = hidden_params["response_cost"]
|
||||
elif self.model_call_details.get("response_cost") is not None:
|
||||
# Preserve response_cost if already calculated (e.g., by pass-through
|
||||
# handlers like Gemini/Vertex which call completion_cost directly)
|
||||
pass
|
||||
else:
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(
|
||||
result=logging_result
|
||||
|
|
@ -2496,23 +2500,29 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
self.model_call_details["async_complete_streaming_response"] = result
|
||||
|
||||
# cost calculation not possible for pass-through
|
||||
self.model_call_details["response_cost"] = None
|
||||
# Only set response_cost to None if not already calculated by
|
||||
# pass-through handlers (e.g. Gemini/Vertex handlers already
|
||||
# compute cost via completion_cost)
|
||||
if self.model_call_details.get("response_cost") is None:
|
||||
self.model_call_details["response_cost"] = None
|
||||
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
result, start_time, end_time
|
||||
)
|
||||
|
||||
# print standard logging payload
|
||||
if (
|
||||
standard_logging_payload := self.model_call_details.get(
|
||||
# Only build standard_logging_object if not already built by
|
||||
# _success_handler_helper_fn
|
||||
if self.model_call_details.get("standard_logging_object") is None:
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
self.model_call_details[
|
||||
"standard_logging_object"
|
||||
] = self._build_standard_logging_payload(
|
||||
result, start_time, end_time
|
||||
)
|
||||
) is not None:
|
||||
emit_standard_logging_payload(standard_logging_payload)
|
||||
|
||||
# print standard logging payload
|
||||
if (
|
||||
standard_logging_payload := self.model_call_details.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
) is not None:
|
||||
emit_standard_logging_payload(standard_logging_payload)
|
||||
callbacks = self.get_combined_callback_list(
|
||||
dynamic_success_callbacks=self.dynamic_async_success_callbacks,
|
||||
global_callbacks=litellm._async_success_callback,
|
||||
|
|
|
|||
|
|
@ -170,9 +170,9 @@ async def test_datadog_logger_not_shadowed_by_llm_obs(monkeypatch):
|
|||
monkeypatch.setenv("DD_API_KEY", "test")
|
||||
monkeypatch.setenv("DD_SITE", "us5.datadoghq.com")
|
||||
|
||||
from litellm.litellm_core_utils import litellm_logging as logging_module
|
||||
from litellm.integrations.datadog.datadog import DataDogLogger
|
||||
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
|
||||
from litellm.litellm_core_utils import litellm_logging as logging_module
|
||||
|
||||
logging_module._in_memory_loggers.clear()
|
||||
|
||||
|
|
@ -205,8 +205,8 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch):
|
|||
monkeypatch.setenv("LOGFIRE_BASE_URL", "https://logfire-api-custom.pydantic.dev") # no trailing slash on purpose
|
||||
|
||||
# Import after env vars are set (important if module-level caching exists)
|
||||
from litellm.litellm_core_utils import litellm_logging as logging_module
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry # logger class
|
||||
from litellm.litellm_core_utils import litellm_logging as logging_module
|
||||
|
||||
logging_module._in_memory_loggers.clear()
|
||||
|
||||
|
|
@ -1589,3 +1589,56 @@ async def test_emit_standard_logging_payload_called_for_non_streaming():
|
|||
assert mock_emit.call_count >= 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_success_handler_preserves_response_cost_for_pass_through_endpoints():
|
||||
"""Regression test: PR #19887 added a pass-through branch in async_success_handler
|
||||
that unconditionally set response_cost=None, overwriting costs already calculated
|
||||
by pass-through handlers (Gemini/Vertex)."""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="gemini-2.5-flash-lite",
|
||||
messages=[{"role": "user", "content": "test"}],
|
||||
stream=False,
|
||||
call_type="pass_through_endpoint",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="test-call-id-cost",
|
||||
function_id="test-function-id-cost",
|
||||
)
|
||||
|
||||
# Simulate what pass-through handlers do: pre-calculate response_cost
|
||||
logging_obj.model_call_details = {
|
||||
"litellm_params": {"metadata": {}, "proxy_server_request": {}},
|
||||
"litellm_call_id": "test-call-id-cost",
|
||||
"response_cost": 0.0000047, # Pre-calculated by pass-through handler
|
||||
"custom_llm_provider": "gemini",
|
||||
}
|
||||
|
||||
result = ModelResponse(
|
||||
id="test-response",
|
||||
model="gemini-2.5-flash-lite",
|
||||
usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
|
||||
)
|
||||
|
||||
start_time = datetime.now()
|
||||
end_time = datetime.now()
|
||||
|
||||
with patch.object(logging_obj, "get_combined_callback_list", return_value=[]):
|
||||
await logging_obj.async_success_handler(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
)
|
||||
|
||||
# response_cost must be preserved, not overwritten to None
|
||||
assert logging_obj.model_call_details.get("response_cost") is not None
|
||||
assert logging_obj.model_call_details["response_cost"] > 0
|
||||
|
||||
# standard_logging_object should also have the cost
|
||||
slo = logging_obj.model_call_details.get("standard_logging_object")
|
||||
assert slo is not None
|
||||
assert slo["response_cost"] > 0
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue