From 66fafa3a7fbefd205a9c3205eb158a443aae8fe8 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 1 Jul 2025 20:17:17 -0700 Subject: [PATCH] [Feat] Polish - add better error validation when users configure prometheus metrics and labels to control cardinality (#12182) * self._pretty_print_invalid_metric_error * docs prometheus.md * test prom validation checks * update metric name * fix _pretty_print_validation_errors * fix linting * test prometheus * test fixes - prometheus --- docs/my-website/docs/proxy/prometheus.md | 16 +- litellm/integrations/prometheus.py | 281 +++++++++++++++++- litellm/proxy/proxy_config.yaml | 13 + litellm/types/integrations/prometheus.py | 44 ++- .../test_prometheus_unit_tests.py | 4 +- .../integrations/test_prometheus.py | 113 ++++++- 6 files changed, 440 insertions(+), 31 deletions(-) diff --git a/docs/my-website/docs/proxy/prometheus.md b/docs/my-website/docs/proxy/prometheus.md index d3fb6eca59..e4ef6f183c 100644 --- a/docs/my-website/docs/proxy/prometheus.md +++ b/docs/my-website/docs/proxy/prometheus.md @@ -64,9 +64,9 @@ Use this for for tracking per [user, key, team, etc.](virtual_keys) | Metric Name | Description | |----------------------|--------------------------------------| | `litellm_spend_metric` | Total Spend, per `"user", "key", "model", "team", "end-user"` | -| `litellm_total_tokens` | input + output tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` | -| `litellm_input_tokens` | input tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` | -| `litellm_output_tokens` | output tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` | +| `litellm_total_tokens_metric` | input + output tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` | +| `litellm_input_tokens_metric` | input tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` | +| `litellm_output_tokens_metric` | output tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` | ### Team - Budget @@ -288,10 +288,11 @@ Control which labels are included for each metric to reduce cardinality: litellm_settings: callbacks: ["prometheus"] prometheus_metrics_config: - - group: "spend_and_tokens" + - group: "token_consumption" metrics: - - "litellm_spend_metric" - - "litellm_total_tokens" + - "litellm_input_tokens_metric" + - "litellm_output_tokens_metric" + - "litellm_total_tokens_metric" include_labels: - "model" - "team" @@ -324,7 +325,6 @@ litellm_settings: # Budget metrics with full label set - group: "budget_tracking" metrics: - - "litellm_spend_metric" - "litellm_remaining_team_budget_metric" include_labels: - "team" @@ -385,7 +385,7 @@ Use these metrics to monitor the health of the DB Transaction Queue. Eg. Monitor -## **🔥 LiteLLM Maintained Grafana Dashboards ** +## 🔥 LiteLLM Maintained Grafana Dashboards Link to Grafana Dashboards maintained by LiteLLM diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 9aea69c34a..0583150b96 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -353,29 +353,289 @@ class PrometheusLogger(CustomLogger): verbose_logger.debug(f"prometheus config: {config}") - label_filters = {} + # Parse and validate all configuration groups + parsed_configs = [] self.enabled_metrics = set() - - # Parse each configuration group + for group_config in config: # Validate configuration using Pydantic if isinstance(group_config, dict): parsed_config = PrometheusMetricsConfig(**group_config) else: parsed_config = group_config - - # Add enabled metrics to the set + + parsed_configs.append(parsed_config) self.enabled_metrics.update(parsed_config.metrics) - # Set label filters for each metric in this group - for metric_name in parsed_config.metrics: - if parsed_config.include_labels: - label_filters[metric_name] = parsed_config.include_labels + # Validate all configurations + validation_results = self._validate_all_configurations(parsed_configs) + + if validation_results.has_errors: + self._pretty_print_validation_errors(validation_results) + error_message = "Configuration validation failed:\n" + "\n".join(validation_results.all_error_messages) + raise ValueError(error_message) + # Build label filters from valid configurations + label_filters = self._build_label_filters(parsed_configs) + # Pretty print the processed configuration self._pretty_print_prometheus_config(label_filters) - return label_filters + + def _validate_all_configurations(self, parsed_configs: List) -> ValidationResults: + """Validate all metric configurations and return collected errors""" + metric_errors = [] + label_errors = [] + + for config in parsed_configs: + for metric_name in config.metrics: + # Validate metric name + metric_error = self._validate_single_metric_name(metric_name) + if metric_error: + metric_errors.append(metric_error) + continue # Skip label validation if metric name is invalid + + # Validate labels if provided + if config.include_labels: + label_error = self._validate_single_metric_labels(metric_name, config.include_labels) + if label_error: + label_errors.append(label_error) + + return ValidationResults(metric_errors=metric_errors, label_errors=label_errors) + + def _validate_single_metric_name(self, metric_name: str) -> Optional[MetricValidationError]: + """Validate a single metric name""" + from typing import get_args + if metric_name not in set(get_args(DEFINED_PROMETHEUS_METRICS)): + return MetricValidationError( + metric_name=metric_name, + valid_metrics=get_args(DEFINED_PROMETHEUS_METRICS) + ) + return None + + def _validate_single_metric_labels(self, metric_name: str, labels: List[str]) -> Optional[LabelValidationError]: + """Validate labels for a single metric""" + from typing import cast + + # Get valid labels for this metric from PrometheusMetricLabels + valid_labels = PrometheusMetricLabels.get_labels(cast(DEFINED_PROMETHEUS_METRICS, metric_name)) + + # Find invalid labels + invalid_labels = [label for label in labels if label not in valid_labels] + + if invalid_labels: + return LabelValidationError( + metric_name=metric_name, + invalid_labels=invalid_labels, + valid_labels=valid_labels + ) + return None + + def _build_label_filters(self, parsed_configs: List) -> Dict[str, List[str]]: + """Build label filters from validated configurations""" + label_filters = {} + + for config in parsed_configs: + for metric_name in config.metrics: + if config.include_labels: + # Only add if metric name is valid (validation already passed) + if self._validate_single_metric_name(metric_name) is None: + label_filters[metric_name] = config.include_labels + + return label_filters + + def _validate_configured_metric_labels(self, metric_name: str, labels: List[str]): + """ + Ensure that all the configured labels are valid for the metric + + Raises ValueError if the metric labels are invalid and pretty prints the error + """ + label_error = self._validate_single_metric_labels(metric_name, labels) + if label_error: + self._pretty_print_invalid_labels_error( + metric_name=label_error.metric_name, + invalid_labels=label_error.invalid_labels, + valid_labels=label_error.valid_labels + ) + raise ValueError(label_error.message) + + return True + + ######################################################### + # Pretty print functions + ######################################################### + + def _pretty_print_validation_errors(self, validation_results: ValidationResults) -> None: + """Pretty print all validation errors using rich""" + try: + from rich.console import Console + from rich.panel import Panel + from rich.table import Table + from rich.text import Text + + console = Console() + + # Create error panel title + title = Text("🚨🚨 Configuration Validation Errors", style="bold red") + + # Print main error panel + console.print("\n") + console.print(Panel(title, border_style="red")) + + # Show invalid metric names if any + if validation_results.metric_errors: + invalid_metrics = [e.metric_name for e in validation_results.metric_errors] + valid_metrics = validation_results.metric_errors[0].valid_metrics # All should have same valid metrics + + metrics_error_text = Text( + f"Invalid Metric Names: {', '.join(invalid_metrics)}", + style="bold red" + ) + console.print(Panel(metrics_error_text, border_style="red")) + + metrics_table = Table( + title="📊 Valid Metric Names", + show_header=True, + header_style="bold green", + title_justify="left", + border_style="green", + ) + metrics_table.add_column("Available Metrics", style="cyan", no_wrap=True) + + for metric in sorted(valid_metrics): + metrics_table.add_row(metric) + + console.print(metrics_table) + + # Show invalid labels if any + if validation_results.label_errors: + for error in validation_results.label_errors: + labels_error_text = Text( + f"Invalid Labels for '{error.metric_name}': {', '.join(error.invalid_labels)}", + style="bold red" + ) + console.print(Panel(labels_error_text, border_style="red")) + + labels_table = Table( + title=f"🏷️ Valid Labels for '{error.metric_name}'", + show_header=True, + header_style="bold green", + title_justify="left", + border_style="green", + ) + labels_table.add_column("Valid Labels", style="cyan", no_wrap=True) + + for label in sorted(error.valid_labels): + labels_table.add_row(label) + + console.print(labels_table) + + console.print("\n") + + except ImportError: + # Fallback to simple logging if rich is not available + for metric_error in validation_results.metric_errors: + verbose_logger.error(metric_error.message) + for label_error in validation_results.label_errors: + verbose_logger.error(label_error.message) + + def _pretty_print_invalid_labels_error( + self, metric_name: str, invalid_labels: List[str], valid_labels: List[str] + ) -> None: + """Pretty print error message for invalid labels using rich""" + try: + from rich.console import Console + from rich.panel import Panel + from rich.table import Table + from rich.text import Text + + console = Console() + + # Create error panel title + title = Text( + f"🚨🚨 Invalid Labels for Metric: '{metric_name}'\nInvalid labels: {', '.join(invalid_labels)}\nPlease specify only valid labels below", + style="bold red" + ) + + # Create valid labels table + labels_table = Table( + title="🏷️ Valid Labels for this Metric", + show_header=True, + header_style="bold green", + title_justify="left", + border_style="green", + ) + labels_table.add_column("Valid Labels", style="cyan", no_wrap=True) + + for label in sorted(valid_labels): + labels_table.add_row(label) + + # Print everything in a nice panel + console.print("\n") + console.print(Panel(title, border_style="red")) + console.print(labels_table) + console.print("\n") + + except ImportError: + # Fallback to simple logging if rich is not available + verbose_logger.error( + f"Invalid labels for metric '{metric_name}': {invalid_labels}. Valid labels: {sorted(valid_labels)}" + ) + + def _pretty_print_invalid_metric_error( + self, invalid_metric_name: str, valid_metrics: tuple + ) -> None: + """Pretty print error message for invalid metric name using rich""" + try: + from rich.console import Console + from rich.panel import Panel + from rich.table import Table + from rich.text import Text + + console = Console() + + # Create error panel title + title = Text(f"🚨🚨 Invalid Metric Name: '{invalid_metric_name}'\nPlease specify one of the allowed metrics below", style="bold red") + + # Create valid metrics table + metrics_table = Table( + title="📊 Valid Metric Names", + show_header=True, + header_style="bold green", + title_justify="left", + border_style="green", + ) + metrics_table.add_column("Available Metrics", style="cyan", no_wrap=True) + + for metric in sorted(valid_metrics): + metrics_table.add_row(metric) + + # Print everything in a nice panel + console.print("\n") + console.print(Panel(title, border_style="red")) + console.print(metrics_table) + console.print("\n") + + except ImportError: + # Fallback to simple logging if rich is not available + verbose_logger.error( + f"Invalid metric name: {invalid_metric_name}. Valid metrics: {sorted(valid_metrics)}" + ) + + ######################################################### + # End of pretty print functions + ######################################################### + + def _valid_metric_name(self, metric_name: str): + """ + Raises ValueError if the metric name is invalid and pretty prints the error + """ + error = self._validate_single_metric_name(metric_name) + if error: + self._pretty_print_invalid_metric_error( + invalid_metric_name=error.metric_name, + valid_metrics=error.valid_metrics) + raise ValueError(error.message) def _pretty_print_prometheus_config( self, label_filters: Dict[str, List[str]] @@ -447,6 +707,7 @@ class PrometheusLogger(CustomLogger): ) verbose_logger.info(f"Label filters: {label_filters}") + def _is_metric_enabled(self, metric_name: str) -> bool: """Check if a metric is enabled based on configuration""" # If no specific configuration is provided, enable all metrics (default behavior) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 8a8fd6794e..01235ca7f7 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -6,3 +6,16 @@ model_list: litellm_params: model: openai/* + +litellm_settings: + callbacks: ["prometheus"] + prometheus_metrics_config: + - group: totals + metrics: + - litellm_spend_metric + - litellm_total_tokens + - litellm_input_tokens_metric + - litellm_output_tokens_metric + include_labels: + - requested_modela + - model \ No newline at end of file diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 6a696345f9..45644f73b5 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -1,11 +1,53 @@ +from dataclasses import dataclass from enum import Enum -from typing import Dict, List, Literal, Optional, Union +from typing import Dict, List, Literal, Optional, Tuple, Union from pydantic import BaseModel, Field from typing_extensions import Annotated import litellm + +@dataclass +class MetricValidationError: + """Error for invalid metric name""" + metric_name: str + valid_metrics: Tuple[str, ...] + + @property + def message(self) -> str: + return f"Invalid metric name: {self.metric_name}" + + +@dataclass +class LabelValidationError: + """Error for invalid labels on a metric""" + metric_name: str + invalid_labels: List[str] + valid_labels: List[str] + + @property + def message(self) -> str: + return f"Invalid labels for metric '{self.metric_name}': {self.invalid_labels}" + + +@dataclass +class ValidationResults: + """Container for all validation results""" + metric_errors: List[MetricValidationError] + label_errors: List[LabelValidationError] + + @property + def has_errors(self) -> bool: + return bool(self.metric_errors or self.label_errors) + + @property + def all_error_messages(self) -> List[str]: + messages = [error.message for error in self.metric_errors] + messages.extend([error.message for error in self.label_errors]) + return messages + + REQUESTED_MODEL = "requested_model" EXCEPTION_STATUS = "exception_status" EXCEPTION_CLASS = "exception_class" diff --git a/tests/logging_callback_tests/test_prometheus_unit_tests.py b/tests/logging_callback_tests/test_prometheus_unit_tests.py index 254ab9f5a5..456cf7fe38 100644 --- a/tests/logging_callback_tests/test_prometheus_unit_tests.py +++ b/tests/logging_callback_tests/test_prometheus_unit_tests.py @@ -1550,7 +1550,7 @@ def test_set_llm_deployment_success_metrics_with_label_filtering(): "litellm_deployment_total_requests", ], include_labels=[ - "requested_model", + "litellm_model_name", "api_provider", "hashed_api_key", ], # Limited labels @@ -1628,7 +1628,7 @@ def test_set_llm_deployment_success_metrics_with_label_filtering(): ) # Should only contain the filtered labels that are supported for this metric - expected_filtered_labels = {"requested_model", "api_provider", "hashed_api_key"} + expected_filtered_labels = {"litellm_model_name", "api_provider", "hashed_api_key"} actual_labels = set(k for k in overhead_labels.keys() if k is not None) # Verify that only expected labels are present (subset of configured labels) diff --git a/tests/test_litellm/integrations/test_prometheus.py b/tests/test_litellm/integrations/test_prometheus.py index 5e3e690fa4..dc8593758d 100644 --- a/tests/test_litellm/integrations/test_prometheus.py +++ b/tests/test_litellm/integrations/test_prometheus.py @@ -230,12 +230,8 @@ def test_prometheus_config_parsing(): "litellm_proxy_total_requests_metric", ], "include_labels": [ - "litellm_model_name", "requested_model", - "api_base", - "api_provider", - "exception_status", - "exception_class", + "team", ], } ] @@ -251,12 +247,8 @@ def test_prometheus_config_parsing(): # Verify label filters exist for each metric expected_labels = [ - "litellm_model_name", "requested_model", - "api_base", - "api_provider", - "exception_status", - "exception_class", + "team", ] expected_metrics = [ @@ -375,6 +367,107 @@ def test_basic_functionality(): print("Basic prometheus configuration test passed!") +# ============================================================================== +# VALIDATION TESTS - Test the new validation logic for metrics and labels +# ============================================================================== + +def test_invalid_metric_name_validation(): + """Test that invalid metric names are caught and raise ValueError""" + # Clear registry before test + clear_prometheus_registry() + + # Set up test configuration with invalid metric name + test_config = [ + { + "group": "service_metrics", + "metrics": [ + "invalid_metric_name_that_does_not_exist", + "litellm_deployment_total_requests", # valid metric + ], + "include_labels": ["litellm_model_name"], + } + ] + + litellm.prometheus_metrics_config = test_config + + # Creating PrometheusLogger should raise ValueError due to invalid metric + with pytest.raises(ValueError) as exc_info: + PrometheusLogger() + + # Verify error message contains information about invalid metric + assert "invalid_metric_name_that_does_not_exist" in str(exc_info.value) + assert "Configuration validation failed" in str(exc_info.value) + + +def test_invalid_labels_validation(): + """Test that invalid labels for metrics are caught and raise ValueError""" + # Clear registry before test + clear_prometheus_registry() + + # Set up test configuration with invalid labels + test_config = [ + { + "group": "service_metrics", + "metrics": ["litellm_deployment_total_requests"], + "include_labels": [ + "litellm_model_name", # valid label + "invalid_label_name", # invalid label + "another_invalid_label", # another invalid label + ], + } + ] + + litellm.prometheus_metrics_config = test_config + + # Creating PrometheusLogger should raise ValueError due to invalid labels + with pytest.raises(ValueError) as exc_info: + PrometheusLogger() + + # Verify error message contains information about invalid labels + assert "invalid_label_name" in str(exc_info.value) + assert "Configuration validation failed" in str(exc_info.value) + + +def test_valid_configuration_passes_validation(): + """Test that valid configuration passes validation without errors""" + # Clear registry before test + clear_prometheus_registry() + + # Set up test configuration with all valid metrics and labels + test_config = [ + { + "group": "service_metrics", + "metrics": [ + "litellm_deployment_total_requests", + "litellm_deployment_failure_responses", + ], + "include_labels": [ + "litellm_model_name", + "api_provider", + "requested_model", + ], + } + ] + + litellm.prometheus_metrics_config = test_config + + # This should not raise any exceptions + try: + logger = PrometheusLogger() + # Verify the logger was created successfully + assert logger is not None + assert hasattr(logger, 'enabled_metrics') + assert 'litellm_deployment_total_requests' in logger.enabled_metrics + assert 'litellm_deployment_failure_responses' in logger.enabled_metrics + except Exception as e: + pytest.fail(f"Valid configuration should not raise exception: {e}") + + +# ============================================================================== +# END VALIDATION TESTS +# ============================================================================== + + # ============================================================================== # SEMANTIC VALIDATION TESTS - Detect logical errors in metric increments # ==============================================================================