From 730fe9c72d7b52193f01a4a58305b5766f3c438b Mon Sep 17 00:00:00 2001 From: TensorNull Date: Wed, 3 Jun 2026 21:40:28 +0800 Subject: [PATCH] fix: address greptile cometapi review comments --- litellm/llms/cometapi/chat/transformation.py | 7 --- .../image_generation/transformation.py | 19 +++++--- litellm/main.py | 48 ++++++++----------- .../chat/test_cometapi_chat_transformation.py | 24 +++++----- .../llms/cometapi/test_cometapi_endpoints.py | 9 +++- 5 files changed, 51 insertions(+), 56 deletions(-) diff --git a/litellm/llms/cometapi/chat/transformation.py b/litellm/llms/cometapi/chat/transformation.py index 45c74740dbb..3bd9334732e 100644 --- a/litellm/llms/cometapi/chat/transformation.py +++ b/litellm/llms/cometapi/chat/transformation.py @@ -44,13 +44,6 @@ class CometAPIConfig(OpenAIGPTConfig): response = super().transform_request( model, messages, optional_params, litellm_params, headers ) - overlapping_keys = set(response).intersection(extra_body) - if overlapping_keys: - raise ValueError( - "CometAPI extra_body cannot override request fields: {}".format( - ", ".join(sorted(overlapping_keys)) - ) - ) response.update(extra_body) return response diff --git a/litellm/llms/cometapi/image_generation/transformation.py b/litellm/llms/cometapi/image_generation/transformation.py index 53cf2107ab2..4e9358a1acb 100644 --- a/litellm/llms/cometapi/image_generation/transformation.py +++ b/litellm/llms/cometapi/image_generation/transformation.py @@ -189,18 +189,25 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig): additional_args={"complete_input_dict": request_data}, original_response=response_data, ) + response_data.update( + { + key: value + for key, value in { + "size": optional_params.get("size"), + "quality": optional_params.get("quality"), + "output_format": optional_params.get( + "output_format", optional_params.get("response_format") + ), + }.items() + if value is not None + } + ) image_response: ImageResponse = convert_to_model_response_object( # type: ignore response_object=response_data, model_response_object=model_response, response_type="image_generation", ) - image_response.size = optional_params.get("size") - image_response.quality = optional_params.get("quality") - image_response.output_format = optional_params.get( - "output_format", optional_params.get("response_format") - ) - return image_response def get_error_class( diff --git a/litellm/main.py b/litellm/main.py index 640c13604c6..2b7b9dd81de 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -108,6 +108,10 @@ from litellm.llms.base_llm.base_model_iterator import ( ) from litellm.llms.bedrock.common_utils import BedrockModelInfo from litellm.llms.cohere.common_utils import CohereModelInfo +from litellm.llms.cometapi.common_utils import ( + get_cometapi_api_base, + require_cometapi_api_key, +) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config from litellm.llms.openai_like.json_loader import JSONProviderRegistry @@ -323,19 +327,6 @@ ovhcloud_transformation = OVHCloudChatConfig() lemonade_transformation = LemonadeChatConfig() -def _get_cometapi_key_and_base( - api_key: Optional[str] = None, api_base: Optional[str] = None -) -> Tuple[str, str]: - from litellm.llms.cometapi.common_utils import ( - get_cometapi_api_base, - require_cometapi_api_key, - ) - - return require_cometapi_api_key( - api_key or litellm.cometapi_key - ), get_cometapi_api_base(api_base) - - MOCK_RESPONSE_TYPE = Union[str, Exception, dict, ModelResponse, ModelResponseStream] ####### COMPLETION ENDPOINTS ################ @@ -2552,9 +2543,8 @@ def completion( # type: ignore # noqa: PLR0915 stream=stream, ) elif custom_llm_provider == "cometapi": - api_key, api_base = _get_cometapi_key_and_base( - api_key=api_key, api_base=api_base - ) + api_key = require_cometapi_api_key(api_key or litellm.cometapi_key) + api_base = get_cometapi_api_base(api_base) ## COMPLETION CALL response = base_llm_http_handler.completion( @@ -5836,9 +5826,8 @@ def embedding( # noqa: PLR0915 litellm_params={}, ) elif custom_llm_provider == "cometapi": - api_key, api_base = _get_cometapi_key_and_base( - api_key=api_key, api_base=api_base - ) + api_key = require_cometapi_api_key(api_key or litellm.cometapi_key) + api_base = get_cometapi_api_base(api_base) response = base_llm_http_handler.embedding( model=model, input=input, @@ -6426,9 +6415,10 @@ def moderation( pass if custom_llm_provider == "cometapi": - api_key, api_base = _get_cometapi_key_and_base( - api_key=api_key or _dynamic_api_key, api_base=api_base or _dynamic_api_base + api_key = require_cometapi_api_key( + api_key or _dynamic_api_key or litellm.cometapi_key ) + api_base = get_cometapi_api_base(api_base or _dynamic_api_base) else: api_key = ( api_key @@ -6491,10 +6481,10 @@ async def amoderation( pass if custom_llm_provider == "cometapi": - api_key, api_base = _get_cometapi_key_and_base( - api_key=api_key or _dynamic_api_key, - api_base=optional_params.api_base or _dynamic_api_base, + api_key = require_cometapi_api_key( + api_key or _dynamic_api_key or litellm.cometapi_key ) + api_base = get_cometapi_api_base(optional_params.api_base or _dynamic_api_base) else: api_key = ( api_key @@ -6753,9 +6743,8 @@ def transcription( # noqa: PLR0915 litellm_params=litellm_params_dict, ) elif custom_llm_provider == "cometapi": - api_key, api_base = _get_cometapi_key_and_base( - api_key=api_key, api_base=api_base - ) + api_key = require_cometapi_api_key(api_key or litellm.cometapi_key) + api_base = get_cometapi_api_base(api_base) response = openai_audio_transcriptions.audio_transcriptions( model=model, audio_file=file, @@ -7005,9 +6994,10 @@ def speech( # noqa: PLR0915 model=model, llm_provider=custom_llm_provider, ) - api_key, api_base = _get_cometapi_key_and_base( - api_key=api_key or dynamic_api_key, api_base=api_base + api_key = require_cometapi_api_key( + api_key or dynamic_api_key or litellm.cometapi_key ) + api_base = get_cometapi_api_base(api_base) headers = headers or litellm.headers response = openai_chat_completions.audio_speech( model=model, diff --git a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py b/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py index 8a237de25e5..7e5483dc8f8 100644 --- a/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py +++ b/tests/test_litellm/llms/cometapi/chat/test_cometapi_chat_transformation.py @@ -139,21 +139,19 @@ class TestCometAPIConfig: {"role": "user", "content": "Hello, world!"} ] - def test_transform_request_extra_body_cannot_override_core_fields(self): - """Test extra_body cannot override the generated request body""" + def test_transform_request_extra_body_can_override_request_fields(self): + """Test extra_body preserves LiteLLM's existing override behavior""" config = CometAPIConfig() - with pytest.raises( - ValueError, - match="CometAPI extra_body cannot override request fields: model", - ): - config.transform_request( - model="cometapi/gpt-5.5", - messages=[{"role": "user", "content": "Hello, world!"}], - optional_params={"extra_body": {"model": "cometapi/gpt-5.5-all"}}, - litellm_params={}, - headers={}, - ) + transformed_request = config.transform_request( + model="cometapi/gpt-5.5", + messages=[{"role": "user", "content": "Hello, world!"}], + optional_params={"extra_body": {"model": "cometapi/gpt-5.5-all"}}, + litellm_params={}, + headers={}, + ) + + assert transformed_request["model"] == "cometapi/gpt-5.5-all" def test_cache_control_flag_removal(self): """Test cache control flag removal from messages""" diff --git a/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py b/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py index 2d01cb41755..b6469aacda3 100644 --- a/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py +++ b/tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py @@ -297,7 +297,11 @@ def test_cometapi_image_generation_normalizes_null_usage_fields(): model_response=ImageResponse(), logging_obj=MagicMock(), request_data={"prompt": "A small comet"}, - optional_params={"size": "1024x1024"}, + optional_params={ + "output_format": "png", + "quality": "high", + "size": "1024x1024", + }, litellm_params={}, encoding=None, ) @@ -308,6 +312,9 @@ def test_cometapi_image_generation_normalizes_null_usage_fields(): assert response.usage.input_tokens == 12 assert response.usage.output_tokens == 100 assert response.usage.total_tokens == 112 + assert response.output_format == "png" + assert response.quality == "high" + assert response.size == "1024x1024" def test_cometapi_image_generation_handles_missing_usage():