mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
add image-genetation support
This commit is contained in:
parent
84bc74a4fd
commit
1b9acac18a
7 changed files with 749 additions and 4 deletions
|
|
@ -334,7 +334,7 @@ def get_llm_provider( # noqa: PLR0915
|
|||
elif endpoint == "dashscope-intl.aliyuncs.com/compatible-mode/v1":
|
||||
custom_llm_provider = "dashscope"
|
||||
dynamic_api_key = get_secret_str("DASHSCOPE_API_KEY")
|
||||
elif endpoint == "api-inference.modelscope.cn/v1":
|
||||
elif endpoint == "https://api-inference.modelscope.cn/v1":
|
||||
custom_llm_provider = "modelscope"
|
||||
dynamic_api_key = get_secret_str("MODELSCOPE_API_KEY")
|
||||
elif endpoint == "api.moonshot.ai/v1":
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
|||
|
||||
|
||||
class ModelScopeChatConfig(OpenAIGPTConfig):
|
||||
DEFAULT_BASE_URL: str = "https://api-inference.modelscope.cn/v1"
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
|
|
@ -49,7 +51,7 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
|
|||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("MODELSCOPE_API_BASE")
|
||||
or "https://api-inference.modelscope.cn/v1"
|
||||
or self.DEFAULT_BASE_URL
|
||||
) # type: ignore
|
||||
dynamic_api_key = api_key or get_secret_str("MODELSCOPE_API_KEY")
|
||||
return api_base, dynamic_api_key
|
||||
|
|
@ -67,7 +69,7 @@ class ModelScopeChatConfig(OpenAIGPTConfig):
|
|||
If api_base is not provided, use the default ModelScope /chat/completions endpoint.
|
||||
"""
|
||||
if not api_base:
|
||||
api_base = "https://api-inference.modelscope.cn/v1"
|
||||
api_base = self.DEFAULT_BASE_URL
|
||||
|
||||
if not api_base.endswith("/chat/completions"):
|
||||
api_base = f"{api_base}/chat/completions"
|
||||
|
|
|
|||
31
litellm/llms/modelscope/image_generation/__init__.py
Normal file
31
litellm/llms/modelscope/image_generation/__init__.py
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
"""
|
||||
ModelScope Image Generation Module
|
||||
|
||||
Factory function for getting the appropriate config class.
|
||||
"""
|
||||
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
|
||||
from .transformation import ModelScopeImageGenerationConfig
|
||||
|
||||
__all__ = [
|
||||
"ModelScopeImageGenerationConfig",
|
||||
"get_modelscope_image_generation_config",
|
||||
]
|
||||
|
||||
|
||||
def get_modelscope_image_generation_config(
|
||||
model: str,
|
||||
) -> BaseImageGenerationConfig:
|
||||
"""
|
||||
Get the ModelScope config for image generation.
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "modelscope/Qwen/Qwen-Image-Edit")
|
||||
|
||||
Returns:
|
||||
BaseImageGenerationConfig instance for ModelScope
|
||||
"""
|
||||
return ModelScopeImageGenerationConfig()
|
||||
246
litellm/llms/modelscope/image_generation/transformation.py
Normal file
246
litellm/llms/modelscope/image_generation/transformation.py
Normal file
|
|
@ -0,0 +1,246 @@
|
|||
"""
|
||||
ModelScope Image Generation Config
|
||||
|
||||
Handles transformation between OpenAI-compatible format and ModelScope API format.
|
||||
|
||||
API Reference: https://modelscope.cn/docs/model-service/API-Inference/intro
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class ModelScopeImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""
|
||||
Configuration for ModelScope image generation.
|
||||
|
||||
Supports text-to-image models like:
|
||||
- Qwen/Qwen-Image-Edit
|
||||
- And other ModelScope-hosted image generation models
|
||||
"""
|
||||
|
||||
DEFAULT_BASE_URL: str = "https://api-inference.modelscope.cn/v1"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIImageGenerationOptionalParams]:
|
||||
"""
|
||||
Return list of OpenAI params supported by ModelScope.
|
||||
|
||||
ModelScope supports standard OpenAI image generation parameters.
|
||||
"""
|
||||
return [
|
||||
"n", # Number of images to generate
|
||||
"size", # Size of the generated images
|
||||
"response_format", # url or b64_json
|
||||
"user", # User identifier
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to ModelScope parameters.
|
||||
|
||||
ModelScope uses the same parameter names as OpenAI.
|
||||
"""
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
if drop_params:
|
||||
non_default_params = {
|
||||
k: v for k, v in non_default_params.items() if k in supported_params
|
||||
}
|
||||
optional_params.update(non_default_params)
|
||||
return optional_params
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete URL for the ModelScope image generation API request.
|
||||
"""
|
||||
base_url: str = (
|
||||
api_base or get_secret_str("MODELSCOPE_API_BASE") or self.DEFAULT_BASE_URL
|
||||
)
|
||||
base_url = base_url.rstrip("/")
|
||||
|
||||
# Return the images endpoint
|
||||
return f"{base_url}/images/generations"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment and set up headers for ModelScope.
|
||||
"""
|
||||
final_api_key: Optional[str] = api_key or get_secret_str("MODELSCOPE_API_KEY")
|
||||
|
||||
if not final_api_key:
|
||||
raise ValueError(
|
||||
"MODELSCOPE_API_KEY is not set. "
|
||||
"Please set it via environment variable or pass api_key parameter."
|
||||
)
|
||||
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {final_api_key}",
|
||||
}
|
||||
|
||||
headers = {**headers, **default_headers}
|
||||
return headers
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform OpenAI-style request to ModelScope request format.
|
||||
|
||||
ModelScope uses the same format as OpenAI for image generation.
|
||||
"""
|
||||
# Build the request body (same as OpenAI)
|
||||
request_data: dict = {
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
}
|
||||
|
||||
# Add optional params
|
||||
for key, value in optional_params.items():
|
||||
if key.startswith("_"):
|
||||
continue
|
||||
request_data[key] = value
|
||||
|
||||
return request_data
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ImageResponse:
|
||||
"""
|
||||
Transform ModelScope response to OpenAI-compatible ImageResponse.
|
||||
|
||||
ModelScope returns the same format as OpenAI:
|
||||
{"created": timestamp, "data": [{"url": "..."}]}
|
||||
"""
|
||||
try:
|
||||
response_data = raw_response.json()
|
||||
except Exception as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"Error parsing ModelScope response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
model=model,
|
||||
)
|
||||
|
||||
# Check for errors in response
|
||||
if "error" in response_data:
|
||||
error_msg = response_data["error"].get(
|
||||
"message", str(response_data["error"])
|
||||
)
|
||||
raise self.get_error_class(
|
||||
error_message=f"ModelScope error: {error_msg}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
model=model,
|
||||
)
|
||||
|
||||
# Extract images from response
|
||||
data_list = response_data.get("data", [])
|
||||
if not model_response.data:
|
||||
model_response.data = []
|
||||
|
||||
for item in data_list:
|
||||
image_obj = ImageObject(
|
||||
url=item.get("url"),
|
||||
b64_json=item.get("b64_json"),
|
||||
revised_prompt=item.get("revised_prompt"),
|
||||
)
|
||||
model_response.data.append(image_obj)
|
||||
|
||||
return model_response
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Union[dict, httpx.Headers],
|
||||
model: Optional[str] = None,
|
||||
) -> Exception:
|
||||
"""Return the appropriate error class for ModelScope."""
|
||||
from litellm.exceptions import (
|
||||
AuthenticationError,
|
||||
BadRequestError,
|
||||
InternalServerError,
|
||||
)
|
||||
|
||||
if status_code == 400:
|
||||
return BadRequestError(
|
||||
message=error_message,
|
||||
model=model or "",
|
||||
llm_provider="modelscope",
|
||||
)
|
||||
elif status_code == 401:
|
||||
return AuthenticationError(
|
||||
message=error_message,
|
||||
model=model or "",
|
||||
llm_provider="modelscope",
|
||||
)
|
||||
elif status_code >= 500:
|
||||
return InternalServerError(
|
||||
message=error_message,
|
||||
model=model or "",
|
||||
llm_provider="modelscope",
|
||||
)
|
||||
else:
|
||||
return BadRequestError(
|
||||
message=error_message,
|
||||
model=model or "",
|
||||
llm_provider="modelscope",
|
||||
)
|
||||
|
|
@ -9298,6 +9298,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return get_dashscope_image_generation_config(model)
|
||||
elif LlmProviders.MODELSCOPE == provider:
|
||||
from litellm.llms.modelscope.image_generation import (
|
||||
get_modelscope_image_generation_config,
|
||||
)
|
||||
|
||||
return get_modelscope_image_generation_config(model)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"""
|
||||
Unit tests for ModelScope configuration.
|
||||
|
||||
These tests validate the DashScopeConfig class which extends OpenAIGPTConfig.
|
||||
These tests validate the ModelScopeChatConfig class which extends OpenAIGPTConfig.
|
||||
ModelScope is an OpenAI-compatible provider with minor customizations.
|
||||
"""
|
||||
|
||||
|
|
@ -0,0 +1,460 @@
|
|||
"""
|
||||
Unit tests for ModelScope image generation configuration.
|
||||
|
||||
These tests validate the ModelScopeImageGenerationConfig class which handles
|
||||
transformation between OpenAI-compatible format and ModelScope API format.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.llms.modelscope.image_generation.transformation import (
|
||||
ModelScopeImageGenerationConfig,
|
||||
)
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
class TestModelScopeImageGenerationTransformation:
|
||||
def setup_method(self):
|
||||
"""Set up test fixtures before each test method."""
|
||||
self.config = ModelScopeImageGenerationConfig()
|
||||
self.model = "modelscope/Qwen/Qwen-Image-Edit"
|
||||
self.logging_obj = MagicMock()
|
||||
|
||||
def test_get_supported_openai_params(self):
|
||||
"""Test that get_supported_openai_params returns correct parameters."""
|
||||
supported_params = self.config.get_supported_openai_params(self.model)
|
||||
|
||||
assert "n" in supported_params
|
||||
assert "size" in supported_params
|
||||
assert "response_format" in supported_params
|
||||
assert "user" in supported_params
|
||||
|
||||
def test_map_openai_params(self):
|
||||
"""Test that map_openai_params correctly passes through parameters."""
|
||||
non_default_params = {
|
||||
"n": 2,
|
||||
"size": "1024x1024",
|
||||
"response_format": "url",
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
result = self.config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=self.model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["n"] == 2
|
||||
assert result["size"] == "1024x1024"
|
||||
assert result["response_format"] == "url"
|
||||
|
||||
def test_map_openai_params_with_user(self):
|
||||
"""Test that map_openai_params correctly passes through user parameter."""
|
||||
non_default_params = {"user": "test-user-123"}
|
||||
optional_params = {}
|
||||
|
||||
result = self.config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=self.model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["user"] == "test-user-123"
|
||||
|
||||
def test_get_complete_url_default(self):
|
||||
"""Test that get_complete_url returns default ModelScope URL."""
|
||||
result = self.config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key="test_key",
|
||||
model=self.model,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert result == "https://api-inference.modelscope.cn/v1/images/generations"
|
||||
|
||||
def test_get_complete_url_with_custom_base(self):
|
||||
"""Test that get_complete_url uses custom api_base."""
|
||||
custom_base = "https://custom.modelscope.cn/v1"
|
||||
|
||||
result = self.config.get_complete_url(
|
||||
api_base=custom_base,
|
||||
api_key="test_key",
|
||||
model=self.model,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert result == f"{custom_base}/images/generations"
|
||||
|
||||
def test_get_complete_url_with_trailing_slash(self):
|
||||
"""Test that get_complete_url strips trailing slashes from base."""
|
||||
custom_base = "https://custom.modelscope.cn/v1/"
|
||||
|
||||
result = self.config.get_complete_url(
|
||||
api_base=custom_base,
|
||||
api_key="test_key",
|
||||
model=self.model,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert result == "https://custom.modelscope.cn/v1/images/generations"
|
||||
|
||||
@patch("litellm.llms.modelscope.image_generation.transformation.get_secret_str")
|
||||
def test_validate_environment_with_api_key(self, mock_get_secret):
|
||||
"""Test that validate_environment correctly sets authorization header."""
|
||||
headers = {}
|
||||
api_key = "test_api_key"
|
||||
|
||||
result = self.config.validate_environment(
|
||||
headers=headers,
|
||||
model=self.model,
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
assert result["Authorization"] == f"Bearer {api_key}"
|
||||
assert result["Content-Type"] == "application/json"
|
||||
mock_get_secret.assert_not_called()
|
||||
|
||||
@patch("litellm.llms.modelscope.image_generation.transformation.get_secret_str")
|
||||
def test_validate_environment_with_secret_key(self, mock_get_secret):
|
||||
"""Test that validate_environment uses secret API key when api_key is None."""
|
||||
mock_get_secret.return_value = "secret_api_key"
|
||||
headers = {}
|
||||
|
||||
result = self.config.validate_environment(
|
||||
headers=headers,
|
||||
model=self.model,
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert result["Authorization"] == "Bearer secret_api_key"
|
||||
mock_get_secret.assert_called_once_with("MODELSCOPE_API_KEY")
|
||||
|
||||
@patch("litellm.llms.modelscope.image_generation.transformation.get_secret_str")
|
||||
def test_validate_environment_no_api_key(self, mock_get_secret):
|
||||
"""Test that validate_environment raises error when no API key is available."""
|
||||
mock_get_secret.return_value = None
|
||||
headers = {}
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
self.config.validate_environment(
|
||||
headers=headers,
|
||||
model=self.model,
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert "MODELSCOPE_API_KEY is not set" in str(exc_info.value)
|
||||
|
||||
def test_transform_image_generation_request_basic(self):
|
||||
"""Test that transform_image_generation_request creates correct request body."""
|
||||
prompt = "A beautiful sunset over mountains"
|
||||
optional_params = {}
|
||||
|
||||
result = self.config.transform_image_generation_request(
|
||||
model=self.model,
|
||||
prompt=prompt,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["model"] == self.model
|
||||
assert result["prompt"] == prompt
|
||||
|
||||
def test_transform_image_generation_request_with_optional_params(self):
|
||||
"""Test that transform_image_generation_request includes optional params."""
|
||||
prompt = "A beautiful sunset"
|
||||
optional_params = {
|
||||
"n": 2,
|
||||
"size": "1024x1024",
|
||||
"response_format": "b64_json",
|
||||
}
|
||||
|
||||
result = self.config.transform_image_generation_request(
|
||||
model=self.model,
|
||||
prompt=prompt,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["model"] == self.model
|
||||
assert result["prompt"] == prompt
|
||||
assert result["n"] == 2
|
||||
assert result["size"] == "1024x1024"
|
||||
assert result["response_format"] == "b64_json"
|
||||
|
||||
def test_transform_image_generation_request_ignores_internal_params(self):
|
||||
"""Test that transform_image_generation_request ignores params starting with _."""
|
||||
prompt = "A beautiful sunset"
|
||||
optional_params = {
|
||||
"n": 2,
|
||||
"_internal_param": "should_be_ignored",
|
||||
}
|
||||
|
||||
result = self.config.transform_image_generation_request(
|
||||
model=self.model,
|
||||
prompt=prompt,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["model"] == self.model
|
||||
assert result["n"] == 2
|
||||
assert "_internal_param" not in result
|
||||
|
||||
def test_transform_image_generation_response_with_url_images(self):
|
||||
"""Test that transform_image_generation_response correctly extracts URL images."""
|
||||
response_data = {
|
||||
"created": 1234567890,
|
||||
"data": [
|
||||
{"url": "https://example.com/image1.png"},
|
||||
{"url": "https://example.com/image2.png"},
|
||||
],
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = response_data
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ImageResponse(data=[])
|
||||
|
||||
result = self.config.transform_image_generation_response(
|
||||
model=self.model,
|
||||
raw_response=mock_response,
|
||||
model_response=model_response,
|
||||
logging_obj=self.logging_obj,
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert len(result.data) == 2
|
||||
assert result.data[0].url == "https://example.com/image1.png"
|
||||
assert result.data[1].url == "https://example.com/image2.png"
|
||||
|
||||
def test_transform_image_generation_response_with_b64_json(self):
|
||||
"""Test that transform_image_generation_response correctly extracts base64 images."""
|
||||
response_data = {
|
||||
"created": 1234567890,
|
||||
"data": [
|
||||
{"b64_json": "iVBORw0KGgoAAAANS"},
|
||||
],
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = response_data
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ImageResponse(data=[])
|
||||
|
||||
result = self.config.transform_image_generation_response(
|
||||
model=self.model,
|
||||
raw_response=mock_response,
|
||||
model_response=model_response,
|
||||
logging_obj=self.logging_obj,
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0].b64_json == "iVBORw0KGgoAAAANS"
|
||||
assert result.data[0].url is None
|
||||
|
||||
def test_transform_image_generation_response_with_revised_prompt(self):
|
||||
"""Test that transform_image_generation_response extracts revised_prompt."""
|
||||
response_data = {
|
||||
"created": 1234567890,
|
||||
"data": [
|
||||
{
|
||||
"url": "https://example.com/image.png",
|
||||
"revised_prompt": "A detailed description of a beautiful sunset",
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = response_data
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ImageResponse(data=[])
|
||||
|
||||
result = self.config.transform_image_generation_response(
|
||||
model=self.model,
|
||||
raw_response=mock_response,
|
||||
model_response=model_response,
|
||||
logging_obj=self.logging_obj,
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert len(result.data) == 1
|
||||
assert (
|
||||
result.data[0].revised_prompt
|
||||
== "A detailed description of a beautiful sunset"
|
||||
)
|
||||
|
||||
def test_transform_image_generation_response_empty_data(self):
|
||||
"""Test that transform_image_generation_response handles empty data array."""
|
||||
response_data = {
|
||||
"created": 1234567890,
|
||||
"data": [],
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = response_data
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ImageResponse(data=[])
|
||||
|
||||
result = self.config.transform_image_generation_response(
|
||||
model=self.model,
|
||||
raw_response=mock_response,
|
||||
model_response=model_response,
|
||||
logging_obj=self.logging_obj,
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert len(result.data) == 0
|
||||
|
||||
def test_transform_image_generation_response_error_handling(self):
|
||||
"""Test that transform_image_generation_response raises error on API error."""
|
||||
response_data = {
|
||||
"error": {
|
||||
"message": "Invalid prompt provided",
|
||||
"type": "invalid_request_error",
|
||||
}
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = response_data
|
||||
mock_response.status_code = 400
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ImageResponse(data=[])
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
self.config.transform_image_generation_response(
|
||||
model=self.model,
|
||||
raw_response=mock_response,
|
||||
model_response=model_response,
|
||||
logging_obj=self.logging_obj,
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert "ModelScope error" in str(exc_info.value)
|
||||
assert "Invalid prompt provided" in str(exc_info.value)
|
||||
|
||||
def test_transform_image_generation_response_json_error(self):
|
||||
"""Test that transform_image_generation_response raises error on invalid JSON."""
|
||||
import json
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.side_effect = json.JSONDecodeError("Invalid JSON", "", 0)
|
||||
mock_response.status_code = 500
|
||||
mock_response.headers = {}
|
||||
|
||||
model_response = ImageResponse(data=[])
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
self.config.transform_image_generation_response(
|
||||
model=self.model,
|
||||
raw_response=mock_response,
|
||||
model_response=model_response,
|
||||
logging_obj=self.logging_obj,
|
||||
request_data={},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert "Error parsing ModelScope response" in str(exc_info.value)
|
||||
|
||||
def test_get_error_class_bad_request(self):
|
||||
"""Test that get_error_class returns BadRequestError for 400 status."""
|
||||
from litellm.exceptions import BadRequestError
|
||||
|
||||
error = self.config.get_error_class(
|
||||
error_message="Bad request",
|
||||
status_code=400,
|
||||
headers={"Content-Type": "application/json"},
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
assert isinstance(error, BadRequestError)
|
||||
|
||||
def test_get_error_class_authentication_error(self):
|
||||
"""Test that get_error_class returns AuthenticationError for 401 status."""
|
||||
from litellm.exceptions import AuthenticationError
|
||||
|
||||
error = self.config.get_error_class(
|
||||
error_message="Invalid API key",
|
||||
status_code=401,
|
||||
headers={"Content-Type": "application/json"},
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
assert isinstance(error, AuthenticationError)
|
||||
|
||||
def test_get_error_class_internal_server_error(self):
|
||||
"""Test that get_error_class returns InternalServerError for 500+ status."""
|
||||
from litellm.exceptions import InternalServerError
|
||||
|
||||
error = self.config.get_error_class(
|
||||
error_message="Internal server error",
|
||||
status_code=500,
|
||||
headers={"Content-Type": "application/json"},
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
assert isinstance(error, InternalServerError)
|
||||
|
||||
def test_get_error_class_default(self):
|
||||
"""Test that get_error_class returns BadRequestError for other status codes."""
|
||||
from litellm.exceptions import BadRequestError
|
||||
|
||||
error = self.config.get_error_class(
|
||||
error_message="Some error",
|
||||
status_code=404,
|
||||
headers={"Content-Type": "application/json"},
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
assert isinstance(error, BadRequestError)
|
||||
Loading…
Add table
Reference in a new issue