feat: enhance cometapi endpoint support

This commit is contained in:
TensorNull 2026-06-03 20:26:10 +08:00
parent d45e9e4d56
commit f2e5e82203
12 changed files with 1324 additions and 719 deletions

View file

@ -288,7 +288,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
| [Codestral (`codestral`)](https://docs.litellm.ai/docs/providers/codestral) | ✅ | ✅ | ✅ | | | | | | | |
| [Cohere (`cohere`)](https://docs.litellm.ai/docs/providers/cohere) | ✅ | ✅ | ✅ | ✅ | | | | | | ✅ |
| [Cohere Chat (`cohere_chat`)](https://docs.litellm.ai/docs/providers/cohere) | ✅ | ✅ | ✅ | | | | | | | |
| [CometAPI (`cometapi`)](https://docs.litellm.ai/docs/providers/cometapi) | ✅ | ✅ | ✅ | ✅ | | | | | | |
| [CometAPI (`cometapi`)](https://docs.litellm.ai/docs/providers/cometapi) | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | | |
| [CompactifAI (`compactifai`)](https://docs.litellm.ai/docs/providers/compactifai) | ✅ | ✅ | ✅ | | | | | | | |
| [Custom (`custom`)](https://docs.litellm.ai/docs/providers/custom_llm_server) | ✅ | ✅ | ✅ | | | | | | | |
| [Custom OpenAI (`custom_openai`)](https://docs.litellm.ai/docs/providers/openai_compatible) | ✅ | ✅ | ✅ | | | ✅ | ✅ | ✅ | ✅ | |

File diff suppressed because one or more lines are too long

View file

@ -152,6 +152,10 @@ def image_generation(
model: Optional[str] = None,
n: Optional[int] = None,
quality: Optional[Union[str, ImageGenerationRequestQuality]] = None,
background: Optional[str] = None,
moderation: Optional[str] = None,
output_compression: Optional[int] = None,
output_format: Optional[str] = None,
response_format: Optional[str] = None,
size: Optional[str] = None,
style: Optional[str] = None,
@ -176,6 +180,10 @@ def image_generation(
model: Optional[str] = None,
n: Optional[int] = None,
quality: Optional[Union[str, ImageGenerationRequestQuality]] = None,
background: Optional[str] = None,
moderation: Optional[str] = None,
output_compression: Optional[int] = None,
output_format: Optional[str] = None,
response_format: Optional[str] = None,
size: Optional[str] = None,
style: Optional[str] = None,
@ -200,6 +208,10 @@ def image_generation( # noqa: PLR0915
model: Optional[str] = None,
n: Optional[int] = None,
quality: Optional[Union[str, ImageGenerationRequestQuality]] = None,
background: Optional[str] = None,
moderation: Optional[str] = None,
output_compression: Optional[int] = None,
output_format: Optional[str] = None,
response_format: Optional[str] = None,
size: Optional[str] = None,
style: Optional[str] = None,
@ -270,6 +282,14 @@ def image_generation( # noqa: PLR0915
non_default_params = {
k: v for k, v in kwargs.items() if k not in default_params
} # model-specific params - pass them straight to the model/provider
for param_name, param_value in {
"background": background,
"moderation": moderation,
"output_compression": output_compression,
"output_format": output_format,
}.items():
if param_value is not None:
non_default_params[param_name] = param_value
image_generation_config: Optional[BaseImageGenerationConfig] = None
if (
@ -404,6 +424,7 @@ def image_generation( # noqa: PLR0915
elif custom_llm_provider in (
litellm.LlmProviders.RECRAFT,
litellm.LlmProviders.AIML,
litellm.LlmProviders.COMETAPI,
litellm.LlmProviders.GEMINI,
litellm.LlmProviders.FAL_AI,
litellm.LlmProviders.STABILITY,
@ -417,9 +438,12 @@ def image_generation( # noqa: PLR0915
f"image generation config is not supported for {custom_llm_provider}"
)
# Resolve api_base from litellm.api_base if not explicitly provided
_api_base = api_base or litellm.api_base
if custom_llm_provider == litellm.LlmProviders.COMETAPI:
_api_base = api_base
api_key = api_key or dynamic_api_key or litellm.cometapi_key
litellm_params_dict["api_base"] = _api_base
litellm_params_dict["api_key"] = api_key
return llm_http_handler.image_generation_handler(
api_key=api_key,
@ -493,7 +517,6 @@ def image_generation( # noqa: PLR0915
):
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers
# Forward OpenAI organization if present (set by proxy pre-call utils)
organization: Optional[str] = kwargs.get("organization", None)
model_response = openai_chat_completions.image_generation(
model=model,

View file

@ -5,17 +5,17 @@ Based on OpenAI-compatible API interface implementation
Documentation: [CometAPI Documentation Link]
"""
from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union
from typing import Any, AsyncIterator, Iterator, List, Optional, Union
import httpx
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, ModelResponseStream
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
from ..common_utils import CometAPIException
from ..common_utils import CometAPIException, get_cometapi_complete_url
class CometAPIConfig(OpenAIGPTConfig):
@ -26,47 +26,6 @@ class CometAPIConfig(OpenAIGPTConfig):
and only need to override necessary methods to handle CometAPI-specific features
"""
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI format parameters to CometAPI format
"""
mapped_openai_params = super().map_openai_params(
non_default_params, optional_params, model, drop_params
)
# CometAPI-specific parameters (if any)
extra_body: dict[str, Any] = {}
# TODO: Add CometAPI-specific parameter handling here
# Example:
# custom_param = non_default_params.pop("custom_param", None)
# if custom_param is not None:
# extra_body["custom_param"] = custom_param
if extra_body:
mapped_openai_params["extra_body"] = extra_body
return mapped_openai_params
def remove_cache_control_flag_from_messages_and_tools(
self,
model: str,
messages: List[AllMessageValues],
tools: Optional[List["ChatCompletionToolParam"]] = None,
) -> Tuple[List[AllMessageValues], Optional[List["ChatCompletionToolParam"]]]:
"""
Remove cache control flags from messages and tools if not supported
"""
# For CometAPI, use default behavior (remove cache control)
return super().remove_cache_control_flag_from_messages_and_tools(
model, messages, tools
)
def transform_request(
self,
model: str,
@ -81,10 +40,17 @@ class CometAPIConfig(OpenAIGPTConfig):
Returns:
dict: The transformed request. Sent as the body of the API call.
"""
extra_body = optional_params.pop("extra_body", {})
extra_body = optional_params.pop("extra_body", {}) or {}
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
@ -103,30 +69,7 @@ class CometAPIConfig(OpenAIGPTConfig):
Returns:
str: The complete URL for the API call.
"""
# Default base
if api_base is None:
api_base = "https://api.cometapi.com/v1"
endpoint = "chat/completions"
# Normalize
api_base = api_base.rstrip("/")
# If endpoint already present, return as-is
if endpoint in api_base:
return api_base
# Ensure we include /v1 prefix when missing
if api_base.endswith("/v1"):
return f"{api_base}/{endpoint}"
if api_base.endswith("/v1/"):
return f"{api_base}{endpoint}"
# If user provided https://api.cometapi.com, add /v1
if api_base == "https://api.cometapi.com":
return f"{api_base}/v1/{endpoint}"
# Generic fallback: if '/v1' not in path, add it
if "/v1" not in api_base.split("//", 1)[-1]:
return f"{api_base}/v1/{endpoint}"
return f"{api_base}/{endpoint}"
return get_cometapi_complete_url(api_base, "chat/completions")
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]

View file

@ -1,4 +1,70 @@
from typing import Optional
from urllib.parse import urlsplit, urlunsplit
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.secret_managers.main import get_secret_str
DEFAULT_COMETAPI_API_BASE = "https://api.cometapi.com/v1"
def get_cometapi_api_key(api_key: Optional[str] = None) -> Optional[str]:
return (
api_key or get_secret_str("COMETAPI_KEY") or get_secret_str("COMETAPI_API_KEY")
)
def get_cometapi_api_base(api_base: Optional[str] = None) -> str:
return (
api_base
or get_secret_str("COMETAPI_BASE_URL")
or get_secret_str("COMETAPI_API_BASE")
or DEFAULT_COMETAPI_API_BASE
)
def get_cometapi_complete_url(api_base: Optional[str], endpoint: str) -> str:
base_url = get_cometapi_api_base(api_base).rstrip("/")
normalized_endpoint = endpoint.strip("/")
parsed_base_url = urlsplit(base_url)
path_segments = [segment for segment in parsed_base_url.path.split("/") if segment]
endpoint_segments = normalized_endpoint.split("/")
invalid_version_segments = [
segment
for segment in path_segments
if segment.startswith("v") and segment[1:2].isdigit() and segment != "v1"
]
if "v1" not in path_segments and invalid_version_segments:
raise ValueError("CometAPI OpenAI-compatible endpoints require a /v1 api_base")
if path_segments[-len(endpoint_segments) :] == endpoint_segments:
if "v1" not in path_segments:
raise ValueError(
"CometAPI OpenAI-compatible endpoints require a /v1 api_base"
)
return base_url
if "v1" in path_segments:
complete_path_segments = path_segments + endpoint_segments
else:
complete_path_segments = path_segments + ["v1"] + endpoint_segments
complete_path = "/" + "/".join(complete_path_segments)
return urlunsplit(
(
parsed_base_url.scheme,
parsed_base_url.netloc,
complete_path,
parsed_base_url.query,
parsed_base_url.fragment,
)
)
def require_cometapi_api_key(api_key: Optional[str] = None) -> str:
final_api_key = get_cometapi_api_key(api_key)
if not final_api_key:
raise ValueError("COMETAPI_KEY or COMETAPI_API_KEY is not set")
return final_api_key
class CometAPIException(BaseLLMException):

View file

@ -9,11 +9,14 @@ import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse, Usage
from ..common_utils import CometAPIException
from ..common_utils import (
CometAPIException,
get_cometapi_complete_url,
require_cometapi_api_key,
)
class CometAPIEmbeddingConfig(BaseEmbeddingConfig):
@ -39,11 +42,7 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig):
"""
Get the complete URL for the CometAPI embedding endpoint.
"""
api_base = (
"https://api.cometapi.com/v1" if api_base is None else api_base.rstrip("/")
)
complete_url = f"{api_base}/embeddings"
return complete_url
return get_cometapi_complete_url(api_base, "embeddings")
def validate_environment(
self,
@ -58,11 +57,10 @@ class CometAPIEmbeddingConfig(BaseEmbeddingConfig):
"""
Validate and set up authentication headers for CometAPI.
"""
if api_key is None:
api_key = get_secret_str("COMETAPI_KEY")
final_api_key = require_cometapi_api_key(api_key)
default_headers = {
"Authorization": f"Bearer {api_key}",
"Authorization": f"Bearer {final_api_key}",
"accept": "application/json",
"Content-Type": "application/json",
}

View file

@ -1,16 +1,23 @@
from typing import TYPE_CHECKING, Any, List, Optional
from typing import TYPE_CHECKING, Any, List, Optional, Union
import httpx
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIImageGenerationOptionalParams,
)
from litellm.types.utils import ImageObject, ImageResponse
from litellm.types.utils import ImageResponse
from litellm.utils import convert_to_model_response_object
from ..common_utils import (
CometAPIException,
get_cometapi_complete_url,
require_cometapi_api_key,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -21,8 +28,28 @@ else:
class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
DEFAULT_BASE_URL: str = "https://api.cometapi.com"
IMAGE_GENERATION_ENDPOINT: str = "v1/images/generations"
@staticmethod
def _normalize_image_usage(response_data: dict) -> None:
usage = response_data.get("usage")
if not isinstance(usage, dict):
return
if usage.get("input_tokens") is None:
usage["input_tokens"] = 0
if usage.get("output_tokens") is None:
usage["output_tokens"] = 0
if usage.get("total_tokens") is None:
usage["total_tokens"] = usage["input_tokens"] + usage["output_tokens"]
input_tokens_details = usage.get("input_tokens_details")
if not isinstance(input_tokens_details, dict):
usage["input_tokens_details"] = {"image_tokens": 0, "text_tokens": 0}
return
if input_tokens_details.get("image_tokens") is None:
input_tokens_details["image_tokens"] = 0
if input_tokens_details.get("text_tokens") is None:
input_tokens_details["text_tokens"] = 0
def get_supported_openai_params(
self, model: str
@ -31,11 +58,16 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
https://api.cometapi.com/v1/images/generations
"""
return [
"background",
"moderation",
"n",
"output_compression",
"output_format",
"quality",
"response_format",
"size",
"style",
"user",
]
def map_openai_params(
@ -50,7 +82,6 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
# CometAPI uses OpenAI-compatible parameters, so we can pass them directly
optional_params[k] = non_default_params[k]
elif drop_params:
pass
@ -73,16 +104,7 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
"""
Get the complete url for the request
"""
complete_url: str = (
api_base
or get_secret_str("COMETAPI_BASE_URL")
or get_secret_str("COMETAPI_API_BASE")
or self.DEFAULT_BASE_URL
)
complete_url = complete_url.rstrip("/")
complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}"
return complete_url
return get_cometapi_complete_url(api_base, "images/generations")
def validate_environment(
self,
@ -94,13 +116,7 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
final_api_key: Optional[str] = (
api_key
or get_secret_str("COMETAPI_KEY")
or get_secret_str("COMETAPI_API_KEY")
)
if not final_api_key:
raise ValueError("COMETAPI_KEY or COMETAPI_API_KEY is not set")
final_api_key = require_cometapi_api_key(api_key)
headers["Authorization"] = f"Bearer {final_api_key}"
headers["Content-Type"] = "application/json"
@ -119,7 +135,6 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
https://api.cometapi.com/v1/images/generations
"""
# CometAPI uses OpenAI-compatible format
request_body = {
"prompt": prompt,
"model": model,
@ -154,17 +169,43 @@ class CometAPIImageGenerationConfig(BaseImageGenerationConfig):
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
if raw_response.status_code >= 400:
error_data = response_data.get("error", response_data)
error_message = (
error_data.get("message")
if isinstance(error_data, dict)
else str(error_data)
)
raise self.get_error_class(
error_message=error_message or str(response_data),
status_code=raw_response.status_code,
headers=raw_response.headers,
)
# CometAPI returns OpenAI-compatible format
# Expected format: {"created": timestamp, "data": [{"url": "...", "b64_json": "..."}]}
if "data" in response_data:
for image_data in response_data["data"]:
image_obj = ImageObject(
b64_json=image_data.get("b64_json"),
url=image_data.get("url"),
)
model_response.data.append(image_obj)
self._normalize_image_usage(response_data)
logging_obj.post_call(
input=request_data.get("prompt", ""),
api_key=api_key,
additional_args={"complete_input_dict": request_data},
original_response=response_data,
)
image_response: ImageResponse = convert_to_model_response_object( # type: ignore
response_object=response_data,
model_response_object=model_response,
response_type="image_generation",
)
return model_response
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(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
return CometAPIException(
message=error_message, status_code=status_code, headers=headers
)

View file

@ -322,6 +322,20 @@ oci_transformation = OCIChatConfig()
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 ################
@ -2538,18 +2552,8 @@ def completion( # type: ignore # noqa: PLR0915
stream=stream,
)
elif custom_llm_provider == "cometapi":
api_key = (
api_key
or litellm.cometapi_key
or get_secret_str("COMETAPI_KEY")
or litellm.api_key
)
api_base = (
api_base
or litellm.api_base
or get_secret_str("COMETAPI_API_BASE")
or "https://api.cometapi.com/v1"
api_key, api_base = _get_cometapi_key_and_base(
api_key=api_key, api_base=api_base
)
## COMPLETION CALL
@ -5832,17 +5836,8 @@ def embedding( # noqa: PLR0915
litellm_params={},
)
elif custom_llm_provider == "cometapi":
api_key = (
api_key
or litellm.cometapi_key
or get_secret_str("COMETAPI_KEY")
or litellm.api_key
)
api_base = (
api_base
or litellm.api_base
or get_secret_str("COMETAPI_API_BASE")
or "https://api.cometapi.com/v1"
api_key, api_base = _get_cometapi_key_and_base(
api_key=api_key, api_base=api_base
)
response = base_llm_http_handler.embedding(
model=model,
@ -6411,16 +6406,38 @@ def adapter_completion(
def moderation(
input: str, model: Optional[str] = None, api_key: Optional[str] = None, **kwargs
) -> OpenAIModerationResponse:
# only supports open ai for now
api_key = (
api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
# Extract api_base from kwargs
api_base = kwargs.get("api_base", None)
custom_llm_provider = kwargs.get("custom_llm_provider", None)
_dynamic_api_key = None
_dynamic_api_base = None
try:
(
model,
custom_llm_provider,
_dynamic_api_key,
_dynamic_api_base,
) = litellm.get_llm_provider(
model=model or "",
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
)
except litellm.BadRequestError:
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
)
else:
api_key = (
api_key
or _dynamic_api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
api_base = api_base or _dynamic_api_base
openai_client = kwargs.get("client", None)
if openai_client is None:
@ -6450,17 +6467,11 @@ async def amoderation(
) -> OpenAIModerationResponse:
from openai import AsyncOpenAI
# only supports open ai for now
api_key = (
api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
optional_params = GenericLiteLLMParams(**kwargs)
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get(
"litellm_logging_obj", None
)
_dynamic_api_key = None
_dynamic_api_base = None
try:
(
@ -6478,6 +6489,21 @@ async def amoderation(
# `model` is optional field for moderation - get_llm_provider will throw BadRequestError if model is not set / not recognized
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,
)
else:
api_key = (
api_key
or _dynamic_api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
api_base = optional_params.api_base or _dynamic_api_base
openai_client = kwargs.get("client", None)
if openai_client is None or not isinstance(openai_client, AsyncOpenAI):
# call helper to get OpenAI client
@ -6485,7 +6511,7 @@ async def amoderation(
_openai_client: AsyncOpenAI = openai_chat_completions._get_openai_client( # type: ignore
is_async=True,
api_key=api_key,
api_base=optional_params.api_base or _dynamic_api_base,
api_base=api_base,
)
else:
_openai_client = openai_client
@ -6587,7 +6613,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
@client
def transcription(
def transcription( # noqa: PLR0915
model: str,
file: FileTypes,
## OPTIONAL OPENAI PARAMS ##
@ -6725,6 +6751,26 @@ def transcription(
max_retries=max_retries,
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
)
response = openai_audio_transcriptions.audio_transcriptions(
model=model,
audio_file=file,
optional_params=optional_params,
model_response=model_response,
atranscription=atranscription,
client=client,
timeout=timeout,
logging_obj=litellm_logging_obj,
max_retries=max_retries,
api_base=api_base,
api_key=api_key,
provider_config=provider_config,
litellm_params=litellm_params_dict,
shared_session=shared_session,
)
elif custom_llm_provider == "openai" or (
custom_llm_provider in litellm.openai_compatible_providers
):
@ -6740,8 +6786,6 @@ def transcription(
or get_secret("OPENAI_ORGANIZATION")
or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
)
# set API KEY
api_key = api_key or litellm.api_key or litellm.openai_key or get_secret("OPENAI_API_KEY") # type: ignore
response = openai_audio_transcriptions.audio_transcriptions(
model=model,
@ -6953,7 +6997,33 @@ def speech( # noqa: PLR0915
Coroutine[Any, Any, HttpxBinaryResponseContent],
None,
] = None
if (
if custom_llm_provider == "cometapi":
if voice is None or not (isinstance(voice, str)):
raise litellm.BadRequestError(
message="'voice' is required to be passed as a string for OpenAI TTS",
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
)
headers = headers or litellm.headers
response = openai_chat_completions.audio_speech(
model=model,
input=input,
voice=voice,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
organization=None,
project=None,
max_retries=max_retries,
timeout=timeout,
client=client, # pass AsyncOpenAI, OpenAI client
aspeech=aspeech,
shared_session=shared_session,
)
elif (
custom_llm_provider == "openai"
or custom_llm_provider in litellm.openai_compatible_providers
):
@ -6970,7 +7040,6 @@ def speech( # noqa: PLR0915
or get_secret("OPENAI_API_BASE")
or "https://api.openai.com/v1"
) # type: ignore
# set API KEY
api_key = (
api_key
or litellm.api_key # for deepinfra/perplexity/anyscale we check in get_llm_provider and pass in the api key from there

View file

@ -3143,6 +3143,10 @@ def get_optional_params_image_gen(
model: Optional[str] = None,
n: Optional[int] = None,
quality: Optional[str] = None,
background: Optional[str] = None,
moderation: Optional[str] = None,
output_compression: Optional[int] = None,
output_format: Optional[str] = None,
response_format: Optional[str] = None,
size: Optional[str] = None,
style: Optional[str] = None,
@ -3177,7 +3181,11 @@ def get_optional_params_image_gen(
passed_params[k] = v
default_params = {
"background": None,
"moderation": None,
"n": None,
"output_compression": None,
"output_format": None,
"quality": None,
"response_format": None,
"size": None,

View file

@ -607,10 +607,10 @@
"messages": true,
"responses": true,
"embeddings": true,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"image_generations": true,
"audio_transcriptions": true,
"audio_speech": true,
"moderations": true,
"batches": false,
"rerank": false,
"a2a": true,

View file

@ -30,7 +30,7 @@ class TestCometAPIChatCompletionStreamingHandler:
chunk = {
"id": "test_id",
"created": 1234567890,
"model": "gpt-3.5-turbo",
"model": "gpt-5.5",
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"choices": [
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
@ -44,7 +44,7 @@ class TestCometAPIChatCompletionStreamingHandler:
assert result.id == "test_id"
assert result.object == "chat.completion.chunk"
assert result.created == 1234567890
assert result.model == "gpt-3.5-turbo"
assert result.model == "gpt-5.5"
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
@ -93,14 +93,14 @@ class TestCometAPIConfig:
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-3.5-turbo",
model="cometapi/gpt-5.5",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={},
litellm_params={},
headers={},
)
assert transformed_request["model"] == "cometapi/gpt-3.5-turbo"
assert transformed_request["model"] == "cometapi/gpt-5.5"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
@ -110,7 +110,7 @@ class TestCometAPIConfig:
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-4",
model="cometapi/gpt-5.5",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={"extra_body": {"custom_param": "custom_value"}},
litellm_params={},
@ -123,12 +123,44 @@ class TestCometAPIConfig:
{"role": "user", "content": "Hello, world!"}
]
def test_transform_request_allows_empty_extra_body(self):
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-5.5",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={"extra_body": None},
litellm_params={},
headers={},
)
assert transformed_request["model"] == "cometapi/gpt-5.5"
assert transformed_request["messages"] == [
{"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"""
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={},
)
def test_cache_control_flag_removal(self):
"""Test cache control flag removal from messages"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-3.5-turbo",
model="cometapi/gpt-5.5",
messages=[
{
"role": "user",
@ -157,7 +189,7 @@ class TestCometAPIConfig:
mapped_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params={},
model="cometapi/gpt-3.5-turbo",
model="cometapi/gpt-5.5",
drop_params=False,
)
@ -179,142 +211,17 @@ class TestCometAPIConfig:
assert error.message == "Test error"
assert error.status_code == 400
def test_get_complete_url(self):
"""Test CometAPI chat endpoint URL normalization"""
config = CometAPIConfig()
# Integration test example (requires real API key)
@pytest.mark.skip(reason="Skipping integration test")
def test_cometapi_integration():
"""
Integration test - requires real API key
Run with: pytest -k test_cometapi_integration -s
"""
import os
from litellm import completion
# Try to get API key from multiple environment variables
api_key = (
os.getenv("COMETAPI_API_KEY")
or os.getenv("COMETAPI_KEY")
or os.getenv("COMET_API_KEY")
)
if not api_key:
pytest.skip("COMETAPI_API_KEY not set - skipping integration test")
response = completion(
model="cometapi/gpt-3.5-turbo",
messages=[{"role": "user", "content": "Say hello in one word"}],
api_key=api_key,
max_tokens=10,
temperature=0.7,
)
# Verify response structure
assert response.choices[0].message.content
assert len(response.choices[0].message.content.strip()) > 0
assert response.model
assert response.usage
assert response.usage.total_tokens > 0
def test_cometapi_streaming_integration():
"""
Integration test for streaming - requires real API key
Run with: pytest -k test_cometapi_streaming_integration -s
"""
import os
from litellm import completion
# Try to get API key from multiple environment variables
api_key = (
os.getenv("COMETAPI_API_KEY")
or os.getenv("COMETAPI_KEY")
or os.getenv("COMET_API_KEY")
)
if not api_key:
pytest.skip("COMETAPI_API_KEY not set - skipping streaming integration test")
try:
print(
f"🔍 Testing streaming with API key: {api_key[:6]}...{api_key[-4:]} (length: {len(api_key)})"
assert (
config.get_complete_url(
api_base="https://api.cometapi.com/v1",
api_key=None,
model="gpt-5.5",
optional_params={},
litellm_params={},
)
== "https://api.cometapi.com/v1/chat/completions"
)
print(f"🔍 API base URL: {os.getenv('COMETAPI_API_BASE', 'default')}")
# test streaming API call
response = completion(
model="cometapi/gpt-3.5-turbo",
messages=[{"role": "user", "content": "Count from 1 to 5"}],
api_key=api_key,
max_tokens=50,
stream=True,
)
# collect streaming response
chunks = []
content_parts = []
for chunk in response:
chunks.append(chunk)
if chunk.choices[0].delta.content:
content_parts.append(chunk.choices[0].delta.content)
# Verify we received at least one chunk and content
assert len(chunks) > 0, "Should receive at least one chunk"
assert len(content_parts) > 0, "Should receive content in chunks"
full_content = "".join(content_parts)
assert len(full_content.strip()) > 0, "Should have non-empty content"
print(f"✅ Received {len(chunks)} chunks")
print(f"✅ Full content: {full_content}")
except Exception as e:
print(f"❌ Streaming integration test error details:")
print(f" Error type: {type(e).__name__}")
print(f" Error message: {str(e)}")
if hasattr(e, "status_code"):
print(f" Status code: {e.status_code}")
if hasattr(e, "response"):
print(f" Response: {e.response}")
# Re-raise with more context for pytest
pytest.fail(f"Streaming integration test failed: {type(e).__name__}: {str(e)}")
def test_cometapi_with_custom_base_url():
"""
Test CometAPI with custom base URL
"""
import os
from litellm import completion
api_key = (
os.getenv("COMETAPI_API_KEY")
or os.getenv("COMETAPI_KEY")
or os.getenv("COMET_API_KEY")
)
custom_base_url = os.getenv("COMETAPI_API_BASE", "https://api.cometapi.com/v1")
if not api_key:
pytest.skip("COMETAPI_API_KEY not set - skipping custom base URL test")
try:
response = completion(
model="cometapi/gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello"}],
api_key=api_key,
api_base=custom_base_url,
max_tokens=5,
)
assert response.choices[0].message.content
print(f"✅ Custom base URL test passed: {response.choices[0].message.content}")
except Exception as e:
pytest.fail(f"Custom base URL test failed: {str(e)}")
if __name__ == "__main__":
# Quick test runner
pytest.main([__file__, "-v"])

View file

@ -0,0 +1,677 @@
import io
import json
from pathlib import Path
from unittest.mock import MagicMock, patch
import httpx
import pytest
import litellm
from litellm import completion
from litellm.images.main import image_generation
from litellm.llms.cometapi.common_utils import (
CometAPIException,
get_cometapi_api_base,
get_cometapi_api_key,
get_cometapi_complete_url,
)
from litellm.llms.cometapi.embed.transformation import CometAPIEmbeddingConfig
from litellm.llms.cometapi.image_generation.transformation import (
CometAPIImageGenerationConfig,
)
from litellm.main import amoderation, embedding, moderation, speech, transcription
from litellm.types.utils import EmbeddingResponse, ImageResponse, TranscriptionResponse
def _clear_cometapi_env(monkeypatch):
monkeypatch.delenv("COMETAPI_KEY", raising=False)
monkeypatch.delenv("COMETAPI_API_KEY", raising=False)
monkeypatch.delenv("COMETAPI_BASE_URL", raising=False)
monkeypatch.delenv("COMETAPI_API_BASE", raising=False)
monkeypatch.setattr(litellm, "cometapi_key", None, raising=False)
def _pollute_openai_globals(monkeypatch):
monkeypatch.setattr(litellm, "api_key", "openai-global-key", raising=False)
monkeypatch.setattr(litellm, "openai_key", "openai-provider-key", raising=False)
monkeypatch.setattr(litellm, "api_base", "https://openai.invalid/v1", raising=False)
@pytest.mark.parametrize(
("api_base", "expected"),
[
(None, "https://api.cometapi.com/v1/embeddings"),
("https://api.cometapi.com", "https://api.cometapi.com/v1/embeddings"),
("https://api.cometapi.com/v1", "https://api.cometapi.com/v1/embeddings"),
(
"https://api.cometapi.com/v1/embeddings",
"https://api.cometapi.com/v1/embeddings",
),
(
"https://proxy.example.com/openai/v1",
"https://proxy.example.com/openai/v1/embeddings",
),
(
"https://proxy.example.com/vertex/openai/v1",
"https://proxy.example.com/vertex/openai/v1/embeddings",
),
(
"https://proxy.example.com/api/v2/openai/v1",
"https://proxy.example.com/api/v2/openai/v1/embeddings",
),
(
"https://proxy.example.com/openai/v1/embeddings",
"https://proxy.example.com/openai/v1/embeddings",
),
],
)
def test_cometapi_embedding_url_normalization(api_base, expected, monkeypatch):
_clear_cometapi_env(monkeypatch)
assert (
CometAPIEmbeddingConfig().get_complete_url(
api_base=api_base,
api_key=None,
model="text-embedding-3-small",
optional_params={},
litellm_params={},
)
== expected
)
def test_cometapi_key_and_base_precedence(monkeypatch):
_clear_cometapi_env(monkeypatch)
monkeypatch.setenv("COMETAPI_KEY", "comet-env-key")
monkeypatch.setenv("COMETAPI_BASE_URL", "https://proxy.example.com/openai/v1")
assert get_cometapi_api_key("explicit-key") == "explicit-key"
assert get_cometapi_api_key() == "comet-env-key"
assert get_cometapi_api_base("https://explicit.example.com/v1") == (
"https://explicit.example.com/v1"
)
assert get_cometapi_api_base() == "https://proxy.example.com/openai/v1"
def test_cometapi_complete_url_preserves_existing_v1_path(monkeypatch):
_clear_cometapi_env(monkeypatch)
assert (
get_cometapi_complete_url("https://proxy.example.com/openai/v1", "embeddings")
== "https://proxy.example.com/openai/v1/embeddings"
)
@pytest.mark.parametrize(
"api_base",
[
"https://proxy.example.com/openai/v1beta",
"https://proxy.example.com/openai/v1beta/embeddings",
"https://proxy.example.com/openai/v10",
"https://proxy.example.com/openai/embeddings",
],
)
def test_cometapi_complete_url_rejects_non_v1_paths(api_base, monkeypatch):
_clear_cometapi_env(monkeypatch)
with pytest.raises(
ValueError,
match="CometAPI OpenAI-compatible endpoints require a /v1 api_base",
):
get_cometapi_complete_url(api_base, "embeddings")
def test_cometapi_embedding_uses_api_key_alias(monkeypatch):
_clear_cometapi_env(monkeypatch)
monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key")
headers = CometAPIEmbeddingConfig().validate_environment(
headers={},
model="text-embedding-3-small",
messages=[],
optional_params={},
litellm_params={},
)
assert headers["Authorization"] == "Bearer comet-env-key"
def test_cometapi_embedding_main_uses_cometapi_key_and_base(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key")
monkeypatch.setenv("COMETAPI_BASE_URL", "https://proxy.example.com/openai/v1")
with patch(
"litellm.main.base_llm_http_handler.embedding",
return_value=EmbeddingResponse(),
) as mock_embedding:
embedding(model="cometapi/text-embedding-3-small", input=["hello"])
call_kwargs = mock_embedding.call_args.kwargs
assert call_kwargs["model"] == "text-embedding-3-small"
assert call_kwargs["api_key"] == "comet-env-key"
assert call_kwargs["api_base"] == "https://proxy.example.com/openai/v1"
def test_cometapi_completion_uses_cometapi_key_and_base(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key")
monkeypatch.setenv("COMETAPI_BASE_URL", "https://proxy.example.com/openai/v1")
mock_response = MagicMock()
with patch(
"litellm.main.base_llm_http_handler.completion",
return_value=mock_response,
) as mock_completion:
response = completion(
model="cometapi/gpt-5.5",
messages=[{"role": "user", "content": "hello"}],
)
call_kwargs = mock_completion.call_args.kwargs
assert response is mock_response
assert call_kwargs["model"] == "gpt-5.5"
assert call_kwargs["api_key"] == "comet-env-key"
assert call_kwargs["api_base"] == "https://proxy.example.com/openai/v1"
assert call_kwargs["custom_llm_provider"] == "cometapi"
def test_cometapi_embedding_missing_key_fails(monkeypatch):
_clear_cometapi_env(monkeypatch)
with pytest.raises(ValueError, match="COMETAPI_KEY or COMETAPI_API_KEY"):
CometAPIEmbeddingConfig().validate_environment(
headers={},
model="text-embedding-3-small",
messages=[],
optional_params={},
litellm_params={},
)
def test_cometapi_image_generation_url_normalization():
config = CometAPIImageGenerationConfig()
assert (
config.get_complete_url(
api_base="https://api.cometapi.com/v1",
api_key=None,
model="gpt-image-2",
optional_params={},
litellm_params={},
)
== "https://api.cometapi.com/v1/images/generations"
)
assert (
config.get_complete_url(
api_base="https://api.cometapi.com/v1/images/generations",
api_key=None,
model="gpt-image-2",
optional_params={},
litellm_params={},
)
== "https://api.cometapi.com/v1/images/generations"
)
def test_cometapi_image_generation_validate_environment_uses_api_key_alias(monkeypatch):
_clear_cometapi_env(monkeypatch)
monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key")
headers = CometAPIImageGenerationConfig().validate_environment(
headers={},
model="gpt-image-2",
messages=[],
optional_params={},
litellm_params={},
)
assert headers["Authorization"] == "Bearer comet-env-key"
assert headers["Content-Type"] == "application/json"
def test_cometapi_image_generation_maps_new_openai_params(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
with patch(
"litellm.images.main.llm_http_handler.image_generation_handler",
return_value=ImageResponse(data=[{"url": "https://example.com/image.png"}]),
) as mock_image_handler:
image_generation(
model="cometapi/gpt-image-2",
prompt="A small comet over a clean API diagram",
api_key="comet-explicit-key",
output_compression=70,
output_format="png",
size="1024x1024",
)
call_kwargs = mock_image_handler.call_args.kwargs
assert call_kwargs["api_key"] == "comet-explicit-key"
assert call_kwargs["custom_llm_provider"] == "cometapi"
assert call_kwargs["litellm_params"]["api_base"] is None
assert (
call_kwargs["image_generation_optional_request_params"]["output_compression"]
== 70
)
assert (
call_kwargs["image_generation_optional_request_params"]["output_format"]
== "png"
)
def test_cometapi_image_generation_normalizes_null_usage_fields():
raw_response = httpx.Response(
200,
json={
"created": 123,
"data": [{"url": "https://example.com/image.png"}],
"usage": {
"input_tokens": 12,
"input_tokens_details": {
"image_tokens": None,
"text_tokens": None,
},
"output_tokens": 100,
"total_tokens": 112,
},
},
request=httpx.Request("POST", "https://api.cometapi.com/v1/images/generations"),
)
response = CometAPIImageGenerationConfig().transform_image_generation_response(
model="gpt-image-2",
raw_response=raw_response,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={"prompt": "A small comet"},
optional_params={"size": "1024x1024"},
litellm_params={},
encoding=None,
)
assert response.data[0].url == "https://example.com/image.png"
assert response.usage.input_tokens_details.image_tokens == 0
assert response.usage.input_tokens_details.text_tokens == 0
assert response.usage.input_tokens == 12
assert response.usage.output_tokens == 100
assert response.usage.total_tokens == 112
def test_cometapi_image_generation_handles_missing_usage():
raw_response = httpx.Response(
200,
json={
"created": 123,
"data": [{"url": "https://example.com/image.png"}],
},
request=httpx.Request("POST", "https://api.cometapi.com/v1/images/generations"),
)
response = CometAPIImageGenerationConfig().transform_image_generation_response(
model="gpt-image-2",
raw_response=raw_response,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={"prompt": "A small comet"},
optional_params={},
litellm_params={},
encoding=None,
)
assert response.data[0].url == "https://example.com/image.png"
assert response.usage.input_tokens == 0
assert response.usage.input_tokens_details.image_tokens == 0
assert response.usage.input_tokens_details.text_tokens == 0
assert response.usage.output_tokens == 0
assert response.usage.total_tokens == 0
def test_cometapi_image_generation_normalizes_null_usage_totals():
raw_response = httpx.Response(
200,
json={
"created": 123,
"data": [{"url": "https://example.com/image.png"}],
"usage": {
"input_tokens": None,
"input_tokens_details": None,
"output_tokens": None,
"total_tokens": None,
},
},
request=httpx.Request("POST", "https://api.cometapi.com/v1/images/generations"),
)
response = CometAPIImageGenerationConfig().transform_image_generation_response(
model="gpt-image-2",
raw_response=raw_response,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={"prompt": "A small comet"},
optional_params={},
litellm_params={},
encoding=None,
)
assert response.usage.input_tokens == 0
assert response.usage.input_tokens_details.image_tokens == 0
assert response.usage.input_tokens_details.text_tokens == 0
assert response.usage.output_tokens == 0
assert response.usage.total_tokens == 0
def test_cometapi_image_generation_raises_provider_error_on_error_response():
raw_response = httpx.Response(
500,
json={
"error": {
"message": "Transparent background is not supported for this model.",
"type": "image_generation_user_error",
"param": "background",
"code": "invalid_value",
}
},
request=httpx.Request("POST", "https://api.cometapi.com/v1/images/generations"),
)
with pytest.raises(
CometAPIException,
match="Transparent background is not supported for this model.",
):
CometAPIImageGenerationConfig().transform_image_generation_response(
model="gpt-image-2",
raw_response=raw_response,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={"prompt": "A small comet"},
optional_params={},
litellm_params={},
encoding=None,
)
def test_cometapi_speech_uses_cometapi_key_and_base(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
monkeypatch.setenv("COMETAPI_API_KEY", "comet-env-key")
with patch("litellm.main.openai_chat_completions.audio_speech") as mock_speech:
mock_speech.return_value = b"audio"
speech(model="cometapi/tts-1", input="hello", voice="alloy")
call_kwargs = mock_speech.call_args.kwargs
assert call_kwargs["model"] == "tts-1"
assert call_kwargs["api_key"] == "comet-env-key"
assert call_kwargs["api_base"] == "https://api.cometapi.com/v1"
def test_cometapi_speech_uses_dynamic_api_key(monkeypatch):
_clear_cometapi_env(monkeypatch)
with (
patch(
"litellm.main.get_llm_provider",
return_value=("tts-1", "cometapi", "dynamic-comet-key", None),
),
patch("litellm.main.openai_chat_completions.audio_speech") as mock_speech,
):
mock_speech.return_value = b"audio"
speech(model="cometapi/tts-1", input="hello", voice="alloy")
assert mock_speech.call_args.kwargs["api_key"] == "dynamic-comet-key"
def test_cometapi_speech_requires_voice(monkeypatch):
_clear_cometapi_env(monkeypatch)
with pytest.raises(litellm.BadRequestError, match="'voice' is required"):
speech(model="cometapi/tts-1", input="hello", api_key="comet-explicit-key")
def test_cometapi_transcription_uses_cometapi_key_and_base(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
audio_file = io.BytesIO(b"not-real-audio")
audio_file.name = "sample.wav"
with patch(
"litellm.main.openai_audio_transcriptions.audio_transcriptions",
return_value=TranscriptionResponse(text="hello"),
) as mock_transcription:
transcription(
model="cometapi/whisper-1",
file=audio_file,
api_key="comet-explicit-key",
)
call_kwargs = mock_transcription.call_args.kwargs
assert call_kwargs["model"] == "whisper-1"
assert call_kwargs["api_key"] == "comet-explicit-key"
assert call_kwargs["api_base"] == "https://api.cometapi.com/v1"
def test_openai_transcription_fallback_is_unchanged(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
audio_file = io.BytesIO(b"not-real-audio")
audio_file.name = "sample.wav"
with patch(
"litellm.main.openai_audio_transcriptions.audio_transcriptions",
return_value=TranscriptionResponse(text="hello"),
) as mock_transcription:
transcription(
model="whisper-1",
file=audio_file,
api_key="openai-explicit-key",
)
call_kwargs = mock_transcription.call_args.kwargs
assert call_kwargs["model"] == "whisper-1"
assert call_kwargs["api_key"] == "openai-explicit-key"
assert call_kwargs["api_base"] == "https://openai.invalid/v1"
def test_openai_speech_fallback_is_unchanged(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
with patch("litellm.main.openai_chat_completions.audio_speech") as mock_speech:
mock_speech.return_value = b"audio"
speech(
model="tts-1",
input="hello",
voice="alloy",
api_key="openai-explicit-key",
)
call_kwargs = mock_speech.call_args.kwargs
assert call_kwargs["model"] == "tts-1"
assert call_kwargs["api_key"] == "openai-explicit-key"
assert call_kwargs["api_base"] == "https://openai.invalid/v1"
assert call_kwargs["organization"] is None
assert call_kwargs["project"] is None
def test_provider_endpoint_matrix_only_updates_cometapi():
support_path = Path(__file__).parents[4] / "provider_endpoints_support.json"
providers = json.loads(support_path.read_text())["providers"]
for endpoint in (
"image_generations",
"audio_transcriptions",
"audio_speech",
"moderations",
):
assert providers["cometapi"]["endpoints"][endpoint] is True
assert providers["a2a"]["endpoints"][endpoint] is False
assert providers["bedrock"]["endpoints"][endpoint] is False
class _ModerationResponse:
def model_dump(self):
return {
"id": "modr-test",
"model": "omni-moderation-latest",
"results": [
{
"flagged": False,
"categories": {},
"category_scores": {},
"category_applied_input_types": {},
}
],
}
class _Moderations:
def __init__(self):
self.calls = []
def create(self, **kwargs):
self.calls.append(kwargs)
return _ModerationResponse()
class _OpenAIClient:
instances = []
def __init__(self, **kwargs):
self.kwargs = kwargs
self.moderations = _Moderations()
self.instances.append(self)
def test_cometapi_moderation_uses_cometapi_key_and_base(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
_OpenAIClient.instances = []
with patch("litellm.main.openai.OpenAI", _OpenAIClient):
moderation(
model="cometapi/omni-moderation-latest",
input="hello",
api_key="comet-explicit-key",
)
assert _OpenAIClient.instances[0].kwargs == {
"api_key": "comet-explicit-key",
"base_url": "https://api.cometapi.com/v1",
}
assert _OpenAIClient.instances[0].moderations.calls[0]["model"] == (
"omni-moderation-latest"
)
def test_cometapi_moderation_accepts_bare_model_with_custom_provider(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
_OpenAIClient.instances = []
with patch("litellm.main.openai.OpenAI", _OpenAIClient):
moderation(
model="omni-moderation-latest",
custom_llm_provider="cometapi",
input="hello",
api_key="comet-explicit-key",
)
assert _OpenAIClient.instances[0].kwargs == {
"api_key": "comet-explicit-key",
"base_url": "https://api.cometapi.com/v1",
}
assert _OpenAIClient.instances[0].moderations.calls[0]["model"] == (
"omni-moderation-latest"
)
def test_openai_moderation_fallback_is_unchanged(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
_OpenAIClient.instances = []
with patch("litellm.main.openai.OpenAI", _OpenAIClient):
moderation(model="omni-moderation-latest", input="hello")
assert _OpenAIClient.instances[0].kwargs == {
"api_key": "openai-global-key",
}
assert _OpenAIClient.instances[0].moderations.calls[0]["model"] == (
"omni-moderation-latest"
)
def test_openai_moderation_without_model_fallback_is_unchanged(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
_OpenAIClient.instances = []
with patch("litellm.main.openai.OpenAI", _OpenAIClient):
moderation(input="hello")
assert _OpenAIClient.instances[0].kwargs == {
"api_key": "openai-global-key",
}
assert "model" not in _OpenAIClient.instances[0].moderations.calls[0]
class _AsyncModerations:
def __init__(self):
self.calls = []
async def create(self, **kwargs):
self.calls.append(kwargs)
return _ModerationResponse()
class _AsyncOpenAIClient:
def __init__(self):
self.moderations = _AsyncModerations()
@pytest.mark.asyncio
async def test_cometapi_amoderation_uses_cometapi_key_and_base(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
fake_client = _AsyncOpenAIClient()
with patch(
"litellm.main.openai_chat_completions._get_openai_client",
return_value=fake_client,
) as mock_get_client:
await amoderation(
model="cometapi/omni-moderation-latest",
input="hello",
api_key="comet-explicit-key",
)
assert mock_get_client.call_args.kwargs["api_key"] == "comet-explicit-key"
assert mock_get_client.call_args.kwargs["api_base"] == "https://api.cometapi.com/v1"
assert fake_client.moderations.calls[0]["model"] == "omni-moderation-latest"
@pytest.mark.asyncio
async def test_openai_amoderation_fallback_is_unchanged(monkeypatch):
_clear_cometapi_env(monkeypatch)
_pollute_openai_globals(monkeypatch)
fake_client = _AsyncOpenAIClient()
with patch(
"litellm.main.openai_chat_completions._get_openai_client",
return_value=fake_client,
) as mock_get_client:
await amoderation(
model="omni-moderation-latest",
input="hello",
api_key="openai-explicit-key",
)
assert mock_get_client.call_args.kwargs["api_key"] == "openai-explicit-key"
assert mock_get_client.call_args.kwargs["api_base"] is None
assert fake_client.moderations.calls[0]["model"] == "omni-moderation-latest"