From 004bde2c45014eb10fc8b73edf0d62220f51d38d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8D=97=E8=BE=B0=E7=87=8F=E7=82=9A?= <95487306+LingXuanYin@users.noreply.github.com> Date: Tue, 20 Jan 2026 11:02:29 +0800 Subject: [PATCH] feat (volcengine) : Support Volcengine responses api (#18508) * Add Volcengine responses adapter * fix llms/volcengine/responses/transformation.py:507:9: F841 Local variable `origin` is assigned to but never used fix llms/volcengine/responses/transformation.py:95: error: Argument "headers" to "VolcEngineError" has incompatible type add more supported optional params removed redundant manual logging/utils fallbacks so litellm/__init__.py uses the registry only. --- litellm/__init__.py | 41 +- litellm/_lazy_imports_registry.py | 3 +- litellm/llms/volcengine/__init__.py | 4 +- .../volcengine/responses/transformation.py | 557 ++++++++++++++++++ litellm/utils.py | 54 +- ...est_volcengine_responses_transformation.py | 274 +++++++++ 6 files changed, 886 insertions(+), 47 deletions(-) create mode 100644 litellm/llms/volcengine/responses/transformation.py create mode 100644 tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index 9eb3f075d5..134f8206d3 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1268,7 +1268,7 @@ if TYPE_CHECKING: from litellm.types.utils import ModelInfo as _ModelInfoType from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.caching.caching import Cache - + # Type stubs for lazy-loaded configs to help mypy from .llms.bedrock.chat.converse_transformation import AmazonConverseConfig as AmazonConverseConfig from .llms.openai_like.chat.handler import OpenAILikeChatConfig as OpenAILikeChatConfig @@ -1374,6 +1374,7 @@ if TYPE_CHECKING: from .llms.azure.responses.o_series_transformation import AzureOpenAIOSeriesResponsesAPIConfig as AzureOpenAIOSeriesResponsesAPIConfig from .llms.xai.responses.transformation import XAIResponsesAPIConfig as XAIResponsesAPIConfig from .llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig as LiteLLMProxyResponsesAPIConfig + from .llms.volcengine.responses.transformation import VolcEngineResponsesAPIConfig as VolcEngineResponsesAPIConfig from .llms.manus.responses.transformation import ManusResponsesAPIConfig as ManusResponsesAPIConfig from .llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig from .llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig as OpenAIOSeriesConfig, OpenAIOSeriesConfig as OpenAIO1Config @@ -1387,7 +1388,7 @@ if TYPE_CHECKING: from .llms.openai.chat.gpt_audio_transformation import OpenAIGPTAudioConfig as OpenAIGPTAudioConfig from .llms.nvidia_nim.chat.transformation import NvidiaNimConfig as NvidiaNimConfig from .llms.nvidia_nim.embed import NvidiaNimEmbeddingConfig as NvidiaNimEmbeddingConfig - + # Type stubs for lazy-loaded config instances openaiOSeriesConfig: OpenAIOSeriesConfig openAIGPTConfig: OpenAIGPTConfig @@ -1395,7 +1396,7 @@ if TYPE_CHECKING: openAIGPT5Config: OpenAIGPT5Config nvidiaNimConfig: NvidiaNimConfig nvidiaNimEmbeddingConfig: NvidiaNimEmbeddingConfig - + # Import config classes that need type stubs (for mypy) - import with _ prefix to avoid circular reference from .llms.vllm.completion.transformation import VLLMConfig as _VLLMConfig from .llms.deepseek.chat.transformation import DeepSeekChatConfig as _DeepSeekChatConfig @@ -1413,7 +1414,7 @@ if TYPE_CHECKING: from .llms.lm_studio.embed.transformation import LmStudioEmbeddingConfig as _LmStudioEmbeddingConfig from .llms.watsonx.embed.transformation import IBMWatsonXEmbeddingConfig as _IBMWatsonXEmbeddingConfig from .llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig as _VertexGeminiConfig - + # Type stubs for lazy-loaded config classes (to help mypy understand types) VLLMConfig: Type[_VLLMConfig] DeepSeekChatConfig: Type[_DeepSeekChatConfig] @@ -1431,7 +1432,7 @@ if TYPE_CHECKING: LmStudioEmbeddingConfig: Type[_LmStudioEmbeddingConfig] IBMWatsonXEmbeddingConfig: Type[_IBMWatsonXEmbeddingConfig] VertexAIConfig: Type[_VertexGeminiConfig] # Alias for VertexGeminiConfig - + from .llms.featherless_ai.chat.transformation import FeatherlessAIConfig as FeatherlessAIConfig from .llms.cerebras.chat import CerebrasConfig as CerebrasConfig from .llms.baseten.chat import BasetenConfig as BasetenConfig @@ -1551,14 +1552,14 @@ if TYPE_CHECKING: # Custom logger class (lazy-loaded) from litellm.integrations.custom_logger import CustomLogger - + # Datadog LLM observability params (lazy-loaded) from litellm.types.integrations.datadog_llm_obs import DatadogLLMObsInitParams - + # Logging callback manager class and instance (lazy-loaded) from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager logging_callback_manager: LoggingCallbackManager - + # provider_list is lazy-loaded from litellm.types.utils import LlmProviders provider_list: List[Union[LlmProviders, str]] @@ -1588,12 +1589,12 @@ def __getattr__(name: str) -> Any: from litellm.llms.custom_httpx.async_client_cleanup import register_async_client_cleanup register_async_client_cleanup() _async_client_cleanup_registered = True - + # Use cached registry from _lazy_imports instead of importing tuples every time from ._lazy_imports import _get_lazy_import_registry - + registry = _get_lazy_import_registry() - + # Check if name is in registry and call the cached handler function if name in registry: handler_func = registry[name] @@ -1608,7 +1609,7 @@ def __getattr__(name: str) -> Any: from .main import encoding as _encoding _globals["encoding"] = _encoding return _globals["encoding"] - + # Lazy load bedrock_tool_name_mappings instance if name == "bedrock_tool_name_mappings": from ._lazy_imports import _get_litellm_globals @@ -1618,7 +1619,7 @@ def __getattr__(name: str) -> Any: from .llms.bedrock.chat.invoke_handler import bedrock_tool_name_mappings as _bedrock_tool_name_mappings _globals["bedrock_tool_name_mappings"] = _bedrock_tool_name_mappings return _globals["bedrock_tool_name_mappings"] - + # Lazy load AzureOpenAIError exception class if name == "AzureOpenAIError": from ._lazy_imports import _get_litellm_globals @@ -1628,7 +1629,7 @@ def __getattr__(name: str) -> Any: from .llms.azure.common_utils import AzureOpenAIError as _AzureOpenAIError _globals["AzureOpenAIError"] = _AzureOpenAIError return _globals["AzureOpenAIError"] - + # Lazy load openaiOSeriesConfig instance if name == "openaiOSeriesConfig": from ._lazy_imports import _get_litellm_globals @@ -1638,7 +1639,7 @@ def __getattr__(name: str) -> Any: config_class = __getattr__("OpenAIOSeriesConfig") _globals["openaiOSeriesConfig"] = config_class() return _globals["openaiOSeriesConfig"] - + # Lazy load other config instances _config_instances = { "openAIGPTConfig": "OpenAIGPTConfig", @@ -1655,11 +1656,11 @@ def __getattr__(name: str) -> Any: config_class = __getattr__(_config_instances[name]) _globals[name] = config_class() return _globals[name] - + # Handle OpenAIO1Config alias if name == "OpenAIO1Config": return __getattr__("OpenAIOSeriesConfig") - + # Lazy load provider_list if name == "provider_list": from ._lazy_imports import _get_litellm_globals @@ -1670,7 +1671,7 @@ def __getattr__(name: str) -> Any: from litellm.types.utils import LlmProviders _globals["provider_list"] = list(LlmProviders) return _globals["provider_list"] - + # Lazy load priority_reservation_settings instance if name == "priority_reservation_settings": from ._lazy_imports import _get_litellm_globals @@ -1681,7 +1682,7 @@ def __getattr__(name: str) -> Any: PriorityReservationSettings = __getattr__("PriorityReservationSettings") _globals["priority_reservation_settings"] = PriorityReservationSettings() return _globals["priority_reservation_settings"] - + # Lazy load logging_callback_manager instance if name == "logging_callback_manager": from ._lazy_imports import _get_litellm_globals @@ -1692,7 +1693,7 @@ def __getattr__(name: str) -> Any: LoggingCallbackManager = __getattr__("LoggingCallbackManager") _globals["logging_callback_manager"] = LoggingCallbackManager() return _globals["logging_callback_manager"] - + # Lazy load _service_logger module if name == "_service_logger": from ._lazy_imports import _get_litellm_globals diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index f37c4dc6d0..3025f1b0c5 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -198,6 +198,7 @@ LLM_CONFIG_NAMES = ( "AzureOpenAIOSeriesResponsesAPIConfig", "XAIResponsesAPIConfig", "LiteLLMProxyResponsesAPIConfig", + "VolcEngineResponsesAPIConfig", "GoogleAIStudioInteractionsConfig", "OpenAIOSeriesConfig", "AnthropicSkillsConfig", @@ -591,6 +592,7 @@ _LLM_CONFIGS_IMPORT_MAP = { "AzureOpenAIOSeriesResponsesAPIConfig": (".llms.azure.responses.o_series_transformation", "AzureOpenAIOSeriesResponsesAPIConfig"), "XAIResponsesAPIConfig": (".llms.xai.responses.transformation", "XAIResponsesAPIConfig"), "LiteLLMProxyResponsesAPIConfig": (".llms.litellm_proxy.responses.transformation", "LiteLLMProxyResponsesAPIConfig"), + "VolcEngineResponsesAPIConfig": (".llms.volcengine.responses.transformation", "VolcEngineResponsesAPIConfig"), "ManusResponsesAPIConfig": (".llms.manus.responses.transformation", "ManusResponsesAPIConfig"), "GoogleAIStudioInteractionsConfig": (".llms.gemini.interactions.transformation", "GoogleAIStudioInteractionsConfig"), "OpenAIOSeriesConfig": (".llms.openai.chat.o_series_transformation", "OpenAIOSeriesConfig"), @@ -774,4 +776,3 @@ __all__ = [ "_LLM_PROVIDER_LOGIC_IMPORT_MAP", "_UTILS_MODULE_IMPORT_MAP", ] - diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py index 0887937bed..fc0098e84d 100644 --- a/litellm/llms/volcengine/__init__.py +++ b/litellm/llms/volcengine/__init__.py @@ -1,6 +1,6 @@ """ Volcengine LLM Provider -Support for Volcengine (ByteDance) chat and embedding models +Support for Volcengine (ByteDance) chat, embedding, and responses models. """ from .chat.transformation import VolcEngineChatConfig @@ -10,6 +10,7 @@ from .common_utils import ( get_volcengine_headers, ) from .embedding import VolcEngineEmbeddingConfig +from .responses.transformation import VolcEngineResponsesAPIConfig # For backward compatibility, keep the old class name VolcEngineConfig = VolcEngineChatConfig @@ -18,6 +19,7 @@ __all__ = [ "VolcEngineChatConfig", "VolcEngineConfig", # backward compatibility "VolcEngineEmbeddingConfig", + "VolcEngineResponsesAPIConfig", "VolcEngineError", "get_volcengine_base_url", "get_volcengine_headers", diff --git a/litellm/llms/volcengine/responses/transformation.py b/litellm/llms/volcengine/responses/transformation.py new file mode 100644 index 0000000000..872c8dcf11 --- /dev/null +++ b/litellm/llms/volcengine/responses/transformation.py @@ -0,0 +1,557 @@ +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Optional, + Tuple, + Union, + get_args, + get_origin, +) + +import httpx +from pydantic import fields as pyd_fields + +import litellm +from litellm._logging import verbose_logger +from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIStreamingResponse +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( + _safe_convert_created_field, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + ResponsesAPIOptionalRequestParams, + ResponsesAPIResponse, +) +from litellm.types.responses.main import DeleteResponseResult +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +from ..common_utils import ( + VolcEngineError, + get_volcengine_base_url, + get_volcengine_headers, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): + _SUPPORTED_OPTIONAL_PARAMS: List[str] = [ + # Doc-listed knobs + "instructions", + "max_output_tokens", + "previous_response_id", + "store", + "reasoning", + "stream", + "temperature", + "top_p", + "text", + "tools", + "tool_choice", + "max_tool_calls", + "thinking", + "caching", + "expire_at", + "context_management", + # LiteLLM-internal metadata (not sent to provider) + "metadata", + # Request plumbing helpers + "extra_headers", + "extra_query", + "extra_body", + "timeout", + ] + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.VOLCENGINE + + def get_supported_openai_params(self, model: str) -> list: + """ + Volcengine Responses API: only documented parameters are supported. + """ + supported = ["input", "model"] + list(self._SUPPORTED_OPTIONAL_PARAMS) + # Do not advertise internal-only metadata to callers; we still accept and drop it before send. + if "metadata" in supported: + supported.remove("metadata") + return supported + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> VolcEngineError: + typed_headers: httpx.Headers = ( + headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers or {}) + ) + return VolcEngineError( + status_code=status_code, + message=error_message, + headers=typed_headers, + ) + + def validate_environment( + self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + """ + Build auth headers for Volcengine Responses API. + """ + if litellm_params is None: + litellm_params = GenericLiteLLMParams() + elif isinstance(litellm_params, dict): + litellm_params = GenericLiteLLMParams(**litellm_params) + + api_key = ( + litellm_params.api_key + or litellm.api_key + or get_secret_str("ARK_API_KEY") + or get_secret_str("VOLCENGINE_API_KEY") + ) + + if api_key is None: + raise ValueError( + "Volcengine API key is required. Set ARK_API_KEY / VOLCENGINE_API_KEY or pass api_key." + ) + + return get_volcengine_headers(api_key=api_key, extra_headers=headers) + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Construct Volcengine Responses API endpoint. + """ + base_url = ( + api_base + or litellm.api_base + or get_secret_str("VOLCENGINE_API_BASE") + or get_secret_str("ARK_API_BASE") + or get_volcengine_base_url() + ) + + base_url = base_url.rstrip("/") + + if base_url.endswith("/responses"): + return base_url + if base_url.endswith("/api/v3"): + return f"{base_url}/responses" + return f"{base_url}/api/v3/responses" + + def map_openai_params( + self, + response_api_optional_params: ResponsesAPIOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """ + Volcengine Responses API aligns with OpenAI parameters. + Remove parameters not supported by the public docs. + """ + params = { + key: value + for key, value in dict(response_api_optional_params).items() + if key in self._SUPPORTED_OPTIONAL_PARAMS + } + + # LiteLLM metadata is internal-only; don't send to provider + params.pop("metadata", None) + + # Volcengine docs do not list parallel_tool_calls; drop it to avoid backend errors. + if "parallel_tool_calls" in params: + verbose_logger.debug( + "Volcengine Responses API: dropping unsupported 'parallel_tool_calls' param." + ) + params.pop("parallel_tool_calls", None) + + return params + + def transform_responses_api_request( + self, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + """ + Volcengine rejects any undocumented fields (including extra_body). Fail fast + with clear errors and re-filter with the documented whitelist before delegating + to the OpenAI base transformer. + """ + allowed = set(self._SUPPORTED_OPTIONAL_PARAMS) + + sanitized_optional = { + k: v for k, v in response_api_optional_request_params.items() if k in allowed + } + # Ensure metadata never reaches provider + sanitized_optional.pop("metadata", None) + sanitized_optional.pop("parallel_tool_calls", None) + + # If extra_body is provided, filter its keys against the same allowlist to avoid + # leaking unsupported params to the provider. + if isinstance(sanitized_optional.get("extra_body"), dict): + filtered_body = { + k: v for k, v in sanitized_optional["extra_body"].items() if k in allowed + } + if filtered_body: + sanitized_optional["extra_body"] = filtered_body + else: + sanitized_optional.pop("extra_body", None) + + return super().transform_responses_api_request( + model=model, + input=input, + response_api_optional_request_params=sanitized_optional, + litellm_params=litellm_params, + headers=headers, + ) + + def transform_streaming_response( + self, + model: str, + parsed_chunk: dict, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIStreamingResponse: + """ + Volcengine may omit required fields; auto-fill them using event model defaults. + """ + chunk = parsed_chunk + + # Patch missing response.output on response.* events + if isinstance(chunk, dict): + resp = chunk.get("response") + if isinstance(resp, dict) and "output" not in resp: + patched_chunk = dict(chunk) + patched_resp = dict(resp) + patched_resp["output"] = [] + patched_chunk["response"] = patched_resp + chunk = patched_chunk + + event_type = str(chunk.get("type")) if isinstance(chunk, dict) else None + event_pydantic_model = OpenAIResponsesAPIConfig.get_event_model_class( + event_type=event_type + ) + + patched_chunk = self._fill_missing_fields(chunk, event_pydantic_model) + + return event_pydantic_model(**patched_chunk) + + def transform_response_api_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + logging_obj.post_call( + original_response=raw_response.text, + additional_args={"complete_input_dict": {}}, + ) + raw_response_json = raw_response.json() + if "created_at" in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["created_at"] + ) + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + try: + response = ResponsesAPIResponse(**raw_response_json) + except Exception: + verbose_logger.debug( + "Volcengine Responses API: falling back to model_construct for response parsing." + ) + response = ResponsesAPIResponse.model_construct(**raw_response_json) + + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + ######################################################### + ########## DELETE RESPONSE API TRANSFORMATION ############## + ######################################################### + def transform_delete_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_delete_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> DeleteResponseResult: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + try: + return DeleteResponseResult(**raw_response_json) + except Exception: + verbose_logger.debug( + "Volcengine Responses API: falling back to model_construct for delete response parsing." + ) + return DeleteResponseResult.model_construct(**raw_response_json) + + ######################################################### + ########## GET RESPONSE API TRANSFORMATION ############### + ######################################################### + def transform_get_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_get_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + response = ResponsesAPIResponse(**raw_response_json) + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + ######################################################### + ########## LIST INPUT ITEMS TRANSFORMATION ############# + ######################################################### + def transform_list_input_items_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + after: Optional[str] = None, + before: Optional[str] = None, + include: Optional[List[str]] = None, + limit: int = 20, + order: Literal["asc", "desc"] = "desc", + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}/input_items" + params: Dict[str, Any] = {} + if after is not None: + params["after"] = after + if before is not None: + params["before"] = before + if include: + params["include"] = ",".join(include) + if limit is not None: + params["limit"] = limit + if order is not None: + params["order"] = order + return url, params + + def transform_list_input_items_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> Dict: + try: + return raw_response.json() + except Exception: + raise VolcEngineError( + 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]: + 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: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + response = ResponsesAPIResponse(**raw_response_json) + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + def should_fake_stream( + self, + model: Optional[str], + stream: Optional[bool], + custom_llm_provider: Optional[str] = None, + ) -> bool: + """ + Volcengine Responses API supports native streaming; never fall back to fake stream. + """ + return False + + @staticmethod + def _fill_missing_fields( + chunk: Any, event_model: Any + ) -> Dict[str, Any]: + """ + Heuristically fill missing required fields with safe defaults based on the + event model's field annotations. This keeps parsing tolerant of providers that + omit non-essential fields. + """ + if not isinstance(chunk, dict) or event_model is None: + return chunk + + patched: Dict[str, Any] = dict(chunk) + fields_map = getattr(event_model, "model_fields", {}) or {} + + for name, field in fields_map.items(): + if name in patched: + patched[name] = VolcEngineResponsesAPIConfig._maybe_fill_nested( + patched[name], field.annotation + ) + continue + + # Explicit default or factory + if field.default is not pyd_fields.PydanticUndefined and field.default is not None: + patched[name] = field.default + continue + if ( + field.default_factory is not None + and field.default_factory is not pyd_fields.PydanticUndefined + ): + patched[name] = field.default_factory() + continue + + # Heuristic defaults for missing required fields + patched[name] = VolcEngineResponsesAPIConfig._default_for_annotation( + field.annotation + ) + + return patched + + @staticmethod + def _default_for_annotation(annotation: Any) -> Any: + origin = get_origin(annotation) + args = get_args(annotation) + + if annotation is int: + return 0 + if annotation is list or origin is list: + return [] + if origin is Union: + # Prefer empty list when any option is a list + if any((arg is list or get_origin(arg) is list) for arg in args): + return [] + if type(None) in args: + return None + if origin is Union and type(None) in args: + return None + + # Fallback to None when no safer guess exists + return None + + @staticmethod + def _maybe_fill_nested(value: Any, annotation: Any) -> Any: + """ + Recursively fill nested dict/list structures based on the annotated model. + """ + model_cls = VolcEngineResponsesAPIConfig._pick_model_class(annotation, value) + args = get_args(annotation) + + if isinstance(value, dict) and model_cls is not None: + return VolcEngineResponsesAPIConfig._fill_missing_fields(value, model_cls) + + if isinstance(value, list): + # Attempt to fill list elements if we know the element annotation + elem_ann: Any = args[0] if args else None + if elem_ann is not None: + return [ + VolcEngineResponsesAPIConfig._maybe_fill_nested(v, elem_ann) + for v in value + ] + + return value + + @staticmethod + def _pick_model_class(annotation: Any, value: Any) -> Optional[Any]: + """ + Choose the best-matching Pydantic model class for a nested dict. + """ + candidates: List[Any] = [] + origin = get_origin(annotation) + + if hasattr(annotation, "model_fields"): + candidates.append(annotation) + if origin is Union: + for arg in get_args(annotation): + if hasattr(arg, "model_fields"): + candidates.append(arg) + + if not candidates: + return None + + # Try to match by literal "type" field when available + if isinstance(value, dict): + v_type = value.get("type") + for candidate in candidates: + try: + type_field = candidate.model_fields.get("type") + if type_field is None: + continue + literal_ann = type_field.annotation + if get_origin(literal_ann) is Literal: + literal_values = get_args(literal_ann) + if v_type in literal_values: + return candidate + except Exception: + continue + + # Fall back to the first candidate + return candidates[0] diff --git a/litellm/utils.py b/litellm/utils.py index ac194e4f33..71f8877aac 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -619,7 +619,7 @@ def load_credentials_from_list(kwargs: dict): """ # Access CredentialAccessor via module to trigger lazy loading if needed CredentialAccessor = getattr(sys.modules[__name__], 'CredentialAccessor') - + credential_name = kwargs.get("litellm_credential_name") if credential_name and litellm.credential_list: credential_accessor = CredentialAccessor.get_credential_values(credential_name) @@ -646,7 +646,7 @@ def _is_gemini_model(model: Optional[str], custom_llm_provider: Optional[str]) - if custom_llm_provider in ["vertex_ai", "vertex_ai_beta"]: return model is not None and "gemini" in model.lower() return True - + # Check if model name contains gemini return model is not None and "gemini" in model.lower() @@ -668,7 +668,7 @@ def _process_assistant_message_tool_calls( """ role = msg_copy.get("role") tool_calls = msg_copy.get("tool_calls") - + if role == "assistant" and isinstance(tool_calls, list): new_tool_calls = [] for tc in tool_calls: @@ -681,17 +681,17 @@ def _process_assistant_message_tool_calls( else: new_tool_calls.append(tc) continue - + # Remove thought signature from ID if present if isinstance(tc_dict.get("id"), str): if thought_signature_separator in tc_dict["id"]: tc_dict["id"] = _remove_thought_signature_from_id( tc_dict["id"], thought_signature_separator ) - + new_tool_calls.append(tc_dict) msg_copy["tool_calls"] = new_tool_calls - + return msg_copy @@ -706,7 +706,7 @@ def _process_tool_message_id(msg_copy: dict, thought_signature_separator: str) - msg_copy["tool_call_id"] = _remove_thought_signature_from_id( msg_copy["tool_call_id"], thought_signature_separator ) - + return msg_copy @@ -717,7 +717,7 @@ def _remove_thought_signatures_from_messages( Remove thought signatures from tool call IDs in all messages. """ processed_messages = [] - + for msg in messages: # Handle Pydantic models (convert to dict) if hasattr(msg, "model_dump"): @@ -728,17 +728,17 @@ def _remove_thought_signatures_from_messages( # Unknown type, keep as is processed_messages.append(msg) continue - + # Process assistant messages with tool_calls msg_dict = _process_assistant_message_tool_calls( msg_dict, thought_signature_separator ) - + # Process tool messages with tool_call_id msg_dict = _process_tool_message_id(msg_dict, thought_signature_separator) - + processed_messages.append(msg_dict) - + return processed_messages @@ -958,7 +958,7 @@ def function_setup( # noqa: PLR0915 input=buffer.getvalue(), model=model, ) - + ### REMOVE THOUGHT SIGNATURES FROM TOOL CALL IDS FOR NON-GEMINI MODELS ### # Gemini models embed thought signatures in tool call IDs. When sending # messages with tool calls to non-Gemini providers, we need to remove these @@ -974,7 +974,7 @@ def function_setup( # noqa: PLR0915 # Get custom_llm_provider to determine target provider custom_llm_provider = kwargs.get("custom_llm_provider") - + # If custom_llm_provider not in kwargs, try to determine it from the model if not custom_llm_provider and model: try: @@ -985,18 +985,18 @@ def function_setup( # noqa: PLR0915 except Exception: # If we can't determine the provider, skip this processing pass - + # Only process if target is NOT a Gemini model if not _is_gemini_model(model, custom_llm_provider): verbose_logger.debug( "Removing thought signatures from tool call IDs for non-Gemini model" ) - + # Process messages to remove thought signatures processed_messages = _remove_thought_signatures_from_messages( messages, THOUGHT_SIGNATURE_SEPARATOR ) - + # Update messages in kwargs or args if "messages" in kwargs: kwargs["messages"] = processed_messages @@ -3035,7 +3035,7 @@ def get_optional_params_embeddings( # noqa: PLR0915 ): # Lazy load get_supported_openai_params get_supported_openai_params = getattr(sys.modules[__name__], 'get_supported_openai_params') - + # retrieve all parameters passed to the function passed_params = locals() custom_llm_provider = passed_params.pop("custom_llm_provider", None) @@ -7084,7 +7084,7 @@ def get_valid_models( # init litellm_params ################################# from litellm.types.router import LiteLLM_Params - + if litellm_params is None: litellm_params = LiteLLM_Params(model="") if api_key is not None: @@ -7618,7 +7618,7 @@ class ProviderConfigManager: @staticmethod def _build_provider_config_map() -> dict[LlmProviders, tuple[Callable, bool]]: """Build the provider-to-config mapping dictionary. - + Returns a dict mapping provider to (factory_function, needs_model_parameter). This avoids expensive inspect.signature() calls at runtime. """ @@ -7784,7 +7784,7 @@ class ProviderConfigManager: ) -> Optional[BaseConfig]: """ Returns the provider config for a given provider. - + Uses O(1) dictionary lookup for fast provider resolution. """ # Check JSON providers FIRST (these override standard mappings) @@ -8015,6 +8015,8 @@ class ProviderConfigManager: # Note: GPT models (gpt-3.5, gpt-4, gpt-5, etc.) support temperature parameter # O-series models (o1, o3) do not contain "gpt" and have different parameter restrictions is_gpt_model = model and "gpt" in model.lower() + is_o_series = model and ("o_series" in model.lower() or (supports_reasoning(model) and not is_gpt_model)) + is_o_series = model and ( "o_series" in model.lower() or (supports_reasoning(model) and not is_gpt_model) @@ -8030,6 +8032,8 @@ class ProviderConfigManager: return litellm.GithubCopilotResponsesAPIConfig() elif litellm.LlmProviders.LITELLM_PROXY == provider: return litellm.LiteLLMProxyResponsesAPIConfig() + elif litellm.LlmProviders.VOLCENGINE == provider: + return litellm.VolcEngineResponsesAPIConfig() elif litellm.LlmProviders.MANUS == provider: return litellm.ManusResponsesAPIConfig() return None @@ -8487,7 +8491,7 @@ class ProviderConfigManager: from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config return get_vertex_ai_ocr_config(model=model) - + MistralOCRConfig = getattr(sys.modules[__name__], 'MistralOCRConfig') PROVIDER_TO_CONFIG_MAP = { litellm.LlmProviders.MISTRAL: MistralOCRConfig, @@ -8925,12 +8929,12 @@ def __getattr__(name: str) -> Any: """Lazy import handler for utils module with cached registry for improved performance.""" # Use cached registry from _lazy_imports instead of importing tuples every time from litellm._lazy_imports import _get_lazy_import_registry - + registry = _get_lazy_import_registry() - + # Check if name is in registry and call the cached handler function if name in registry: handler_func = registry[name] return handler_func(name) - + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py new file mode 100644 index 0000000000..16930bcb99 --- /dev/null +++ b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py @@ -0,0 +1,274 @@ +""" +Tests for Volcengine Responses API transformation. +""" +import os +import sys + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm +from litellm.llms.volcengine.responses.transformation import ( + VolcEngineResponsesAPIConfig, +) +from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams +from litellm.types.responses.main import DeleteResponseResult +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + + +class TestVolcengineResponsesAPITransformation: + """Test Volcengine Responses API configuration and transformations.""" + + def test_provider_config_registration(self): + """Provider registry should return VolcEngineResponsesAPIConfig.""" + config = ProviderConfigManager.get_provider_responses_api_config( + model="volcengine/demo-model", + provider=LlmProviders.VOLCENGINE, + ) + + assert config is not None, "Config should not be None for Volcengine provider" + assert isinstance( + config, VolcEngineResponsesAPIConfig + ), f"Expected VolcEngineResponsesAPIConfig, got {type(config)}" + assert ( + config.custom_llm_provider == LlmProviders.VOLCENGINE + ), "custom_llm_provider should be VOLCENGINE" + + def test_parallel_tool_calls_dropped(self): + """Volcengine does not list parallel_tool_calls; ensure it is removed.""" + config = VolcEngineResponsesAPIConfig() + params = ResponsesAPIOptionalRequestParams( + parallel_tool_calls=True, + temperature=0.5, + metadata={"k": "v"}, + ) + + mapped = config.map_openai_params( + response_api_optional_params=params, + model="volcengine/demo-model", + drop_params=False, + ) + + assert "parallel_tool_calls" not in mapped, "parallel_tool_calls must be dropped" + assert mapped.get("temperature") == 0.5 + assert "metadata" not in mapped, "Undocumented params should not be included" + + def test_unsupported_params_are_dropped(self): + """Unknown fields should be dropped before send, including nested extra_body.""" + config = VolcEngineResponsesAPIConfig() + + request = config.transform_responses_api_request( + model="volcengine/demo-model", + input="hi", + response_api_optional_request_params={ + "unsupported_custom_param": 0.1, + "temperature": 0.2, + "metadata": {"k": "v"}, + "extra_body": {"unsupported_custom_param": 1, "temperature": 0.3}, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert "unsupported_custom_param" not in request + assert request["temperature"] == 0.2 + assert "metadata" not in request + assert "extra_body" in request + assert "unsupported_custom_param" not in request["extra_body"] + assert request["extra_body"]["temperature"] == 0.3 + + def test_get_complete_url_variants(self): + """Ensure Volcengine endpoint construction handles different bases.""" + config = VolcEngineResponsesAPIConfig() + + default_url = config.get_complete_url(api_base=None, litellm_params={}) + assert default_url == "https://ark.cn-beijing.volces.com/api/v3/responses" + + api_base_with_api = config.get_complete_url( + api_base="https://custom.volc.com/api/v3", litellm_params={} + ) + assert api_base_with_api == "https://custom.volc.com/api/v3/responses" + + api_base_full = config.get_complete_url( + api_base="https://custom.volc.com/api/v3/responses", litellm_params={} + ) + assert api_base_full == "https://custom.volc.com/api/v3/responses" + + @pytest.mark.parametrize( + "litellm_params, expected_key", + [ + ({"api_key": "dict-key"}, "dict-key"), + (GenericLiteLLMParams(api_key="attr-key"), "attr-key"), + ], + ) + def test_validate_environment_uses_api_key( + self, monkeypatch, litellm_params, expected_key + ): + """validate_environment should pull api key from params/env and attach headers.""" + config = VolcEngineResponsesAPIConfig() + + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.delenv("ARK_API_KEY", raising=False) + monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False) + + headers = config.validate_environment( + headers={}, model="volcengine/demo-model", litellm_params=litellm_params + ) + + assert headers.get("Authorization") == f"Bearer {expected_key}" + assert headers.get("Content-Type") == "application/json" + + def test_validate_environment_raises_without_key(self, monkeypatch): + """validate_environment should error when no key is available.""" + config = VolcEngineResponsesAPIConfig() + + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.delenv("ARK_API_KEY", raising=False) + monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False) + + with pytest.raises(ValueError): + config.validate_environment( + headers={}, model="volcengine/demo", litellm_params={} + ) + + def test_unsupported_params_are_dropped_with_extra_body(self): + """Unknown fields (including extra_body) should be dropped before send.""" + config = VolcEngineResponsesAPIConfig() + + request = config.transform_responses_api_request( + model="volcengine/demo-model", + input="hi", + response_api_optional_request_params={ + "unsupported_custom_param": 0.1, + "temperature": 0.2, + "metadata": {"k": "v"}, + "extra_body": {"unsupported_custom_param": 1, "temperature": 0.3}, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert "unsupported_custom_param" not in request + assert "metadata" not in request + assert request["temperature"] == 0.2 + assert "extra_body" in request + assert "unsupported_custom_param" not in request["extra_body"] + assert request["extra_body"]["temperature"] == 0.3 + + def test_valid_thinking_caching_and_expire_at_pass(self): + """Documented params should pass through without validation errors.""" + config = VolcEngineResponsesAPIConfig() + request = config.transform_responses_api_request( + model="volcengine/demo-model", + input="hi", + response_api_optional_request_params={ + "instructions": "do X", + "thinking": {"type": "enabled"}, + "caching": {"type": "enabled"}, + "expire_at": 1234567890, + "temperature": 0.5, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert request["thinking"]["type"] == "enabled" + assert request["caching"]["type"] == "enabled" + assert request["expire_at"] == 1234567890 + assert request["instructions"] == "do X" + + def test_supported_params_limited_to_docs(self): + """Supported params should match documented Volcengine surface.""" + config = VolcEngineResponsesAPIConfig() + supported = set(config.get_supported_openai_params("volcengine/demo-model")) + + expected = { + "input", + "model", + "instructions", + "max_output_tokens", + "previous_response_id", + "store", + "reasoning", + "stream", + "temperature", + "top_p", + "text", + "tools", + "tool_choice", + "max_tool_calls", + "thinking", + "caching", + "expire_at", + "extra_headers", + "extra_query", + "extra_body", + "timeout", + } + + assert supported == expected + + def test_error_class_returns_volcengine_error(self): + """Errors should be wrapped with VolcEngineError for consistent handling.""" + config = VolcEngineResponsesAPIConfig() + error = config.get_error_class("bad request", 400, headers={"x": "y"}) + from litellm.llms.volcengine.common_utils import VolcEngineError + + assert isinstance(error, VolcEngineError) + assert error.status_code == 400 + assert error.message == "bad request" + assert error.headers.get("x") == "y" + + def test_transform_response_api_response_sets_headers_and_created_at(self): + """Responses should include processed headers and keep created_at intact.""" + config = VolcEngineResponsesAPIConfig() + response_payload = { + "id": "resp_123", + "object": "response", + "created_at": 123, + "status": "completed", + "output": [], + "model": "demo-model", + "usage": {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}, + } + http_response = httpx.Response( + status_code=200, + json=response_payload, + request=httpx.Request("POST", "https://example.com/responses"), + headers={"x-test": "1"}, + ) + + result = config.transform_response_api_response( + model="volcengine/demo-model", + raw_response=http_response, + logging_obj=type( + "Logger", + (), + {"post_call": staticmethod(lambda **kwargs: None)}, + ), + ) + + assert result.created_at == 123 + assert result._hidden_params["headers"].get("x-test") == "1" + assert "additional_headers" in result._hidden_params + + def test_transform_delete_response_api_response_parses_json(self): + """DELETE response parsing should return DeleteResponseResult.""" + config = VolcEngineResponsesAPIConfig() + http_response = httpx.Response( + status_code=200, + json={"id": "resp_123", "deleted": True}, + request=httpx.Request("DELETE", "https://example.com/responses/resp_123"), + ) + + result = config.transform_delete_response_api_response( + raw_response=http_response, + logging_obj=None, + ) + + assert isinstance(result, DeleteResponseResult) + assert result.deleted is True