Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 15 additions & 59 deletions litellm/integrations/prometheus.py
Original file line number Diff line number Diff line change
Expand Up @@ -1055,16 +1055,16 @@ async def async_log_success_event(self, kwargs, response_obj, start_time, end_ti
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,
Expand All @@ -1086,13 +1086,6 @@ def _increment_token_metrics(
):
_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"
Expand Down Expand Up @@ -1655,49 +1648,12 @@ async def async_post_call_success_hook(
):
"""
Proxy level tracking - triggered when the proxy responds with a success response to the client
"""
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
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).
"""
pass

def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any:
"""Get value from dict or Pydantic model."""
Expand Down
18 changes: 12 additions & 6 deletions litellm/proxy/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -450,18 +450,24 @@ def get_proxy_hook(self, hook: str) -> Optional[CustomLogger]:
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

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
Comment thread
shivamrawat1 marked this conversation as resolved.
litellm.logging_callback_manager.add_litellm_callback(initialized_callback)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

string_callbacks_to_replace is never populated — callbacks silently break

The string_callbacks_to_replace dict is created at line 457 but never receives any entries. The first commit (0f9f0da4) correctly had:

if initialized_callback is not None:
    string_callbacks_to_replace[idx] = initialized_callback

The second commit (b672cc0c) removed this assignment without replacing it. As a result, the replacement loop at lines 468-470 iterates over an empty dict, and no string callbacks are ever initialized or replaced. This means the entire fix for the double-counting bug is ineffective — string callbacks like "prometheus" remain as raw strings in litellm.callbacks and are never turned into PrometheusLogger instances via this code path.

Additionally, the else branch that registered pre-existing CustomLogger instances with add_litellm_callback was also removed, so non-string callbacks are no longer registered with the callback manager either.

Suggested change
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
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
litellm.logging_callback_manager.add_litellm_callback(initialized_callback)
for idx, callback in enumerate(litellm.callbacks):
if isinstance(callback, str):
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 initialized_callback is not None:
string_callbacks_to_replace[idx] = initialized_callback
else:
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
litellm.logging_callback_manager.add_litellm_callback(initialized_callback)


async def update_request_status(
self, litellm_call_id: str, status: Literal["success", "fail"]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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:
Expand Down Expand Up @@ -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."
)
Expand Down
175 changes: 175 additions & 0 deletions tests/litellm/proxy/test_init_litellm_callbacks.py
Original file line number Diff line number Diff line change
@@ -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