mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
feat: enhance cometapi endpoint support
This commit is contained in:
parent
d45e9e4d56
commit
f2e5e82203
12 changed files with 1324 additions and 719 deletions
|
|
@ -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) | ✅ | ✅ | ✅ | | | ✅ | ✅ | ✅ | ✅ | |
|
||||
|
|
|
|||
689
cookbook/LiteLLM_CometAPI.ipynb
vendored
689
cookbook/LiteLLM_CometAPI.ipynb
vendored
File diff suppressed because one or more lines are too long
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
159
litellm/main.py
159
litellm/main.py
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
677
tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py
Normal file
677
tests/test_litellm/llms/cometapi/test_cometapi_endpoints.py
Normal 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"
|
||||
Loading…
Add table
Reference in a new issue