fix: fix test

This commit is contained in:
Krrish Dholakia 2025-09-13 13:36:05 -07:00
parent d75c0ff2c4
commit 459e66a9cd
2 changed files with 134 additions and 116 deletions

View file

@ -1040,13 +1040,6 @@ class PrometheusLogger(CustomLogger):
# increment total LLM requests and spend metric
self._increment_top_level_request_and_spend_metrics(
end_user_id=end_user_id,
user_api_key=user_api_key,
user_api_key_alias=user_api_key_alias,
model=model,
user_api_team=user_api_team,
user_api_team_alias=user_api_team_alias,
user_id=user_id,
response_cost=response_cost,
enum_values=enum_values,
)
@ -1213,13 +1206,6 @@ class PrometheusLogger(CustomLogger):
def _increment_top_level_request_and_spend_metrics(
self,
end_user_id: Optional[str],
user_api_key: Optional[str],
user_api_key_alias: Optional[str],
model: Optional[str],
user_api_team: Optional[str],
user_api_team_alias: Optional[str],
user_id: Optional[str],
response_cost: float,
enum_values: UserAPIKeyLabelValues,
):

View file

@ -560,13 +560,6 @@ def test_increment_top_level_request_and_spend_metrics(prometheus_logger):
prometheus_logger.litellm_spend_metric = MagicMock()
prometheus_logger._increment_top_level_request_and_spend_metrics(
end_user_id="user1",
user_api_key="key1",
user_api_key_alias="alias1",
model="gpt-3.5-turbo",
user_api_team="team1",
user_api_team_alias="team_alias1",
user_id="user1",
response_cost=0.1,
enum_values=enum_values,
)
@ -584,7 +577,13 @@ def test_increment_top_level_request_and_spend_metrics(prometheus_logger):
prometheus_logger.litellm_requests_metric.labels().inc.assert_called_once()
prometheus_logger.litellm_spend_metric.labels.assert_called_once_with(
"user1", "key1", "alias1", "gpt-3.5-turbo", "team1", "team_alias1", "user1"
end_user=None,
hashed_api_key="test_hash",
api_key_alias="test_alias",
model="gpt-3.5-turbo",
team="test_team",
team_alias="test_team_alias",
user=None,
)
prometheus_logger.litellm_spend_metric.labels().inc.assert_called_once_with(0.1)
@ -1141,22 +1140,28 @@ def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch):
# Configure tags with wildcard patterns
monkeypatch.setattr(
"litellm.custom_prometheus_tags",
["User-Agent: curl/*", "User-Agent: python-requests/*", "Environment: prod*", "Service: api-gateway*", "exact-match"]
"litellm.custom_prometheus_tags",
[
"User-Agent: curl/*",
"User-Agent: python-requests/*",
"Environment: prod*",
"Service: api-gateway*",
"exact-match",
],
)
# Test tags that should match the wildcard patterns
tags = [
"User-Agent: curl/7.68.0",
"User-Agent: python-requests/2.28.1",
"User-Agent: curl/7.68.0",
"User-Agent: python-requests/2.28.1",
"Environment: production",
"Service: api-gateway-v2",
"exact-match",
"other-tag"
"other-tag",
]
result = get_custom_labels_from_tags(tags)
expected = {
"tag_User_Agent__curl__": "true", # matches "User-Agent: curl/*"
"tag_User_Agent__python_requests__": "true", # matches "User-Agent: python-requests/*"
@ -1164,7 +1169,7 @@ def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch):
"tag_Service__api_gateway_": "true", # matches "Service: api-gateway*"
"tag_exact_match": "true", # exact match
}
assert result == expected
@ -1174,26 +1179,26 @@ def test_get_custom_labels_from_tags_wildcard_no_matches(monkeypatch):
# Configure tags with wildcard patterns
monkeypatch.setattr(
"litellm.custom_prometheus_tags",
["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"]
"litellm.custom_prometheus_tags",
["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"],
)
# Test tags that should NOT match the wildcard patterns
tags = [
"User-Agent: curl/7.68.0", # doesn't match "User-Agent: firefox/*"
"Environment: production", # doesn't match "Environment: dev*"
"Environment: production", # doesn't match "Environment: dev*"
"Service: api-gateway-v2", # doesn't match "Service: web-app*"
"other-tag"
"other-tag",
]
result = get_custom_labels_from_tags(tags)
expected = {
"tag_User_Agent__firefox__": "false", # no match for "User-Agent: firefox/*"
"tag_Environment__dev_": "false", # no match for "Environment: dev*"
"tag_Service__web_app_": "false", # no match for "Service: web-app*"
}
assert result == expected
@ -1204,48 +1209,69 @@ def test_tag_matches_wildcard_configured_pattern():
)
# Test cases that should match
assert _tag_matches_wildcard_configured_pattern(
tags=["User-Agent: curl/7.68.0", "prod", "other"],
configured_tag="User-Agent: curl/*"
) is True
assert _tag_matches_wildcard_configured_pattern(
tags=["User-Agent: python-requests/2.28.1", "test"],
configured_tag="User-Agent: python-requests/*"
) is True
assert _tag_matches_wildcard_configured_pattern(
tags=["Environment: production", "debug"],
configured_tag="Environment: prod*"
) is True
assert (
_tag_matches_wildcard_configured_pattern(
tags=["User-Agent: curl/7.68.0", "prod", "other"],
configured_tag="User-Agent: curl/*",
)
is True
)
assert (
_tag_matches_wildcard_configured_pattern(
tags=["User-Agent: python-requests/2.28.1", "test"],
configured_tag="User-Agent: python-requests/*",
)
is True
)
assert (
_tag_matches_wildcard_configured_pattern(
tags=["Environment: production", "debug"],
configured_tag="Environment: prod*",
)
is True
)
# Test exact match (no wildcard)
assert _tag_matches_wildcard_configured_pattern(
tags=["prod", "test"],
configured_tag="prod"
) is True
assert (
_tag_matches_wildcard_configured_pattern(
tags=["prod", "test"], configured_tag="prod"
)
is True
)
# Test cases that should NOT match
assert _tag_matches_wildcard_configured_pattern(
tags=["User-Agent: firefox/98.0", "prod"],
configured_tag="User-Agent: curl/*"
) is False
assert _tag_matches_wildcard_configured_pattern(
tags=["Environment: development", "test"],
configured_tag="Environment: prod*"
) is False
assert _tag_matches_wildcard_configured_pattern(
tags=["staging", "test"],
configured_tag="prod"
) is False
assert (
_tag_matches_wildcard_configured_pattern(
tags=["User-Agent: firefox/98.0", "prod"],
configured_tag="User-Agent: curl/*",
)
is False
)
assert (
_tag_matches_wildcard_configured_pattern(
tags=["Environment: development", "test"],
configured_tag="Environment: prod*",
)
is False
)
assert (
_tag_matches_wildcard_configured_pattern(
tags=["staging", "test"], configured_tag="prod"
)
is False
)
# Test with empty tags
assert _tag_matches_wildcard_configured_pattern(
tags=[],
configured_tag="User-Agent: curl/*"
) is False
assert (
_tag_matches_wildcard_configured_pattern(
tags=[], configured_tag="User-Agent: curl/*"
)
is False
)
@pytest.mark.asyncio(scope="session")
@ -1908,12 +1934,12 @@ def test_set_llm_deployment_success_metrics_with_label_filtering():
async def test_prometheus_token_metrics_with_prometheus_config():
"""
Test that validates the renamed token metrics are incremented correctly with a prometheus config.
This test ensures that after the metric renaming (git diff):
- litellm_total_tokens -> litellm_total_tokens_metric
- litellm_input_tokens -> litellm_input_tokens_metric
- litellm_input_tokens -> litellm_input_tokens_metric
- litellm_output_tokens -> litellm_output_tokens_metric
All three metrics should be properly incremented when making a successful completion request.
"""
from prometheus_client import CollectorRegistry, Counter
@ -1925,39 +1951,39 @@ async def test_prometheus_token_metrics_with_prometheus_config():
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
REGISTRY.unregister(collector)
# Set up prometheus configuration that includes the token metrics
config = [
PrometheusMetricsConfig(
group="token_metrics_test",
metrics=[
"litellm_total_tokens_metric",
"litellm_input_tokens_metric",
"litellm_input_tokens_metric",
"litellm_output_tokens_metric",
"litellm_requests_metric"
"litellm_requests_metric",
],
include_labels=[
"model",
"hashed_api_key",
"hashed_api_key",
"api_key_alias",
"team",
"team_alias"
"team_alias",
],
)
]
# Mock litellm.prometheus_metrics_config
with patch("litellm.prometheus_metrics_config", config):
# Create PrometheusLogger with the configuration
prometheus_logger = PrometheusLogger()
# Test data with specific token counts
standard_logging_payload = create_standard_logging_payload()
standard_logging_payload["total_tokens"] = 1500
standard_logging_payload["prompt_tokens"] = 900
standard_logging_payload["completion_tokens"] = 600
standard_logging_payload["response_cost"] = 0.075
kwargs = {
"model": "gpt-3.5-turbo",
"stream": False,
@ -1971,7 +1997,7 @@ async def test_prometheus_token_metrics_with_prometheus_config():
}
},
"start_time": datetime.now() - timedelta(seconds=2),
"completion_start_time": datetime.now() - timedelta(seconds=1),
"completion_start_time": datetime.now() - timedelta(seconds=1),
"api_call_start_time": datetime.now() - timedelta(seconds=1.5),
"end_time": datetime.now(),
"standard_logging_object": standard_logging_payload,
@ -1987,69 +2013,75 @@ async def test_prometheus_token_metrics_with_prometheus_config():
print("final registry values", REGISTRY._collector_to_names)
# Get metric collectors directly from registry
# Get metric collectors directly from registry
metric_collectors = {}
for collector, names in REGISTRY._collector_to_names.items():
metric_name = names[0] # First name is the base metric name
metric_collectors[metric_name] = collector
print("=== Final Metric Values (Direct Access) ===")
# Expected values
# Expected values
expected_values = {
"litellm_total_tokens_metric": 1500.0,
"litellm_input_tokens_metric": 900.0,
"litellm_output_tokens_metric": 600.0,
"litellm_requests_metric": 1.0
"litellm_requests_metric": 1.0,
}
expected_label_values = {
'api_key_alias': 'test_alias',
'hashed_api_key': 'test_hash',
'model': 'gpt-3.5-turbo',
'team': 'test_team',
'team_alias': 'test_team_alias'
"api_key_alias": "test_alias",
"hashed_api_key": "test_hash",
"model": "gpt-3.5-turbo",
"team": "test_team",
"team_alias": "test_team_alias",
}
# Validate each metric directly
for metric_name, expected_value in expected_values.items():
if metric_name in metric_collectors:
collector = metric_collectors[metric_name]
# Get all samples for this metric
samples = list(collector.collect())[0].samples
# Find the _total sample (the actual counter value)
total_sample = None
for sample in samples:
if sample.name.endswith('_total'):
if sample.name.endswith("_total"):
total_sample = sample
break
if total_sample:
actual_value = total_sample.value
actual_labels = total_sample.labels
print(f"✓ {metric_name}: expected={expected_value}, actual={actual_value}")
print(
f"✓ {metric_name}: expected={expected_value}, actual={actual_value}"
)
print(f" Labels: {actual_labels}")
# Validate the value
assert actual_value == expected_value, f"Expected {expected_value}, got {actual_value} for {metric_name}"
assert (
actual_value == expected_value
), f"Expected {expected_value}, got {actual_value} for {metric_name}"
# Validate the labels
for label_key, expected_label_value in expected_label_values.items():
for (
label_key,
expected_label_value,
) in expected_label_values.items():
actual_label_value = actual_labels.get(label_key)
assert actual_label_value == expected_label_value, f"Expected label {label_key}={expected_label_value}, got {actual_label_value}"
assert (
actual_label_value == expected_label_value
), f"Expected label {label_key}={expected_label_value}, got {actual_label_value}"
print(f" ✓ {metric_name} VALIDATED")
else:
raise AssertionError(f"No _total sample found for {metric_name}")
else:
raise AssertionError(f"Metric {metric_name} not found in registry")
print("✓ All token metrics validated successfully!")
# check final value of metrics in registry