diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py index d4964b9667..8db0fcf752 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py @@ -119,6 +119,7 @@ class PagerDutyAlerting(SlackAlerting): user_api_key_end_user_id=_meta.get("user_api_key_end_user_id"), user_api_key_user_email=_meta.get("user_api_key_user_email"), user_api_key_request_route=_meta.get("user_api_key_request_route"), + user_api_key_auth_metadata=_meta.get("user_api_key_auth_metadata"), ) ) @@ -196,7 +197,11 @@ class PagerDutyAlerting(SlackAlerting): user_api_key_alias=user_api_key_dict.key_alias, user_api_key_spend=user_api_key_dict.spend, user_api_key_max_budget=user_api_key_dict.max_budget, - user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None, + user_api_key_budget_reset_at=( + user_api_key_dict.budget_reset_at.isoformat() + if user_api_key_dict.budget_reset_at + else None + ), user_api_key_org_id=user_api_key_dict.org_id, user_api_key_team_id=user_api_key_dict.team_id, user_api_key_user_id=user_api_key_dict.user_id, @@ -204,6 +209,7 @@ class PagerDutyAlerting(SlackAlerting): user_api_key_end_user_id=user_api_key_dict.end_user_id, user_api_key_user_email=user_api_key_dict.user_email, user_api_key_request_route=user_api_key_dict.request_route, + user_api_key_auth_metadata=user_api_key_dict.metadata, ) ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 018b339d01..08e6540f76 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -51,7 +51,11 @@ class _ProxyDBLogger(CustomLogger): user_api_key_alias=user_api_key_dict.key_alias, user_api_key_spend=user_api_key_dict.spend, user_api_key_max_budget=user_api_key_dict.max_budget, - user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None, + user_api_key_budget_reset_at=( + user_api_key_dict.budget_reset_at.isoformat() + if user_api_key_dict.budget_reset_at + else None + ), user_api_key_user_email=user_api_key_dict.user_email, user_api_key_user_id=user_api_key_dict.user_id, user_api_key_team_id=user_api_key_dict.team_id, @@ -59,15 +63,16 @@ class _ProxyDBLogger(CustomLogger): user_api_key_team_alias=user_api_key_dict.team_alias, user_api_key_end_user_id=user_api_key_dict.end_user_id, user_api_key_request_route=user_api_key_dict.request_route, + user_api_key_auth_metadata=user_api_key_dict.metadata, ) ) _metadata["user_api_key"] = user_api_key_dict.api_key _metadata["status"] = "failure" - _metadata[ - "error_information" - ] = StandardLoggingPayloadSetup.get_error_information( - original_exception=original_exception, - traceback_str=traceback_str, + _metadata["error_information"] = ( + StandardLoggingPayloadSetup.get_error_information( + original_exception=original_exception, + traceback_str=traceback_str, + ) ) existing_metadata: dict = request_data.get("metadata", None) or {} diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 005b474424..227ea4db9f 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -96,9 +96,13 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona # langfuse requires b64 encoded headers - we construct that here _langfuse_public_key = custom_headers["LANGFUSE_PUBLIC_KEY"] _langfuse_secret_key = custom_headers["LANGFUSE_SECRET_KEY"] - if isinstance(_langfuse_public_key, str) and _langfuse_public_key.startswith("os.environ/"): + if isinstance( + _langfuse_public_key, str + ) and _langfuse_public_key.startswith("os.environ/"): _langfuse_public_key = get_secret_str(_langfuse_public_key) - if isinstance(_langfuse_secret_key, str) and _langfuse_secret_key.startswith("os.environ/"): + if isinstance( + _langfuse_secret_key, str + ) and _langfuse_secret_key.startswith("os.environ/"): _langfuse_secret_key = get_secret_str(_langfuse_secret_key) headers["Authorization"] = "Basic " + b64encode( f"{_langfuse_public_key}:{_langfuse_secret_key}".encode("utf-8") @@ -107,7 +111,9 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona # for all other headers headers[key] = value if isinstance(value, str) and "os.environ/" in value: - verbose_proxy_logger.debug("pass through endpoint - looking up 'os.environ/' variable") + verbose_proxy_logger.debug( + "pass through endpoint - looking up 'os.environ/' variable" + ) # get string section that is os.environ/ start_index = value.find("os.environ/") _variable_name = value[start_index:] @@ -200,7 +206,9 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 # skip router if user passed their key if "api_key" in data: llm_response = asyncio.create_task(litellm.aadapter_completion(**data)) - elif llm_router is not None and data["model"] in router_model_names: # model in router model list + elif ( + llm_router is not None and data["model"] in router_model_names + ): # model in router model list llm_response = asyncio.create_task(llm_router.aadapter_completion(**data)) elif ( llm_router is not None @@ -214,8 +222,8 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 llm_response = asyncio.create_task( llm_router.aadapter_completion(**data, specific_deployment=True) ) - elif ( - llm_router is not None and llm_router.has_model_id(data["model"]) + elif llm_router is not None and llm_router.has_model_id( + data["model"] ): # model in router model list llm_response = asyncio.create_task(llm_router.aadapter_completion(**data)) elif ( @@ -229,7 +237,10 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 else: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "completion: Invalid model name passed in model=" + data.get("model", "")}, + detail={ + "error": "completion: Invalid model name passed in model=" + + data.get("model", "") + }, ) # Await the llm_response task @@ -243,7 +254,9 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 ### ALERTING ### asyncio.create_task( - proxy_logging_obj.update_request_status(litellm_call_id=data.get("litellm_call_id", ""), status="success") + proxy_logging_obj.update_request_status( + litellm_call_id=data.get("litellm_call_id", ""), status="success" + ) ) verbose_proxy_logger.debug("final response: %s", response) @@ -265,7 +278,11 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 await proxy_logging_obj.post_call_failure_hook( user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data ) - verbose_proxy_logger.exception("litellm.proxy.proxy_server.completion(): Exception occured - {}".format(str(e))) + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.completion(): Exception occured - {}".format( + str(e) + ) + ) error_msg = f"{str(e)}" raise ProxyException( message=getattr(e, "message", error_msg), @@ -284,7 +301,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): ) -> dict: excluded_headers = {"transfer-encoding", "content-encoding"} - return_headers = {key: value for key, value in headers.items() if key.lower() not in excluded_headers} + return_headers = { + key: value + for key, value in headers.items() + if key.lower() not in excluded_headers + } if litellm_call_id: return_headers["x-litellm-call-id"] = litellm_call_id if custom_headers: @@ -411,8 +432,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): for field_name, field_value in form_data.items(): if isinstance(field_value, (StarletteUploadFile, UploadFile)): - files[field_name] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file( - upload_file=field_value + files[field_name] = ( + await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file( + upload_file=field_value + ) ) else: form_data_dict[field_name] = field_value @@ -462,8 +485,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): user_api_key_spend=user_api_key_dict.spend, user_api_key_max_budget=user_api_key_dict.max_budget, user_api_key_budget_reset_at=( - user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None + user_api_key_dict.budget_reset_at.isoformat() + if user_api_key_dict.budget_reset_at + else None ), + user_api_key_auth_metadata=user_api_key_dict.metadata, ) ) @@ -496,12 +522,16 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): "passthrough_logging_payload": passthrough_logging_payload, } - logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload + logging_obj.model_call_details["passthrough_logging_payload"] = ( + passthrough_logging_payload + ) return kwargs @staticmethod - def construct_target_url_with_subpath(base_target: str, subpath: str, include_subpath: Optional[bool]) -> str: + def construct_target_url_with_subpath( + base_target: str, subpath: str, include_subpath: Optional[bool] + ) -> str: """ Helper function to construct the full target URL with subpath handling. @@ -604,7 +634,9 @@ async def pass_through_request( # noqa: PLR0915 ).encode("ascii") ) - endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url)) + endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type( + str(url) + ) if custom_body: _parsed_body = custom_body @@ -665,13 +697,15 @@ async def pass_through_request( # noqa: PLR0915 logging_obj.model_call_details["litellm_call_id"] = litellm_call_id # combine url with query params for logging - requested_query_params: Optional[dict] = ( - query_params or dict(request.query_params) + requested_query_params: Optional[dict] = query_params or dict( + request.query_params ) requested_query_params_str = None if requested_query_params: - requested_query_params_str = "&".join(f"{k}={v}" for k, v in requested_query_params.items()) + requested_query_params_str = "&".join( + f"{k}={v}" for k, v in requested_query_params.items() + ) logging_url = str(url) if requested_query_params_str: @@ -689,9 +723,11 @@ async def pass_through_request( # noqa: PLR0915 "headers": headers, }, ) - stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( - parsed_body=_parsed_body, - stream=stream, + stream = ( + HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( + parsed_body=_parsed_body, + stream=stream, + ) ) if stream: @@ -708,7 +744,9 @@ async def pass_through_request( # noqa: PLR0915 try: response.raise_for_status() except httpx.HTTPStatusError as e: - raise HTTPException(status_code=e.response.status_code, detail=await e.response.aread()) + raise HTTPException( + status_code=e.response.status_code, detail=await e.response.aread() + ) return StreamingResponse( PassThroughStreamingHandler.chunk_processor( @@ -730,16 +768,20 @@ async def pass_through_request( # noqa: PLR0915 verbose_proxy_logger.debug("request method: {}".format(request.method)) verbose_proxy_logger.debug("request url: {}".format(url)) verbose_proxy_logger.debug("request headers: {}".format(headers)) - verbose_proxy_logger.debug("requested_query_params={}".format(requested_query_params)) + verbose_proxy_logger.debug( + "requested_query_params={}".format(requested_query_params) + ) verbose_proxy_logger.debug("request body: {}".format(_parsed_body)) - response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler( - request=request, - async_client=async_client, - url=url, - headers=headers, - requested_query_params=requested_query_params, - _parsed_body=_parsed_body, + response = ( + await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler( + request=request, + async_client=async_client, + url=url, + headers=headers, + requested_query_params=requested_query_params, + _parsed_body=_parsed_body, + ) ) verbose_proxy_logger.debug("response.headers= %s", response.headers) @@ -747,7 +789,9 @@ async def pass_through_request( # noqa: PLR0915 try: response.raise_for_status() except httpx.HTTPStatusError as e: - raise HTTPException(status_code=e.response.status_code, detail=await e.response.aread()) + raise HTTPException( + status_code=e.response.status_code, detail=await e.response.aread() + ) return StreamingResponse( PassThroughStreamingHandler.chunk_processor( @@ -769,7 +813,9 @@ async def pass_through_request( # noqa: PLR0915 try: response.raise_for_status() except httpx.HTTPStatusError as e: - raise HTTPException(status_code=e.response.status_code, detail=e.response.text) + raise HTTPException( + status_code=e.response.status_code, detail=e.response.text + ) if response.status_code >= 300: raise HTTPException(status_code=response.status_code, detail=response.text) @@ -822,7 +868,9 @@ async def pass_through_request( # noqa: PLR0915 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)) + "litellm.proxy.proxy_server.pass_through_endpoint(): Exception occured - {}".format( + str(e) + ) ) ######################################################### @@ -921,12 +969,16 @@ def create_pass_through_route( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), query_params: Optional[dict] = None, custom_body: Optional[dict] = None, - stream: Optional[bool] = None, # if pass-through endpoint is a streaming request + stream: Optional[ + bool + ] = None, # if pass-through endpoint is a streaming request subpath: str = "", # captures sub-paths when include_subpath=True ): # Construct the full target URL with subpath if needed - full_target = HttpPassThroughEndpointHelpers.construct_target_url_with_subpath( - base_target=target, subpath=subpath, include_subpath=include_subpath + full_target = ( + HttpPassThroughEndpointHelpers.construct_target_url_with_subpath( + base_target=target, subpath=subpath, include_subpath=include_subpath + ) ) return await pass_through_request( # type: ignore @@ -1078,7 +1130,9 @@ async def websocket_passthrough_request( # noqa: PLR0915 # Create a dummy request object for WebSocket connections to maintain compatibility # with the existing _init_kwargs_for_pass_through_endpoint function class DummyRequest: - def __init__(self, url: str, method: str = "WEBSOCKET", headers: Optional[dict] = None): + def __init__( + self, url: str, method: str = "WEBSOCKET", headers: Optional[dict] = None + ): self.url = url self.method = method self.headers = headers or {} @@ -1183,9 +1237,9 @@ async def websocket_passthrough_request( # noqa: PLR0915 ) if extracted_model: kwargs["model"] = extracted_model - kwargs[ - "custom_llm_provider" - ] = "vertex_ai-language-models" + kwargs["custom_llm_provider"] = ( + "vertex_ai-language-models" + ) # Update logging object with correct model logging_obj.model = extracted_model logging_obj.model_call_details[ @@ -1251,9 +1305,9 @@ async def websocket_passthrough_request( # noqa: PLR0915 # Update logging object with correct model logging_obj.model = extracted_model logging_obj.model_call_details["model"] = extracted_model - logging_obj.model_call_details[ - "custom_llm_provider" - ] = "vertex_ai_language_models" + logging_obj.model_call_details["custom_llm_provider"] = ( + "vertex_ai_language_models" + ) verbose_proxy_logger.debug( f"WebSocket passthrough ({endpoint}): Successfully extracted model '{extracted_model}' and set provider to 'vertex_ai' from server setup response" ) @@ -1597,11 +1651,15 @@ class InitPassThroughEndpointHelpers: def remove_endpoint_routes(endpoint_id: str): """Remove all routes for a specific endpoint ID from the registry""" keys_to_remove = [ - key for key, value in _registered_pass_through_routes.items() if value["endpoint_id"] == endpoint_id + key + for key, value in _registered_pass_through_routes.items() + if value["endpoint_id"] == endpoint_id ] for key in keys_to_remove: del _registered_pass_through_routes[key] - verbose_proxy_logger.debug("Removed pass-through route from registry: %s", key) + verbose_proxy_logger.debug( + "Removed pass-through route from registry: %s", key + ) @staticmethod def is_registered_pass_through_route(route: str) -> bool: @@ -1625,11 +1683,13 @@ class InitPassThroughEndpointHelpers: if len(parts) == 3: route_type = parts[1] registered_path = parts[2] - + if route_type == "exact" and route == registered_path: return True elif route_type == "subpath": - if route == registered_path or route.startswith(registered_path + "/"): + if route == registered_path or route.startswith( + registered_path + "/" + ): return True return False @@ -1669,7 +1729,9 @@ async def initialize_pass_through_endpoints( if _path is None: raise ValueError("Path is required for pass-through endpoint") _custom_headers = endpoint.get("headers", None) - _custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers) + _custom_headers = await set_env_variables_in_header( + custom_headers=_custom_headers + ) _forward_headers = endpoint.get("forward_headers", None) _merge_query_params = endpoint.get("merge_query_params", None) _auth = endpoint.get("auth", None) @@ -1688,7 +1750,9 @@ async def initialize_pass_through_endpoints( continue # Add exact path route - verbose_proxy_logger.debug("Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id) + verbose_proxy_logger.debug( + "Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id + ) InitPassThroughEndpointHelpers.add_exact_path_route( app=app, path=_path, @@ -1715,7 +1779,9 @@ async def initialize_pass_through_endpoints( endpoint_id=endpoint_id, ) - verbose_proxy_logger.debug("Added new pass through endpoint: %s (ID: %s)", _path, endpoint_id) + verbose_proxy_logger.debug( + "Added new pass through endpoint: %s (ID: %s)", _path, endpoint_id + ) async def _get_pass_through_endpoints_from_db( @@ -1819,7 +1885,11 @@ async def update_pass_through_endpoints( # Find the index for updating the list endpoint_index = None for idx, endpoint in enumerate(pass_through_endpoint_data): - _endpoint = PassThroughGenericEndpoint(**endpoint) if isinstance(endpoint, dict) else endpoint + _endpoint = ( + PassThroughGenericEndpoint(**endpoint) + if isinstance(endpoint, dict) + else endpoint + ) if _endpoint.id == endpoint_id: endpoint_index = idx break @@ -1827,7 +1897,9 @@ async def update_pass_through_endpoints( if endpoint_index is None: raise HTTPException( status_code=404, - detail={"error": f"Could not find index for endpoint with ID '{endpoint_id}'"}, + detail={ + "error": f"Could not find index for endpoint with ID '{endpoint_id}'" + }, ) # Get the update data as dict, excluding None values for partial updates @@ -1858,9 +1930,13 @@ async def update_pass_through_endpoints( field_value=pass_through_endpoint_data, config_type="general_settings", ) - await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict) + await update_config_general_settings( + data=updated_data, user_api_key_dict=user_api_key_dict + ) - return PassThroughEndpointResponse(endpoints=[updated_endpoint] if updated_endpoint else []) + return PassThroughEndpointResponse( + endpoints=[updated_endpoint] if updated_endpoint else [] + ) @router.post( @@ -1887,7 +1963,9 @@ async def create_pass_through_endpoints( field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict ) except Exception: - response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None) + response = ConfigFieldInfo( + field_name="pass_through_endpoints", field_value=None + ) ## Auto-generate ID if not provided data_dict = data.model_dump() @@ -1905,7 +1983,9 @@ async def create_pass_through_endpoints( field_value=response.field_value, config_type="general_settings", ) - await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict) + await update_config_general_settings( + data=updated_data, user_api_key_dict=user_api_key_dict + ) # Return the created endpoint with the generated ID created_endpoint = PassThroughGenericEndpoint(**data_dict) @@ -1938,7 +2018,9 @@ async def delete_pass_through_endpoints( field_name="pass_through_endpoints", user_api_key_dict=user_api_key_dict ) except Exception: - response = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None) + response = ConfigFieldInfo( + field_name="pass_through_endpoints", field_value=None + ) ## Update field by removing endpoint pass_through_endpoint_data: Optional[List] = response.field_value @@ -1954,13 +2036,21 @@ async def delete_pass_through_endpoints( if found_endpoint is None: raise HTTPException( status_code=400, - detail={"error": "Endpoint with ID '{}' was not found in pass-through endpoint list.".format(endpoint_id)}, + detail={ + "error": "Endpoint with ID '{}' was not found in pass-through endpoint list.".format( + endpoint_id + ) + }, ) # Find the index for deleting from the list endpoint_index = None for idx, endpoint in enumerate(pass_through_endpoint_data): - _endpoint = PassThroughGenericEndpoint(**endpoint) if isinstance(endpoint, dict) else endpoint + _endpoint = ( + PassThroughGenericEndpoint(**endpoint) + if isinstance(endpoint, dict) + else endpoint + ) if _endpoint.id == endpoint_id: endpoint_index = idx break @@ -1968,7 +2058,9 @@ async def delete_pass_through_endpoints( if endpoint_index is None: raise HTTPException( status_code=400, - detail={"error": f"Could not find index for endpoint with ID '{endpoint_id}'"}, + detail={ + "error": f"Could not find index for endpoint with ID '{endpoint_id}'" + }, ) # Remove the endpoint @@ -1984,7 +2076,9 @@ async def delete_pass_through_endpoints( field_value=pass_through_endpoint_data, config_type="general_settings", ) - await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict) + await update_config_general_settings( + data=updated_data, user_api_key_dict=user_api_key_dict + ) return PassThroughEndpointResponse(endpoints=[response_obj]) @@ -2022,4 +2116,6 @@ async def initialize_pass_through_endpoints_in_db(): Gets all pass-through endpoints from db and initializes them in the proxy server. """ pass_through_endpoints = await _get_pass_through_endpoints_from_db() - await initialize_pass_through_endpoints(pass_through_endpoints=pass_through_endpoints) + await initialize_pass_through_endpoints( + pass_through_endpoints=pass_through_endpoints + )