diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 7f131efb043..fcc13509ceb 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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), diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 1994e27ecf5..f5bcc6ba11c 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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()) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index b13d614678a..a13b0dc216e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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, )