From 9379e3d0472860045f00057c9137185d1147527c Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 20 Apr 2024 16:13:11 -0700 Subject: [PATCH] fix(lowest_tpm_rpm_v2.py): use a combined tpm+rpm query in async get cache, to reduce redis client calls in high traffic --- litellm/integrations/prometheus.py | 2 +- litellm/integrations/prometheus_services.py | 53 +++++++++++++++----- litellm/proxy/_new_secret_config.yaml | 15 ++---- litellm/router_strategy/lowest_tpm_rpm_v2.py | 12 +++-- 4 files changed, 54 insertions(+), 28 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 74632d49a06..30a1188fe9c 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -19,7 +19,7 @@ class PrometheusLogger: **kwargs, ): try: - verbose_logger.debug(f"in init prometheus metrics") + print(f"in init prometheus metrics") from prometheus_client import Counter self.litellm_llm_api_failed_requests_metric = Counter( diff --git a/litellm/integrations/prometheus_services.py b/litellm/integrations/prometheus_services.py index 548d0a2a3af..45f70a8c1ad 100644 --- a/litellm/integrations/prometheus_services.py +++ b/litellm/integrations/prometheus_services.py @@ -44,9 +44,18 @@ class PrometheusServicesLogger: ) # store the prometheus histogram/counter we need to call for each field in payload for service in self.services: - histogram = self.create_histogram(service) - counter = self.create_counter(service) - self.payload_to_prometheus_map[service] = [histogram, counter] + histogram = self.create_histogram(service, type_of_request="latency") + counter_failed_request = self.create_counter( + service, type_of_request="failed_requests" + ) + counter_total_requests = self.create_counter( + service, type_of_request="total_requests" + ) + self.payload_to_prometheus_map[service] = [ + histogram, + counter_failed_request, + counter_total_requests, + ] self.prometheus_to_amount_map: dict = ( {} @@ -74,26 +83,26 @@ class PrometheusServicesLogger: return metric return None - def create_histogram(self, label: str): - metric_name = "litellm_{}_latency".format(label) + def create_histogram(self, service: str, type_of_request: str): + metric_name = "litellm_{}_{}".format(service, type_of_request) is_registered = self.is_metric_registered(metric_name) if is_registered: return self.get_metric(metric_name) return self.Histogram( metric_name, - "Latency for {} service".format(label), - labelnames=[label], + "Latency for {} service".format(service), + labelnames=[service], ) - def create_counter(self, label: str): - metric_name = "litellm_{}_failed_requests".format(label) + def create_counter(self, service: str, type_of_request: str): + metric_name = "litellm_{}_{}".format(service, type_of_request) is_registered = self.is_metric_registered(metric_name) if is_registered: return self.get_metric(metric_name) return self.Counter( metric_name, - "Total failed requests for {} service".format(label), - labelnames=[label], + "Total {} for {} service".format(type_of_request, service), + labelnames=[service], ) def observe_histogram( @@ -120,6 +129,8 @@ class PrometheusServicesLogger: if self.mock_testing: self.mock_testing_success_calls += 1 + print(f"payload call type: {payload.call_type}") + if payload.service.value in self.payload_to_prometheus_map: prom_objects = self.payload_to_prometheus_map[payload.service.value] for obj in prom_objects: @@ -129,11 +140,19 @@ class PrometheusServicesLogger: labels=payload.service.value, amount=payload.duration, ) + elif isinstance(obj, self.Counter) and "total_requests" in obj._name: + self.increment_counter( + counter=obj, + labels=payload.service.value, + amount=1, # LOG TOTAL REQUESTS TO PROMETHEUS + ) def service_failure_hook(self, payload: ServiceLoggerPayload): if self.mock_testing: self.mock_testing_failure_calls += 1 + print(f"payload call type: {payload.call_type}") + if payload.service.value in self.payload_to_prometheus_map: prom_objects = self.payload_to_prometheus_map[payload.service.value] for obj in prom_objects: @@ -141,7 +160,7 @@ class PrometheusServicesLogger: self.increment_counter( counter=obj, labels=payload.service.value, - amount=1, # LOG ERROR COUNT TO PROMETHEUS + amount=1, # LOG ERROR COUNT / TOTAL REQUESTS TO PROMETHEUS ) async def async_service_success_hook(self, payload: ServiceLoggerPayload): @@ -151,6 +170,8 @@ class PrometheusServicesLogger: if self.mock_testing: self.mock_testing_success_calls += 1 + print(f"payload call type: {payload.call_type}") + if payload.service.value in self.payload_to_prometheus_map: prom_objects = self.payload_to_prometheus_map[payload.service.value] for obj in prom_objects: @@ -160,12 +181,20 @@ class PrometheusServicesLogger: labels=payload.service.value, amount=payload.duration, ) + elif isinstance(obj, self.Counter) and "total_requests" in obj._name: + self.increment_counter( + counter=obj, + labels=payload.service.value, + amount=1, # LOG TOTAL REQUESTS TO PROMETHEUS + ) async def async_service_failure_hook(self, payload: ServiceLoggerPayload): print(f"received error payload: {payload.error}") if self.mock_testing: self.mock_testing_failure_calls += 1 + print(f"payload call type: {payload.call_type}") + if payload.service.value in self.payload_to_prometheus_map: prom_objects = self.payload_to_prometheus_map[payload.service.value] for obj in prom_objects: diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 53c59ff8a7e..d717dc15958 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -4,14 +4,12 @@ model_list: model: openai/my-fake-model api_key: my-fake-key api_base: https://openai-function-calling-workers.tasslexyz.workers.dev/ - # api_base: http://0.0.0.0:8080 stream_timeout: 0.001 - model_name: fake-openai-endpoint litellm_params: model: openai/my-fake-model-2 api_key: my-fake-key api_base: https://openai-function-calling-workers.tasslexyz.workers.dev/ - # api_base: http://0.0.0.0:8080 stream_timeout: 0.001 - litellm_params: model: azure/chatgpt-v-2 @@ -30,15 +28,8 @@ model_list: # api_key: my-fake-key # api_base: https://exampleopenaiendpoint-production.up.railway.app/ -# litellm_settings: -# success_callback: ["prometheus"] -# failure_callback: ["prometheus"] -# service_callback: ["prometheus_system"] -# upperbound_key_generate_params: -# max_budget: os.environ/LITELLM_UPPERBOUND_KEYS_MAX_BUDGET - router_settings: - # routing_strategy: usage-based-routing-v2 + routing_strategy: usage-based-routing-v2 # redis_url: "os.environ/REDIS_URL" redis_host: os.environ/REDIS_HOST redis_port: os.environ/REDIS_PORT @@ -48,6 +39,10 @@ router_settings: litellm_settings: num_retries: 3 # retry call 3 times on each model_name allowed_fails: 3 # cooldown model if it fails > 1 call in a minute. + success_callback: ["prometheus"] + failure_callback: ["prometheus"] + service_callback: ["prometheus_system"] + general_settings: alerting: ["slack"] diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index b2b6df42bfe..39dbcd9d059 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -407,13 +407,15 @@ class LowestTPMLoggingHandler_v2(CustomLogger): tpm_keys.append(tpm_key) rpm_keys.append(rpm_key) - tpm_values = await self.router_cache.async_batch_get_cache( - keys=tpm_keys - ) # [1, 2, None, ..] - rpm_values = await self.router_cache.async_batch_get_cache( - keys=rpm_keys + combined_tpm_rpm_keys = tpm_keys + rpm_keys + + combined_tpm_rpm_values = await self.router_cache.async_batch_get_cache( + keys=combined_tpm_rpm_keys ) # [1, 2, None, ..] + tpm_values = combined_tpm_rpm_values[: len(tpm_keys)] + rpm_values = combined_tpm_rpm_values[len(tpm_keys) :] + return self._common_checks_available_deployment( model_group=model_group, healthy_deployments=healthy_deployments,