fix: address greptile cometapi review comments

This commit is contained in:
TensorNull 2026-06-03 21:40:28 +08:00
parent bdecd5497c
commit 730fe9c72d
5 changed files with 51 additions and 56 deletions

View file

@ -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

View file

@ -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(

View file

@ -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,

View file

@ -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"""

View file

@ -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():