diff --git a/.github/workflows/interpret_load_test.py b/.github/workflows/interpret_load_test.py index 2eea46778a..9d95c768fc 100644 --- a/.github/workflows/interpret_load_test.py +++ b/.github/workflows/interpret_load_test.py @@ -78,7 +78,7 @@ if __name__ == "__main__": existing_release_body + "\n\n" + "### Don't want to maintain your internal proxy? get in touch 🎉" - + "Hosted Proxy Alpha: https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat" + + "\nHosted Proxy Alpha: https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat" + "\n\n" + "## Load Test LiteLLM Proxy Results" + "\n\n" diff --git a/README.md b/README.md index 6c81181f3a..3caeb830bb 100644 --- a/README.md +++ b/README.md @@ -5,7 +5,7 @@

Call all LLM APIs using the OpenAI format [Bedrock, Huggingface, VertexAI, TogetherAI, Azure, OpenAI, etc.]

-

OpenAI Proxy Server | Enterprise Tier

+

OpenAI Proxy Server | Hosted Proxy (Preview) | Enterprise Tier

PyPI Version @@ -128,7 +128,9 @@ response = completion(model="gpt-3.5-turbo", messages=[{"role": "user", "content # OpenAI Proxy - ([Docs](https://docs.litellm.ai/docs/simple_proxy)) -Set Budgets & Rate limits across multiple projects +Track spend + Load Balance across multiple projects + +[Hosted Proxy (Preview)](https://docs.litellm.ai/docs/hosted) The proxy provides: diff --git a/docs/my-website/docs/hosted.md b/docs/my-website/docs/hosted.md new file mode 100644 index 0000000000..9be6e775dc --- /dev/null +++ b/docs/my-website/docs/hosted.md @@ -0,0 +1,49 @@ +import Image from '@theme/IdealImage'; + +# Hosted LiteLLM Proxy + +LiteLLM maintains the proxy, so you can focus on your core products. + +## [**Get Onboarded**](https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat) + +This is in alpha. Schedule a call with us, and we'll give you a hosted proxy within 30 minutes. + +[**🚨 Schedule Call**](https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat) + +### **Status**: Alpha + +Our proxy is already used in production by customers. + +See our status page for [**live reliability**](https://status.litellm.ai/) + +### **Benefits** +- **No Maintenance, No Infra**: We'll maintain the proxy, and spin up any additional infrastructure (e.g.: separate server for spend logs) to make sure you can load balance + track spend across multiple LLM projects. +- **Reliable**: Our hosted proxy is tested on 1k requests per second, making it reliable for high load. +- **Secure**: LiteLLM is currently undergoing SOC-2 compliance, to make sure your data is as secure as possible. + +### Pricing + +Pricing is based on usage. We can figure out a price that works for your team, on the call. + +[**🚨 Schedule Call**](https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat) + +## **Screenshots** + +### 1. Create keys + + + +### 2. Add Models + + + +### 3. Track spend + + + + +### 4. Configure load balancing + + + +#### [**🚨 Schedule Call**](https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat) \ No newline at end of file diff --git a/docs/my-website/docs/proxy/prod.md b/docs/my-website/docs/proxy/prod.md index 7f1109dd50..980bba5426 100644 --- a/docs/my-website/docs/proxy/prod.md +++ b/docs/my-website/docs/proxy/prod.md @@ -16,7 +16,7 @@ Expected Performance in Production | `/chat/completions` Requests/hour | `126K` | -## 1. Switch of Debug Logging +## 1. Switch off Debug Logging Remove `set_verbose: True` from your config.yaml ```yaml @@ -40,7 +40,7 @@ Use this Docker `CMD`. This will start the proxy with 1 Uvicorn Async Worker CMD ["--port", "4000", "--config", "./proxy_server_config.yaml"] ``` -## 2. Batch write spend updates every 60s +## 3. Batch write spend updates every 60s The default proxy batch write is 10s. This is to make it easy to see spend when debugging locally. @@ -49,11 +49,35 @@ In production, we recommend using a longer interval period of 60s. This reduces ```yaml general_settings: master_key: sk-1234 - proxy_batch_write_at: 5 # 👈 Frequency of batch writing logs to server (in seconds) + proxy_batch_write_at: 60 # 👈 Frequency of batch writing logs to server (in seconds) ``` +## 4. use Redis 'port','host', 'password'. NOT 'redis_url' -## 3. Move spend logs to separate server +When connecting to Redis use redis port, host, and password params. Not 'redis_url'. We've seen a 80 RPS difference between these 2 approaches when using the async redis client. + +This is still something we're investigating. Keep track of it [here](https://github.com/BerriAI/litellm/issues/3188) + +Recommended to do this for prod: + +```yaml +router_settings: + routing_strategy: usage-based-routing-v2 + # redis_url: "os.environ/REDIS_URL" + redis_host: os.environ/REDIS_HOST + redis_port: os.environ/REDIS_PORT + redis_password: os.environ/REDIS_PASSWORD +``` + +## 5. Switch off resetting budgets + +Add this to your config.yaml. (Only spend per Key, User and Team will be tracked - spend per API Call will not be written to the LiteLLM Database) +```yaml +general_settings: + disable_reset_budget: true +``` + +## 6. Move spend logs to separate server (BETA) Writing each spend log to the db can slow down your proxy. In testing we saw a 70% improvement in median response time, by moving writing spend logs to a separate server. @@ -141,24 +165,6 @@ A t2.micro should be sufficient to handle 1k logs / minute on this server. This consumes at max 120MB, and <0.1 vCPU. -## 4. Switch off resetting budgets - -Add this to your config.yaml. (Only spend per Key, User and Team will be tracked - spend per API Call will not be written to the LiteLLM Database) -```yaml -general_settings: - disable_spend_logs: true - disable_reset_budget: true -``` - -## 5. Switch of `litellm.telemetry` - -Switch of all telemetry tracking done by litellm - -```yaml -litellm_settings: - telemetry: False -``` - ## Machine Specifications to Deploy LiteLLM | Service | Spec | CPUs | Memory | Architecture | Version| diff --git a/docs/my-website/docs/proxy/quick_start.md b/docs/my-website/docs/proxy/quick_start.md index a7ca4743bc..050d9b5983 100644 --- a/docs/my-website/docs/proxy/quick_start.md +++ b/docs/my-website/docs/proxy/quick_start.md @@ -348,6 +348,29 @@ query_result = embeddings.embed_query(text) print(f"TITAN EMBEDDINGS") print(query_result[:5]) +``` + + + +This is **not recommended**. There is duplicate logic as the proxy also uses the sdk, which might lead to unexpected errors. + +```python +from litellm import completion + +response = completion( + model="openai/gpt-3.5-turbo", + messages = [ + { + "role": "user", + "content": "this is a test request, write a short poem" + } + ], + api_key="anything", + base_url="http://0.0.0.0:4000" + ) + +print(response) + ``` diff --git a/docs/my-website/docs/routing.md b/docs/my-website/docs/routing.md index c10d804990..5d9b38cc1f 100644 --- a/docs/my-website/docs/routing.md +++ b/docs/my-website/docs/routing.md @@ -279,7 +279,7 @@ router_settings: ``` - + **Default** Picks a deployment based on the provided **Requests per minute (rpm) or Tokens per minute (tpm)** diff --git a/docs/my-website/docusaurus.config.js b/docs/my-website/docusaurus.config.js index 0dadd71d6f..235af3f28c 100644 --- a/docs/my-website/docusaurus.config.js +++ b/docs/my-website/docusaurus.config.js @@ -105,6 +105,12 @@ const config = { label: 'Enterprise', to: "docs/enterprise" }, + { + sidebarId: 'tutorialSidebar', + position: 'left', + label: '🚀 Hosted', + to: "docs/hosted" + }, { href: 'https://github.com/BerriAI/litellm', label: 'GitHub', diff --git a/docs/my-website/img/litellm_hosted_ui_add_models.png b/docs/my-website/img/litellm_hosted_ui_add_models.png new file mode 100644 index 0000000000..207e952297 Binary files /dev/null and b/docs/my-website/img/litellm_hosted_ui_add_models.png differ diff --git a/docs/my-website/img/litellm_hosted_ui_create_key.png b/docs/my-website/img/litellm_hosted_ui_create_key.png new file mode 100644 index 0000000000..039d265806 Binary files /dev/null and b/docs/my-website/img/litellm_hosted_ui_create_key.png differ diff --git a/docs/my-website/img/litellm_hosted_ui_router.png b/docs/my-website/img/litellm_hosted_ui_router.png new file mode 100644 index 0000000000..9f20dd4ab5 Binary files /dev/null and b/docs/my-website/img/litellm_hosted_ui_router.png differ diff --git a/docs/my-website/img/litellm_hosted_usage_dashboard.png b/docs/my-website/img/litellm_hosted_usage_dashboard.png new file mode 100644 index 0000000000..8513551d3e Binary files /dev/null and b/docs/my-website/img/litellm_hosted_usage_dashboard.png differ diff --git a/litellm/__init__.py b/litellm/__init__.py index 21f98e8b36..b9d9891ca2 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -16,11 +16,24 @@ dotenv.load_dotenv() if set_verbose == True: _turn_on_debug() ############################################# +### Callbacks /Logging / Success / Failure Handlers ### input_callback: List[Union[str, Callable]] = [] success_callback: List[Union[str, Callable]] = [] failure_callback: List[Union[str, Callable]] = [] service_callback: List[Union[str, Callable]] = [] callbacks: List[Callable] = [] +_langfuse_default_tags: Optional[ + List[ + Literal[ + "user_api_key_alias", + "user_api_key_user_id", + "user_api_key_user_email", + "user_api_key_team_alias", + "semantic-similarity", + "proxy_base_url", + ] + ] +] = None _async_input_callback: List[Callable] = ( [] ) # internal variable - async custom callbacks are routed here. @@ -32,6 +45,8 @@ _async_failure_callback: List[Callable] = ( ) # internal variable - async custom callbacks are routed here. pre_call_rules: List[Callable] = [] post_call_rules: List[Callable] = [] +## end of callbacks ############# + email: Optional[str] = ( None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 ) diff --git a/litellm/_redis.py b/litellm/_redis.py index 69ff6f3f2c..d7789472c1 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -32,6 +32,25 @@ def _get_redis_kwargs(): return available_args +def _get_redis_url_kwargs(client=None): + if client is None: + client = redis.Redis.from_url + arg_spec = inspect.getfullargspec(redis.Redis.from_url) + + # Only allow primitive arguments + exclude_args = { + "self", + "connection_pool", + "retry", + } + + include_args = ["url"] + + available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args + + return available_args + + def _get_redis_env_kwarg_mapping(): PREFIX = "REDIS_" @@ -91,27 +110,39 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("password", None) elif "host" not in redis_kwargs or redis_kwargs["host"] is None: raise ValueError("Either 'host' or 'url' must be specified for redis.") - litellm.print_verbose(f"redis_kwargs: {redis_kwargs}") + # litellm.print_verbose(f"redis_kwargs: {redis_kwargs}") return redis_kwargs def get_redis_client(**env_overrides): redis_kwargs = _get_redis_client_logic(**env_overrides) if "url" in redis_kwargs and redis_kwargs["url"] is not None: - redis_kwargs.pop( - "connection_pool", None - ) # redis.from_url doesn't support setting your own connection pool - return redis.Redis.from_url(**redis_kwargs) + args = _get_redis_url_kwargs() + url_kwargs = {} + for arg in redis_kwargs: + if arg in args: + url_kwargs[arg] = redis_kwargs[arg] + + return redis.Redis.from_url(**url_kwargs) return redis.Redis(**redis_kwargs) def get_redis_async_client(**env_overrides): redis_kwargs = _get_redis_client_logic(**env_overrides) if "url" in redis_kwargs and redis_kwargs["url"] is not None: - redis_kwargs.pop( - "connection_pool", None - ) # redis.from_url doesn't support setting your own connection pool - return async_redis.Redis.from_url(**redis_kwargs) + args = _get_redis_url_kwargs(client=async_redis.Redis.from_url) + url_kwargs = {} + for arg in redis_kwargs: + if arg in args: + url_kwargs[arg] = redis_kwargs[arg] + else: + litellm.print_verbose( + "REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format( + arg + ) + ) + return async_redis.Redis.from_url(**url_kwargs) + return async_redis.Redis( socket_timeout=5, **redis_kwargs, @@ -124,4 +155,9 @@ def get_redis_connection_pool(**env_overrides): return async_redis.BlockingConnectionPool.from_url( timeout=5, url=redis_kwargs["url"] ) + connection_class = async_redis.Connection + if "ssl" in redis_kwargs and redis_kwargs["ssl"] is not None: + connection_class = async_redis.SSLConnection + redis_kwargs.pop("ssl", None) + redis_kwargs["connection_class"] = connection_class return async_redis.BlockingConnectionPool(timeout=5, **redis_kwargs) diff --git a/litellm/caching.py b/litellm/caching.py index d73112d21c..86e3ef40d1 100644 --- a/litellm/caching.py +++ b/litellm/caching.py @@ -149,18 +149,19 @@ class RedisCache(BaseCache): if password is not None: redis_kwargs["password"] = password + ### HEALTH MONITORING OBJECT ### + if kwargs.get("service_logger_obj", None) is not None and isinstance( + kwargs["service_logger_obj"], ServiceLogging + ): + self.service_logger_obj = kwargs.pop("service_logger_obj") + else: + self.service_logger_obj = ServiceLogging() + redis_kwargs.update(kwargs) self.redis_client = get_redis_client(**redis_kwargs) self.redis_kwargs = redis_kwargs self.async_redis_conn_pool = get_redis_connection_pool(**redis_kwargs) - if "url" in redis_kwargs and redis_kwargs["url"] is not None: - parsed_kwargs = redis.connection.parse_url(redis_kwargs["url"]) - redis_kwargs.update(parsed_kwargs) - self.redis_kwargs.update(parsed_kwargs) - # pop url - self.redis_kwargs.pop("url") - # redis namespaces self.namespace = namespace # for high traffic, we store the redis results in memory and then batch write to redis @@ -172,8 +173,15 @@ class RedisCache(BaseCache): except Exception as e: pass - ### HEALTH MONITORING OBJECT ### - self.service_logger_obj = ServiceLogging() + ### ASYNC HEALTH PING ### + try: + # asyncio.get_running_loop().create_task(self.ping()) + result = asyncio.get_running_loop().create_task(self.ping()) + except Exception: + pass + + ### SYNC HEALTH PING ### + self.redis_client.ping() def init_async_client(self): from ._redis import get_redis_async_client @@ -601,15 +609,72 @@ class RedisCache(BaseCache): print_verbose(f"Error occurred in pipeline read - {str(e)}") return key_value_dict - async def ping(self): + def sync_ping(self) -> bool: + """ + Tests if the sync redis client is correctly setup. + """ + print_verbose(f"Pinging Sync Redis Cache") + start_time = time.time() + try: + response = self.redis_client.ping() + print_verbose(f"Redis Cache PING: {response}") + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + self.service_logger_obj.service_success_hook( + service=ServiceTypes.REDIS, + duration=_duration, + call_type="sync_ping", + ) + return response + except Exception as e: + # NON blocking - notify users Redis is throwing an exception + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + self.service_logger_obj.service_failure_hook( + service=ServiceTypes.REDIS, + duration=_duration, + error=e, + call_type="sync_ping", + ) + print_verbose( + f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}" + ) + traceback.print_exc() + raise e + + async def ping(self) -> bool: _redis_client = self.init_async_client() + start_time = time.time() async with _redis_client as redis_client: print_verbose(f"Pinging Async Redis Cache") try: response = await redis_client.ping() - print_verbose(f"Redis Cache PING: {response}") + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=_duration, + call_type="async_ping", + ) + ) + return response except Exception as e: # NON blocking - notify users Redis is throwing an exception + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=_duration, + error=e, + call_type="async_ping", + ) + ) print_verbose( f"LiteLLM Redis Cache PING: - Got exception from REDIS : {str(e)}" ) @@ -1216,7 +1281,7 @@ class DualCache(BaseCache): self.in_memory_cache.set_cache(key, redis_result[key], **kwargs) for key, value in redis_result.items(): - result[sublist_keys.index(key)] = value + result[keys.index(key)] = value print_verbose(f"async batch get cache: cache result: {result}") return result @@ -1266,7 +1331,6 @@ class DualCache(BaseCache): keys, **kwargs ) - print_verbose(f"in_memory_result: {in_memory_result}") if in_memory_result is not None: result = in_memory_result if None in result and self.redis_cache is not None and local_only == False: @@ -1290,9 +1354,9 @@ class DualCache(BaseCache): key, redis_result[key], **kwargs ) for key, value in redis_result.items(): - result[sublist_keys.index(key)] = value + index = keys.index(key) + result[index] = value - print_verbose(f"async batch get cache: cache result: {result}") return result except Exception as e: traceback.print_exc() diff --git a/litellm/integrations/langfuse.py b/litellm/integrations/langfuse.py index 3b13446a6a..38ab9c994b 100644 --- a/litellm/integrations/langfuse.py +++ b/litellm/integrations/langfuse.py @@ -280,13 +280,13 @@ class LangFuseLogger: clean_metadata = {} if isinstance(metadata, dict): for key, value in metadata.items(): - # generate langfuse tags - if key in [ - "user_api_key_alias", - "user_api_key_user_id", - "user_api_key_team_alias", - "semantic-similarity", - ]: + + # generate langfuse tags - Default Tags sent to Langfuse from LiteLLM Proxy + if ( + litellm._langfuse_default_tags is not None + and isinstance(litellm._langfuse_default_tags, list) + and key in litellm._langfuse_default_tags + ): tags.append(f"{key}:{value}") # clean litellm metadata before logging @@ -300,6 +300,15 @@ class LangFuseLogger: else: clean_metadata[key] = value + if ( + litellm._langfuse_default_tags is not None + and isinstance(litellm._langfuse_default_tags, list) + and "proxy_base_url" in litellm._langfuse_default_tags + ): + proxy_base_url = os.environ.get("PROXY_BASE_URL", None) + if proxy_base_url is not None: + tags.append(f"proxy_base_url:{proxy_base_url}") + api_base = litellm_params.get("api_base", None) if api_base: clean_metadata["api_base"] = api_base diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 7943d5dba9..30a1188fe9 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( @@ -67,7 +67,7 @@ class PrometheusLogger: # unpack kwargs model = kwargs.get("model", "") - response_cost = kwargs.get("response_cost", 0.0) + response_cost = kwargs.get("response_cost", 0.0) or 0 litellm_params = kwargs.get("litellm_params", {}) or {} proxy_server_request = litellm_params.get("proxy_server_request") or {} end_user_id = proxy_server_request.get("body", {}).get("user", None) diff --git a/litellm/integrations/prometheus_services.py b/litellm/integrations/prometheus_services.py index 4171593baf..45f70a8c1a 100644 --- a/litellm/integrations/prometheus_services.py +++ b/litellm/integrations/prometheus_services.py @@ -30,7 +30,6 @@ class PrometheusServicesLogger: raise Exception( "Missing prometheus_client. Run `pip install prometheus-client`" ) - print("INITIALIZES PROMETHEUS SERVICE LOGGER!") self.Histogram = Histogram self.Counter = Counter @@ -45,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 = ( {} @@ -75,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( @@ -121,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: @@ -130,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: @@ -142,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): @@ -152,7 +170,8 @@ class PrometheusServicesLogger: if self.mock_testing: self.mock_testing_success_calls += 1 - print(f"LOGS SUCCESSFUL CALL TO PROMETHEUS - payload={payload}") + 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: @@ -162,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/llms/ollama.py b/litellm/llms/ollama.py index 670c565904..740747c8e4 100644 --- a/litellm/llms/ollama.py +++ b/litellm/llms/ollama.py @@ -253,7 +253,7 @@ def get_ollama_response( model_response["choices"][0]["message"]["content"] = response_json["response"] model_response["created"] = int(time.time()) model_response["model"] = "ollama/" + model - prompt_tokens = response_json.get("prompt_eval_count", len(encoding.encode(prompt))) # type: ignore + prompt_tokens = response_json.get("prompt_eval_count", len(encoding.encode(prompt, disallowed_special=()))) # type: ignore completion_tokens = response_json.get("eval_count", len(response_json.get("message",dict()).get("content", ""))) model_response["usage"] = litellm.Usage( prompt_tokens=prompt_tokens, @@ -355,7 +355,7 @@ async def ollama_acompletion(url, data, model_response, encoding, logging_obj): ] model_response["created"] = int(time.time()) model_response["model"] = "ollama/" + data["model"] - prompt_tokens = response_json.get("prompt_eval_count", len(encoding.encode(data["prompt"]))) # type: ignore + prompt_tokens = response_json.get("prompt_eval_count", len(encoding.encode(data["prompt"], disallowed_special=()))) # type: ignore completion_tokens = response_json.get("eval_count", len(response_json.get("message",dict()).get("content", ""))) model_response["usage"] = litellm.Usage( prompt_tokens=prompt_tokens, diff --git a/litellm/llms/ollama_chat.py b/litellm/llms/ollama_chat.py index aea00a303f..917336d05c 100644 --- a/litellm/llms/ollama_chat.py +++ b/litellm/llms/ollama_chat.py @@ -148,7 +148,7 @@ class OllamaChatConfig: if param == "top_p": optional_params["top_p"] = value if param == "frequency_penalty": - optional_params["repeat_penalty"] = param + optional_params["repeat_penalty"] = value if param == "stop": optional_params["stop"] = value if param == "response_format" and value["type"] == "json_object": diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index 52589d2de4..38478d931b 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -481,10 +481,11 @@ def construct_tool_use_system_prompt( ): # from https://github.com/anthropics/anthropic-cookbook/blob/main/function_calling/function_calling.ipynb tool_str_list = [] for tool in tools: + tool_function = get_attribute_or_key(tool, "function") tool_str = construct_format_tool_for_claude_prompt( - tool["function"]["name"], - tool["function"].get("description", ""), - tool["function"].get("parameters", {}), + get_attribute_or_key(tool_function, "name"), + get_attribute_or_key(tool_function, "description", ""), + get_attribute_or_key(tool_function, "parameters", {}), ) tool_str_list.append(tool_str) tool_use_system_prompt = ( @@ -608,7 +609,8 @@ def convert_to_anthropic_tool_result_xml(message: dict) -> str: """ name = message.get("name") - content = message.get("content") + content = message.get("content", "") + content = content.replace("<", "<").replace(">", ">").replace("&", "&") # We can't determine from openai message format whether it's a successful or # error call result so default to the successful result template @@ -629,13 +631,15 @@ def convert_to_anthropic_tool_result_xml(message: dict) -> str: def convert_to_anthropic_tool_invoke_xml(tool_calls: list) -> str: invokes = "" for tool in tool_calls: - if tool["type"] != "function": + if get_attribute_or_key(tool, "type") != "function": continue - tool_name = tool["function"]["name"] + tool_function = get_attribute_or_key(tool,"function") + tool_name = get_attribute_or_key(tool_function, "name") + tool_arguments = get_attribute_or_key(tool_function, "arguments") parameters = "".join( f"<{param}>{val}\n" - for param, val in json.loads(tool["function"]["arguments"]).items() + for param, val in json.loads(tool_arguments).items() ) invokes += ( "\n" @@ -689,7 +693,7 @@ def anthropic_messages_pt_xml(messages: list): { "type": "text", "text": ( - convert_to_anthropic_tool_result(messages[msg_i]) + convert_to_anthropic_tool_result_xml(messages[msg_i]) if messages[msg_i]["role"] == "tool" else messages[msg_i]["content"] ), @@ -710,7 +714,7 @@ def anthropic_messages_pt_xml(messages: list): if messages[msg_i].get( "tool_calls", [] ): # support assistant tool invoke convertion - assistant_text += convert_to_anthropic_tool_invoke( # type: ignore + assistant_text += convert_to_anthropic_tool_invoke_xml( # type: ignore messages[msg_i]["tool_calls"] ) @@ -822,12 +826,12 @@ def convert_to_anthropic_tool_invoke(tool_calls: list) -> list: anthropic_tool_invoke = [ { "type": "tool_use", - "id": tool["id"], - "name": tool["function"]["name"], - "input": json.loads(tool["function"]["arguments"]), + "id": get_attribute_or_key(tool, "id"), + "name": get_attribute_or_key(get_attribute_or_key(tool, "function"), "name"), + "input": json.loads(get_attribute_or_key(get_attribute_or_key(tool, "function"), "arguments")), } for tool in tool_calls - if tool["type"] == "function" + if get_attribute_or_key(tool, "type") == "function" ] return anthropic_tool_invoke @@ -1048,7 +1052,8 @@ def cohere_message_pt(messages: list): tool_result = convert_openai_message_to_cohere_tool_result(message) tool_results.append(tool_result) else: - prompt += message["content"] + prompt += message["content"] + "\n\n" + prompt = prompt.rstrip() return prompt, tool_results @@ -1122,12 +1127,6 @@ def _gemini_vision_convert_messages(messages: list): Returns: tuple: A tuple containing the prompt (a string) and the processed images (a list of objects representing the images). """ - try: - from PIL import Image - except: - raise Exception( - "gemini image conversion failed please run `pip install Pillow`" - ) try: # given messages for gpt-4 vision, convert them for gemini @@ -1154,6 +1153,12 @@ def _gemini_vision_convert_messages(messages: list): image = _load_image_from_url(img) processed_images.append(image) else: + try: + from PIL import Image + except: + raise Exception( + "gemini image conversion failed please run `pip install Pillow`" + ) # Case 2: Image filepath (e.g. temp.jpeg) given image = Image.open(img) processed_images.append(image) @@ -1370,3 +1375,8 @@ def prompt_factory( return default_pt( messages=messages ) # default that covers Bloom, T-5, any non-chat tuned model (e.g. base Llama2) + +def get_attribute_or_key(tool_or_function, attribute, default=None): + if hasattr(tool_or_function, attribute): + return getattr(tool_or_function, attribute) + return tool_or_function.get(attribute, default) diff --git a/litellm/llms/vertex_ai_anthropic.py b/litellm/llms/vertex_ai_anthropic.py index 9bce746dd6..34709e0c56 100644 --- a/litellm/llms/vertex_ai_anthropic.py +++ b/litellm/llms/vertex_ai_anthropic.py @@ -123,7 +123,7 @@ class VertexAIAnthropicConfig: """ -- Run client init +- Run client init - Support async completion, streaming """ @@ -236,17 +236,19 @@ def completion( if client is None: if vertex_credentials is not None and isinstance(vertex_credentials, str): import google.oauth2.service_account - - json_obj = json.loads(vertex_credentials) - creds = ( google.oauth2.service_account.Credentials.from_service_account_info( - json_obj, + json.loads(vertex_credentials), scopes=["https://www.googleapis.com/auth/cloud-platform"], ) ) ### CHECK IF ACCESS access_token = refresh_auth(credentials=creds) + else: + import google.auth + creds, _ = google.auth.default() + ### CHECK IF ACCESS + access_token = refresh_auth(credentials=creds) vertex_ai_client = AnthropicVertex( project_id=vertex_project, diff --git a/litellm/main.py b/litellm/main.py index 65696b3c0c..87942f7040 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -609,6 +609,7 @@ def completion( "client", "rpm", "tpm", + "max_parallel_requests", "input_cost_per_token", "output_cost_per_token", "input_cost_per_second", @@ -2560,6 +2561,7 @@ def embedding( client = kwargs.pop("client", None) rpm = kwargs.pop("rpm", None) tpm = kwargs.pop("tpm", None) + max_parallel_requests = kwargs.pop("max_parallel_requests", None) model_info = kwargs.get("model_info", None) metadata = kwargs.get("metadata", None) encoding_format = kwargs.get("encoding_format", None) @@ -2617,6 +2619,7 @@ def embedding( "client", "rpm", "tpm", + "max_parallel_requests", "input_cost_per_token", "output_cost_per_token", "input_cost_per_second", @@ -3476,6 +3479,7 @@ def image_generation( "client", "rpm", "tpm", + "max_parallel_requests", "input_cost_per_token", "output_cost_per_token", "hf_model_name", diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 0f7c24576e..d717dc1595 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -3,10 +3,14 @@ model_list: litellm_params: 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 + api_base: https://openai-function-calling-workers.tasslexyz.workers.dev/ + 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/ stream_timeout: 0.001 - rpm: 10 - litellm_params: model: azure/chatgpt-v-2 api_base: os.environ/AZURE_API_BASE @@ -24,30 +28,23 @@ model_list: # api_key: my-fake-key # api_base: https://exampleopenaiendpoint-production.up.railway.app/ +router_settings: + routing_strategy: usage-based-routing-v2 + # redis_url: "os.environ/REDIS_URL" + redis_host: os.environ/REDIS_HOST + redis_port: os.environ/REDIS_PORT + redis_password: os.environ/REDIS_PASSWORD + enable_pre_call_checks: True + 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"] - upperbound_key_generate_params: - max_budget: os.environ/LITELLM_UPPERBOUND_KEYS_MAX_BUDGET -router_settings: - routing_strategy: usage-based-routing-v2 - redis_host: os.environ/REDIS_HOST - redis_password: os.environ/REDIS_PASSWORD - redis_port: os.environ/REDIS_PORT - enable_pre_call_checks: True general_settings: - master_key: sk-1234 - allow_user_auth: true alerting: ["slack"] - store_model_in_db: True // set via environment variable - os.environ["STORE_MODEL_IN_DB"] = "True" - proxy_batch_write_at: 5 # 👈 Frequency of batch writing logs to server (in seconds) - enable_jwt_auth: True - alerting: ["slack"] - litellm_jwtauth: - admin_jwt_scope: "litellm_proxy_admin" - public_key_ttl: os.environ/LITELLM_PUBLIC_KEY_TTL - user_id_jwt_field: "sub" - org_id_jwt_field: "azp" \ No newline at end of file + alerting_threshold: 300 # sends alerts if requests hang for 5min+ and responses take 5min+ + proxy_batch_write_at: 60 # Frequency of batch writing logs to server (in seconds) \ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b697b6e976..ca9926cef0 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -87,6 +87,14 @@ class LiteLLMRoutes(enum.Enum): "/v2/key/info", ] + sso_only_routes: List = [ + "/key/generate", + "/key/update", + "/key/delete", + "/global/spend/logs", + "/global/predict/spend/logs", + ] + management_routes: List = [ # key "/key/generate", "/key/update", diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index d1bf53a6bf..cd62556493 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -13,4 +13,8 @@ model_list: general_settings: store_model_in_db: true master_key: sk-1234 - alerting: ["slack"] \ No newline at end of file + alerting: ["slack"] + +litellm_settings: + success_callback: ["langfuse"] + _langfuse_default_tags: ["user_api_key_alias", "user_api_key_user_id", "user_api_key_user_email", "user_api_key_team_alias", "semantic-similarity", "proxy_base_url"] \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index db85b7ba10..7f8d7d4a36 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1053,6 +1053,11 @@ async def user_api_key_auth( status_code=status.HTTP_403_FORBIDDEN, detail="key not allowed to access this team's info", ) + elif ( + _has_user_setup_sso() + and route in LiteLLMRoutes.sso_only_routes.value + ): + pass else: raise Exception( f"Only master key can be used to generate, delete, update info for new keys/users/teams. Route={route}" @@ -1102,6 +1107,13 @@ async def user_api_key_auth( return UserAPIKeyAuth( api_key=api_key, user_role="proxy_admin", **valid_token_dict ) + elif ( + _has_user_setup_sso() + and route in LiteLLMRoutes.sso_only_routes.value + ): + return UserAPIKeyAuth( + api_key=api_key, user_role="app_owner", **valid_token_dict + ) else: raise Exception( f"This key is made for LiteLLM UI, Tried to access route: {route}. Not allowed" @@ -4166,6 +4178,14 @@ async def audio_transcriptions( file.filename is not None ) # make sure filename passed in (needed for type) + _original_filename = file.filename + file_extension = os.path.splitext(file.filename)[1] + # rename the file to a random hash file name -> we eventuall remove the file and don't want to remove any local files + file.filename = f"tmp-request" + str(uuid.uuid4()) + file_extension + + # IMP - Asserts that we've renamed the uploaded file, since we run os.remove(file.filename), we should rename the original file + assert file.filename != _original_filename + with open(file.filename, "wb+") as f: f.write(await file.read()) try: @@ -5713,6 +5733,20 @@ async def new_user(data: NewUserRequest): "user" # only create a user, don't create key if 'auto_create_key' set to False ) response = await generate_key_helper_fn(**data_json) + + # Admin UI Logic + # if team_id passed add this user to the team + if data_json.get("team_id", None) is not None: + await team_member_add( + data=TeamMemberAddRequest( + team_id=data_json.get("team_id", None), + member=Member( + user_id=data_json.get("user_id", None), + role="user", + user_email=data_json.get("user_email", None), + ), + ) + ) return NewUserResponse( key=response.get("token", ""), expires=response.get("expires", None), @@ -6518,13 +6552,20 @@ async def team_member_add( existing_team_row = await prisma_client.get_data( # type: ignore team_id=data.team_id, table_name="team", query_type="find_unique" ) + if existing_team_row is None: + raise HTTPException( + status_code=404, + detail={ + "error": f"Team not found for team_id={getattr(data, 'team_id', None)}" + }, + ) new_member = data.member existing_team_row.members_with_roles.append(new_member) complete_team_data = LiteLLM_TeamTable( - **existing_team_row.model_dump(), + **_get_pydantic_json_dict(existing_team_row), ) team_row = await prisma_client.update_data( @@ -8112,36 +8153,33 @@ async def auth_callback(request: Request): } user_role = getattr(user_info, "user_role", None) - else: - ## check if user-email in db ## - user_info = await prisma_client.db.litellm_usertable.find_first( - where={"user_email": user_email} - ) - if user_info is not None: - user_defined_values = { - "models": getattr(user_info, "models", user_id_models), - "user_id": getattr(user_info, "user_id", user_id), - "user_email": getattr(user_info, "user_id", user_email), - "user_role": getattr(user_info, "user_role", None), - } - user_role = getattr(user_info, "user_role", None) + ## check if user-email in db ## + user_info = await prisma_client.db.litellm_usertable.find_first( + where={"user_email": user_email} + ) + if user_info is not None: + user_defined_values = { + "models": getattr(user_info, "models", user_id_models), + "user_id": getattr(user_info, "user_id", user_id), + "user_email": getattr(user_info, "user_id", user_email), + "user_role": getattr(user_info, "user_role", None), + } + user_role = getattr(user_info, "user_role", None) - # update id - await prisma_client.db.litellm_usertable.update_many( - where={"user_email": user_email}, data={"user_id": user_id} # type: ignore - ) - elif litellm.default_user_params is not None and isinstance( - litellm.default_user_params, dict - ): - user_defined_values = { - "models": litellm.default_user_params.get( - "models", user_id_models - ), - "user_id": litellm.default_user_params.get("user_id", user_id), - "user_email": litellm.default_user_params.get( - "user_email", user_email - ), - } + # update id + await prisma_client.db.litellm_usertable.update_many( + where={"user_email": user_email}, data={"user_id": user_id} # type: ignore + ) + elif litellm.default_user_params is not None and isinstance( + litellm.default_user_params, dict + ): + user_defined_values = { + "models": litellm.default_user_params.get("models", user_id_models), + "user_id": litellm.default_user_params.get("user_id", user_id), + "user_email": litellm.default_user_params.get( + "user_email", user_email + ), + } except Exception as e: pass diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 02e8a41668..18f1b837f4 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -238,7 +238,10 @@ class ProxyLogging: litellm_params = kwargs.get("litellm_params", {}) model = kwargs.get("model", "") api_base = litellm.get_api_base(model=model, optional_params=litellm_params) - messages = kwargs.get("messages", "") + messages = kwargs.get("messages", None) + # if messages does not exist fallback to "input" + if messages is None: + messages = kwargs.get("input", None) # only use first 100 chars for alerting _messages = str(messages)[:100] @@ -282,7 +285,10 @@ class ProxyLogging: ): if request_data is not None: model = request_data.get("model", "") - messages = request_data.get("messages", "") + messages = request_data.get("messages", None) + if messages is None: + # if messages does not exist fallback to "input" + messages = request_data.get("input", None) trace_id = request_data.get("metadata", {}).get( "trace_id", None ) # get langfuse trace id diff --git a/litellm/router.py b/litellm/router.py index 8145ef619e..fda53eb4fa 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -26,7 +26,12 @@ from litellm.llms.custom_httpx.azure_dall_e_2 import ( CustomHTTPTransport, AsyncCustomHTTPTransport, ) -from litellm.utils import ModelResponse, CustomStreamWrapper, get_utc_datetime +from litellm.utils import ( + ModelResponse, + CustomStreamWrapper, + get_utc_datetime, + calculate_max_parallel_requests, +) import copy from litellm._logging import verbose_router_logger import logging @@ -61,6 +66,7 @@ class Router: num_retries: int = 0, timeout: Optional[float] = None, default_litellm_params={}, # default params for Router.chat.completion.create + default_max_parallel_requests: Optional[int] = None, set_verbose: bool = False, debug_level: Literal["DEBUG", "INFO"] = "INFO", fallbacks: List = [], @@ -198,6 +204,7 @@ class Router: ) # use a dual cache (Redis+In-Memory) for tracking cooldowns, usage, etc. self.default_deployment = None # use this to track the users default deployment, when they want to use model = * + self.default_max_parallel_requests = default_max_parallel_requests if model_list: model_list = copy.deepcopy(model_list) @@ -213,6 +220,7 @@ class Router: ) # cache to track failed call per deployment, if num failed calls within 1 minute > allowed fails, then add it to cooldown self.num_retries = num_retries or litellm.num_retries or 0 self.timeout = timeout or litellm.request_timeout + self.retry_after = retry_after self.routing_strategy = routing_strategy self.fallbacks = fallbacks or litellm.fallbacks @@ -298,7 +306,7 @@ class Router: else: litellm.failure_callback = [self.deployment_callback_on_failure] verbose_router_logger.info( - f"Intialized router with Routing strategy: {self.routing_strategy}\n\nRouting fallbacks: {self.fallbacks}\n\nRouting context window fallbacks: {self.context_window_fallbacks}" + f"Intialized router with Routing strategy: {self.routing_strategy}\n\nRouting fallbacks: {self.fallbacks}\n\nRouting context window fallbacks: {self.context_window_fallbacks}\n\nRouter Redis Caching={self.cache.redis_cache}" ) self.routing_strategy_args = routing_strategy_args @@ -496,7 +504,9 @@ class Router: ) rpm_semaphore = self._get_client( - deployment=deployment, kwargs=kwargs, client_type="rpm_client" + deployment=deployment, + kwargs=kwargs, + client_type="max_parallel_requests", ) if rpm_semaphore is not None and isinstance( @@ -681,7 +691,9 @@ class Router: ### CONCURRENCY-SAFE RPM CHECKS ### rpm_semaphore = self._get_client( - deployment=deployment, kwargs=kwargs, client_type="rpm_client" + deployment=deployment, + kwargs=kwargs, + client_type="max_parallel_requests", ) if rpm_semaphore is not None and isinstance( @@ -803,7 +815,9 @@ class Router: ### CONCURRENCY-SAFE RPM CHECKS ### rpm_semaphore = self._get_client( - deployment=deployment, kwargs=kwargs, client_type="rpm_client" + deployment=deployment, + kwargs=kwargs, + client_type="max_parallel_requests", ) if rpm_semaphore is not None and isinstance( @@ -1049,7 +1063,9 @@ class Router: ) rpm_semaphore = self._get_client( - deployment=deployment, kwargs=kwargs, client_type="rpm_client" + deployment=deployment, + kwargs=kwargs, + client_type="max_parallel_requests", ) if rpm_semaphore is not None and isinstance( @@ -1243,7 +1259,9 @@ class Router: ### CONCURRENCY-SAFE RPM CHECKS ### rpm_semaphore = self._get_client( - deployment=deployment, kwargs=kwargs, client_type="rpm_client" + deployment=deployment, + kwargs=kwargs, + client_type="max_parallel_requests", ) if rpm_semaphore is not None and isinstance( @@ -1862,17 +1880,23 @@ class Router: model_id = model["model_info"]["id"] # ### IF RPM SET - initialize a semaphore ### rpm = litellm_params.get("rpm", None) - if rpm: - semaphore = asyncio.Semaphore(rpm) - cache_key = f"{model_id}_rpm_client" + tpm = litellm_params.get("tpm", None) + max_parallel_requests = litellm_params.get("max_parallel_requests", None) + calculated_max_parallel_requests = calculate_max_parallel_requests( + rpm=rpm, + max_parallel_requests=max_parallel_requests, + tpm=tpm, + default_max_parallel_requests=self.default_max_parallel_requests, + ) + if calculated_max_parallel_requests: + semaphore = asyncio.Semaphore(calculated_max_parallel_requests) + cache_key = f"{model_id}_max_parallel_requests_client" self.cache.set_cache( key=cache_key, value=semaphore, local_only=True, ) - # print("STORES SEMAPHORE IN CACHE") - #### for OpenAI / Azure we need to initalize the Client for High Traffic ######## custom_llm_provider = litellm_params.get("custom_llm_provider") custom_llm_provider = custom_llm_provider or model_name.split("/", 1)[0] or "" @@ -2537,8 +2561,8 @@ class Router: The appropriate client based on the given client_type and kwargs. """ model_id = deployment["model_info"]["id"] - if client_type == "rpm_client": - cache_key = "{}_rpm_client".format(model_id) + if client_type == "max_parallel_requests": + cache_key = "{}_max_parallel_requests_client".format(model_id) client = self.cache.get_cache(key=cache_key, local_only=True) return client elif client_type == "async": @@ -2778,6 +2802,7 @@ class Router: """ if ( self.routing_strategy != "usage-based-routing-v2" + and self.routing_strategy != "simple-shuffle" ): # prevent regressions for other routing strategies, that don't have async get available deployments implemented. return self.get_available_deployment( model=model, @@ -2828,6 +2853,25 @@ class Router: messages=messages, input=input, ) + elif self.routing_strategy == "simple-shuffle": + # if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm + ############## Check if we can do a RPM/TPM based weighted pick ################# + rpm = healthy_deployments[0].get("litellm_params").get("rpm", None) + if rpm is not None: + # use weight-random pick if rpms provided + rpms = [m["litellm_params"].get("rpm", 0) for m in healthy_deployments] + verbose_router_logger.debug(f"\nrpms {rpms}") + total_rpm = sum(rpms) + weights = [rpm / total_rpm for rpm in rpms] + verbose_router_logger.debug(f"\n weights {weights}") + # Perform weighted random pick + selected_index = random.choices(range(len(rpms)), weights=weights)[0] + verbose_router_logger.debug(f"\n selected index, {selected_index}") + deployment = healthy_deployments[selected_index] + verbose_router_logger.info( + f"get_available_deployment for model: {model}, Selected deployment: {self.print_deployment(deployment) or deployment[0]} for model: {model}" + ) + return deployment or deployment[0] if deployment is None: verbose_router_logger.info( diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 4f6364c2ba..39dbcd9d05 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -187,6 +187,7 @@ class LowestTPMLoggingHandler_v2(CustomLogger): request=httpx.Request(method="tpm_rpm_limits", url="https://github.com/BerriAI/litellm"), # type: ignore ), ) + return deployment except Exception as e: if isinstance(e, litellm.RateLimitError): @@ -406,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, diff --git a/litellm/tests/example_config_yaml/cache_with_params.yaml b/litellm/tests/example_config_yaml/cache_with_params.yaml index d43c1d0336..068e2cc4a2 100644 --- a/litellm/tests/example_config_yaml/cache_with_params.yaml +++ b/litellm/tests/example_config_yaml/cache_with_params.yaml @@ -8,4 +8,4 @@ litellm_settings: cache_params: type: "redis" supported_call_types: ["embedding", "aembedding"] - host: "localhost" \ No newline at end of file + host: "os.environ/REDIS_HOST" \ No newline at end of file diff --git a/litellm/tests/test_amazing_vertex_completion.py b/litellm/tests/test_amazing_vertex_completion.py index a27edd06f4..14bdde493c 100644 --- a/litellm/tests/test_amazing_vertex_completion.py +++ b/litellm/tests/test_amazing_vertex_completion.py @@ -615,8 +615,11 @@ def test_gemini_pro_function_calling(): model="gemini-pro", messages=messages, tools=tools, tool_choice="auto" ) print(f"completion: {completion}") - assert completion.choices[0].message.content is None - assert len(completion.choices[0].message.tool_calls) == 1 + # assert completion.choices[0].message.content is None ## GEMINI PRO is very chatty. + if hasattr(completion.choices[0].message, "tool_calls") and isinstance( + completion.choices[0].message.tool_calls, list + ): + assert len(completion.choices[0].message.tool_calls) == 1 except litellm.APIError as e: pass except litellm.RateLimitError as e: diff --git a/litellm/tests/test_bedrock_completion.py b/litellm/tests/test_bedrock_completion.py index 4b1781cd93..ca2ffea5f5 100644 --- a/litellm/tests/test_bedrock_completion.py +++ b/litellm/tests/test_bedrock_completion.py @@ -269,6 +269,30 @@ def test_bedrock_claude_3_tool_calling(): assert isinstance( response.choices[0].message.tool_calls[0].function.arguments, str ) + messages.append( + response.choices[0].message.model_dump() + ) # Add assistant tool invokes + tool_result = ( + '{"location": "Boston", "temperature": "72", "unit": "fahrenheit"}' + ) + # Add user submitted tool results in the OpenAI format + messages.append( + { + "tool_call_id": response.choices[0].message.tool_calls[0].id, + "role": "tool", + "name": response.choices[0].message.tool_calls[0].function.name, + "content": tool_result, + } + ) + # In the second response, Claude should deduce answer from tool results + second_response = completion( + model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0", + messages=messages, + tools=tools, + tool_choice="auto", + ) + print(f"second response: {second_response}") + assert isinstance(second_response.choices[0].message.content, str) except RateLimitError: pass except Exception as e: diff --git a/litellm/tests/test_caching.py b/litellm/tests/test_caching.py index ae7bb42893..16f1b33804 100644 --- a/litellm/tests/test_caching.py +++ b/litellm/tests/test_caching.py @@ -33,6 +33,51 @@ def generate_random_word(length=4): messages = [{"role": "user", "content": "who is ishaan 5222"}] +@pytest.mark.asyncio +async def test_dual_cache_async_batch_get_cache(): + """ + Unit testing for Dual Cache async_batch_get_cache() + - 2 item query + - in_memory result has a partial hit (1/2) + - hit redis for the other -> expect to return None + - expect result = [in_memory_result, None] + """ + from litellm.caching import DualCache, InMemoryCache, RedisCache + + in_memory_cache = InMemoryCache() + redis_cache = RedisCache() # get credentials from environment + dual_cache = DualCache(in_memory_cache=in_memory_cache, redis_cache=redis_cache) + + in_memory_cache.set_cache(key="test_value", value="hello world") + + result = await dual_cache.async_batch_get_cache(keys=["test_value", "test_value_2"]) + + assert result[0] == "hello world" + assert result[1] == None + + +def test_dual_cache_batch_get_cache(): + """ + Unit testing for Dual Cache batch_get_cache() + - 2 item query + - in_memory result has a partial hit (1/2) + - hit redis for the other -> expect to return None + - expect result = [in_memory_result, None] + """ + from litellm.caching import DualCache, InMemoryCache, RedisCache + + in_memory_cache = InMemoryCache() + redis_cache = RedisCache() # get credentials from environment + dual_cache = DualCache(in_memory_cache=in_memory_cache, redis_cache=redis_cache) + + in_memory_cache.set_cache(key="test_value", value="hello world") + + result = dual_cache.batch_get_cache(keys=["test_value", "test_value_2"]) + + assert result[0] == "hello world" + assert result[1] == None + + # @pytest.mark.skip(reason="") def test_caching_dynamic_args(): # test in memory cache try: @@ -390,6 +435,7 @@ async def test_embedding_caching_azure_individual_items_reordered(): @pytest.mark.asyncio async def test_embedding_caching_base_64(): """ """ + litellm.set_verbose = True litellm.cache = Cache( type="redis", host=os.environ["REDIS_HOST"], @@ -408,6 +454,8 @@ async def test_embedding_caching_base_64(): caching=True, encoding_format="base64", ) + await asyncio.sleep(5) + print("\n\nCALL2\n\n") embedding_val_2 = await aembedding( model="azure/azure-embedding-model", input=inputs, @@ -1094,10 +1142,6 @@ def test_custom_redis_cache_params(): port=os.environ["REDIS_PORT"], password=os.environ["REDIS_PASSWORD"], db=0, - ssl=True, - ssl_certfile="./redis_user.crt", - ssl_keyfile="./redis_user_private.key", - ssl_ca_certs="./redis_ca.pem", ) print(litellm.cache.cache.redis_client) @@ -1105,7 +1149,7 @@ def test_custom_redis_cache_params(): litellm.success_callback = [] litellm._async_success_callback = [] except Exception as e: - pytest.fail(f"Error occurred:", e) + pytest.fail(f"Error occurred: {str(e)}") def test_get_cache_key(): diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index fff3464254..09053cf17b 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -2489,7 +2489,7 @@ def test_completion_deep_infra_mistral(): # Gemini tests def test_completion_gemini(): litellm.set_verbose = True - model_name = "gemini/gemini-pro" + model_name = "gemini/gemini-1.5-pro-latest" messages = [{"role": "user", "content": "Hey, how's it going?"}] try: response = completion(model=model_name, messages=messages) diff --git a/litellm/tests/test_function_calling.py b/litellm/tests/test_function_calling.py index f76a082f60..aedc21665b 100644 --- a/litellm/tests/test_function_calling.py +++ b/litellm/tests/test_function_calling.py @@ -221,6 +221,9 @@ def test_parallel_function_call_stream(): # test_parallel_function_call_stream() +@pytest.mark.skip( + reason="Flaky test. Groq function calling is not reliable for ci/cd testing." +) def test_groq_parallel_function_call(): litellm.set_verbose = True try: diff --git a/litellm/tests/test_key_generate_prisma.py b/litellm/tests/test_key_generate_prisma.py index fdb7649d52..08618c9889 100644 --- a/litellm/tests/test_key_generate_prisma.py +++ b/litellm/tests/test_key_generate_prisma.py @@ -120,6 +120,15 @@ async def test_new_user_response(prisma_client): await litellm.proxy.proxy_server.prisma_client.connect() from litellm.proxy.proxy_server import user_api_key_cache + await new_team( + NewTeamRequest( + team_id="ishaan-special-team", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role="proxy_admin", api_key="sk-1234", user_id="1234" + ), + ) + _response = await new_user( data=NewUserRequest( models=["azure-gpt-3.5"], @@ -999,10 +1008,32 @@ def test_generate_and_update_key(prisma_client): async def test(): await litellm.proxy.proxy_server.prisma_client.connect() + + # create team "litellm-core-infra@gmail.com"" + print("creating team litellm-core-infra@gmail.com") + await new_team( + NewTeamRequest( + team_id="litellm-core-infra@gmail.com", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role="proxy_admin", api_key="sk-1234", user_id="1234" + ), + ) + + await new_team( + NewTeamRequest( + team_id="ishaan-special-team", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role="proxy_admin", api_key="sk-1234", user_id="1234" + ), + ) + request = NewUserRequest( - metadata={"team": "litellm-team3", "project": "litellm-project3"}, + metadata={"project": "litellm-project3"}, team_id="litellm-core-infra@gmail.com", ) + key = await new_user(request) print(key) @@ -1015,7 +1046,6 @@ def test_generate_and_update_key(prisma_client): print("\n info for key=", result["info"]) assert result["info"]["max_parallel_requests"] == None assert result["info"]["metadata"] == { - "team": "litellm-team3", "project": "litellm-project3", } assert result["info"]["team_id"] == "litellm-core-infra@gmail.com" @@ -1037,7 +1067,7 @@ def test_generate_and_update_key(prisma_client): # update the team id response2 = await update_key_fn( request=Request, - data=UpdateKeyRequest(key=generated_key, team_id="ishaan"), + data=UpdateKeyRequest(key=generated_key, team_id="ishaan-special-team"), ) print("response2=", response2) @@ -1048,11 +1078,10 @@ def test_generate_and_update_key(prisma_client): print("\n info for key=", result["info"]) assert result["info"]["max_parallel_requests"] == None assert result["info"]["metadata"] == { - "team": "litellm-team3", "project": "litellm-project3", } assert result["info"]["models"] == ["ada", "babbage", "curie", "davinci"] - assert result["info"]["team_id"] == "ishaan" + assert result["info"]["team_id"] == "ishaan-special-team" # cleanup - delete key delete_key_request = KeyRequest(keys=[generated_key]) @@ -1941,6 +1970,15 @@ async def test_master_key_hashing(prisma_client): await litellm.proxy.proxy_server.prisma_client.connect() from litellm.proxy.proxy_server import user_api_key_cache + await new_team( + NewTeamRequest( + team_id="ishaans-special-team", + ), + user_api_key_dict=UserAPIKeyAuth( + user_role="proxy_admin", api_key="sk-1234", user_id="1234" + ), + ) + _response = await new_user( data=NewUserRequest( models=["azure-gpt-3.5"], diff --git a/litellm/tests/test_prometheus_service.py b/litellm/tests/test_prometheus_service.py index 63ff347d3a..9e3441abb5 100644 --- a/litellm/tests/test_prometheus_service.py +++ b/litellm/tests/test_prometheus_service.py @@ -65,23 +65,18 @@ async def test_completion_with_caching_bad_call(): - Assert failure callback gets called """ litellm.set_verbose = True - sl = ServiceLogging(mock_testing=True) + try: - litellm.cache = Cache(type="redis", host="hello-world") + from litellm.caching import RedisCache + litellm.service_callback = ["prometheus_system"] + sl = ServiceLogging(mock_testing=True) - litellm.cache.cache.service_logger_obj = sl - - messages = [{"role": "user", "content": "Hey, how's it going?"}] - response1 = await acompletion( - model="gpt-3.5-turbo", messages=messages, caching=True - ) - response1 = await acompletion( - model="gpt-3.5-turbo", messages=messages, caching=True - ) + RedisCache(host="hello-world", service_logger_obj=sl) except Exception as e: - pass + print(f"Receives exception = {str(e)}") + await asyncio.sleep(5) assert sl.mock_testing_async_failure_hook > 0 assert sl.mock_testing_async_success_hook == 0 assert sl.mock_testing_sync_success_hook == 0 @@ -144,64 +139,3 @@ async def test_router_with_caching(): except Exception as e: pytest.fail(f"An exception occured - {str(e)}") - - -@pytest.mark.asyncio -async def test_router_with_caching_bad_call(): - """ - - Run completion with caching (incorrect credentials) - - Assert failure callback gets called - """ - try: - - def get_azure_params(deployment_name: str): - params = { - "model": f"azure/{deployment_name}", - "api_key": os.environ["AZURE_API_KEY"], - "api_version": os.environ["AZURE_API_VERSION"], - "api_base": os.environ["AZURE_API_BASE"], - } - return params - - model_list = [ - { - "model_name": "azure/gpt-4", - "litellm_params": get_azure_params("chatgpt-v-2"), - "tpm": 100, - }, - { - "model_name": "azure/gpt-4", - "litellm_params": get_azure_params("chatgpt-v-2"), - "tpm": 1000, - }, - ] - - router = litellm.Router( - model_list=model_list, - set_verbose=True, - debug_level="DEBUG", - routing_strategy="usage-based-routing-v2", - redis_host="hello world", - redis_port=os.environ["REDIS_PORT"], - redis_password=os.environ["REDIS_PASSWORD"], - ) - - litellm.service_callback = ["prometheus_system"] - - sl = ServiceLogging(mock_testing=True) - sl.prometheusServicesLogger.mock_testing = True - router.cache.redis_cache.service_logger_obj = sl - - messages = [{"role": "user", "content": "Hey, how's it going?"}] - try: - response1 = await router.acompletion(model="azure/gpt-4", messages=messages) - response1 = await router.acompletion(model="azure/gpt-4", messages=messages) - except Exception as e: - pass - - assert sl.mock_testing_async_failure_hook > 0 - assert sl.mock_testing_async_success_hook == 0 - assert sl.mock_testing_sync_success_hook == 0 - - except Exception as e: - pytest.fail(f"An exception occured - {str(e)}") diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index d58cf7c2f8..052646db81 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -362,7 +362,9 @@ def test_load_router_config(): ] # init with all call types except Exception as e: - pytest.fail("Proxy: Got exception reading config", e) + pytest.fail( + f"Proxy: Got exception reading config: {str(e)}\n{traceback.format_exc()}" + ) # test_load_router_config() diff --git a/litellm/tests/test_router_caching.py b/litellm/tests/test_router_caching.py index 1fb699c177..ebace161c9 100644 --- a/litellm/tests/test_router_caching.py +++ b/litellm/tests/test_router_caching.py @@ -15,6 +15,61 @@ from litellm import Router ## 2. 2 models - openai, azure - 2 diff model groups, 1 caching group +@pytest.mark.asyncio +async def test_router_async_caching_with_ssl_url(): + """ + Tests when a redis url is passed to the router, if caching is correctly setup + """ + try: + router = Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": { + "model": "gpt-3.5-turbo-0613", + "api_key": os.getenv("OPENAI_API_KEY"), + }, + "tpm": 100000, + "rpm": 10000, + }, + ], + redis_url=os.getenv("REDIS_SSL_URL"), + ) + + response = await router.cache.redis_cache.ping() + print(f"response: {response}") + assert response == True + except Exception as e: + pytest.fail(f"An exception occurred - {str(e)}") + + +def test_router_sync_caching_with_ssl_url(): + """ + Tests when a redis url is passed to the router, if caching is correctly setup + """ + try: + router = Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": { + "model": "gpt-3.5-turbo-0613", + "api_key": os.getenv("OPENAI_API_KEY"), + }, + "tpm": 100000, + "rpm": 10000, + }, + ], + redis_url=os.getenv("REDIS_SSL_URL"), + ) + + response = router.cache.redis_cache.sync_ping() + print(f"response: {response}") + assert response == True + except Exception as e: + pytest.fail(f"An exception occurred - {str(e)}") + + @pytest.mark.asyncio async def test_acompletion_caching_on_router(): # tests acompletion + caching on router diff --git a/litellm/tests/test_router_debug_logs.py b/litellm/tests/test_router_debug_logs.py index a768864aeb..0bc711b157 100644 --- a/litellm/tests/test_router_debug_logs.py +++ b/litellm/tests/test_router_debug_logs.py @@ -81,7 +81,7 @@ def test_async_fallbacks(caplog): # Define the expected log messages # - error request, falling back notice, success notice expected_logs = [ - "Intialized router with Routing strategy: simple-shuffle\n\nRouting fallbacks: [{'gpt-3.5-turbo': ['azure/gpt-3.5-turbo']}]\n\nRouting context window fallbacks: None", + "Intialized router with Routing strategy: simple-shuffle\n\nRouting fallbacks: [{'gpt-3.5-turbo': ['azure/gpt-3.5-turbo']}]\n\nRouting context window fallbacks: None\n\nRouter Redis Caching=None", "litellm.acompletion(model=gpt-3.5-turbo)\x1b[31m Exception OpenAIException - Error code: 401 - {'error': {'message': 'Incorrect API key provided: bad-key. You can find your API key at https://platform.openai.com/account/api-keys.', 'type': 'invalid_request_error', 'param': None, 'code': 'invalid_api_key'}}\x1b[0m", "Falling back to model_group = azure/gpt-3.5-turbo", "litellm.acompletion(model=azure/chatgpt-v-2)\x1b[32m 200 OK\x1b[0m", diff --git a/litellm/tests/test_router_max_parallel_requests.py b/litellm/tests/test_router_max_parallel_requests.py new file mode 100644 index 0000000000..f9cac6aafb --- /dev/null +++ b/litellm/tests/test_router_max_parallel_requests.py @@ -0,0 +1,115 @@ +# What is this? +## Unit tests for the max_parallel_requests feature on Router +import sys, os, time, inspect, asyncio, traceback +from datetime import datetime +import pytest + +sys.path.insert(0, os.path.abspath("../..")) +import litellm +from litellm.utils import calculate_max_parallel_requests +from typing import Optional + +""" +- only rpm +- only tpm +- only max_parallel_requests +- max_parallel_requests + rpm +- max_parallel_requests + tpm +- max_parallel_requests + tpm + rpm +""" + + +max_parallel_requests_values = [None, 10] +tpm_values = [None, 20, 300000] +rpm_values = [None, 30] +default_max_parallel_requests = [None, 40] + + +@pytest.mark.parametrize( + "max_parallel_requests, tpm, rpm, default_max_parallel_requests", + [ + (mp, tp, rp, dmp) + for mp in max_parallel_requests_values + for tp in tpm_values + for rp in rpm_values + for dmp in default_max_parallel_requests + ], +) +def test_scenario(max_parallel_requests, tpm, rpm, default_max_parallel_requests): + calculated_max_parallel_requests = calculate_max_parallel_requests( + max_parallel_requests=max_parallel_requests, + rpm=rpm, + tpm=tpm, + default_max_parallel_requests=default_max_parallel_requests, + ) + if max_parallel_requests is not None: + assert max_parallel_requests == calculated_max_parallel_requests + elif rpm is not None: + assert rpm == calculated_max_parallel_requests + elif tpm is not None: + calculated_rpm = int(tpm / 1000 / 6) + if calculated_rpm == 0: + calculated_rpm = 1 + print( + f"test calculated_rpm: {calculated_rpm}, calculated_max_parallel_requests={calculated_max_parallel_requests}" + ) + assert calculated_rpm == calculated_max_parallel_requests + elif default_max_parallel_requests is not None: + assert calculated_max_parallel_requests == default_max_parallel_requests + else: + assert calculated_max_parallel_requests is None + + +@pytest.mark.parametrize( + "max_parallel_requests, tpm, rpm, default_max_parallel_requests", + [ + (mp, tp, rp, dmp) + for mp in max_parallel_requests_values + for tp in tpm_values + for rp in rpm_values + for dmp in default_max_parallel_requests + ], +) +def test_setting_mpr_limits_per_model( + max_parallel_requests, tpm, rpm, default_max_parallel_requests +): + deployment = { + "model_name": "gpt-3.5-turbo", + "litellm_params": { + "model": "gpt-3.5-turbo", + "max_parallel_requests": max_parallel_requests, + "tpm": tpm, + "rpm": rpm, + }, + "model_info": {"id": "my-unique-id"}, + } + + router = litellm.Router( + model_list=[deployment], + default_max_parallel_requests=default_max_parallel_requests, + ) + + mpr_client: Optional[asyncio.Semaphore] = router._get_client( + deployment=deployment, + kwargs={}, + client_type="max_parallel_requests", + ) + + if max_parallel_requests is not None: + assert max_parallel_requests == mpr_client._value + elif rpm is not None: + assert rpm == mpr_client._value + elif tpm is not None: + calculated_rpm = int(tpm / 1000 / 6) + if calculated_rpm == 0: + calculated_rpm = 1 + print( + f"test calculated_rpm: {calculated_rpm}, calculated_max_parallel_requests={mpr_client._value}" + ) + assert calculated_rpm == mpr_client._value + elif default_max_parallel_requests is not None: + assert mpr_client._value == default_max_parallel_requests + else: + assert mpr_client is None + + # raise Exception("it worked!") diff --git a/litellm/tests/test_tpm_rpm_routing copy.py b/litellm/tests/test_tpm_rpm_routing copy.py deleted file mode 100644 index 8fe30cfcc0..0000000000 --- a/litellm/tests/test_tpm_rpm_routing copy.py +++ /dev/null @@ -1,385 +0,0 @@ -#### What this tests #### -# This tests the router's ability to pick deployment with lowest tpm - -import sys, os, asyncio, time, random -from datetime import datetime -import traceback -from dotenv import load_dotenv - -load_dotenv() -import os - -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path -import pytest -from litellm import Router -import litellm -from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler -from litellm.caching import DualCache - -### UNIT TESTS FOR TPM/RPM ROUTING ### - - -def test_tpm_rpm_updated(): - test_cache = DualCache() - model_list = [] - lowest_tpm_logger = LowestTPMLoggingHandler( - router_cache=test_cache, model_list=model_list - ) - model_group = "gpt-3.5-turbo" - deployment_id = "1234" - kwargs = { - "litellm_params": { - "metadata": { - "model_group": "gpt-3.5-turbo", - "deployment": "azure/chatgpt-v-2", - }, - "model_info": {"id": deployment_id}, - } - } - start_time = time.time() - response_obj = {"usage": {"total_tokens": 50}} - end_time = time.time() - lowest_tpm_logger.log_success_event( - response_obj=response_obj, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - current_minute = datetime.now().strftime("%H-%M") - tpm_count_api_key = f"{model_group}:tpm:{current_minute}" - rpm_count_api_key = f"{model_group}:rpm:{current_minute}" - assert ( - response_obj["usage"]["total_tokens"] - == test_cache.get_cache(key=tpm_count_api_key)[deployment_id] - ) - assert 1 == test_cache.get_cache(key=rpm_count_api_key)[deployment_id] - - -# test_tpm_rpm_updated() - - -def test_get_available_deployments(): - test_cache = DualCache() - model_list = [ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "azure/chatgpt-v-2"}, - "model_info": {"id": "1234"}, - }, - { - "model_name": "gpt-3.5-turbo", - "litellm_params": {"model": "azure/chatgpt-v-2"}, - "model_info": {"id": "5678"}, - }, - ] - lowest_tpm_logger = LowestTPMLoggingHandler( - router_cache=test_cache, model_list=model_list - ) - model_group = "gpt-3.5-turbo" - ## DEPLOYMENT 1 ## - deployment_id = "1234" - kwargs = { - "litellm_params": { - "metadata": { - "model_group": "gpt-3.5-turbo", - "deployment": "azure/chatgpt-v-2", - }, - "model_info": {"id": deployment_id}, - } - } - start_time = time.time() - response_obj = {"usage": {"total_tokens": 50}} - end_time = time.time() - lowest_tpm_logger.log_success_event( - response_obj=response_obj, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - ## DEPLOYMENT 2 ## - deployment_id = "5678" - kwargs = { - "litellm_params": { - "metadata": { - "model_group": "gpt-3.5-turbo", - "deployment": "azure/chatgpt-v-2", - }, - "model_info": {"id": deployment_id}, - } - } - start_time = time.time() - response_obj = {"usage": {"total_tokens": 20}} - end_time = time.time() - lowest_tpm_logger.log_success_event( - response_obj=response_obj, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - - ## CHECK WHAT'S SELECTED ## - print( - lowest_tpm_logger.get_available_deployments( - model_group=model_group, - healthy_deployments=model_list, - input=["Hello world"], - ) - ) - assert ( - lowest_tpm_logger.get_available_deployments( - model_group=model_group, - healthy_deployments=model_list, - input=["Hello world"], - )["model_info"]["id"] - == "5678" - ) - - -# test_get_available_deployments() - - -def test_router_get_available_deployments(): - """ - Test if routers 'get_available_deployments' returns the least busy deployment - """ - model_list = [ - { - "model_name": "azure-model", - "litellm_params": { - "model": "azure/gpt-turbo", - "api_key": "os.environ/AZURE_FRANCE_API_KEY", - "api_base": "https://openai-france-1234.openai.azure.com", - "rpm": 1440, - }, - "model_info": {"id": 1}, - }, - { - "model_name": "azure-model", - "litellm_params": { - "model": "azure/gpt-35-turbo", - "api_key": "os.environ/AZURE_EUROPE_API_KEY", - "api_base": "https://my-endpoint-europe-berri-992.openai.azure.com", - "rpm": 6, - }, - "model_info": {"id": 2}, - }, - ] - router = Router( - model_list=model_list, - routing_strategy="usage-based-routing", - set_verbose=False, - num_retries=3, - ) # type: ignore - - print(f"router id's: {router.get_model_ids()}") - ## DEPLOYMENT 1 ## - deployment_id = 1 - kwargs = { - "litellm_params": { - "metadata": { - "model_group": "azure-model", - }, - "model_info": {"id": 1}, - } - } - start_time = time.time() - response_obj = {"usage": {"total_tokens": 50}} - end_time = time.time() - router.lowesttpm_logger.log_success_event( - response_obj=response_obj, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - ## DEPLOYMENT 2 ## - deployment_id = 2 - kwargs = { - "litellm_params": { - "metadata": { - "model_group": "azure-model", - }, - "model_info": {"id": 2}, - } - } - start_time = time.time() - response_obj = {"usage": {"total_tokens": 20}} - end_time = time.time() - router.lowesttpm_logger.log_success_event( - response_obj=response_obj, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - - ## CHECK WHAT'S SELECTED ## - # print(router.lowesttpm_logger.get_available_deployments(model_group="azure-model")) - assert ( - router.get_available_deployment(model="azure-model")["model_info"]["id"] == "2" - ) - - -# test_get_available_deployments() -# test_router_get_available_deployments() - - -def test_router_skip_rate_limited_deployments(): - """ - Test if routers 'get_available_deployments' raises No Models Available error if max tpm would be reached by message - """ - model_list = [ - { - "model_name": "azure-model", - "litellm_params": { - "model": "azure/gpt-turbo", - "api_key": "os.environ/AZURE_FRANCE_API_KEY", - "api_base": "https://openai-france-1234.openai.azure.com", - "tpm": 1440, - }, - "model_info": {"id": 1}, - }, - ] - router = Router( - model_list=model_list, - routing_strategy="usage-based-routing", - set_verbose=False, - num_retries=3, - ) # type: ignore - - ## DEPLOYMENT 1 ## - deployment_id = 1 - kwargs = { - "litellm_params": { - "metadata": { - "model_group": "azure-model", - }, - "model_info": {"id": deployment_id}, - } - } - start_time = time.time() - response_obj = {"usage": {"total_tokens": 1439}} - end_time = time.time() - router.lowesttpm_logger.log_success_event( - response_obj=response_obj, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) - - ## CHECK WHAT'S SELECTED ## - # print(router.lowesttpm_logger.get_available_deployments(model_group="azure-model")) - try: - router.get_available_deployment( - model="azure-model", - messages=[{"role": "user", "content": "Hey, how's it going?"}], - ) - pytest.fail(f"Should have raised No Models Available error") - except Exception as e: - print(f"An exception occurred! {str(e)}") - - -def test_single_deployment_tpm_zero(): - import litellm - import os - from datetime import datetime - - model_list = [ - { - "model_name": "gpt-3.5-turbo", - "litellm_params": { - "model": "gpt-3.5-turbo", - "api_key": os.getenv("OPENAI_API_KEY"), - "tpm": 0, - }, - } - ] - - router = litellm.Router( - model_list=model_list, - routing_strategy="usage-based-routing", - cache_responses=True, - ) - - model = "gpt-3.5-turbo" - messages = [{"content": "Hello, how are you?", "role": "user"}] - try: - router.get_available_deployment( - model=model, - messages=[{"role": "user", "content": "Hey, how's it going?"}], - ) - pytest.fail(f"Should have raised No Models Available error") - except Exception as e: - print(f"it worked - {str(e)}! \n{traceback.format_exc()}") - - -@pytest.mark.asyncio -async def test_router_completion_streaming(): - messages = [ - {"role": "user", "content": "Hello, can you generate a 500 words poem?"} - ] - model = "azure-model" - model_list = [ - { - "model_name": "azure-model", - "litellm_params": { - "model": "azure/gpt-turbo", - "api_key": "os.environ/AZURE_FRANCE_API_KEY", - "api_base": "https://openai-france-1234.openai.azure.com", - "rpm": 1440, - }, - "model_info": {"id": 1}, - }, - { - "model_name": "azure-model", - "litellm_params": { - "model": "azure/gpt-35-turbo", - "api_key": "os.environ/AZURE_EUROPE_API_KEY", - "api_base": "https://my-endpoint-europe-berri-992.openai.azure.com", - "rpm": 6, - }, - "model_info": {"id": 2}, - }, - ] - router = Router( - model_list=model_list, - routing_strategy="usage-based-routing", - set_verbose=False, - ) # type: ignore - - ### Make 3 calls, test if 3rd call goes to lowest tpm deployment - - ## CALL 1+2 - tasks = [] - response = None - final_response = None - for _ in range(2): - tasks.append(router.acompletion(model=model, messages=messages)) - response = await asyncio.gather(*tasks) - - if response is not None: - ## CALL 3 - await asyncio.sleep(1) # let the token update happen - current_minute = datetime.now().strftime("%H-%M") - picked_deployment = router.lowesttpm_logger.get_available_deployments( - model_group=model, - healthy_deployments=router.healthy_deployments, - messages=messages, - ) - final_response = await router.acompletion(model=model, messages=messages) - print(f"min deployment id: {picked_deployment}") - tpm_key = f"{model}:tpm:{current_minute}" - rpm_key = f"{model}:rpm:{current_minute}" - - tpm_dict = router.cache.get_cache(key=tpm_key) - print(f"tpm_dict: {tpm_dict}") - rpm_dict = router.cache.get_cache(key=rpm_key) - print(f"rpm_dict: {rpm_dict}") - print(f"model id: {final_response._hidden_params['model_id']}") - assert ( - final_response._hidden_params["model_id"] - == picked_deployment["model_info"]["id"] - ) - - -# asyncio.run(test_router_completion_streaming()) diff --git a/litellm/tests/test_tpm_rpm_routing_v2.py b/litellm/tests/test_tpm_rpm_routing_v2.py index 4a0256f6a6..9a43ae3ca1 100644 --- a/litellm/tests/test_tpm_rpm_routing_v2.py +++ b/litellm/tests/test_tpm_rpm_routing_v2.py @@ -23,6 +23,10 @@ from litellm.caching import DualCache ### UNIT TESTS FOR TPM/RPM ROUTING ### +""" +- Given 2 deployments, make sure it's shuffling deployments correctly. +""" + def test_tpm_rpm_updated(): test_cache = DualCache() diff --git a/litellm/utils.py b/litellm/utils.py index 1f23fac981..055f4afdb2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4207,9 +4207,7 @@ def supports_vision(model: str): return True return False else: - raise Exception( - f"Model not in model_prices_and_context_window.json. You passed model={model}." - ) + return False def supports_parallel_function_calling(model: str): @@ -5397,6 +5395,49 @@ def get_optional_params( return optional_params +def calculate_max_parallel_requests( + max_parallel_requests: Optional[int], + rpm: Optional[int], + tpm: Optional[int], + default_max_parallel_requests: Optional[int], +) -> Optional[int]: + """ + Returns the max parallel requests to send to a deployment. + + Used in semaphore for async requests on router. + + Parameters: + - max_parallel_requests - Optional[int] - max_parallel_requests allowed for that deployment + - rpm - Optional[int] - requests per minute allowed for that deployment + - tpm - Optional[int] - tokens per minute allowed for that deployment + - default_max_parallel_requests - Optional[int] - default_max_parallel_requests allowed for any deployment + + Returns: + - int or None (if all params are None) + + Order: + max_parallel_requests > rpm > tpm / 6 (azure formula) > default max_parallel_requests + + Azure RPM formula: + 6 rpm per 1000 TPM + https://learn.microsoft.com/en-us/azure/ai-services/openai/quotas-limits + + + """ + if max_parallel_requests is not None: + return max_parallel_requests + elif rpm is not None: + return rpm + elif tpm is not None: + calculated_rpm = int(tpm / 1000 / 6) + if calculated_rpm == 0: + calculated_rpm = 1 + return calculated_rpm + elif default_max_parallel_requests is not None: + return default_max_parallel_requests + return None + + def get_api_base(model: str, optional_params: dict) -> Optional[str]: """ Returns the api base used for calling the model. @@ -7886,6 +7927,8 @@ def exception_type( elif ( "429 Quota exceeded" in error_str or "IndexError: list index out of range" in error_str + or "429 Unable to submit request because the service is temporarily out of capacity." + in error_str ): exception_mapping_worked = True raise RateLimitError( diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index f69b7ef1a2..7c2d742672 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -55,6 +55,20 @@ model_list: api_base: https://openai-function-calling-workers.tasslexyz.workers.dev/ stream_timeout: 0.001 rpm: 1 + - model_name: fake-openai-endpoint-3 + litellm_params: + model: openai/my-fake-model + api_key: my-fake-key + api_base: https://openai-function-calling-workers.tasslexyz.workers.dev/ + stream_timeout: 0.001 + rpm: 10 + - model_name: fake-openai-endpoint-3 + litellm_params: + model: openai/my-fake-model-2 + api_key: my-fake-key + api_base: https://openai-function-calling-workers.tasslexyz.workers.dev/ + stream_timeout: 0.001 + rpm: 10 - model_name: "*" litellm_params: model: openai/* @@ -82,9 +96,9 @@ litellm_settings: router_settings: routing_strategy: usage-based-routing-v2 - redis_host: os.environ/REDIS_HOST - redis_password: os.environ/REDIS_PASSWORD - redis_port: os.environ/REDIS_PORT + # redis_host: os.environ/REDIS_HOST + # redis_password: os.environ/REDIS_PASSWORD + # redis_port: os.environ/REDIS_PORT enable_pre_call_checks: true general_settings: diff --git a/pyproject.toml b/pyproject.toml index 8cfbf88fdd..a5de973741 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "1.35.15" +version = "1.35.17" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT" @@ -80,7 +80,7 @@ requires = ["poetry-core", "wheel"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "1.35.15" +version = "1.35.17" version_files = [ "pyproject.toml:^version" ] diff --git a/tests/test_keys.py b/tests/test_keys.py index 39787eb97f..f21c50c0dd 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -14,6 +14,24 @@ sys.path.insert( import litellm +async def generate_team(session): + url = "http://0.0.0.0:4000/team/new" + headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + data = { + "team_id": "litellm-dashboard", + } + + async with session.post(url, headers=headers, json=data) as response: + status = response.status + response_text = await response.text() + + print(f"Response (Status code: {status}):") + print(response_text) + print() + _json_response = await response.json() + return _json_response + + async def generate_user( session, user_role="app_owner", @@ -668,7 +686,7 @@ async def test_key_rate_limit(): @pytest.mark.asyncio -async def test_key_delete(): +async def test_key_delete_ui(): """ Admin UI flow - DO NOT DELETE -> Create a key with user_id = "ishaan" @@ -680,6 +698,8 @@ async def test_key_delete(): key = key_gen["key"] # generate a admin UI key + team = await generate_team(session=session) + print("generated team: ", team) admin_ui_key = await generate_user(session=session, user_role="proxy_admin") print( "trying to delete key=", diff --git a/tests/test_openai_endpoints.py b/tests/test_openai_endpoints.py index f6bf218aec..c77eeba5b0 100644 --- a/tests/test_openai_endpoints.py +++ b/tests/test_openai_endpoints.py @@ -102,6 +102,47 @@ async def chat_completion(session, key, model="gpt-4"): return await response.json() +async def chat_completion_with_headers(session, key, model="gpt-4"): + url = "http://0.0.0.0:4000/chat/completions" + headers = { + "Authorization": f"Bearer {key}", + "Content-Type": "application/json", + } + data = { + "model": model, + "messages": [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello!"}, + ], + } + + async with session.post(url, headers=headers, json=data) as response: + status = response.status + response_text = await response.text() + + print(response_text) + print() + + if status != 200: + raise Exception(f"Request did not return a 200 status code: {status}") + + response_header_check( + response + ) # calling the function to check response headers + + raw_headers = response.raw_headers + raw_headers_json = {} + + for ( + item + ) in ( + response.raw_headers + ): # ((b'date', b'Fri, 19 Apr 2024 21:17:29 GMT'), (), ) + raw_headers_json[item[0].decode("utf-8")] = item[1].decode("utf-8") + + return raw_headers_json + + async def completion(session, key): url = "http://0.0.0.0:4000/completions" headers = { @@ -218,6 +259,39 @@ async def test_chat_completion_ratelimit(): try: await asyncio.gather(*tasks) pytest.fail("Expected at least 1 call to fail") + except Exception as e: + if "Request did not return a 200 status code: 429" in str(e): + pass + else: + pytest.fail(f"Wrong error received - {str(e)}") + + +@pytest.mark.asyncio +async def test_chat_completion_different_deployments(): + """ + - call model group with 2 deployments + - make 5 calls + - expect 2 unique deployments + """ + async with aiohttp.ClientSession() as session: + # key_gen = await generate_key(session=session) + key = "sk-1234" + results = [] + for _ in range(5): + results.append( + await chat_completion_with_headers( + session=session, key=key, model="fake-openai-endpoint-3" + ) + ) + try: + print(f"results: {results}") + init_model_id = results[0]["x-litellm-model-id"] + deployments_shuffled = False + for result in results[1:]: + if init_model_id != result["x-litellm-model-id"]: + deployments_shuffled = True + if deployments_shuffled == False: + pytest.fail("Expected at least 1 shuffled call") except Exception as e: pass diff --git a/ui/litellm-dashboard/src/components/create_key_button.tsx b/ui/litellm-dashboard/src/components/create_key_button.tsx index 8dde3fb001..d8716d304e 100644 --- a/ui/litellm-dashboard/src/components/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/create_key_button.tsx @@ -2,7 +2,7 @@ import React, { useState, useEffect, useRef } from "react"; import { Button, TextInput, Grid, Col } from "@tremor/react"; -import { Card, Metric, Text, Title, Subtitle } from "@tremor/react"; +import { Card, Metric, Text, Title, Subtitle, Accordion, AccordionHeader, AccordionBody, } from "@tremor/react"; import { CopyToClipboard } from 'react-copy-to-clipboard'; import { Button as Button2, @@ -248,16 +248,143 @@ const CreateKey: React.FC = ({ ) : ( <> - + - - + - - + + + + + + Optional Settings + + + { + if (value && team && team.max_budget !== null && value > team.max_budget) { + throw new Error(`Budget cannot exceed team max budget: $${team.max_budget}`); + } + }, + }, + ]} + > + + + + + + { + if (value && team && team.tpm_limit !== null && value > team.tpm_limit) { + throw new Error(`TPM limit cannot exceed team TPM limit: ${team.tpm_limit}`); + } + }, + }, + ]} + > + + + { + if (value && team && team.rpm_limit !== null && value > team.rpm_limit) { + throw new Error(`RPM limit cannot exceed team RPM limit: ${team.rpm_limit}`); + } + }, + }, + ]} + > + + + + + + + + + + + + )}
diff --git a/ui/litellm-dashboard/src/components/dashboard_default_team.tsx b/ui/litellm-dashboard/src/components/dashboard_default_team.tsx index b3976912b2..c845ef1508 100644 --- a/ui/litellm-dashboard/src/components/dashboard_default_team.tsx +++ b/ui/litellm-dashboard/src/components/dashboard_default_team.tsx @@ -4,6 +4,7 @@ import { Select, SelectItem, Text, Title } from "@tremor/react"; interface DashboardTeamProps { teams: Object[] | null; setSelectedTeam: React.Dispatch>; + userRole: string | null; } type TeamInterface = { @@ -15,6 +16,7 @@ type TeamInterface = { const DashboardTeam: React.FC = ({ teams, setSelectedTeam, + userRole, }) => { const defaultTeam: TeamInterface = { models: [], @@ -25,19 +27,27 @@ const DashboardTeam: React.FC = ({ const [value, setValue] = useState(defaultTeam); - const updatedTeams = teams ? [...teams, defaultTeam] : [defaultTeam]; - + let updatedTeams; + if (userRole === "App User") { + // Non-Admin SSO users should only see their own team - they should not see "Default Team" + updatedTeams = teams; + } else { + updatedTeams = teams ? [...teams, defaultTeam] : [defaultTeam]; + } return (
Select Team - - If you belong to multiple teams, this setting controls which team is - used by default when creating new API Keys. - - - Default Team: If no team_id is set for a key, it will be grouped under here. - + {userRole !== "App User" && ( + <> + + If you belong to multiple teams, this setting controls which team is used by default when creating new API Keys. + + + Default Team: If no team_id is set for a key, it will be grouped under here. + + + )} {updatedTeams && updatedTeams.length > 0 ? (