mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fixed double counting
This commit is contained in:
parent
af3acdda18
commit
0f9f0da413
5 changed files with 225 additions and 91 deletions
|
|
@ -1055,16 +1055,16 @@ class PrometheusLogger(CustomLogger):
|
|||
enum_values=enum_values,
|
||||
)
|
||||
|
||||
if (
|
||||
standard_logging_payload["stream"] is True
|
||||
): # log successful streaming requests from logging event hook.
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
|
||||
# increment litellm_proxy_total_requests_metric for all successful requests
|
||||
# (both streaming and non-streaming) in this single location to prevent
|
||||
# double-counting that occurs when async_post_call_success_hook also increments
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
|
||||
|
||||
def _increment_token_metrics(
|
||||
self,
|
||||
|
|
@ -1086,13 +1086,6 @@ class PrometheusLogger(CustomLogger):
|
|||
):
|
||||
_tags = standard_logging_payload["request_tags"]
|
||||
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_total_tokens_metric"
|
||||
|
|
@ -1655,49 +1648,12 @@ class PrometheusLogger(CustomLogger):
|
|||
):
|
||||
"""
|
||||
Proxy level tracking - triggered when the proxy responds with a success response to the client
|
||||
|
||||
Note: litellm_proxy_total_requests_metric is NOT incremented here to avoid
|
||||
double-counting. It is incremented in async_log_success_event which fires
|
||||
for all successful requests (both streaming and non-streaming).
|
||||
"""
|
||||
try:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup,
|
||||
)
|
||||
|
||||
if self._should_skip_metrics_for_invalid_key(
|
||||
user_api_key_dict=user_api_key_dict
|
||||
):
|
||||
return
|
||||
|
||||
_metadata = data.get("metadata", {}) or {}
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
end_user=user_api_key_dict.end_user_id,
|
||||
hashed_api_key=user_api_key_dict.api_key,
|
||||
api_key_alias=user_api_key_dict.key_alias,
|
||||
requested_model=data.get("model", ""),
|
||||
team=user_api_key_dict.team_id,
|
||||
team_alias=user_api_key_dict.team_alias,
|
||||
user=user_api_key_dict.user_id,
|
||||
user_email=user_api_key_dict.user_email,
|
||||
status_code="200",
|
||||
route=user_api_key_dict.request_route,
|
||||
tags=StandardLoggingPayloadSetup._get_request_tags(
|
||||
litellm_params=data,
|
||||
proxy_server_request=data.get("proxy_server_request", {}),
|
||||
),
|
||||
client_ip=_metadata.get("requester_ip_address"),
|
||||
user_agent=_metadata.get("user_agent"),
|
||||
)
|
||||
_labels = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(
|
||||
metric_name="litellm_proxy_total_requests_metric"
|
||||
),
|
||||
enum_values=enum_values,
|
||||
)
|
||||
self.litellm_proxy_total_requests_metric.labels(**_labels).inc()
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"prometheus Layer Error(): Exception occured - {}".format(str(e))
|
||||
)
|
||||
pass
|
||||
pass
|
||||
|
||||
def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
|
||||
"""Get value from dict or Pydantic model."""
|
||||
|
|
|
|||
|
|
@ -450,18 +450,32 @@ class ProxyLogging:
|
|||
def _init_litellm_callbacks(self, llm_router: Optional[Router] = None):
|
||||
self._add_proxy_hooks(llm_router)
|
||||
litellm.logging_callback_manager.add_litellm_callback(self.service_logging_obj) # type: ignore
|
||||
for callback in litellm.callbacks:
|
||||
|
||||
# Track string callbacks and their initialized instances so we can
|
||||
# replace them in-place, preventing duplicates (string + instance) in
|
||||
# litellm.callbacks which caused double-counting of metrics.
|
||||
string_callbacks_to_replace: Dict[int, CustomLogger] = {}
|
||||
|
||||
for idx, callback in enumerate(litellm.callbacks):
|
||||
if isinstance(callback, str):
|
||||
callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class( # type: ignore
|
||||
initialized_callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class(
|
||||
cast(_custom_logger_compatible_callbacks_literal, callback),
|
||||
internal_usage_cache=self.internal_usage_cache.dual_cache,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if callback is None:
|
||||
continue
|
||||
if initialized_callback is not None:
|
||||
string_callbacks_to_replace[idx] = initialized_callback
|
||||
else:
|
||||
# Only add non-string callbacks to the manager; string
|
||||
# callbacks will be replaced in-place below and are already
|
||||
# present in litellm.callbacks, so adding them to the manager
|
||||
# (which appends to litellm.callbacks) would create duplicates.
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback)
|
||||
# Replace string entries in litellm.callbacks with initialized instances
|
||||
for idx, initialized_callback in string_callbacks_to_replace.items():
|
||||
litellm.callbacks[idx] = initialized_callback
|
||||
|
||||
async def update_request_status(
|
||||
self, litellm_call_id: str, status: Literal["success", "fail"]
|
||||
|
|
|
|||
|
|
@ -792,7 +792,8 @@ async def test_async_post_call_success_hook(prometheus_logger):
|
|||
"""
|
||||
Test for the async_post_call_success_hook method
|
||||
|
||||
it should increment the litellm_proxy_total_requests_metric
|
||||
litellm_proxy_total_requests_metric is NOT incremented here to avoid double-counting.
|
||||
It is incremented in async_log_success_event instead.
|
||||
"""
|
||||
# Mock the prometheus metric
|
||||
prometheus_logger.litellm_proxy_total_requests_metric = MagicMock()
|
||||
|
|
@ -817,23 +818,8 @@ async def test_async_post_call_success_hook(prometheus_logger):
|
|||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
# Assert total requests metric was incremented with correct labels
|
||||
prometheus_logger.litellm_proxy_total_requests_metric.labels.assert_called_once_with(
|
||||
end_user=None,
|
||||
hashed_api_key="test_key",
|
||||
api_key_alias="test_alias",
|
||||
requested_model="gpt-3.5-turbo",
|
||||
team="test_team",
|
||||
team_alias="test_team_alias",
|
||||
user="test_user",
|
||||
status_code="200",
|
||||
user_email=None,
|
||||
route=user_api_key_dict.request_route,
|
||||
model_id=None,
|
||||
client_ip=None,
|
||||
user_agent=None,
|
||||
)
|
||||
prometheus_logger.litellm_proxy_total_requests_metric.labels().inc.assert_called_once()
|
||||
# Assert total requests metric was NOT incremented (moved to async_log_success_event)
|
||||
prometheus_logger.litellm_proxy_total_requests_metric.labels.assert_not_called()
|
||||
|
||||
|
||||
def test_set_llm_deployment_success_metrics(prometheus_logger):
|
||||
|
|
|
|||
|
|
@ -545,6 +545,9 @@ async def test_request_counter_semantic_validation(mock_prometheus_logger):
|
|||
CRITICAL TEST: Validates that request counters are incremented by 1, not by token count.
|
||||
This test specifically catches the bug where litellm_proxy_total_requests_metric
|
||||
is incorrectly incremented by total_tokens instead of 1.
|
||||
|
||||
The metric is now ONLY incremented in async_log_success_event (for both streaming
|
||||
and non-streaming) to prevent double-counting.
|
||||
"""
|
||||
from datetime import datetime, timedelta
|
||||
from unittest.mock import MagicMock
|
||||
|
|
@ -583,18 +586,18 @@ async def test_request_counter_semantic_validation(mock_prometheus_logger):
|
|||
},
|
||||
}
|
||||
|
||||
# Call the success event
|
||||
# Call the success event - should increment for both streaming and non-streaming
|
||||
await mock_prometheus_logger.async_log_success_event(
|
||||
kwargs, None, kwargs["start_time"], kwargs["end_time"]
|
||||
)
|
||||
|
||||
# CRITICAL ASSERTION: Request counter should not be incremented
|
||||
# CRITICAL ASSERTION: Request counter should be incremented by 1
|
||||
total_requests_metric = mock_prometheus_logger.litellm_proxy_total_requests_metric
|
||||
assert (
|
||||
len(total_requests_metric.inc_calls) == 0
|
||||
), "Request metric should not be incremented"
|
||||
len(total_requests_metric.inc_calls) == 1
|
||||
), "Request metric should be incremented once in async_log_success_event"
|
||||
|
||||
# Call the post-call logging hook
|
||||
# Call the post-call logging hook - should NOT increment (to prevent double-counting)
|
||||
await mock_prometheus_logger.async_post_call_success_hook(
|
||||
data={},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
|
|
@ -607,11 +610,11 @@ async def test_request_counter_semantic_validation(mock_prometheus_logger):
|
|||
response=MagicMock(),
|
||||
)
|
||||
|
||||
# CRITICAL ASSERTION: Request counter be incremented by 1
|
||||
# CRITICAL ASSERTION: Request counter should still be 1 (not incremented again)
|
||||
total_requests_metric = mock_prometheus_logger.litellm_proxy_total_requests_metric
|
||||
assert (
|
||||
len(total_requests_metric.inc_calls) == 1
|
||||
), "Request metric should not be incremented"
|
||||
), "Request metric should not be incremented again in async_post_call_success_hook"
|
||||
|
||||
# Check that ALL request counter increments are by 1 (not by token count)
|
||||
for inc_value in total_requests_metric.inc_calls:
|
||||
|
|
@ -684,8 +687,8 @@ async def test_multiple_requests_counter_semantics(mock_prometheus_logger):
|
|||
expected_total_tokens = num_requests * tokens_per_request # 3 * 500 = 1500
|
||||
|
||||
# With the bug, total_request_increments would be 1500 instead of 3
|
||||
assert total_request_increments == 0, (
|
||||
f"SEMANTIC BUG: Request counter total increments = 0, "
|
||||
assert total_request_increments == num_requests, (
|
||||
f"SEMANTIC BUG: Request counter total increments = {total_request_increments}, "
|
||||
f"expected {num_requests}. This suggests request counters are being incremented "
|
||||
f"by token counts instead of request counts."
|
||||
)
|
||||
|
|
|
|||
175
tests/litellm/proxy/test_init_litellm_callbacks.py
Normal file
175
tests/litellm/proxy/test_init_litellm_callbacks.py
Normal file
|
|
@ -0,0 +1,175 @@
|
|||
"""
|
||||
Unit tests for ProxyLogging._init_litellm_callbacks.
|
||||
|
||||
Validates that string callbacks in litellm.callbacks are replaced in-place
|
||||
with their initialized instances, preventing duplicate entries (string + instance)
|
||||
that caused double-counting of metrics like litellm_proxy_total_requests_metric.
|
||||
"""
|
||||
|
||||
from typing import List, Union
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class FakeCustomLogger(CustomLogger):
|
||||
"""A minimal CustomLogger subclass for testing."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class TestInitLitellmCallbacks:
|
||||
"""Tests for ProxyLogging._init_litellm_callbacks."""
|
||||
|
||||
def _make_proxy_logging(self):
|
||||
"""Create a ProxyLogging instance with mocked dependencies."""
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
mock_cache = MagicMock()
|
||||
proxy_logging = ProxyLogging(user_api_key_cache=mock_cache)
|
||||
return proxy_logging
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.utils.ProxyLogging._add_proxy_hooks",
|
||||
new_callable=lambda: lambda self, *a, **kw: None,
|
||||
)
|
||||
def test_should_replace_string_callback_with_instance(self, _mock_hooks):
|
||||
"""
|
||||
When litellm.callbacks contains a string callback (e.g. "lago"),
|
||||
_init_litellm_callbacks should replace the string with the initialized
|
||||
CustomLogger instance, not leave both the string and instance in the list.
|
||||
"""
|
||||
fake_logger = FakeCustomLogger()
|
||||
|
||||
# Start with a string callback in litellm.callbacks
|
||||
litellm.callbacks = ["lago"] # type: ignore
|
||||
|
||||
proxy_logging = self._make_proxy_logging()
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class",
|
||||
return_value=fake_logger,
|
||||
):
|
||||
proxy_logging._init_litellm_callbacks(llm_router=None)
|
||||
|
||||
# The string "lago" should be replaced by the instance, not appended
|
||||
string_entries = [c for c in litellm.callbacks if isinstance(c, str)]
|
||||
instance_entries = [
|
||||
c for c in litellm.callbacks if isinstance(c, FakeCustomLogger)
|
||||
]
|
||||
|
||||
assert len(string_entries) == 0, (
|
||||
f"String callbacks should have been replaced, but found: {string_entries}"
|
||||
)
|
||||
assert len(instance_entries) == 1, (
|
||||
f"Expected exactly one FakeCustomLogger instance, found {len(instance_entries)}"
|
||||
)
|
||||
assert instance_entries[0] is fake_logger
|
||||
|
||||
# Clean up
|
||||
litellm.callbacks = [] # type: ignore
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.utils.ProxyLogging._add_proxy_hooks",
|
||||
new_callable=lambda: lambda self, *a, **kw: None,
|
||||
)
|
||||
def test_should_not_duplicate_existing_instance_callbacks(self, _mock_hooks):
|
||||
"""
|
||||
When litellm.callbacks already contains a CustomLogger instance (not a string),
|
||||
_init_litellm_callbacks should not create a duplicate.
|
||||
"""
|
||||
existing_logger = FakeCustomLogger()
|
||||
|
||||
litellm.callbacks = [existing_logger] # type: ignore
|
||||
|
||||
proxy_logging = self._make_proxy_logging()
|
||||
|
||||
proxy_logging._init_litellm_callbacks(llm_router=None)
|
||||
|
||||
# Count how many FakeCustomLogger instances are in litellm.callbacks
|
||||
instance_count = sum(
|
||||
1 for c in litellm.callbacks if isinstance(c, FakeCustomLogger)
|
||||
)
|
||||
assert instance_count == 1, (
|
||||
f"Expected exactly 1 FakeCustomLogger instance, found {instance_count}. "
|
||||
f"litellm.callbacks = {litellm.callbacks}"
|
||||
)
|
||||
|
||||
# Clean up
|
||||
litellm.callbacks = [] # type: ignore
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.utils.ProxyLogging._add_proxy_hooks",
|
||||
new_callable=lambda: lambda self, *a, **kw: None,
|
||||
)
|
||||
def test_should_handle_unrecognized_string_callback(self, _mock_hooks):
|
||||
"""
|
||||
When _init_custom_logger_compatible_class returns None for a string callback,
|
||||
the string should remain in litellm.callbacks (not crash).
|
||||
"""
|
||||
litellm.callbacks = ["unknown_callback"] # type: ignore
|
||||
|
||||
proxy_logging = self._make_proxy_logging()
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class",
|
||||
return_value=None,
|
||||
):
|
||||
proxy_logging._init_litellm_callbacks(llm_router=None)
|
||||
|
||||
# The unknown string callback should still be there (not replaced, not crashed)
|
||||
assert "unknown_callback" in litellm.callbacks
|
||||
|
||||
# Clean up
|
||||
litellm.callbacks = [] # type: ignore
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.utils.ProxyLogging._add_proxy_hooks",
|
||||
new_callable=lambda: lambda self, *a, **kw: None,
|
||||
)
|
||||
def test_should_replace_multiple_string_callbacks(self, _mock_hooks):
|
||||
"""
|
||||
When litellm.callbacks contains multiple string callbacks,
|
||||
each should be replaced with its corresponding initialized instance.
|
||||
"""
|
||||
fake_logger_a = FakeCustomLogger()
|
||||
fake_logger_b = FakeCustomLogger()
|
||||
|
||||
litellm.callbacks = ["callback_a", "callback_b"] # type: ignore
|
||||
|
||||
proxy_logging = self._make_proxy_logging()
|
||||
|
||||
call_count = 0
|
||||
|
||||
def mock_init_class(callback_name, **kwargs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return fake_logger_a
|
||||
return fake_logger_b
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class",
|
||||
side_effect=mock_init_class,
|
||||
):
|
||||
proxy_logging._init_litellm_callbacks(llm_router=None)
|
||||
|
||||
string_entries = [c for c in litellm.callbacks if isinstance(c, str)]
|
||||
instance_entries = [
|
||||
c for c in litellm.callbacks if isinstance(c, FakeCustomLogger)
|
||||
]
|
||||
|
||||
assert len(string_entries) == 0, (
|
||||
f"All string callbacks should have been replaced: {string_entries}"
|
||||
)
|
||||
assert len(instance_entries) == 2, (
|
||||
f"Expected 2 FakeCustomLogger instances, found {len(instance_entries)}"
|
||||
)
|
||||
assert instance_entries[0] is fake_logger_a
|
||||
assert instance_entries[1] is fake_logger_b
|
||||
|
||||
# Clean up
|
||||
litellm.callbacks = [] # type: ignore
|
||||
Loading…
Add table
Reference in a new issue