mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-03 04:22:22 +00:00
fix: fix linting errors
This commit is contained in:
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -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 {}
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user