mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(stability): add Stability AI image generation support (#17894)
Add direct Stability AI REST API support for image generation endpoints. This enables using Stability's SD3, SD3.5, and Stable Image models via LiteLLM's OpenAI-compatible interface. Changes: - Add STABILITY provider to LlmProviders enum - Create StabilityImageGenerationConfig with multipart/form-data support - Add OpenAI size to Stability aspect_ratio mapping - Register provider in ProviderConfigManager - Add 9 Stability models to model_prices_and_context_window.json - Add documentation at docs/providers/stability.md - Add 25 unit tests Supported models: - stability/sd3, sd3-large, sd3-large-turbo, sd3-medium - stability/sd3.5-large, sd3.5-large-turbo, sd3.5-medium - stability/stable-image-ultra, stable-image-core
This commit is contained in:
parent
5262896d62
commit
bd1a075a89
15 changed files with 1178 additions and 12 deletions
181
docs/my-website/docs/providers/stability.md
Normal file
181
docs/my-website/docs/providers/stability.md
Normal file
|
|
@ -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.
|
||||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
0
litellm/llms/stability/__init__.py
Normal file
0
litellm/llms/stability/__init__.py
Normal file
37
litellm/llms/stability/image_generation/__init__.py
Normal file
37
litellm/llms/stability/image_generation/__init__.py
Normal file
|
|
@ -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()
|
||||
274
litellm/llms/stability/image_generation/transformation.py
Normal file
274
litellm/llms/stability/image_generation/transformation.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
212
litellm/types/llms/stability.py
Normal file
212
litellm/types/llms/stability.py
Normal file
|
|
@ -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",
|
||||
}
|
||||
|
|
@ -2996,6 +2996,7 @@ class LlmProviders(str, Enum):
|
|||
HYPERBOLIC = "hyperbolic"
|
||||
RECRAFT = "recraft"
|
||||
FAL_AI = "fal_ai"
|
||||
STABILITY = "stability"
|
||||
HEROKU = "heroku"
|
||||
AIML = "aiml"
|
||||
COMETAPI = "cometapi"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
0
tests/test_litellm/llms/stability/__init__.py
Normal file
0
tests/test_litellm/llms/stability/__init__.py
Normal file
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue