(feat) Add Tag-based budgets on litellm router / proxy (#7236)

* add BudgetConfig

* add _get_tags_from_request_kwargs

* test_tag_budgets_e2e_test_expect_to_fail

* add a check for request tags

* fix _async_get_cache_keys_for_router_budget_limiting

* fix test

* fix _sync_in_memory_spend_with_redis

* _async_get_cache_keys_for_router_budget_limiting

* fix _init_tag_budgets

* fix type casting

* docs show error for tag budget limit hit

* fix _get_tags_from_request_kwargs

* fix undo change
This commit is contained in:
Ishaan Jaff 2024-12-14 17:28:36 -08:00 • committed by GitHub
parent 7b7023789d
commit 7103198805
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 323 additions and 36 deletions

View file

@ -5,6 +5,7 @@ import TabItem from '@theme/TabItem';
LiteLLM Supports setting the following budgets:
- Provider budget - $100/day for OpenAI, $100/day for Azure.
- Model budget - $100/day for gpt-4 https://api-base-1, $100/day for gpt-4o https://api-base-2
- Tag budget - $10/day for tag=`product:chat-bot`, $100/day for tag=`product:chat-bot-2`
## Provider Budgets
@ -269,6 +270,90 @@ Expected response on failure
</Tabs>
## Tag Budgets
Use this to set budgets for tags - example $10/day for tag=`product:chat-bot`, $100/day for tag=`product:chat-bot-2`
### Quick Start
Set tag budgets by setting `tag_budget_config` in your `proxy_config.yaml` file
```yaml
model_list:
- model_name: gpt-4o
litellm_params:
model: openai/gpt-4o
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
tag_budget_config:
product:chat-bot: # (Tag)
max_budget: 0.000000000001 # (USD)
budget_duration: 1d # (Duration)
product:chat-bot-2: # (Tag)
max_budget: 100 # (USD)
budget_duration: 1d # (Duration)
```
#### Make a test request
We expect the first request to succeed, and the second request to fail since we cross the budget for `openai/gpt-4o`
**[Langchain, OpenAI SDK Usage Examples](../proxy/user_keys#request-format)**
<Tabs>
<TabItem label="Successful Call " value = "allowed">
```shell
curl -i http://localhost:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-d '{
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "hi my name is test request"}
],
"metadata": {"tags": ["product:chat-bot"]}
}'
```
</TabItem>
<TabItem label="Unsuccessful call" value = "not-allowed">
Expect this to fail since since we cross the budget for tag=`product:chat-bot`
```shell
curl -i http://localhost:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer sk-1234" \
-d '{
"model": "gpt-4o",
"messages": [
{"role": "user", "content": "hi my name is test request"}
],
"metadata": {"tags": ["product:chat-bot"]}
}
```
Expected response on failure
```json
{
"error": {
"message": "No deployments available - crossed budget: Exceeded budget for tag='product:chat-bot', tag_spend=0.00015250000000000002, tag_budget_limit=1e-12",
"type": "None",
"param": "None",
"code": "429"
}
}
```
</TabItem>
</Tabs>
## Multi-instance setup
If you are using a multi-instance setup, you will need to set the Redis host, port, and password in the `proxy_config.yaml` file. Redis is used to sync the spend across LiteLLM instances.

View file

@ -9,6 +9,7 @@ from typing import Callable, List, Optional, Dict, Union, Any, Literal, get_args
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.caching.caching import Cache, DualCache, RedisCache, InMemoryCache
from litellm.types.llms.bedrock import COHERE_EMBEDDING_INPUT_TYPES
from litellm.types.utils import ImageObject, BudgetConfig
from litellm._logging import (
set_verbose,
_turn_on_debug,
@ -291,6 +292,7 @@ max_user_budget: Optional[float] = None
default_max_internal_user_budget: Optional[float] = None
max_internal_user_budget: Optional[float] = None
internal_user_budget_duration: Optional[str] = None
tag_budget_config: Optional[Dict[str, BudgetConfig]] = None
max_end_user_budget: Optional[float] = None
disable_end_user_cost_tracking: Optional[bool] = None
disable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
@ -989,7 +991,6 @@ ALL_LITELLM_RESPONSE_TYPES = [
TextCompletionResponse,
]
from .types.utils import ImageObject
from .llms.custom_llm import CustomLLM
from .llms.openai_like.chat.handler import OpenAILikeChatConfig
from .llms.galadriel.chat.transformation import GaladrielChatConfig

View file

@ -3,12 +3,12 @@ model_list:
litellm_params:
model: openai/gpt-4o
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
tag_budget_config:
product:chat-bot: # (Tag)
max_budget: 0.000000000001 # (USD)
budget_duration: 1d # (Duration)
- model_name: gpt-4o-mini
litellm_params:
model: openai/gpt-4o-mini
api_key: os.environ/OPENAI_API_KEY
product:chat-bot-2: # (Tag)
max_budget: 100 # (USD)
budget_duration: 1d # (Duration)
budget_duration: 1d # (Duration)

View file

@ -29,6 +29,7 @@ from litellm.caching.redis_cache import RedisPipelineIncrementOperation
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs
from litellm.router_utils.cooldown_callbacks import (
_get_prometheus_logger_from_callbacks,
)
@ -39,7 +40,7 @@ from litellm.types.router import (
LiteLLM_Params,
RouterErrors,
)
from litellm.types.utils import StandardLoggingPayload
from litellm.types.utils import BudgetConfig, StandardLoggingPayload
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -67,8 +68,10 @@ class RouterBudgetLimiting(CustomLogger):
provider_budget_config
)
self.deployment_budget_config: Optional[GenericBudgetConfigType] = None
self.tag_budget_config: Optional[GenericBudgetConfigType] = None
self._init_provider_budgets()
self._init_deployment_budgets(model_list=model_list)
self._init_tag_budgets()
# Add self to litellm callbacks if it's a list
if isinstance(litellm.callbacks, list):
@ -104,33 +107,12 @@ class RouterBudgetLimiting(CustomLogger):
request_kwargs
)
# Build combined cache keys for both provider and deployment budgets
cache_keys = []
provider_configs: Dict[str, GenericBudgetInfo] = {}
deployment_configs: Dict[str, GenericBudgetInfo] = {}
for deployment in healthy_deployments:
# Check provider budgets
if self.provider_budget_config:
provider = self._get_llm_provider_for_deployment(deployment)
if provider is not None:
budget_config = self._get_budget_config_for_provider(provider)
if budget_config is not None:
provider_configs[provider] = budget_config
cache_keys.append(
f"provider_spend:{provider}:{budget_config.time_period}"
)
# Check deployment budgets
if self.deployment_budget_config:
model_id = deployment.get("model_info", {}).get("id")
if model_id is not None:
budget_config = self._get_budget_config_for_deployment(model_id)
if budget_config is not None:
deployment_configs[model_id] = budget_config
cache_keys.append(
f"deployment_spend:{model_id}:{budget_config.time_period}"
)
cache_keys, provider_configs, deployment_configs = (
await self._async_get_cache_keys_for_router_budget_limiting(
healthy_deployments=healthy_deployments,
request_kwargs=request_kwargs,
)
)
# Single cache read for all spend values
if len(cache_keys) > 0:
@ -152,6 +134,9 @@ class RouterBudgetLimiting(CustomLogger):
deployment_configs=deployment_configs,
spend_map=spend_map,
potential_deployments=potential_deployments,
request_tags=_get_tags_from_request_kwargs(
request_kwargs=request_kwargs
),
)
)
@ -171,13 +156,14 @@ class RouterBudgetLimiting(CustomLogger):
provider_configs: Dict[str, GenericBudgetInfo],
deployment_configs: Dict[str, GenericBudgetInfo],
spend_map: Dict[str, float],
request_tags: List[str],
) -> Tuple[List[Dict[str, Any]], str]:
"""
Filter out deployments that have exceeded their budget limit.
Follow budget checks are run here:
- Provider budget
- Deployment budget
- Request tags budget
Returns:
Tuple[List[Dict[str, Any]], str]:
- A tuple containing the filtered deployments
@ -226,11 +212,79 @@ class RouterBudgetLimiting(CustomLogger):
is_within_budget = False
continue
# Check tag budget
if self.tag_budget_config and is_within_budget:
for _tag in request_tags:
_tag_budget_config = self._get_budget_config_for_tag(_tag)
if _tag_budget_config:
_tag_spend = spend_map.get(
f"tag_spend:{_tag}:{_tag_budget_config.time_period}", 0.0
)
if _tag_spend >= _tag_budget_config.budget_limit:
debug_msg = f"Exceeded budget for tag='{_tag}', tag_spend={_tag_spend}, tag_budget_limit={_tag_budget_config.budget_limit}"
verbose_router_logger.debug(debug_msg)
deployment_above_budget_info += f"{debug_msg}\n"
is_within_budget = False
continue
if is_within_budget:
potential_deployments.append(deployment)
return potential_deployments, deployment_above_budget_info
async def _async_get_cache_keys_for_router_budget_limiting(
self,
healthy_deployments: List[Dict[str, Any]],
request_kwargs: Optional[Dict] = None,
) -> Tuple[List[str], Dict[str, GenericBudgetInfo], Dict[str, GenericBudgetInfo]]:
"""
Returns list of cache keys to fetch from router cache for budget limiting and provider and deployment configs
Returns:
Tuple[List[str], Dict[str, GenericBudgetInfo], Dict[str, GenericBudgetInfo]]:
- List of cache keys to fetch from router cache for budget limiting
- Dict of provider budget configs `provider_configs`
- Dict of deployment budget configs `deployment_configs`
"""
cache_keys: List[str] = []
provider_configs: Dict[str, GenericBudgetInfo] = {}
deployment_configs: Dict[str, GenericBudgetInfo] = {}
for deployment in healthy_deployments:
# Check provider budgets
if self.provider_budget_config:
provider = self._get_llm_provider_for_deployment(deployment)
if provider is not None:
budget_config = self._get_budget_config_for_provider(provider)
if budget_config is not None:
provider_configs[provider] = budget_config
cache_keys.append(
f"provider_spend:{provider}:{budget_config.time_period}"
)
# Check deployment budgets
if self.deployment_budget_config:
model_id = deployment.get("model_info", {}).get("id")
if model_id is not None:
budget_config = self._get_budget_config_for_deployment(model_id)
if budget_config is not None:
deployment_configs[model_id] = budget_config
cache_keys.append(
f"deployment_spend:{model_id}:{budget_config.time_period}"
)
# Check tag budgets
if self.tag_budget_config:
request_tags = _get_tags_from_request_kwargs(
request_kwargs=request_kwargs
)
for _tag in request_tags:
_tag_budget_config = self._get_budget_config_for_tag(_tag)
if _tag_budget_config:
cache_keys.append(
f"tag_spend:{_tag}:{_tag_budget_config.time_period}"
)
return cache_keys, provider_configs, deployment_configs
async def _get_or_set_budget_start_time(
self, start_time_key: str, current_time: float, ttl_seconds: int
) -> float:
@ -344,6 +398,22 @@ class RouterBudgetLimiting(CustomLogger):
response_cost=response_cost,
)
request_tags = _get_tags_from_request_kwargs(kwargs)
if len(request_tags) > 0:
for _tag in request_tags:
_tag_budget_config = self._get_budget_config_for_tag(_tag)
if _tag_budget_config:
_tag_spend_key = (
f"tag_spend:{_tag}:{_tag_budget_config.time_period}"
)
_tag_start_time_key = f"tag_budget_start_time:{_tag}"
await self._increment_spend_for_key(
budget_config=_tag_budget_config,
spend_key=_tag_spend_key,
start_time_key=_tag_start_time_key,
response_cost=response_cost,
)
async def _increment_spend_for_key(
self,
budget_config: GenericBudgetInfo,
@ -478,6 +548,12 @@ class RouterBudgetLimiting(CustomLogger):
f"deployment_spend:{model_id}:{config.time_period}"
)
if self.tag_budget_config is not None:
for tag, config in self.tag_budget_config.items():
if config is None:
continue
cache_keys.append(f"tag_spend:{tag}:{config.time_period}")
# Batch fetch current spend values from Redis
redis_values = await self.router_cache.redis_cache.async_batch_get_cache(
key_list=cache_keys
@ -514,6 +590,11 @@ class RouterBudgetLimiting(CustomLogger):
return None
return self.provider_budget_config.get(provider, None)
def _get_budget_config_for_tag(self, tag: str) -> Optional[GenericBudgetInfo]:
if self.tag_budget_config is None:
return None
return self.tag_budget_config.get(tag, None)
def _get_llm_provider_for_deployment(self, deployment: Dict) -> Optional[str]:
try:
_litellm_params: LiteLLM_Params = LiteLLM_Params(
@ -632,10 +713,14 @@ class RouterBudgetLimiting(CustomLogger):
Either:
- provider_budget_config is set
- budgets are set for deployments in the model_list
- tag_budget_config is set
"""
if provider_budget_config is not None:
return True
if litellm.tag_budget_config is not None:
return True
if model_list is None:
return False
@ -707,3 +792,23 @@ class RouterBudgetLimiting(CustomLogger):
verbose_router_logger.debug(
f"Initialized Deployment Budget Config: {self.deployment_budget_config}"
)
def _init_tag_budgets(self):
if litellm.tag_budget_config is None:
return
if self.tag_budget_config is None:
self.tag_budget_config = {}
for _tag, _tag_budget_config in litellm.tag_budget_config.items():
if isinstance(_tag_budget_config, dict):
_tag_budget_config = BudgetConfig(**_tag_budget_config)
_generic_budget_config = GenericBudgetInfo(
time_period=_tag_budget_config.budget_duration,
budget_limit=_tag_budget_config.max_budget,
)
self.tag_budget_config[_tag] = _generic_budget_config
verbose_router_logger.debug(
f"Initialized Tag Budget Config: {self.tag_budget_config}"
)

View file

@ -108,3 +108,27 @@ async def get_deployments_for_tag(
healthy_deployments,
)
return healthy_deployments
def _get_tags_from_request_kwargs(
request_kwargs: Optional[Dict[Any, Any]] = None
) -> List[str]:
"""
Helper to get tags from request kwargs
Args:
request_kwargs: The request kwargs to get tags from
Returns:
List[str]: The tags from the request kwargs
"""
if request_kwargs is None:
return []
if "metadata" in request_kwargs:
metadata = request_kwargs["metadata"]
return metadata.get("tags", [])
elif "litellm_params" in request_kwargs:
litellm_params = request_kwargs["litellm_params"]
_metadata = litellm_params.get("metadata", {})
return _metadata.get("tags", [])
return []

View file

@ -1662,6 +1662,11 @@ class StandardKeyGenerationConfig(TypedDict, total=False):
personal_key_generation: PersonalUIKeyGenerationConfig
class BudgetConfig(BaseModel):
max_budget: float
budget_duration: str
class LlmProviders(str, Enum):
OPENAI = "openai"
OPENAI_LIKE = "openai_like" # embedding only

View file

@ -22,6 +22,7 @@ import logging
from litellm._logging import verbose_router_logger
import litellm
from datetime import timezone, timedelta
from litellm.types.utils import BudgetConfig
verbose_router_logger.setLevel(logging.DEBUG)
@ -46,6 +47,9 @@ def cleanup_redis():
for key in redis_client.scan_iter("deployment_spend:*"):
print("deleting key", key)
redis_client.delete(key)
for key in redis_client.scan_iter("tag_spend:*"):
print("deleting key", key)
redis_client.delete(key)
except Exception as e:
print(f"Error cleaning up Redis: {str(e)}")
@ -674,3 +678,66 @@ async def test_deployment_budgets_e2e_test_expect_to_fail():
# Verify the error is related to budget exceeded
assert "Exceeded budget for deployment" in str(exc_info.value)
@pytest.mark.asyncio
async def test_tag_budgets_e2e_test_expect_to_fail():
"""
Expected behavior:
- first request passes, all subsequent requests fail
"""
cleanup_redis()
TAG_NAME = "product:chat-bot"
TAG_NAME_2 = "product:chat-bot-2"
litellm.tag_budget_config = {
TAG_NAME: BudgetConfig(max_budget=0.000000000001, budget_duration="1d"),
TAG_NAME_2: BudgetConfig(max_budget=100, budget_duration="1d"),
}
router = Router(
model_list=[
{
"model_name": "openai/gpt-4o-mini", # openai model name
"litellm_params": {
"model": "openai/gpt-4o-mini",
},
},
],
redis_host=os.getenv("REDIS_HOST"),
redis_port=int(os.getenv("REDIS_PORT")),
redis_password=os.getenv("REDIS_PASSWORD"),
)
response = await router.acompletion(
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="openai/gpt-4o-mini",
metadata={"tags": [TAG_NAME]},
)
print(response)
await asyncio.sleep(2.5)
for _ in range(3):
with pytest.raises(Exception) as exc_info:
response = await router.acompletion(
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="openai/gpt-4o-mini",
metadata={"tags": [TAG_NAME]},
)
print(response)
print("response.hidden_params", response._hidden_params)
await asyncio.sleep(0.5)
# Verify the error is related to budget exceeded
assert f"Exceeded budget for tag='{TAG_NAME}'" in str(exc_info.value)
# test with tag-2 expect to pass
for _ in range(2):
response = await router.acompletion(
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="openai/gpt-4o-mini",
metadata={"tags": [TAG_NAME_2]},
)
print(response)