mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
commit
37f2712bba
6 changed files with 232 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
43
litellm/proxy/auth/auth_utils.py
Normal file
43
litellm/proxy/auth/auth_utils.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
103
litellm/tests/test_spend_calculate_endpoint.py
Normal file
103
litellm/tests/test_spend_calculate_endpoint.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue