diff --git a/litellm/llms/azure/responses/transformation.py b/litellm/llms/azure/responses/transformation.py index 488a711669d..0050bd163d1 100644 --- a/litellm/llms/azure/responses/transformation.py +++ b/litellm/llms/azure/responses/transformation.py @@ -1,5 +1,7 @@ from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple +import httpx + from litellm._logging import verbose_logger from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig @@ -194,3 +196,66 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): params["order"] = order verbose_logger.debug(f"list input items url={url}") return url, params + + ######################################################### + ########## CANCEL RESPONSE API TRANSFORMATION ########## + ######################################################### + def transform_cancel_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + """ + Transform the cancel response API request into a URL and data + + Azure OpenAI API expects the following request: + - POST /openai/responses/{response_id}/cancel?api-version=xxx + + This function handles URLs with query parameters by inserting the response_id + at the correct location (before any query parameters). + """ + from urllib.parse import urlparse, urlunparse + + # Parse the URL to separate its components + parsed_url = urlparse(api_base) + + # Insert the response_id and /cancel at the end of the path component + # Remove trailing slash if present to avoid double slashes + path = parsed_url.path.rstrip("/") + new_path = f"{path}/{response_id}/cancel" + + # Reconstruct the URL with all original components but with the modified path + cancel_url = urlunparse( + ( + parsed_url.scheme, # http, https + parsed_url.netloc, # domain name, port + new_path, # path with response_id and /cancel added + parsed_url.params, # parameters + parsed_url.query, # query string + parsed_url.fragment, # fragment + ) + ) + + data: Dict = {} + verbose_logger.debug(f"cancel response url={cancel_url}") + return cancel_url, data + + def transform_cancel_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + """ + Transform the cancel response API response into a ResponsesAPIResponse + """ + try: + raw_response_json = raw_response.json() + except Exception: + from litellm.llms.azure.chat.gpt_transformation import AzureOpenAIError + + raise AzureOpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + return ResponsesAPIResponse(**raw_response_json) diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 4da4f7652e0..facabbda72a 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -217,3 +217,28 @@ class BaseResponsesAPIConfig(ABC): ) -> bool: """Returns True if litellm should fake a stream for the given model and stream value""" return False + + ######################################################### + ########## CANCEL RESPONSE API TRANSFORMATION ########## + ######################################################### + @abstractmethod + def transform_cancel_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + pass + + @abstractmethod + def transform_cancel_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + pass + + ######################################################### + ########## END CANCEL RESPONSE API TRANSFORMATION ####### + ######################################################### diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index d691549bc6b..8b925a375a1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2200,6 +2200,7 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, optional_params={}, ) + if _is_async: return self.async_create_file( transformed_request=transformed_request, @@ -2216,7 +2217,6 @@ class BaseLLMHTTPHandler: sync_httpx_client = _get_httpx_client() else: sync_httpx_client = client - if isinstance(transformed_request, dict) and "method" in transformed_request: # Handle pre-signed requests (e.g., from Bedrock S3 uploads) @@ -2283,11 +2283,11 @@ class BaseLLMHTTPHandler: e=e, provider_config=provider_config, ) - + # Store the upload URL in litellm_params for the transformation method litellm_params_with_url = dict(litellm_params) litellm_params_with_url["upload_url"] = api_base - + return provider_config.transform_create_file_response( model=None, raw_response=upload_response, @@ -2423,7 +2423,7 @@ class BaseLLMHTTPHandler: # get config from model, custom llm provider if model is None: raise ValueError("model is required for create_batch") - + headers = provider_config.validate_environment( api_key=api_key, headers=headers, @@ -2606,6 +2606,159 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params_with_request, ) + def cancel_response_api_handler( + self, + response_id: str, + responses_api_provider_config: BaseResponsesAPIConfig, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + custom_llm_provider: Optional[str], + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + _is_async: bool = False, + ) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: + """ + Async version of the responses API handler. + Uses async HTTP client to make requests. + """ + if _is_async: + return self.async_cancel_response_api_handler( + response_id=response_id, + responses_api_provider_config=responses_api_provider_config, + litellm_params=litellm_params, + logging_obj=logging_obj, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + client=client, + ) + if client is None or not isinstance(client, HTTPHandler): + sync_httpx_client = _get_httpx_client( + params={"ssl_verify": litellm_params.get("ssl_verify", None)} + ) + else: + sync_httpx_client = client + + headers = responses_api_provider_config.validate_environment( + headers=extra_headers or {}, model="None", litellm_params=litellm_params + ) + + if extra_headers: + headers.update(extra_headers) + + api_base = responses_api_provider_config.get_complete_url( + api_base=litellm_params.api_base, + litellm_params=dict(litellm_params), + ) + + url, data = responses_api_provider_config.transform_cancel_response_api_request( + response_id=response_id, + api_base=api_base, + litellm_params=litellm_params, + headers=headers, + ) + + ## LOGGING + logging_obj.pre_call( + input=response_id, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": url, + "headers": headers, + }, + ) + + try: + response = sync_httpx_client.post( + url=url, headers=headers, json=data, timeout=timeout + ) + + except Exception as e: + raise self._handle_error( + e=e, + provider_config=responses_api_provider_config, + ) + + return responses_api_provider_config.transform_cancel_response_api_response( + raw_response=response, + logging_obj=logging_obj, + ) + + async def async_cancel_response_api_handler( + self, + response_id: str, + responses_api_provider_config: BaseResponsesAPIConfig, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + custom_llm_provider: Optional[str], + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + _is_async: bool = False, + ) -> ResponsesAPIResponse: + """ + Async version of the cancel response API handler. + Uses async HTTP client to make requests. + """ + if client is None or not isinstance(client, AsyncHTTPHandler): + async_httpx_client = get_async_httpx_client( + llm_provider=litellm.LlmProviders(custom_llm_provider), + params={"ssl_verify": litellm_params.get("ssl_verify", None)}, + ) + else: + async_httpx_client = client + + headers = responses_api_provider_config.validate_environment( + headers=extra_headers or {}, model="None", litellm_params=litellm_params + ) + + if extra_headers: + headers.update(extra_headers) + + api_base = responses_api_provider_config.get_complete_url( + api_base=litellm_params.api_base, + litellm_params=dict(litellm_params), + ) + + url, data = responses_api_provider_config.transform_cancel_response_api_request( + response_id=response_id, + api_base=api_base, + litellm_params=litellm_params, + headers=headers, + ) + + ## LOGGING + logging_obj.pre_call( + input=response_id, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": url, + "headers": headers, + }, + ) + + try: + response = await async_httpx_client.post( + url=url, headers=headers, json=data, timeout=timeout + ) + + except Exception as e: + raise self._handle_error( + e=e, + provider_config=responses_api_provider_config, + ) + + return responses_api_provider_config.transform_cancel_response_api_response( + raw_response=response, + logging_obj=logging_obj, + ) + def list_files(self): """ Lists all files @@ -2766,10 +2919,7 @@ class BaseLLMHTTPHandler: _is_async: bool = False, fake_stream: bool = False, litellm_metadata: Optional[Dict[str, Any]] = None, - ) -> Union[ - ImageResponse, - Coroutine[Any, Any, ImageResponse], - ]: + ) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]: """ Handles image edit requests. @@ -2959,10 +3109,7 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, litellm_metadata: Optional[Dict[str, Any]] = None, api_key: Optional[str] = None, - ) -> Union[ - ImageResponse, - Coroutine[Any, Any, ImageResponse], - ]: + ) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]: """ Handles image generation requests. When _is_async=True, returns a coroutine instead of making the call directly. @@ -3196,15 +3343,16 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url, request_body = ( - vector_store_provider_config.transform_search_vector_store_request( - vector_store_id=vector_store_id, - query=query, - vector_store_search_optional_params=vector_store_search_optional_params, - api_base=api_base, - litellm_logging_obj=logging_obj, - litellm_params=dict(litellm_params), - ) + ( + url, + request_body, + ) = vector_store_provider_config.transform_search_vector_store_request( + vector_store_id=vector_store_id, + query=query, + vector_store_search_optional_params=vector_store_search_optional_params, + api_base=api_base, + litellm_logging_obj=logging_obj, + litellm_params=dict(litellm_params), ) all_optional_params: Dict[str, Any] = dict(litellm_params) all_optional_params.update(vector_store_search_optional_params or {}) @@ -3295,15 +3443,16 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url, request_body = ( - vector_store_provider_config.transform_search_vector_store_request( - vector_store_id=vector_store_id, - query=query, - vector_store_search_optional_params=vector_store_search_optional_params, - api_base=api_base, - litellm_logging_obj=logging_obj, - litellm_params=dict(litellm_params), - ) + ( + url, + request_body, + ) = vector_store_provider_config.transform_search_vector_store_request( + vector_store_id=vector_store_id, + query=query, + vector_store_search_optional_params=vector_store_search_optional_params, + api_base=api_base, + litellm_logging_obj=logging_obj, + litellm_params=dict(litellm_params), ) all_optional_params: Dict[str, Any] = dict(litellm_params) @@ -3377,11 +3526,12 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url, request_body = ( - vector_store_provider_config.transform_create_vector_store_request( - vector_store_create_optional_params=vector_store_create_optional_params, - api_base=api_base, - ) + ( + url, + request_body, + ) = vector_store_provider_config.transform_create_vector_store_request( + vector_store_create_optional_params=vector_store_create_optional_params, + api_base=api_base, ) logging_obj.pre_call( @@ -3452,11 +3602,12 @@ class BaseLLMHTTPHandler: litellm_params=dict(litellm_params), ) - url, request_body = ( - vector_store_provider_config.transform_create_vector_store_request( - vector_store_create_optional_params=vector_store_create_optional_params, - api_base=api_base, - ) + ( + url, + request_body, + ) = vector_store_provider_config.transform_create_vector_store_request( + vector_store_create_optional_params=vector_store_create_optional_params, + api_base=api_base, ) logging_obj.pre_call( @@ -3535,13 +3686,14 @@ class BaseLLMHTTPHandler: sync_httpx_client = client # Get headers and URL from the provider config - headers, api_base = ( - generate_content_provider_config.sync_get_auth_token_and_url( - api_base=litellm_params.api_base, - model=model, - litellm_params=dict(litellm_params), - stream=stream, - ) + ( + headers, + api_base, + ) = generate_content_provider_config.sync_get_auth_token_and_url( + api_base=litellm_params.api_base, + model=model, + litellm_params=dict(litellm_params), + stream=stream, ) if extra_headers: @@ -3641,13 +3793,14 @@ class BaseLLMHTTPHandler: async_httpx_client = client # Get headers and URL from the provider config - headers, api_base = ( - await generate_content_provider_config.get_auth_token_and_url( - model=model, - litellm_params=dict(litellm_params), - stream=stream, - api_base=litellm_params.api_base, - ) + ( + headers, + api_base, + ) = await generate_content_provider_config.get_auth_token_and_url( + model=model, + litellm_params=dict(litellm_params), + stream=stream, + api_base=litellm_params.api_base, ) if extra_headers: diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 1d52f74b7b9..25078267571 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -425,3 +425,39 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): raise OpenAIError( message=raw_response.text, status_code=raw_response.status_code ) + + ######################################################### + ########## CANCEL RESPONSE API TRANSFORMATION ########## + ######################################################### + def transform_cancel_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + """ + Transform the cancel response API request into a URL and data + + OpenAI API expects the following request + - POST /v1/responses/{response_id}/cancel + """ + url = f"{api_base}/{response_id}/cancel" + data: Dict = {} + return url, data + + def transform_cancel_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + """ + Transform the cancel response API response into a ResponsesAPIResponse + """ + try: + raw_response_json = raw_response.json() + except Exception: + raise OpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + return ResponsesAPIResponse(**raw_response_json) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e900975f1cc..5739e652043 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -259,6 +259,7 @@ class ProxyBaseLLMRequestProcessing: "_arealtime", "aget_responses", "adelete_responses", + "acancel_responses", "acreate_batch", "aretrieve_batch", "afile_content", @@ -355,6 +356,7 @@ class ProxyBaseLLMRequestProcessing: "_arealtime", "aget_responses", "adelete_responses", + "acancel_responses", "atext_completion", "aimage_edit", "alist_input_items", diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 18481f11e2f..c87690854f3 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -285,3 +285,75 @@ async def get_response_input_items( proxy_logging_obj=proxy_logging_obj, version=version, ) + + +@router.post( + "/v1/responses/{response_id}/cancel", + dependencies=[Depends(user_api_key_auth)], + tags=["responses"], +) +@router.post( + "/responses/{response_id}/cancel", + dependencies=[Depends(user_api_key_auth)], + tags=["responses"], +) +async def cancel_response( + response_id: str, + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Cancel a response by ID. + + Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/cancel + + ```bash + curl -X POST http://localhost:4000/v1/responses/resp_abc123/cancel \ + -H "Authorization: Bearer sk-1234" + ``` + """ + from litellm.proxy.proxy_server import ( + _read_request_body, + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + data = await _read_request_body(request=request) + data["response_id"] = response_id + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="acancel_responses", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=select_data_generator, + model=None, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index cdeea0094a6..2a4281d6357 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -24,6 +24,7 @@ ROUTE_ENDPOINT_MAPPING = { "aresponses": "/responses", "alist_input_items": "/responses/{response_id}/input_items", "aimage_edit": "/images/edits", + "acancel_responses": "/responses/{response_id}/cancel", } @@ -70,6 +71,8 @@ async def route_request( "aresponses", "aget_responses", "adelete_responses", + "acancel_responses", + "acreate_response_reply", "alist_input_items", "_arealtime", # private function for realtime API "aimage_edit", @@ -86,6 +89,11 @@ async def route_request( team_id = get_team_id_from_data(data) router_model_names = llm_router.model_names if llm_router is not None else [] + # Preprocess Google GenAI generate content requests + if route_type in ["agenerate_content", "agenerate_content_stream"]: + # Map generationConfig to config parameter for Google GenAI compatibility + if "generationConfig" in data and "config" not in data: + data["config"] = data.pop("generationConfig") if "api_key" in data or "api_base" in data: if llm_router is not None: return getattr(llm_router, f"{route_type}")(**data) @@ -149,6 +157,7 @@ async def route_request( "amoderation", "aget_responses", "adelete_responses", + "acancel_responses", "alist_input_items", "avector_store_create", "avector_store_search", diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 04ee2b343f5..d3cb5a7de2a 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -167,13 +167,17 @@ async def aresponses_api_with_mcp( # Process MCP tools through the complete pipeline (fetch + filter + deduplicate + transform) user_api_key_auth = kwargs.get("user_api_key_auth") - + # Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods - original_mcp_tools = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( - user_api_key_auth=user_api_key_auth, - mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy + original_mcp_tools = ( + await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + user_api_key_auth=user_api_key_auth, + mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, + ) + ) + openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( + original_mcp_tools ) - openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(original_mcp_tools) # Combine with other tools all_tools = openai_tools + other_tools if (openai_tools or other_tools) else None @@ -212,15 +216,15 @@ async def aresponses_api_with_mcp( from litellm.responses.mcp.mcp_streaming_iterator import ( create_mcp_list_tools_events, ) - + base_item_id = f"mcp_{uuid.uuid4().hex[:8]}" mcp_discovery_events = await create_mcp_list_tools_events( mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, user_api_key_auth=user_api_key_auth, base_item_id=base_item_id, - pre_processed_mcp_tools=original_mcp_tools + pre_processed_mcp_tools=original_mcp_tools, ) - + return LiteLLM_Proxy_MCP_Handler._create_mcp_streaming_response( input=input, model=model, @@ -229,23 +233,21 @@ async def aresponses_api_with_mcp( mcp_discovery_events=mcp_discovery_events, call_params=call_params, previous_response_id=previous_response_id, - **kwargs + **kwargs, ) - + # Determine if we should auto-execute tools - should_auto_execute = ( - bool(mcp_tools_with_litellm_proxy) - and LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( - mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy - ) + should_auto_execute = bool( + mcp_tools_with_litellm_proxy + ) and LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( + mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy ) - + # Prepare parameters for the initial call initial_call_params = LiteLLM_Proxy_MCP_Handler._prepare_initial_call_params( - call_params=call_params, - should_auto_execute=should_auto_execute + call_params=call_params, should_auto_execute=should_auto_execute ) - + ######################################################### # Make initial response API call ######################################################### @@ -263,9 +265,8 @@ async def aresponses_api_with_mcp( # Auto-Execute Tools Handling # If auto-execute tools is True, then we need to execute the tool calls ######################################################### - if ( - should_auto_execute - and isinstance(response, ResponsesAPIResponse) + if should_auto_execute and isinstance( + response, ResponsesAPIResponse ): # type: ignore tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_response( response=response @@ -285,19 +286,21 @@ async def aresponses_api_with_mcp( ) # Prepare parameters for follow-up call (restores original stream setting) - follow_up_call_params = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params( - call_params=call_params, - original_stream_setting=stream or False + follow_up_call_params = ( + LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params( + call_params=call_params, original_stream_setting=stream or False + ) ) - + # Create tool execution events for streaming if needed tool_execution_events = [] if stream: - tool_execution_events = LiteLLM_Proxy_MCP_Handler._create_tool_execution_events( - tool_calls=tool_calls, - tool_results=tool_results + tool_execution_events = ( + LiteLLM_Proxy_MCP_Handler._create_tool_execution_events( + tool_calls=tool_calls, tool_results=tool_results + ) ) - + final_response = await LiteLLM_Proxy_MCP_Handler._make_follow_up_call( follow_up_input=follow_up_input, model=model, @@ -307,13 +310,20 @@ async def aresponses_api_with_mcp( ) # If streaming and we have tool execution events, wrap the response - if stream and tool_execution_events and (hasattr(final_response, '__aiter__') or hasattr(final_response, '__iter__')): + if ( + stream + and tool_execution_events + and ( + hasattr(final_response, "__aiter__") + or hasattr(final_response, "__iter__") + ) + ): from litellm.responses.mcp.mcp_streaming_iterator import ( MCPEnhancedStreamingIterator, ) + final_response = MCPEnhancedStreamingIterator( - base_iterator=final_response, - mcp_events=tool_execution_events + base_iterator=final_response, mcp_events=tool_execution_events ) # Add custom output elements to the final response (for non-streaming) @@ -321,7 +331,7 @@ async def aresponses_api_with_mcp( # Fetch MCP tools again for output elements (without OpenAI transformation) mcp_tools_for_output = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( user_api_key_auth=user_api_key_auth, - mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy + mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, ) final_response = ( LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response( @@ -1150,4 +1160,163 @@ def list_input_items( original_exception=e, completion_kwargs=local_vars, extra_kwargs=kwargs, - ) \ No newline at end of file + ) + + +@client +async def acancel_responses( + response_id: str, + # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. + # The extra values given here take precedence over values defined on the client or passed to this method. + extra_headers: Optional[Dict[str, Any]] = None, + extra_query: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + # LiteLLM specific params, + custom_llm_provider: Optional[str] = None, + **kwargs, +) -> ResponsesAPIResponse: + """ + Async version of the POST Cancel Responses API + + POST /v1/responses/{response_id}/cancel endpoint in the responses API + + """ + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["acancel_responses"] = True + + # get custom llm provider from response_id + decoded_response_id: DecodedResponseId = ( + ResponsesAPIRequestUtils._decode_responses_api_response_id( + response_id=response_id, + ) + ) + response_id = decoded_response_id.get("response_id") or response_id + custom_llm_provider = ( + decoded_response_id.get("custom_llm_provider") or custom_llm_provider + ) + + func = partial( + cancel_responses, + response_id=response_id, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + extra_query=extra_query, + extra_body=extra_body, + timeout=timeout, + **kwargs, + ) + + ctx = contextvars.copy_context() + func_with_context = partial(ctx.run, func) + init_response = await loop.run_in_executor(None, func_with_context) + + if asyncio.iscoroutine(init_response): + response = await init_response + else: + response = init_response + return response + except Exception as e: + raise litellm.exception_type( + model=None, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def cancel_responses( + response_id: str, + # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. + # The extra values given here take precedence over values defined on the client or passed to this method. + extra_headers: Optional[Dict[str, Any]] = None, + extra_query: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + # LiteLLM specific params, + custom_llm_provider: Optional[str] = None, + **kwargs, +) -> Union[ResponsesAPIResponse, Coroutine[Any, Any, ResponsesAPIResponse]]: + """ + Synchronous version of the POST Responses API + + POST /v1/responses/{response_id}/cancel endpoint in the responses API + + """ + local_vars = locals() + try: + litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + _is_async = kwargs.pop("acancel_responses", False) is True + + # get llm provider logic + litellm_params = GenericLiteLLMParams(**kwargs) + + # get custom llm provider from response_id + decoded_response_id: DecodedResponseId = ( + ResponsesAPIRequestUtils._decode_responses_api_response_id( + response_id=response_id, + ) + ) + response_id = decoded_response_id.get("response_id") or response_id + custom_llm_provider = ( + decoded_response_id.get("custom_llm_provider") or custom_llm_provider + ) + + if custom_llm_provider is None: + raise ValueError("custom_llm_provider is required but passed as None") + + # get provider config + responses_api_provider_config: Optional[ + BaseResponsesAPIConfig + ] = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=litellm.LlmProviders(custom_llm_provider), + ) + + if responses_api_provider_config is None: + raise ValueError( + f"CANCEL responses is not supported for {custom_llm_provider}" + ) + + local_vars.update(kwargs) + + # Pre Call logging + litellm_logging_obj.update_environment_variables( + model=None, + optional_params={ + "response_id": response_id, + }, + litellm_params={ + "litellm_call_id": litellm_call_id, + }, + custom_llm_provider=custom_llm_provider, + ) + + # Call the handler with _is_async flag instead of directly calling the async handler + response = base_llm_http_handler.cancel_response_api_handler( + response_id=response_id, + custom_llm_provider=custom_llm_provider, + responses_api_provider_config=responses_api_provider_config, + litellm_params=litellm_params, + logging_obj=litellm_logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout or request_timeout, + _is_async=_is_async, + client=kwargs.get("client"), + ) + + return response + except Exception as e: + raise litellm.exception_type( + model=None, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) diff --git a/litellm/router.py b/litellm/router.py index 519c7797daf..1978c14aafb 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -359,9 +359,9 @@ class Router: ) # names of models under litellm_params. ex. azure/chatgpt-v-2 self.deployment_latency_map = {} ### CACHING ### - cache_type: Literal["local", "redis", "redis-semantic", "s3", "disk"] = ( - "local" # default to an in-memory cache - ) + cache_type: Literal[ + "local", "redis", "redis-semantic", "s3", "disk" + ] = "local" # default to an in-memory cache redis_cache = None cache_config: Dict[str, Any] = {} @@ -403,9 +403,9 @@ class Router: self.default_max_parallel_requests = default_max_parallel_requests self.provider_default_deployment_ids: List[str] = [] self.pattern_router = PatternMatchRouter() - self.team_pattern_routers: Dict[str, PatternMatchRouter] = ( - {} - ) # {"TEAM_ID": PatternMatchRouter} + self.team_pattern_routers: Dict[ + str, PatternMatchRouter + ] = {} # {"TEAM_ID": PatternMatchRouter} self.auto_routers: Dict[str, "AutoRouter"] = {} if model_list is not None: @@ -587,9 +587,9 @@ class Router: ) ) - self.model_group_retry_policy: Optional[Dict[str, RetryPolicy]] = ( - model_group_retry_policy - ) + self.model_group_retry_policy: Optional[ + Dict[str, RetryPolicy] + ] = model_group_retry_policy self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None if allowed_fails_policy is not None: @@ -782,6 +782,9 @@ class Router: self.aget_responses = self.factory_function( litellm.aget_responses, call_type="aget_responses" ) + self.acancel_responses = self.factory_function( + litellm.acancel_responses, call_type="acancel_responses" + ) self.adelete_responses = self.factory_function( litellm.adelete_responses, call_type="adelete_responses" ) @@ -873,7 +876,6 @@ class Router: def add_optional_pre_call_checks( self, optional_pre_call_checks: Optional[OptionalPreCallChecks] ): - if optional_pre_call_checks is not None: for pre_call_check in optional_pre_call_checks: _callback: Optional[CustomLogger] = None @@ -1209,10 +1211,7 @@ class Router: async def _acompletion( self, model: str, messages: List[Dict[str, str]], **kwargs - ) -> Union[ - ModelResponse, - CustomStreamWrapper, - ]: + ) -> Union[ModelResponse, CustomStreamWrapper,]: """ - Get an available deployment - call it with a semaphore over the call @@ -2713,7 +2712,6 @@ class Router: passthrough_on_no_deployment = kwargs.pop("passthrough_on_no_deployment", False) function_name = "_ageneric_api_call_with_fallbacks" try: - parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) try: deployment = await self.async_get_available_deployment( @@ -3046,7 +3044,7 @@ class Router: from litellm.router_utils.common_utils import add_model_file_id_mappings verbose_router_logger.debug( - f"Inside _acreate_file()- model: {model}; kwargs: {kwargs}" + f"Inside _atext_completion()- model: {model}; kwargs: {kwargs}" ) parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) healthy_deployments = await self.async_get_healthy_deployments( @@ -3157,9 +3155,9 @@ class Router: healthy_deployments=healthy_deployments, responses=responses ) returned_response = cast(OpenAIFileObject, responses[0]) - returned_response._hidden_params["model_file_id_mapping"] = ( - model_file_id_mapping - ) + returned_response._hidden_params[ + "model_file_id_mapping" + ] = model_file_id_mapping return returned_response except Exception as e: verbose_router_logger.exception( @@ -3485,6 +3483,7 @@ class Router: "moderation", "anthropic_messages", "aresponses", + "acancel_responses", "responses", "aget_responses", "adelete_responses", @@ -3578,6 +3577,7 @@ class Router: ) elif call_type in ( "aget_responses", + "acancel_responses", "adelete_responses", "alist_input_items", ): @@ -3625,7 +3625,7 @@ class Router: """ Initialize the Responses API endpoints on the router. - GET, DELETE Responses API Requests encode the model_id in the response_id, this function decodes the response_id and sets the model to the model_id. + GET, DELETE, CANCEL Responses API Requests encode the model_id in the response_id, this function decodes the response_id and sets the model to the model_id. """ from litellm.responses.utils import ResponsesAPIRequestUtils @@ -3720,11 +3720,11 @@ class Router: if isinstance(e, litellm.ContextWindowExceededError): if context_window_fallbacks is not None: - context_window_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=context_window_fallbacks, - model_group=model_group, - ) + context_window_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=context_window_fallbacks, + model_group=model_group, ) if context_window_fallback_model_group is None: raise original_exception @@ -3756,11 +3756,11 @@ class Router: e.message += "\n{}".format(error_message) elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: - content_policy_fallback_model_group: Optional[List[str]] = ( - self._get_fallback_model_group_from_fallbacks( - fallbacks=content_policy_fallbacks, - model_group=model_group, - ) + content_policy_fallback_model_group: Optional[ + List[str] + ] = self._get_fallback_model_group_from_fallbacks( + fallbacks=content_policy_fallbacks, + model_group=model_group, ) if content_policy_fallback_model_group is None: raise original_exception @@ -4414,7 +4414,7 @@ class Router: return tpm_key except Exception as e: - verbose_router_logger.debug( + verbose_router_logger.exception( "litellm.router.Router::deployment_callback_on_success(): Exception occured - {}".format( str(e) ) @@ -4992,26 +4992,26 @@ class Router: """ from litellm.router_strategy.auto_router.auto_router import AutoRouter - auto_router_config_path: Optional[str] = ( - deployment.litellm_params.auto_router_config_path - ) + auto_router_config_path: Optional[ + str + ] = deployment.litellm_params.auto_router_config_path auto_router_config: Optional[str] = deployment.litellm_params.auto_router_config if auto_router_config_path is None and auto_router_config is None: raise ValueError( "auto_router_config_path or auto_router_config is required for auto-router deployments. Please set it in the litellm_params" ) - default_model: Optional[str] = ( - deployment.litellm_params.auto_router_default_model - ) + default_model: Optional[ + str + ] = deployment.litellm_params.auto_router_default_model if default_model is None: raise ValueError( "auto_router_default_model is required for auto-router deployments. Please set it in the litellm_params" ) - embedding_model: Optional[str] = ( - deployment.litellm_params.auto_router_embedding_model - ) + embedding_model: Optional[ + str + ] = deployment.litellm_params.auto_router_embedding_model if embedding_model is None: raise ValueError( "auto_router_embedding_model is required for auto-router deployments. Please set it in the litellm_params" diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index fc6983520fd..5ed4fbbb7b8 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -590,3 +590,60 @@ class BaseResponsesAPITest(ABC): assert function_call_item["status"] == "completed", "status value should be preserved" print("✅ OpenAI Responses API dict input filtering test passed") + + @pytest.mark.parametrize("sync_mode", [False, True]) + @pytest.mark.flaky(retries=3, delay=2) + @pytest.mark.asyncio + async def test_basic_openai_responses_cancel_endpoint(self, sync_mode): + litellm._turn_on_debug() + litellm.set_verbose = True + base_completion_call_args = self.get_base_completion_call_args() + if sync_mode: + response = litellm.responses( + input="Basic ping", max_output_tokens=20, background=True, **base_completion_call_args + ) + + # cancel the response + if isinstance(response, ResponsesAPIResponse): + cancel_result = litellm.cancel_responses( + response_id=response.id, **base_completion_call_args + ) + assert cancel_result is not None + assert hasattr(cancel_result, "id") + # The actual response structure depends on the provider implementation + assert isinstance(cancel_result, ResponsesAPIResponse) + else: + raise ValueError("response is not a ResponsesAPIResponse") + else: + response = await litellm.aresponses( + input="Basic ping", max_output_tokens=20, background=True, **base_completion_call_args + ) + + # async cancel the response + if isinstance(response, ResponsesAPIResponse): + cancel_result = await litellm.acancel_responses( + response_id=response.id, **base_completion_call_args + ) + assert cancel_result is not None + assert hasattr(cancel_result, "id") + # The actual response structure depends on the provider implementation + assert isinstance(cancel_result, ResponsesAPIResponse) + else: + raise ValueError("response is not a ResponsesAPIResponse") + + @pytest.mark.parametrize("sync_mode", [False, True]) + @pytest.mark.asyncio + async def test_cancel_responses_invalid_response_id(self, sync_mode): + """Test cancel_responses with invalid response ID should raise appropriate error""" + base_completion_call_args = self.get_base_completion_call_args() + + if sync_mode: + with pytest.raises(Exception): + litellm.cancel_responses( + response_id="invalid_response_id_12345", **base_completion_call_args + ) + else: + with pytest.raises(Exception): + await litellm.acancel_responses( + response_id="invalid_response_id_12345", **base_completion_call_args + ) \ No newline at end of file diff --git a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py index e637066d2f9..7e7def0ee03 100644 --- a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py +++ b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py @@ -128,3 +128,57 @@ def test_anthropic_with_responses_api(): previous_response_id="hi", ) print("anthropic response=", response) + + +def test_cancel_response(): + client = get_test_client() + from litellm.types.llms.openai import ResponsesAPIResponse + response = client.responses.create( + model="gpt-4o", input="just respond with the word 'ping'", background=True + ) + print("basic response=", response) + + # cancel the response + cancel_response = client.responses.cancel(response.id) + print("CANCEL response=", cancel_response) + + # verify cancel response structure + assert hasattr(cancel_response, "id") + # Note: Cancel response returns ResponsesAPIResponse, not DeleteResponseResult + # The actual response structure depends on the provider implementation + assert isinstance(cancel_response, ResponsesAPIResponse) + + +def test_cancel_streaming_response(): + client = get_test_client() + from litellm.types.llms.openai import ResponsesAPIResponse + stream = client.responses.create( + model="gpt-4o", input="just respond with the word 'ping'", stream=True, background=True + ) + + collected_chunks = [] + response_id = None + for chunk in stream: + print("stream chunk=", chunk) + collected_chunks.append(chunk) + # Extract response ID from the first chunk that has it + if response_id is None and hasattr(chunk, 'response') and hasattr(chunk.response, 'id'): + response_id = chunk.response.id + + assert len(collected_chunks) > 0 + + # cancel the response if we got a response ID + if response_id: + cancel_response = client.responses.cancel(response_id) + print("CANCEL streaming response=", cancel_response) + assert hasattr(cancel_response, "id") + # Note: Cancel response returns ResponsesAPIResponse, not DeleteResponseResult + # The actual response structure depends on the provider implementation + assert isinstance(cancel_response, ResponsesAPIResponse) + + +def test_cancel_invalid_response_id(): + client = get_test_client() + with pytest.raises(Exception): + # Try to cancel a non-existent response ID + client.responses.cancel("invalid_response_id_12345") \ No newline at end of file diff --git a/tests/test_litellm/llms/azure/response/test_azure_transformation.py b/tests/test_litellm/llms/azure/response/test_azure_transformation.py index 5a0db987eff..124f0e93db8 100644 --- a/tests/test_litellm/llms/azure/response/test_azure_transformation.py +++ b/tests/test_litellm/llms/azure/response/test_azure_transformation.py @@ -293,3 +293,55 @@ class TestAzureResponsesAPIConfig: litellm_params={"api_version": None}, ) assert result_none_version == expected_url + + def test_azure_cancel_response_api_request(self): + """Test Azure cancel response API request transformation""" + from litellm.types.router import GenericLiteLLMParams + + response_id = "resp_test123" + api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview" + litellm_params = GenericLiteLLMParams(api_version="2024-05-01-preview") + headers = {"Authorization": "Bearer test-key"} + + url, data = self.config.transform_cancel_response_api_request( + response_id=response_id, + api_base=api_base, + litellm_params=litellm_params, + headers=headers, + ) + + expected_url = "https://test.openai.azure.com/openai/responses/resp_test123/cancel?api-version=2024-05-01-preview" + assert url == expected_url + assert data == {} + + def test_azure_cancel_response_api_response(self): + """Test Azure cancel response API response transformation""" + from unittest.mock import Mock + from litellm.types.llms.openai import ResponsesAPIResponse + + # Mock response + mock_response = Mock() + mock_response.json.return_value = { + "id": "resp_test123", + "object": "response", + "created_at": 1234567890, + "output": [], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "top_p": 1.0, + "status": "cancelled" + } + mock_response.text = "test response" + mock_response.status_code = 200 + + # Mock logging object + mock_logging_obj = Mock() + + result = self.config.transform_cancel_response_api_response( + raw_response=mock_response, + logging_obj=mock_logging_obj, + ) + + assert isinstance(result, ResponsesAPIResponse) + assert result.id == "resp_test123" \ No newline at end of file