add image-genetation support

This commit is contained in:
yrk 2026-05-19 09:48:40 +08:00
parent 84bc74a4fd
commit 1b9acac18a
7 changed files with 749 additions and 4 deletions

View file

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

View file

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

View 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()

View 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",
)

View file

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

View file

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

View file

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