Merge pull request #4389 from BerriAI/litellm_allow_user_to_define_public_routes

[Feat-Enterprise] - Allow setting custom public routes
This commit is contained in:
Ishaan Jaff 2024-06-24 20:23:35 -07:00 • committed by GitHub
commit 37f2712bba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 232 additions and 8 deletions

View file

@ -1627,3 +1627,9 @@ class CommonProxyErrors(enum.Enum):
no_llm_router = "No models configured on proxy"
not_allowed_access = "Admin-only endpoint. Not allowed to access this."
not_premium_user = "You must be a LiteLLM Enterprise user to use this feature. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://calendly.com/d/4mp-gd3-k5k/litellm-1-1-onboarding-chat"
class SpendCalculateRequest(LiteLLMBase):
model: Optional[str] = None
messages: Optional[List] = None
completion_response: Optional[dict] = None

View file

@ -0,0 +1,43 @@
from litellm._logging import verbose_proxy_logger
def route_in_additonal_public_routes(current_route: str):
"""
Helper to check if the user defined public_routes on config.yaml
Parameters:
- current_route: str - the route the user is trying to call
Returns:
- bool - True if the route is defined in public_routes
- bool - False if the route is not defined in public_routes
In order to use this the litellm config.yaml should have the following in general_settings:
```yaml
general_settings:
master_key: sk-1234
public_routes: ["LiteLLMRoutes.public_routes", "/spend/calculate"]
```
"""
# check if user is premium_user - if not do nothing
from litellm.proxy._types import LiteLLMRoutes
from litellm.proxy.proxy_server import general_settings, premium_user
try:
if premium_user is not True:
return False
# check if this is defined on the config
if general_settings is None:
return False
routes_defined = general_settings.get("public_routes", [])
if current_route in routes_defined:
return True
return False
except Exception as e:
verbose_proxy_logger.error(f"route_in_additonal_public_routes: {str(e)}")
return False

View file

@ -56,6 +56,7 @@ from litellm.proxy.auth.auth_checks import (
get_user_object,
log_to_opentelemetry,
)
from litellm.proxy.auth.auth_utils import route_in_additonal_public_routes
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.utils import _to_ns
@ -137,7 +138,10 @@ async def user_api_key_auth(
"""
route: str = request.url.path
if route in LiteLLMRoutes.public_routes.value:
if (
route in LiteLLMRoutes.public_routes.value
or route_in_additonal_public_routes(current_route=route)
):
# check if public endpoint
return UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY)

View file

@ -21,6 +21,8 @@ model_list:
general_settings:
master_key: sk-1234
alerting: ["slack", "email"]
public_routes: ["LiteLLMRoutes.public_routes", "/spend/calculate"]
litellm_settings:
success_callback: ["prometheus"]

View file

@ -1199,7 +1199,7 @@ async def _get_spend_report_for_time_range(
}
},
)
async def calculate_spend(request: Request):
async def calculate_spend(request: SpendCalculateRequest):
"""
Accepts all the params of completion_cost.
@ -1248,14 +1248,80 @@ async def calculate_spend(request: Request):
}'
```
"""
from litellm import completion_cost
try:
from litellm import completion_cost
from litellm.cost_calculator import CostPerToken
from litellm.proxy.proxy_server import llm_router
data = await request.json()
if "completion_response" in data:
data["completion_response"] = litellm.ModelResponse(
**data["completion_response"]
_cost = None
if request.model is not None:
if request.messages is None:
raise HTTPException(
status_code=400,
detail="Bad Request - messages must be provided if 'model' is provided",
)
# check if model in llm_router
_model_in_llm_router = None
cost_per_token: Optional[CostPerToken] = None
if llm_router is not None:
for model in llm_router.model_list:
if model.get("model_name") == request.model:
_model_in_llm_router = model
"""
3 cases for /spend/calculate
1. user passes model, and model is defined on litellm config.yaml or in DB. use info on config or in DB in this case
2. user passes model, and model is not defined on litellm config.yaml or in DB. Pass model as is to litellm.completion_cost
3. user passes completion_response
"""
if _model_in_llm_router is not None:
_litellm_params = _model_in_llm_router.get("litellm_params")
_litellm_model_name = _litellm_params.get("model")
input_cost_per_token = _litellm_params.get("input_cost_per_token")
output_cost_per_token = _litellm_params.get("output_cost_per_token")
if (
input_cost_per_token is not None
or output_cost_per_token is not None
):
cost_per_token = CostPerToken(
input_cost_per_token=input_cost_per_token,
output_cost_per_token=output_cost_per_token,
)
_cost = completion_cost(
model=_litellm_model_name,
messages=request.messages,
custom_cost_per_token=cost_per_token,
)
else:
_cost = completion_cost(model=request.model, messages=request.messages)
elif request.completion_response is not None:
_completion_response = litellm.ModelResponse(**request.completion_response)
_cost = completion_cost(completion_response=_completion_response)
else:
raise HTTPException(
status_code=400,
detail="Bad Request - Either 'model' or 'completion_response' must be provided",
)
return {"cost": _cost}
except Exception as e:
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "detail", str(e)),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
)
error_msg = f"{str(e)}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
)
return {"cost": completion_cost(**data)}
@router.get(

View file

@ -0,0 +1,103 @@
import os
import sys
import pytest
from dotenv import load_dotenv
from fastapi import Request
from fastapi.routing import APIRoute
import litellm
from litellm.proxy._types import SpendCalculateRequest
from litellm.proxy.spend_tracking.spend_management_endpoints import calculate_spend
from litellm.router import Router
# this file is to test litellm/proxy
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
@pytest.mark.asyncio
async def test_spend_calc_model_messages():
cost_obj = await calculate_spend(
request=SpendCalculateRequest(
model="gpt-3.5-turbo",
messages=[
{"role": "user", "content": "What is the capital of France?"},
],
)
)
print("calculated cost", cost_obj)
cost = cost_obj["cost"]
assert cost > 0.0
@pytest.mark.asyncio
async def test_spend_calc_model_on_router_messages():
from litellm.proxy.proxy_server import llm_router as init_llm_router
temp_llm_router = Router(
model_list=[
{
"model_name": "special-llama-model",
"litellm_params": {
"model": "groq/llama3-8b-8192",
},
}
]
)
setattr(litellm.proxy.proxy_server, "llm_router", temp_llm_router)
cost_obj = await calculate_spend(
request=SpendCalculateRequest(
model="special-llama-model",
messages=[
{"role": "user", "content": "What is the capital of France?"},
],
)
)
print("calculated cost", cost_obj)
_cost = cost_obj["cost"]
assert _cost > 0.0
# set router to init value
setattr(litellm.proxy.proxy_server, "llm_router", init_llm_router)
@pytest.mark.asyncio
async def test_spend_calc_using_response():
cost_obj = await calculate_spend(
request=SpendCalculateRequest(
completion_response={
"id": "chatcmpl-3bc7abcd-f70b-48ab-a16c-dfba0b286c86",
"choices": [
{
"finish_reason": "stop",
"index": 0,
"message": {
"content": "Yooo! What's good?",
"role": "assistant",
},
}
],
"created": "1677652288",
"model": "groq/llama3-8b-8192",
"object": "chat.completion",
"system_fingerprint": "fp_873a560973",
"usage": {
"completion_tokens": 8,
"prompt_tokens": 12,
"total_tokens": 20,
},
}
)
)
print("calculated cost", cost_obj)
cost = cost_obj["cost"]
assert cost > 0.0