mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(prometheus): expose MCP tool metadata in Prometheus metrics (#31899)
Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
parent
64dc5080b9
commit
85db18e618
3 changed files with 362 additions and 0 deletions
|
|
@ -593,6 +593,21 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=[],
|
||||
)
|
||||
|
||||
########################################
|
||||
# MCP Tool Call Metrics
|
||||
########################################
|
||||
self.litellm_mcp_tool_calls_total = self._counter_factory(
|
||||
name="litellm_mcp_tool_calls_total",
|
||||
documentation="Total MCP tool calls, segmented by tool and server name",
|
||||
labelnames=self.get_labels_for_metric("litellm_mcp_tool_calls_total"),
|
||||
)
|
||||
|
||||
self.litellm_mcp_tool_call_spend_metric = self._counter_factory(
|
||||
name="litellm_mcp_tool_call_spend_metric",
|
||||
documentation="Total spend on MCP tool calls, segmented by tool and server name",
|
||||
labelnames=self.get_labels_for_metric("litellm_mcp_tool_call_spend_metric"),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
print_verbose(f"Got exception on init prometheus client {str(e)}")
|
||||
raise e
|
||||
|
|
@ -1300,6 +1315,13 @@ class PrometheusLogger(CustomLogger):
|
|||
label_context=label_context,
|
||||
)
|
||||
|
||||
# MCP tool call metrics
|
||||
self._increment_mcp_tool_call_metrics(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=response_cost,
|
||||
)
|
||||
|
||||
# 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
|
||||
|
|
@ -1521,6 +1543,49 @@ class PrometheusLogger(CustomLogger):
|
|||
amount=float(provider_cache_creation_tokens),
|
||||
)
|
||||
|
||||
def _increment_mcp_tool_call_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
response_cost: float,
|
||||
) -> None:
|
||||
metadata = standard_logging_payload.get("metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
mcp_meta = metadata.get("mcp_tool_call_metadata")
|
||||
if not isinstance(mcp_meta, dict):
|
||||
return
|
||||
|
||||
mcp_enum_values = UserAPIKeyLabelValues(
|
||||
mcp_tool_name=mcp_meta.get("name"),
|
||||
mcp_server_name=mcp_meta.get("mcp_server_name"),
|
||||
hashed_api_key=enum_values.hashed_api_key,
|
||||
api_key_alias=enum_values.api_key_alias,
|
||||
team=enum_values.team,
|
||||
team_alias=enum_values.team_alias,
|
||||
user=enum_values.user,
|
||||
end_user=enum_values.end_user,
|
||||
)
|
||||
mcp_label_context = PrometheusLabelFactoryContext(mcp_enum_values)
|
||||
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
self,
|
||||
self.litellm_mcp_tool_calls_total,
|
||||
"litellm_mcp_tool_calls_total",
|
||||
mcp_enum_values,
|
||||
label_context=mcp_label_context,
|
||||
)
|
||||
|
||||
if response_cost > 0:
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
self,
|
||||
self.litellm_mcp_tool_call_spend_metric,
|
||||
"litellm_mcp_tool_call_spend_metric",
|
||||
mcp_enum_values,
|
||||
label_context=mcp_label_context,
|
||||
amount=response_cost,
|
||||
)
|
||||
|
||||
async def _increment_remaining_budget_metrics(
|
||||
self,
|
||||
user_api_team: Optional[str],
|
||||
|
|
|
|||
|
|
@ -188,6 +188,8 @@ class UserAPIKeyLabelNames(Enum):
|
|||
STREAM = "stream"
|
||||
ORG_ID = "org_id"
|
||||
ORG_ALIAS = "org_alias"
|
||||
MCP_TOOL_NAME = "mcp_tool_name"
|
||||
MCP_SERVER_NAME = "mcp_server_name"
|
||||
|
||||
|
||||
DEFINED_PROMETHEUS_METRICS = Literal[
|
||||
|
|
@ -264,6 +266,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_check_batch_cost_jobs_processed_total",
|
||||
"litellm_check_batch_cost_errors_total",
|
||||
"litellm_check_batch_cost_last_run_timestamp",
|
||||
# MCP tool call metrics
|
||||
"litellm_mcp_tool_calls_total",
|
||||
"litellm_mcp_tool_call_spend_metric",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -737,6 +742,20 @@ class PrometheusMetricLabels:
|
|||
|
||||
litellm_check_batch_cost_last_run_timestamp: List[str] = []
|
||||
|
||||
# MCP tool call metrics
|
||||
litellm_mcp_tool_calls_total: list[str] = [
|
||||
UserAPIKeyLabelNames.MCP_TOOL_NAME.value,
|
||||
UserAPIKeyLabelNames.MCP_SERVER_NAME.value,
|
||||
UserAPIKeyLabelNames.API_KEY_HASH.value,
|
||||
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
|
||||
UserAPIKeyLabelNames.TEAM.value,
|
||||
UserAPIKeyLabelNames.TEAM_ALIAS.value,
|
||||
UserAPIKeyLabelNames.USER.value,
|
||||
UserAPIKeyLabelNames.END_USER.value,
|
||||
]
|
||||
|
||||
litellm_mcp_tool_call_spend_metric: list[str] = list(litellm_mcp_tool_calls_total)
|
||||
|
||||
@staticmethod
|
||||
def get_labels(label_name: DEFINED_PROMETHEUS_METRICS) -> List[str]:
|
||||
default_labels = getattr(PrometheusMetricLabels, label_name)
|
||||
|
|
@ -840,6 +859,8 @@ class UserAPIKeyLabelValues:
|
|||
stream: Optional[str] = None
|
||||
org_id: Optional[str] = None
|
||||
org_alias: Optional[str] = None
|
||||
mcp_tool_name: Optional[str] = None
|
||||
mcp_server_name: Optional[str] = None
|
||||
|
||||
# Added for test compatibility.
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,276 @@
|
|||
"""
|
||||
Unit tests for MCP tool call Prometheus metrics (LIT-3765).
|
||||
|
||||
These metrics expose ``mcp_tool_call_metadata`` in Prometheus so Grafana
|
||||
dashboards can break down MCP usage by server and tool name.
|
||||
|
||||
Run with:
|
||||
uv run pytest tests/test_litellm/integrations/test_prometheus_mcp_tool_metrics.py -v
|
||||
"""
|
||||
|
||||
from typing import get_args
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.types.integrations.prometheus import (
|
||||
DEFINED_PROMETHEUS_METRICS,
|
||||
PrometheusMetricLabels,
|
||||
UserAPIKeyLabelNames,
|
||||
UserAPIKeyLabelValues,
|
||||
)
|
||||
|
||||
|
||||
MCP_METRICS = (
|
||||
"litellm_mcp_tool_calls_total",
|
||||
"litellm_mcp_tool_call_spend_metric",
|
||||
)
|
||||
|
||||
|
||||
def _make_mock_logger():
|
||||
logger = MagicMock()
|
||||
for name in MCP_METRICS:
|
||||
setattr(logger, name, MagicMock())
|
||||
logger.get_labels_for_metric = MagicMock(
|
||||
return_value=PrometheusMetricLabels.litellm_mcp_tool_calls_total,
|
||||
)
|
||||
return logger
|
||||
|
||||
|
||||
def _make_enum_values(
|
||||
*,
|
||||
mcp_tool_name: str = "get_weather",
|
||||
mcp_server_name: str = "weather-server",
|
||||
) -> UserAPIKeyLabelValues:
|
||||
return UserAPIKeyLabelValues(
|
||||
mcp_tool_name=mcp_tool_name,
|
||||
mcp_server_name=mcp_server_name,
|
||||
hashed_api_key="sk-hash-123",
|
||||
api_key_alias="test-key",
|
||||
team="team-1",
|
||||
team_alias="Test Team",
|
||||
user="user-1",
|
||||
end_user="end-user-1",
|
||||
)
|
||||
|
||||
|
||||
def _make_payload(
|
||||
*,
|
||||
mcp_tool_name: str = "get_weather",
|
||||
mcp_server_name: str = "weather-server",
|
||||
response_cost: float = 0.005,
|
||||
) -> dict:
|
||||
return {
|
||||
"model": "gpt-4o",
|
||||
"model_group": "gpt-4o",
|
||||
"model_id": "model-123",
|
||||
"api_base": "https://api.openai.com",
|
||||
"custom_llm_provider": "openai",
|
||||
"response_cost": response_cost,
|
||||
"completion_tokens": 50,
|
||||
"prompt_tokens": 100,
|
||||
"total_tokens": 150,
|
||||
"request_tags": [],
|
||||
"stream": False,
|
||||
"metadata": {
|
||||
"user_api_key_hash": "sk-hash-123",
|
||||
"user_api_key_alias": "test-key",
|
||||
"user_api_key_team_id": "team-1",
|
||||
"user_api_key_team_alias": "Test Team",
|
||||
"user_api_key_user_id": "user-1",
|
||||
"user_api_key_user_email": None,
|
||||
"user_api_key_org_id": None,
|
||||
"user_api_key_org_alias": None,
|
||||
"mcp_tool_call_metadata": {
|
||||
"name": mcp_tool_name,
|
||||
"mcp_server_name": mcp_server_name,
|
||||
"namespaced_tool_name": f"{mcp_server_name}/{mcp_tool_name}",
|
||||
"arguments": {"city": "SF"},
|
||||
"result": {"temp": 72},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class TestMCPMetricRegistration:
|
||||
def test_metrics_in_defined_prometheus_metrics(self):
|
||||
defined = get_args(DEFINED_PROMETHEUS_METRICS)
|
||||
for name in MCP_METRICS:
|
||||
assert name in defined, f"{name} missing from DEFINED_PROMETHEUS_METRICS"
|
||||
|
||||
def test_metric_labels_defined(self):
|
||||
for name in MCP_METRICS:
|
||||
assert hasattr(PrometheusMetricLabels, name), f"{name} missing from PrometheusMetricLabels"
|
||||
|
||||
def test_mcp_labels_include_tool_and_server_name(self):
|
||||
labels = PrometheusMetricLabels.litellm_mcp_tool_calls_total
|
||||
assert UserAPIKeyLabelNames.MCP_TOOL_NAME.value in labels
|
||||
assert UserAPIKeyLabelNames.MCP_SERVER_NAME.value in labels
|
||||
|
||||
def test_spend_metric_shares_label_set_with_calls_metric(self):
|
||||
assert (
|
||||
PrometheusMetricLabels.litellm_mcp_tool_call_spend_metric
|
||||
== PrometheusMetricLabels.litellm_mcp_tool_calls_total
|
||||
)
|
||||
assert (
|
||||
PrometheusMetricLabels.litellm_mcp_tool_call_spend_metric
|
||||
is not PrometheusMetricLabels.litellm_mcp_tool_calls_total
|
||||
)
|
||||
|
||||
def test_enum_values_accept_mcp_fields(self):
|
||||
vals = _make_enum_values()
|
||||
assert vals.mcp_tool_name == "get_weather"
|
||||
assert vals.mcp_server_name == "weather-server"
|
||||
|
||||
def test_enum_values_default_mcp_fields_to_none(self):
|
||||
vals = UserAPIKeyLabelValues(user="u1")
|
||||
assert vals.mcp_tool_name is None
|
||||
assert vals.mcp_server_name is None
|
||||
|
||||
|
||||
class TestIncrementMCPToolCallMetrics:
|
||||
def test_increments_calls_counter_when_mcp_metadata_present(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload()
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.005,
|
||||
)
|
||||
|
||||
logger.litellm_mcp_tool_calls_total.labels.assert_called_once()
|
||||
logger.litellm_mcp_tool_calls_total.labels().inc.assert_called_once_with(1.0)
|
||||
|
||||
def test_increments_spend_counter_when_cost_positive(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload(response_cost=0.01)
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.01,
|
||||
)
|
||||
|
||||
logger.litellm_mcp_tool_call_spend_metric.labels.assert_called_once()
|
||||
logger.litellm_mcp_tool_call_spend_metric.labels().inc.assert_called_once_with(0.01)
|
||||
|
||||
def test_skips_spend_counter_when_cost_zero(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload(response_cost=0.0)
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.0,
|
||||
)
|
||||
|
||||
logger.litellm_mcp_tool_calls_total.labels.assert_called_once()
|
||||
logger.litellm_mcp_tool_call_spend_metric.labels.assert_not_called()
|
||||
|
||||
def test_noop_when_no_mcp_metadata(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload()
|
||||
payload["metadata"]["mcp_tool_call_metadata"] = None
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.005,
|
||||
)
|
||||
|
||||
for name in MCP_METRICS:
|
||||
getattr(logger, name).labels.assert_not_called()
|
||||
|
||||
def test_noop_when_metadata_missing(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = {"metadata": None}
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.005,
|
||||
)
|
||||
|
||||
for name in MCP_METRICS:
|
||||
getattr(logger, name).labels.assert_not_called()
|
||||
|
||||
def test_label_values_carry_tool_and_server_name(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload(
|
||||
mcp_tool_name="search_docs",
|
||||
mcp_server_name="docs-mcp",
|
||||
)
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.005,
|
||||
)
|
||||
|
||||
labels_passed = logger.litellm_mcp_tool_calls_total.labels.call_args
|
||||
assert labels_passed.kwargs["mcp_tool_name"] == "search_docs"
|
||||
assert labels_passed.kwargs["mcp_server_name"] == "docs-mcp"
|
||||
|
||||
def test_label_values_carry_team_and_key_from_parent(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload()
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
hashed_api_key="sk-parent-key",
|
||||
api_key_alias="parent-alias",
|
||||
team="parent-team",
|
||||
team_alias="Parent Team",
|
||||
user="parent-user",
|
||||
end_user="parent-end-user",
|
||||
)
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.005,
|
||||
)
|
||||
|
||||
labels_passed = logger.litellm_mcp_tool_calls_total.labels.call_args
|
||||
assert labels_passed.kwargs["hashed_api_key"] == "sk-parent-key"
|
||||
assert labels_passed.kwargs["team"] == "parent-team"
|
||||
assert labels_passed.kwargs["team_alias"] == "Parent Team"
|
||||
assert labels_passed.kwargs["user"] == "parent-user"
|
||||
|
||||
def test_handles_missing_server_name_gracefully(self):
|
||||
logger = _make_mock_logger()
|
||||
payload = _make_payload()
|
||||
payload["metadata"]["mcp_tool_call_metadata"] = {
|
||||
"name": "standalone_tool",
|
||||
"arguments": {},
|
||||
"result": {},
|
||||
}
|
||||
enum_values = _make_enum_values()
|
||||
|
||||
PrometheusLogger._increment_mcp_tool_call_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=enum_values,
|
||||
response_cost=0.0,
|
||||
)
|
||||
|
||||
labels_passed = logger.litellm_mcp_tool_calls_total.labels.call_args
|
||||
assert labels_passed.kwargs["mcp_tool_name"] == "standalone_tool"
|
||||
assert labels_passed.kwargs["mcp_server_name"] is None
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
Loading…
Add table
Reference in a new issue