Merge pull request #9439 from BerriAI/litellm_dev_03_20_2025_p2

support returning api-base on pass-through endpoints +  consistently return 404 if team not found in DB
This commit is contained in:
Krish Dholakia 2025-03-21 10:52:36 -07:00 • committed by GitHub
commit ea1b282512
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 35 additions and 9 deletions

View file

@ -57,7 +57,9 @@ class ProxyBaseLLMRequestProcessing:
"x-litellm-call-id": call_id,
"x-litellm-model-id": model_id,
"x-litellm-cache-key": cache_key,
"x-litellm-model-api-base": api_base,
"x-litellm-model-api-base": (
api_base.split("?")[0] if api_base else None
), # don't include query params, risk of leaking sensitive info
"x-litellm-version": version,
"x-litellm-model-region": model_region,
"x-litellm-response-cost": str(response_cost),

View file

@ -470,7 +470,7 @@ async def update_team(
if existing_team_row is None:
raise HTTPException(
status_code=400,
status_code=404,
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
@ -1137,14 +1137,16 @@ async def delete_team(
team_rows: List[LiteLLM_TeamTable] = []
for team_id in data.team_ids:
try:
team_row_base: BaseModel = (
team_row_base: Optional[BaseModel] = (
await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
)
if team_row_base is None:
raise Exception
except Exception:
raise HTTPException(
status_code=400,
status_code=404,
detail={"error": f"Team not found, passed team_id={team_id}"},
)
team_row_pydantic = LiteLLM_TeamTable(**team_row_base.model_dump())

View file

@ -1,6 +1,7 @@
import ast
import asyncio
import json
import uuid
from base64 import b64encode
from datetime import datetime
from typing import Dict, List, Optional, Union
@ -284,7 +285,9 @@ class HttpPassThroughEndpointHelpers:
@staticmethod
def get_response_headers(
headers: httpx.Headers, litellm_call_id: Optional[str] = None
headers: httpx.Headers,
litellm_call_id: Optional[str] = None,
custom_headers: Optional[dict] = None,
) -> dict:
excluded_headers = {"transfer-encoding", "content-encoding"}
@ -295,6 +298,8 @@ class HttpPassThroughEndpointHelpers:
}
if litellm_call_id:
return_headers["x-litellm-call-id"] = litellm_call_id
if custom_headers:
return_headers.update(custom_headers)
return return_headers
@ -365,8 +370,9 @@ async def pass_through_request( # noqa: PLR0915
query_params: Optional[dict] = None,
stream: Optional[bool] = None,
):
litellm_call_id = str(uuid.uuid4())
url: Optional[httpx.URL] = None
try:
import uuid
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.proxy_server import proxy_logging_obj
@ -416,8 +422,6 @@ async def pass_through_request( # noqa: PLR0915
)
async_client = async_client_obj.client
litellm_call_id = str(uuid.uuid4())
# create logging object
start_time = datetime.now()
logging_obj = Logging(
@ -596,15 +600,31 @@ async def pass_through_request( # noqa: PLR0915
)
)
## CUSTOM HEADERS - `x-litellm-*`
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=litellm_call_id,
model_id=None,
cache_key=None,
api_base=str(url._uri_reference),
)
return Response(
content=content,
status_code=response.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
litellm_call_id=litellm_call_id,
custom_headers=custom_headers,
),
)
except Exception as e:
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=litellm_call_id,
model_id=None,
cache_key=None,
api_base=str(url._uri_reference) if url else None,
)
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - {}".format(
str(e)
@ -616,6 +636,7 @@ async def pass_through_request( # noqa: PLR0915
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST),
headers=custom_headers,
)
else:
error_msg = f"{str(e)}"
@ -624,6 +645,7 @@ async def pass_through_request( # noqa: PLR0915
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
headers=custom_headers,
)