diff --git a/docs/my-website/docs/providers/stability.md b/docs/my-website/docs/providers/stability.md new file mode 100644 index 00000000000..49773fffdb3 --- /dev/null +++ b/docs/my-website/docs/providers/stability.md @@ -0,0 +1,181 @@ +# Stability AI +https://stability.ai/ + +## Overview + +| Property | Details | +|-------|-------| +| Description | Stability AI creates open AI models for image, video, audio, and 3D generation. Known for Stable Diffusion. | +| Provider Route on LiteLLM | `stability/` | +| Link to Provider Doc | [Stability AI API ↗](https://platform.stability.ai/docs/api-reference) | +| Supported Operations | [`/images/generations`](#image-generation) | + +LiteLLM supports Stability AI Image Generation calls via the Stability AI REST API (not via Bedrock). + +## API Key + +```python +# env variable +os.environ['STABILITY_API_KEY'] = "your-api-key" +``` + +Get your API key from the [Stability AI Platform](https://platform.stability.ai/). + +## Image Generation + +### Usage - LiteLLM Python SDK + +```python showLineNumbers +from litellm import image_generation +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +# Stability AI image generation call +response = image_generation( + model="stability/sd3.5-large", + prompt="A beautiful sunset over a calm ocean", +) +print(response) +``` + +### Usage - LiteLLM Proxy Server + +#### 1. Setup config.yaml + +```yaml showLineNumbers +model_list: + - model_name: sd3 + litellm_params: + model: stability/sd3.5-large + api_key: os.environ/STABILITY_API_KEY + model_info: + mode: image_generation + +general_settings: + master_key: sk-1234 +``` + +#### 2. Start the proxy + +```bash showLineNumbers +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +#### 3. Test it + +```bash showLineNumbers +curl --location 'http://0.0.0.0:4000/v1/images/generations' \ +--header 'Content-Type: application/json' \ +--header 'Authorization: Bearer sk-1234' \ +--data '{ + "model": "sd3", + "prompt": "A beautiful sunset over a calm ocean" +}' +``` + +### Advanced Usage - With Additional Parameters + +```python showLineNumbers +from litellm import image_generation +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +response = image_generation( + model="stability/sd3.5-large", + prompt="A beautiful sunset over a calm ocean", + size="1792x1024", # Maps to aspect_ratio 16:9 + negative_prompt="blurry, low quality", # Stability-specific + seed=12345, # For reproducibility +) +print(response) +``` + +### Supported Parameters + +Stability AI supports the following OpenAI-compatible parameters: + +| Parameter | Type | Description | Example | +|-----------|------|-------------|---------| +| `size` | string | Image dimensions (mapped to aspect_ratio) | `"1024x1024"` | +| `n` | integer | Number of images (note: Stability returns 1 per request) | `1` | +| `response_format` | string | Format of response (`b64_json` only for Stability) | `"b64_json"` | + +### Size to Aspect Ratio Mapping + +The `size` parameter is automatically mapped to Stability's `aspect_ratio`: + +| OpenAI Size | Stability Aspect Ratio | +|-------------|----------------------| +| `1024x1024` | `1:1` | +| `1792x1024` | `16:9` | +| `1024x1792` | `9:16` | +| `512x512` | `1:1` | +| `256x256` | `1:1` | + +### Using Stability-Specific Parameters + +You can pass parameters that are specific to Stability AI directly in your request: + +```python showLineNumbers +from litellm import image_generation +import os + +os.environ['STABILITY_API_KEY'] = "your-api-key" + +response = image_generation( + model="stability/sd3.5-large", + prompt="A beautiful sunset over a calm ocean", + # Stability-specific parameters + negative_prompt="blurry, watermark, text", + aspect_ratio="16:9", # Use directly instead of size + seed=42, + output_format="png", # png, jpeg, or webp +) +print(response) +``` + +### Supported Image Generation Models + +| Model Name | Function Call | Description | +|------------|---------------|-------------| +| sd3 | `image_generation(model="stability/sd3", ...)` | Stable Diffusion 3 | +| sd3-large | `image_generation(model="stability/sd3-large", ...)` | SD3 Large | +| sd3-large-turbo | `image_generation(model="stability/sd3-large-turbo", ...)` | SD3 Large Turbo (faster) | +| sd3-medium | `image_generation(model="stability/sd3-medium", ...)` | SD3 Medium | +| sd3.5-large | `image_generation(model="stability/sd3.5-large", ...)` | SD 3.5 Large (recommended) | +| sd3.5-large-turbo | `image_generation(model="stability/sd3.5-large-turbo", ...)` | SD 3.5 Large Turbo | +| sd3.5-medium | `image_generation(model="stability/sd3.5-medium", ...)` | SD 3.5 Medium | +| stable-image-ultra | `image_generation(model="stability/stable-image-ultra", ...)` | Stable Image Ultra | +| stable-image-core | `image_generation(model="stability/stable-image-core", ...)` | Stable Image Core | + +For more details on available models and features, see: https://platform.stability.ai/docs/api-reference + +## Response Format + +Stability AI returns images in base64 format. The response is OpenAI-compatible: + +```python +{ + "created": 1234567890, + "data": [ + { + "b64_json": "iVBORw0KGgo..." # Base64 encoded image + } + ] +} +``` + +## Comparing with Bedrock + +LiteLLM supports Stability AI models via two routes: + +| Route | Provider | Use Case | +|-------|----------|----------| +| `stability/` | Stability AI Direct API | Direct access, all latest models | +| `bedrock/stability.*` | AWS Bedrock | AWS integration, enterprise features | + +Use `stability/` for direct API access. Use `bedrock/stability.*` if you're already using AWS Bedrock. diff --git a/litellm/images/main.py b/litellm/images/main.py index 770b16c1ed2..4aae96bf715 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -346,6 +346,7 @@ def image_generation( # noqa: PLR0915 litellm.LlmProviders.AIML, litellm.LlmProviders.GEMINI, litellm.LlmProviders.FAL_AI, + litellm.LlmProviders.STABILITY, litellm.LlmProviders.RUNWAYML, litellm.LlmProviders.VERTEX_AI, ): diff --git a/litellm/llms/base_llm/image_generation/transformation.py b/litellm/llms/base_llm/image_generation/transformation.py index fc8db8c65c7..151e2893d1c 100644 --- a/litellm/llms/base_llm/image_generation/transformation.py +++ b/litellm/llms/base_llm/image_generation/transformation.py @@ -103,3 +103,11 @@ class BaseImageGenerationConfig(ABC): raise NotImplementedError( "ImageVariationConfig implements 'transform_response_image_variation' for image variation models" ) + + def use_multipart_form_data(self) -> bool: + """ + Returns True if this provider requires multipart/form-data instead of JSON. + + Override this method in subclasses that need form-data (e.g., Stability AI). + """ + return False diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 381d94f0186..4a7789a181f 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -3965,12 +3965,24 @@ class BaseLLMHTTPHandler: ) try: - response = sync_httpx_client.post( - url=api_base, - headers=headers, - json=data, - timeout=timeout, - ) + # Check if provider requires multipart/form-data (e.g., Stability AI) + if image_generation_provider_config.use_multipart_form_data(): + # Use form-data: pass files={} to force multipart encoding + response = sync_httpx_client.post( + url=api_base, + headers=headers, + data=data, + files={"none": ""}, # Forces multipart/form-data + timeout=timeout, + ) + else: + # Use JSON (default) + response = sync_httpx_client.post( + url=api_base, + headers=headers, + json=data, + timeout=timeout, + ) except Exception as e: raise self._handle_error( @@ -4063,12 +4075,24 @@ class BaseLLMHTTPHandler: ) try: - response = await async_httpx_client.post( - url=api_base, - headers=headers, - json=data, - timeout=timeout, - ) + # Check if provider requires multipart/form-data (e.g., Stability AI) + if image_generation_provider_config.use_multipart_form_data(): + # Use form-data: pass files={} to force multipart encoding + response = await async_httpx_client.post( + url=api_base, + headers=headers, + data=data, + files={"none": ""}, # Forces multipart/form-data + timeout=timeout, + ) + else: + # Use JSON (default) + response = await async_httpx_client.post( + url=api_base, + headers=headers, + json=data, + timeout=timeout, + ) except Exception as e: raise self._handle_error( diff --git a/litellm/llms/stability/__init__.py b/litellm/llms/stability/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/stability/image_generation/__init__.py b/litellm/llms/stability/image_generation/__init__.py new file mode 100644 index 00000000000..391fec6ddca --- /dev/null +++ b/litellm/llms/stability/image_generation/__init__.py @@ -0,0 +1,37 @@ +""" +Stability AI Image Generation Module + +Factory function for getting the appropriate config class. +""" + +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) + +from .transformation import StabilityImageGenerationConfig + +__all__ = [ + "StabilityImageGenerationConfig", + "get_stability_image_generation_config", +] + + +def get_stability_image_generation_config(model: str) -> BaseImageGenerationConfig: + """ + Get the appropriate Stability AI config for the given model. + + Currently all models use the same config class, but this factory + allows for model-specific configs in the future. + + Args: + model: The model name (e.g., "stability/sd3", "stability/stable-image-ultra") + + Returns: + BaseImageGenerationConfig instance for Stability AI + """ + # For now, all models use the same config + # In the future, we could have model-specific configs: + # - StabilitySD3Config for SD3 models + # - StabilityUltraConfig for Ultra models + # - etc. + return StabilityImageGenerationConfig() diff --git a/litellm/llms/stability/image_generation/transformation.py b/litellm/llms/stability/image_generation/transformation.py new file mode 100644 index 00000000000..d69dd399b2c --- /dev/null +++ b/litellm/llms/stability/image_generation/transformation.py @@ -0,0 +1,274 @@ +""" +Stability AI Image Generation Config + +Handles transformation between OpenAI-compatible format and Stability AI API format. + +API Reference: https://platform.stability.ai/docs/api-reference +""" + +from typing import TYPE_CHECKING, Any, List, Optional + +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.llms.stability import ( + OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO, + STABILITY_GENERATION_MODELS, + StabilityImageGenerationRequest, +) +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 StabilityImageGenerationConfig(BaseImageGenerationConfig): + """ + Configuration for Stability AI image generation. + + Supports: + - Stable Diffusion 3 (SD3, SD3.5) + - Stable Image Ultra + - Stable Image Core + """ + + DEFAULT_BASE_URL: str = "https://api.stability.ai" + + def get_supported_openai_params( + self, model: str + ) -> List[OpenAIImageGenerationOptionalParams]: + """ + Return list of OpenAI params supported by Stability AI. + + https://platform.stability.ai/docs/api-reference + """ + return [ + "n", # Number of images (Stability always returns 1, we can loop) + "size", # Maps to aspect_ratio + "response_format", # b64_json or url (Stability only returns b64) + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to Stability AI parameters. + + OpenAI -> Stability mappings: + - size -> aspect_ratio + - n -> (handled separately, Stability returns 1 image per request) + """ + supported_params = self.get_supported_openai_params(model) + + for k, v in non_default_params.items(): + if k not in optional_params: + if k in supported_params: + # Map size to aspect_ratio + if k == "size" and v in OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO: + optional_params["aspect_ratio"] = ( + OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO[v] + ) + elif k == "n": + # Store n for later, but don't pass to Stability + optional_params["_n"] = v + elif k == "response_format": + # Stability only returns base64, store for response handling + optional_params["_response_format"] = v + else: + optional_params[k] = v + elif drop_params: + pass + else: + raise ValueError( + f"Parameter {k} is not supported for model {model}. " + f"Supported parameters are {supported_params}. " + f"Set drop_params=True to drop unsupported parameters." + ) + + return optional_params + + def _get_model_endpoint(self, model: str) -> str: + """ + Get the API endpoint for a given model. + """ + # Remove "stability/" prefix if present + model_name = model.lower() + if model_name.startswith("stability/"): + model_name = model_name[10:] # Remove "stability/" prefix + + # Check if model is in our mapping + for key, endpoint in STABILITY_GENERATION_MODELS.items(): + if key in model_name: + return endpoint + + # Default to SD3 endpoint + return "/v2beta/stable-image/generate/sd3" + + 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 Stability AI API request. + """ + base_url: str = ( + api_base + or get_secret_str("STABILITY_API_BASE") + or self.DEFAULT_BASE_URL + ) + base_url = base_url.rstrip("/") + + endpoint = self._get_model_endpoint(model) + return f"{base_url}{endpoint}" + + 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 Stability AI. + """ + final_api_key: Optional[str] = api_key or get_secret_str("STABILITY_API_KEY") + + if not final_api_key: + raise ValueError( + "STABILITY_API_KEY is not set. " + "Please set it via environment variable or pass api_key parameter." + ) + + headers["Authorization"] = f"Bearer {final_api_key}" + headers["Accept"] = "application/json" + 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 Stability AI request format. + + Note: Stability AI uses multipart/form-data, but the HTTP handler + will handle the conversion from dict to form data. + """ + # Build Stability request + stability_request: StabilityImageGenerationRequest = { + "prompt": prompt, + "output_format": "png", # Default to PNG + } + + # Add optional params (already mapped in map_openai_params) + for key, value in optional_params.items(): + # Skip internal params (prefixed with _) + if key.startswith("_"): + continue + # Add supported Stability params + if key in [ + "negative_prompt", + "aspect_ratio", + "seed", + "output_format", + "model", + "mode", + "strength", + "style_preset", + ]: + stability_request[key] = value # type: ignore + + return dict(stability_request) + + 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 Stability AI response to OpenAI-compatible ImageResponse. + + Stability returns: {"image": "base64...", "finish_reason": "SUCCESS", "seed": 123} + OpenAI expects: {"data": [{"b64_json": "base64..."}], "created": timestamp} + """ + try: + response_data = raw_response.json() + except Exception as e: + raise self.get_error_class( + error_message=f"Error parsing Stability AI response: {e}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check for errors in response + if "errors" in response_data: + raise self.get_error_class( + error_message=f"Stability AI error: {response_data['errors']}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + # Check finish_reason + finish_reason = response_data.get("finish_reason", "") + if finish_reason == "CONTENT_FILTERED": + raise self.get_error_class( + error_message="Content was filtered by Stability AI safety systems", + status_code=400, + headers=raw_response.headers, + ) + + if not model_response.data: + model_response.data = [] + + # Extract image from response + image_b64 = response_data.get("image") + if image_b64: + model_response.data.append( + ImageObject( + b64_json=image_b64, + url=None, + revised_prompt=None, + ) + ) + + return model_response + + def use_multipart_form_data(self) -> bool: + """ + Stability AI requires multipart/form-data for image generation. + """ + return True diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 1ca26164431..c584deb683a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -23758,6 +23758,60 @@ "max_tokens": 8000, "mode": "chat" }, + "stability/sd3": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.065, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3-large": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.065, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3-large-turbo": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.04, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3-medium": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.035, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3.5-large": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.065, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3.5-large-turbo": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.04, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3.5-medium": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.035, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/stable-image-ultra": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.08, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/stable-image-core": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.03, + "supported_endpoints": ["/v1/images/generations"] + }, "stability.sd3-5-large-v1:0": { "litellm_provider": "bedrock", "max_input_tokens": 77, diff --git a/litellm/types/llms/stability.py b/litellm/types/llms/stability.py new file mode 100644 index 00000000000..33199ff769d --- /dev/null +++ b/litellm/types/llms/stability.py @@ -0,0 +1,212 @@ +""" +Type definitions for Stability AI API + +API Reference: https://platform.stability.ai/docs/api-reference +""" + +from typing import List, Literal, Optional + +from typing_extensions import TypedDict + + +class StabilityImageGenerationRequest(TypedDict, total=False): + """ + Base request parameters for Stability AI image generation. + + Used for endpoints: + - /v2beta/stable-image/generate/sd3 + - /v2beta/stable-image/generate/ultra + - /v2beta/stable-image/generate/core + """ + prompt: str # Required - text prompt for image generation + negative_prompt: Optional[str] # What to avoid in the image + aspect_ratio: Optional[str] # e.g., "1:1", "16:9", "9:16", "4:3", "3:4", "21:9", "9:21" + seed: Optional[int] # Random seed for reproducibility (0 to 4294967294) + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + model: Optional[str] # Model variant (e.g., "sd3.5-large", "sd3.5-medium") + mode: Optional[Literal["text-to-image", "image-to-image"]] # Generation mode + image: Optional[str] # Base64-encoded image for image-to-image + strength: Optional[float] # How much to transform the image (0-1) + style_preset: Optional[str] # Style preset name + + +class StabilityImageGenerationResponse(TypedDict, total=False): + """ + Response from Stability AI image generation endpoints. + """ + image: str # Base64-encoded image + finish_reason: str # "SUCCESS", "CONTENT_FILTERED", etc. + seed: int # The seed used for generation + + +class StabilityUpscaleRequest(TypedDict, total=False): + """ + Request parameters for Stability AI upscale endpoints. + + Used for endpoints: + - /v2beta/stable-image/upscale/fast + - /v2beta/stable-image/upscale/conservative + - /v2beta/stable-image/upscale/creative + """ + image: str # Required - Base64-encoded image to upscale + prompt: Optional[str] # Text prompt (required for creative upscale) + negative_prompt: Optional[str] # What to avoid + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + seed: Optional[int] # Random seed + creativity: Optional[float] # Creativity level for creative upscale (0-0.35) + + +class StabilityInpaintRequest(TypedDict, total=False): + """ + Request parameters for Stability AI inpaint endpoint. + + Endpoint: /v2beta/stable-image/edit/inpaint + """ + image: str # Required - Base64-encoded image to edit + prompt: str # Required - Description of desired changes + mask: Optional[str] # Base64-encoded mask (white = edit, black = keep) + negative_prompt: Optional[str] # What to avoid + seed: Optional[int] # Random seed + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + grow_mask: Optional[int] # Pixels to grow the mask by (0-100) + + +class StabilityOutpaintRequest(TypedDict, total=False): + """ + Request parameters for Stability AI outpaint endpoint. + + Endpoint: /v2beta/stable-image/edit/outpaint + """ + image: str # Required - Base64-encoded image to expand + prompt: Optional[str] # Description of content to generate + negative_prompt: Optional[str] # What to avoid + left: Optional[int] # Pixels to expand left (0-2000) + right: Optional[int] # Pixels to expand right (0-2000) + up: Optional[int] # Pixels to expand up (0-2000) + down: Optional[int] # Pixels to expand down (0-2000) + seed: Optional[int] # Random seed + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + creativity: Optional[float] # How creative to be (0-1) + + +class StabilityEraseRequest(TypedDict, total=False): + """ + Request parameters for Stability AI erase endpoint. + + Endpoint: /v2beta/stable-image/edit/erase + """ + image: str # Required - Base64-encoded image + mask: Optional[str] # Base64-encoded mask (white = erase) + seed: Optional[int] # Random seed + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + grow_mask: Optional[int] # Pixels to grow the mask by (0-100) + + +class StabilitySearchReplaceRequest(TypedDict, total=False): + """ + Request parameters for Stability AI search-and-replace endpoint. + + Endpoint: /v2beta/stable-image/edit/search-and-replace + """ + image: str # Required - Base64-encoded image + prompt: str # Required - Description of object to add + search_prompt: str # Required - Description of object to find and replace + negative_prompt: Optional[str] # What to avoid + seed: Optional[int] # Random seed + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + grow_mask: Optional[int] # Pixels to grow detected mask + + +class StabilityRemoveBackgroundRequest(TypedDict, total=False): + """ + Request parameters for Stability AI remove-background endpoint. + + Endpoint: /v2beta/stable-image/edit/remove-background + """ + image: str # Required - Base64-encoded image + output_format: Optional[Literal["png", "webp"]] # Output format (no jpeg - needs transparency) + + +class StabilityControlRequest(TypedDict, total=False): + """ + Request parameters for Stability AI control endpoints. + + Used for endpoints: + - /v2beta/stable-image/control/sketch + - /v2beta/stable-image/control/structure + - /v2beta/stable-image/control/style + """ + image: str # Required - Base64-encoded control image (sketch/structure/style reference) + prompt: str # Required - Description of desired output + negative_prompt: Optional[str] # What to avoid + control_strength: Optional[float] # How strongly to follow the control (0-1) + seed: Optional[int] # Random seed + output_format: Optional[Literal["jpeg", "png", "webp"]] # Output format + + +class StabilityEditResponse(TypedDict, total=False): + """ + Response from Stability AI edit/upscale/control endpoints. + """ + image: str # Base64-encoded result image + finish_reason: str # "SUCCESS", "CONTENT_FILTERED", etc. + seed: int # The seed used + + +# Mapping of OpenAI size to Stability aspect_ratio +OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO = { + "1024x1024": "1:1", + "1792x1024": "16:9", + "1024x1792": "9:16", + "512x512": "1:1", + "256x256": "1:1", +} + +# Stability AI supported aspect ratios +STABILITY_ASPECT_RATIOS = [ + "1:1", + "16:9", + "9:16", + "4:3", + "3:4", + "21:9", + "9:21", + "3:2", + "2:3", + "5:4", + "4:5", +] + +# Stability AI model endpoints +STABILITY_GENERATION_MODELS = { + "sd3": "/v2beta/stable-image/generate/sd3", + "sd3.5-large": "/v2beta/stable-image/generate/sd3", + "sd3.5-large-turbo": "/v2beta/stable-image/generate/sd3", + "sd3.5-medium": "/v2beta/stable-image/generate/sd3", + "sd3-large": "/v2beta/stable-image/generate/sd3", + "sd3-large-turbo": "/v2beta/stable-image/generate/sd3", + "sd3-medium": "/v2beta/stable-image/generate/sd3", + "stable-image-ultra": "/v2beta/stable-image/generate/ultra", + "stable-image-core": "/v2beta/stable-image/generate/core", +} + +STABILITY_EDIT_ENDPOINTS = { + "inpaint": "/v2beta/stable-image/edit/inpaint", + "outpaint": "/v2beta/stable-image/edit/outpaint", + "erase": "/v2beta/stable-image/edit/erase", + "search-and-replace": "/v2beta/stable-image/edit/search-and-replace", + "search-and-recolor": "/v2beta/stable-image/edit/search-and-recolor", + "remove-background": "/v2beta/stable-image/edit/remove-background", +} + +STABILITY_UPSCALE_ENDPOINTS = { + "fast": "/v2beta/stable-image/upscale/fast", + "conservative": "/v2beta/stable-image/upscale/conservative", + "creative": "/v2beta/stable-image/upscale/creative", +} + +STABILITY_CONTROL_ENDPOINTS = { + "sketch": "/v2beta/stable-image/control/sketch", + "structure": "/v2beta/stable-image/control/structure", + "style": "/v2beta/stable-image/control/style", +} diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 9d4dd7e6601..3ca803e017f 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2996,6 +2996,7 @@ class LlmProviders(str, Enum): HYPERBOLIC = "hyperbolic" RECRAFT = "recraft" FAL_AI = "fal_ai" + STABILITY = "stability" HEROKU = "heroku" AIML = "aiml" COMETAPI = "cometapi" diff --git a/litellm/utils.py b/litellm/utils.py index 9c3fa503ee3..e22109712f7 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7801,6 +7801,12 @@ class ProviderConfigManager: ) return get_fal_ai_image_generation_config(model) + elif LlmProviders.STABILITY == provider: + from litellm.llms.stability.image_generation import ( + get_stability_image_generation_config, + ) + + return get_stability_image_generation_config(model) elif LlmProviders.RUNWAYML == provider: from litellm.llms.runwayml.image_generation import ( get_runwayml_image_generation_config, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 1ca26164431..c584deb683a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -23758,6 +23758,60 @@ "max_tokens": 8000, "mode": "chat" }, + "stability/sd3": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.065, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3-large": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.065, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3-large-turbo": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.04, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3-medium": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.035, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3.5-large": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.065, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3.5-large-turbo": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.04, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/sd3.5-medium": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.035, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/stable-image-ultra": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.08, + "supported_endpoints": ["/v1/images/generations"] + }, + "stability/stable-image-core": { + "litellm_provider": "stability", + "mode": "image_generation", + "output_cost_per_image": 0.03, + "supported_endpoints": ["/v1/images/generations"] + }, "stability.sd3-5-large-v1:0": { "litellm_provider": "bedrock", "max_input_tokens": 77, diff --git a/tests/test_litellm/llms/stability/__init__.py b/tests/test_litellm/llms/stability/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/stability/image_generation/__init__.py b/tests/test_litellm/llms/stability/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/stability/image_generation/test_stability_image_generation.py b/tests/test_litellm/llms/stability/image_generation/test_stability_image_generation.py new file mode 100644 index 00000000000..85fe9552f00 --- /dev/null +++ b/tests/test_litellm/llms/stability/image_generation/test_stability_image_generation.py @@ -0,0 +1,314 @@ +""" +Tests for Stability AI Image Generation transformation + +Tests the transformation of OpenAI-compatible requests/responses to Stability AI format. +""" + +import json +from unittest.mock import MagicMock + +import httpx +import pytest + +from litellm.llms.stability.image_generation import ( + StabilityImageGenerationConfig, + get_stability_image_generation_config, +) +from litellm.types.llms.stability import ( + OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO, + STABILITY_GENERATION_MODELS, +) +from litellm.types.utils import ImageResponse + + +class TestStabilityImageGenerationConfig: + """Test the StabilityImageGenerationConfig class""" + + def setup_method(self): + """Set up test fixtures""" + self.config = StabilityImageGenerationConfig() + + def test_get_supported_openai_params(self): + """Test that supported OpenAI params are returned""" + params = self.config.get_supported_openai_params("stability/sd3") + assert "n" in params + assert "size" in params + assert "response_format" in params + + def test_map_openai_params_size_to_aspect_ratio(self): + """Test that OpenAI size is mapped to Stability aspect_ratio""" + non_default_params = {"size": "1024x1024"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="stability/sd3", + drop_params=False, + ) + + assert result.get("aspect_ratio") == "1:1" + + def test_map_openai_params_size_16_9(self): + """Test that 1792x1024 maps to 16:9 aspect ratio""" + non_default_params = {"size": "1792x1024"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="stability/sd3", + drop_params=False, + ) + + assert result.get("aspect_ratio") == "16:9" + + def test_map_openai_params_n_stored_internally(self): + """Test that n parameter is stored with underscore prefix""" + non_default_params = {"n": 2} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="stability/sd3", + drop_params=False, + ) + + assert result.get("_n") == 2 + assert "n" not in result + + def test_map_openai_params_unsupported_raises_error(self): + """Test that unsupported params raise ValueError when drop_params=False""" + non_default_params = {"unsupported_param": "value"} + optional_params = {} + + with pytest.raises(ValueError) as exc_info: + self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="stability/sd3", + drop_params=False, + ) + + assert "unsupported_param" in str(exc_info.value) + + def test_map_openai_params_unsupported_dropped(self): + """Test that unsupported params are dropped when drop_params=True""" + non_default_params = {"unsupported_param": "value", "size": "1024x1024"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="stability/sd3", + drop_params=True, + ) + + assert "unsupported_param" not in result + assert result.get("aspect_ratio") == "1:1" + + def test_get_model_endpoint_sd3(self): + """Test that SD3 model gets correct endpoint""" + endpoint = self.config._get_model_endpoint("stability/sd3") + assert endpoint == "/v2beta/stable-image/generate/sd3" + + def test_get_model_endpoint_sd35_large(self): + """Test that SD3.5 Large model gets correct endpoint""" + endpoint = self.config._get_model_endpoint("stability/sd3.5-large") + assert endpoint == "/v2beta/stable-image/generate/sd3" + + def test_get_model_endpoint_ultra(self): + """Test that Stable Image Ultra model gets correct endpoint""" + endpoint = self.config._get_model_endpoint("stability/stable-image-ultra") + assert endpoint == "/v2beta/stable-image/generate/ultra" + + def test_get_model_endpoint_core(self): + """Test that Stable Image Core model gets correct endpoint""" + endpoint = self.config._get_model_endpoint("stability/stable-image-core") + assert endpoint == "/v2beta/stable-image/generate/core" + + def test_get_complete_url(self): + """Test that complete URL is constructed correctly""" + url = self.config.get_complete_url( + api_base=None, + api_key="test-key", + model="stability/sd3", + optional_params={}, + litellm_params={}, + ) + + assert url == "https://api.stability.ai/v2beta/stable-image/generate/sd3" + + def test_get_complete_url_with_custom_base(self): + """Test that custom api_base is used when provided""" + url = self.config.get_complete_url( + api_base="https://custom.stability.ai", + api_key="test-key", + model="stability/sd3", + optional_params={}, + litellm_params={}, + ) + + assert url == "https://custom.stability.ai/v2beta/stable-image/generate/sd3" + + def test_validate_environment_sets_headers(self): + """Test that validate_environment sets correct headers""" + headers = self.config.validate_environment( + headers={}, + model="stability/sd3", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test-api-key", + ) + + assert headers["Authorization"] == "Bearer test-api-key" + assert headers["Accept"] == "application/json" + + def test_validate_environment_raises_without_api_key(self): + """Test that validate_environment raises error without API key""" + with pytest.raises(ValueError) as exc_info: + self.config.validate_environment( + headers={}, + model="stability/sd3", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + assert "STABILITY_API_KEY" in str(exc_info.value) + + def test_transform_image_generation_request(self): + """Test transformation of request to Stability format""" + result = self.config.transform_image_generation_request( + model="stability/sd3", + prompt="A beautiful sunset", + optional_params={"aspect_ratio": "16:9", "negative_prompt": "blurry"}, + litellm_params={}, + headers={}, + ) + + assert result["prompt"] == "A beautiful sunset" + assert result["output_format"] == "png" + assert result["aspect_ratio"] == "16:9" + assert result["negative_prompt"] == "blurry" + + def test_transform_image_generation_request_ignores_internal_params(self): + """Test that internal params (prefixed with _) are not included""" + result = self.config.transform_image_generation_request( + model="stability/sd3", + prompt="Test", + optional_params={"_n": 2, "_response_format": "url", "aspect_ratio": "1:1"}, + litellm_params={}, + headers={}, + ) + + assert "_n" not in result + assert "_response_format" not in result + assert result["aspect_ratio"] == "1:1" + + def test_transform_image_generation_response(self): + """Test transformation of Stability response to OpenAI format""" + # Mock the raw response + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "image": "base64encodedimage==", + "finish_reason": "SUCCESS", + "seed": 12345, + } + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + mock_logging = MagicMock() + + result = self.config.transform_image_generation_response( + model="stability/sd3", + raw_response=mock_response, + model_response=model_response, + logging_obj=mock_logging, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert len(result.data) == 1 + assert result.data[0].b64_json == "base64encodedimage==" + assert result.data[0].url is None + + def test_transform_image_generation_response_content_filtered(self): + """Test that content filtered response raises error""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "finish_reason": "CONTENT_FILTERED", + } + mock_response.status_code = 200 + mock_response.headers = {} + + model_response = ImageResponse(data=[]) + mock_logging = MagicMock() + + with pytest.raises(Exception) as exc_info: + self.config.transform_image_generation_response( + model="stability/sd3", + raw_response=mock_response, + model_response=model_response, + logging_obj=mock_logging, + request_data={}, + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert "filtered" in str(exc_info.value).lower() + + +class TestFactoryFunction: + """Test the factory function""" + + def test_get_stability_image_generation_config(self): + """Test that factory returns correct config type""" + config = get_stability_image_generation_config("stability/sd3") + assert isinstance(config, StabilityImageGenerationConfig) + + def test_factory_returns_config_for_any_model(self): + """Test that factory works for any model name""" + config = get_stability_image_generation_config("stability/custom-model") + assert isinstance(config, StabilityImageGenerationConfig) + + +class TestOpenAISizeMapping: + """Test the size to aspect ratio mapping""" + + def test_all_sizes_have_mappings(self): + """Test that standard OpenAI sizes have mappings""" + expected_sizes = ["1024x1024", "1792x1024", "1024x1792", "512x512", "256x256"] + for size in expected_sizes: + assert size in OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO + + def test_square_sizes_map_to_1_1(self): + """Test that square sizes map to 1:1""" + square_sizes = ["1024x1024", "512x512", "256x256"] + for size in square_sizes: + assert OPENAI_SIZE_TO_STABILITY_ASPECT_RATIO[size] == "1:1" + + +class TestStabilityGenerationModels: + """Test the model endpoint mappings""" + + def test_sd3_models_use_sd3_endpoint(self): + """Test that SD3 models use the SD3 endpoint""" + sd3_models = ["sd3", "sd3-large", "sd3-medium", "sd3.5-large"] + for model in sd3_models: + assert STABILITY_GENERATION_MODELS[model] == "/v2beta/stable-image/generate/sd3" + + def test_ultra_model_uses_ultra_endpoint(self): + """Test that Ultra model uses ultra endpoint""" + assert STABILITY_GENERATION_MODELS["stable-image-ultra"] == "/v2beta/stable-image/generate/ultra" + + def test_core_model_uses_core_endpoint(self): + """Test that Core model uses core endpoint""" + assert STABILITY_GENERATION_MODELS["stable-image-core"] == "/v2beta/stable-image/generate/core"