diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 7b9e575e030..ef13215495e 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -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": diff --git a/litellm/llms/modelscope/chat/transformation.py b/litellm/llms/modelscope/chat/transformation.py index b657bc8950f..257959685e8 100644 --- a/litellm/llms/modelscope/chat/transformation.py +++ b/litellm/llms/modelscope/chat/transformation.py @@ -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" diff --git a/litellm/llms/modelscope/image_generation/__init__.py b/litellm/llms/modelscope/image_generation/__init__.py new file mode 100644 index 00000000000..8b28ea962ce --- /dev/null +++ b/litellm/llms/modelscope/image_generation/__init__.py @@ -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() diff --git a/litellm/llms/modelscope/image_generation/transformation.py b/litellm/llms/modelscope/image_generation/transformation.py new file mode 100644 index 00000000000..2716a8057ef --- /dev/null +++ b/litellm/llms/modelscope/image_generation/transformation.py @@ -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", + ) diff --git a/litellm/utils.py b/litellm/utils.py index 8790feadae0..0c89eeec325 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/tests/test_litellm/llms/modelscope/test_modelscope_chat_transformation.py b/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py similarity index 97% rename from tests/test_litellm/llms/modelscope/test_modelscope_chat_transformation.py rename to tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py index c868eaaacaa..89c521533a1 100644 --- a/tests/test_litellm/llms/modelscope/test_modelscope_chat_transformation.py +++ b/tests/test_litellm/llms/modelscope/chat/test_modelscope_chat_transformation.py @@ -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. """ diff --git a/tests/test_litellm/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py b/tests/test_litellm/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py new file mode 100644 index 00000000000..69d50237f72 --- /dev/null +++ b/tests/test_litellm/llms/modelscope/image_generation/test_modelscope_image_gen_transformation.py @@ -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)