[Feat] Add FAL AI Image Generations on LiteLLM (#16067)

* add fal-ai provider

* fix image_generation_handler

* init FalAIImageGenerationConfig

* init cost_calculator

* init FAL AI

* TestFAL_AI_ImageGeneration

* fix load_custom_provider_entrypoints

* TestFAL_AI_ImageGeneration

* add imagen4 transform FAL AI

* add FAL AI imagen 4 transform

* BaseImageGenTest

* test_fal_ai_image_generation_basic

* add BRIA + Recraft img gen

* add recraft + BRIA

* test_fal_ai_image_generation_basic

* tests for flux PRO v11

* Add FAL AI SD

* test FAL AI SD

* docs FAL AI

* docs fal ai

* Using Model-Specific Parameters

* add fal ai model prices

* add fall_ai JPG logo

* ui fixes FAL AI

* fix linting

* fix linting

* fix bedrock test_get_request_body_stability3

* test_custom_llm_provider_entrypoint
This commit is contained in:
Ishaan Jaff 2025-10-29 13:10:51 -07:00 • committed by GitHub
parent 4939793ade
commit 99feefd614
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 2037 additions and 31 deletions

View file

@ -0,0 +1,310 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Fal AI
Fal AI provides fast, scalable access to state-of-the-art image generation models including FLUX, Stable Diffusion, Imagen, and more.
## Overview
| Property | Details |
|----------|---------|
| Description | Fal AI offers optimized infrastructure for running image generation models at scale with low latency. |
| Provider Route on LiteLLM | `fal_ai/` |
| Provider Doc | [Fal AI Documentation ↗](https://fal.ai/models) |
| Supported Operations | [`/images/generations`](#image-generation) |
## Setup
### API Key
```python showLineNumbers
import os
# Set your Fal AI API key
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
```
Get your API key from [fal.ai](https://fal.ai/).
## Supported Models
| Model Name | Description | Documentation |
|------------|-------------|---------------|
| `fal_ai/fal-ai/flux-pro/v1.1-ultra` | FLUX Pro v1.1 Ultra - High-quality image generation | [Docs ↗](https://fal.ai/models/fal-ai/flux-pro/v1.1-ultra) |
| `fal_ai/fal-ai/imagen4/preview` | Google's Imagen 4 - Highest quality model | [Docs ↗](https://fal.ai/models/fal-ai/imagen4/preview) |
| `fal_ai/fal-ai/recraft/v3/text-to-image` | Recraft v3 - Multiple style options | [Docs ↗](https://fal.ai/models/fal-ai/recraft/v3/text-to-image) |
| `fal_ai/fal-ai/stable-diffusion-v35-medium` | Stable Diffusion v3.5 Medium | [Docs ↗](https://fal.ai/models/fal-ai/stable-diffusion-v35-medium) |
| `fal_ai/bria/text-to-image/3.2` | Bria 3.2 - Commercial-grade generation | [Docs ↗](https://fal.ai/models/bria/text-to-image/3.2) |
## Image Generation
### Usage - LiteLLM Python SDK
<Tabs>
<TabItem value="basic" label="Basic Usage">
```python showLineNumbers title="Basic Image Generation"
import litellm
import os
# Set your API key
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
# Generate an image
response = litellm.image_generation(
model="fal_ai/fal-ai/flux-pro/v1.1-ultra",
prompt="A serene mountain landscape at sunset with vibrant colors"
)
print(response.data[0].url)
```
</TabItem>
<TabItem value="imagen4" label="Imagen 4">
```python showLineNumbers title="Google Imagen 4 Generation"
import litellm
import os
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
# Generate with Imagen 4
response = litellm.image_generation(
model="fal_ai/fal-ai/imagen4/preview",
prompt="A vintage 1960s kitchen with flour package on countertop",
aspect_ratio="16:9",
num_images=1
)
print(response.data[0].url)
```
</TabItem>
<TabItem value="recraft" label="Recraft v3">
```python showLineNumbers title="Recraft v3 with Style"
import litellm
import os
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
# Generate with specific style
response = litellm.image_generation(
model="fal_ai/fal-ai/recraft/v3/text-to-image",
prompt="A red panda eating bamboo",
style="realistic_image",
image_size="landscape_4_3"
)
print(response.data[0].url)
```
</TabItem>
<TabItem value="async" label="Async Usage">
```python showLineNumbers title="Async Image Generation"
import litellm
import asyncio
import os
async def generate_image():
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
response = await litellm.aimage_generation(
model="fal_ai/fal-ai/stable-diffusion-v35-medium",
prompt="A cyberpunk cityscape with neon lights",
guidance_scale=7.5,
num_inference_steps=50
)
print(response.data[0].url)
return response
asyncio.run(generate_image())
```
</TabItem>
<TabItem value="advanced" label="Advanced Parameters">
```python showLineNumbers title="Advanced FLUX Pro Generation"
import litellm
import os
os.environ["FAL_AI_API_KEY"] = "your-fal-api-key"
# Generate with advanced parameters
response = litellm.image_generation(
model="fal_ai/fal-ai/flux-pro/v1.1-ultra",
prompt="A majestic dragon soaring over mountains",
n=2,
size="1792x1024", # Maps to aspect_ratio="16:9"
seed=42,
safety_tolerance="2",
enhance_prompt=True
)
for image in response.data:
print(f"Generated image: {image.url}")
```
</TabItem>
</Tabs>
### Usage - LiteLLM Proxy Server
#### 1. Configure your config.yaml
```yaml showLineNumbers title="Fal AI Image Generation Configuration"
model_list:
- model_name: flux-ultra
litellm_params:
model: fal_ai/fal-ai/flux-pro/v1.1-ultra
api_key: os.environ/FAL_AI_API_KEY
model_info:
mode: image_generation
- model_name: imagen4
litellm_params:
model: fal_ai/fal-ai/imagen4/preview
api_key: os.environ/FAL_AI_API_KEY
model_info:
mode: image_generation
- model_name: stable-diffusion
litellm_params:
model: fal_ai/fal-ai/stable-diffusion-v35-medium
api_key: os.environ/FAL_AI_API_KEY
model_info:
mode: image_generation
general_settings:
master_key: sk-1234
```
#### 2. Start LiteLLM Proxy Server
```bash showLineNumbers title="Start Proxy Server"
litellm --config /path/to/config.yaml
# RUNNING on http://0.0.0.0:4000
```
#### 3. Make requests
<Tabs>
<TabItem value="openai-sdk" label="OpenAI SDK">
```python showLineNumbers title="Generate via Proxy - OpenAI SDK"
from openai import OpenAI
client = OpenAI(
base_url="http://localhost:4000",
api_key="sk-1234"
)
response = client.images.generate(
model="flux-ultra",
prompt="A beautiful sunset over the ocean",
n=1,
size="1024x1024"
)
print(response.data[0].url)
```
</TabItem>
<TabItem value="litellm-sdk" label="LiteLLM SDK">
```python showLineNumbers title="Generate via Proxy - LiteLLM SDK"
import litellm
response = litellm.image_generation(
model="litellm_proxy/imagen4",
prompt="A cozy coffee shop interior",
api_base="http://localhost:4000",
api_key="sk-1234"
)
print(response.data[0].url)
```
</TabItem>
<TabItem value="curl" label="cURL">
```bash showLineNumbers title="Generate via Proxy - cURL"
curl --location 'http://localhost:4000/v1/images/generations' \
--header 'Content-Type: application/json' \
--header 'Authorization: Bearer sk-1234' \
--data '{
"model": "stable-diffusion",
"prompt": "A serene Japanese garden with cherry blossoms",
"n": 1,
"size": "1024x1024"
}'
```
</TabItem>
</Tabs>
## Using Model-Specific Parameters
LiteLLM forwards any additional parameters directly to the Fal AI API. You can pass model-specific parameters in your request and they will be sent to Fal AI.
```python showLineNumbers title="Pass Model-Specific Parameters"
import litellm
# Any parameters beyond the standard ones are forwarded to Fal AI
response = litellm.image_generation(
model="fal_ai/fal-ai/flux-pro/v1.1-ultra",
prompt="A beautiful sunset",
# Model-specific Fal AI parameters
aspect_ratio="16:9",
safety_tolerance="2",
enhance_prompt=True,
seed=42
)
```
For the complete list of parameters supported by each model, see:
- [FLUX Pro v1.1-ultra Parameters ↗](https://fal.ai/models/fal-ai/flux-pro/v1.1-ultra/api)
- [Imagen 4 Parameters ↗](https://fal.ai/models/fal-ai/imagen4/preview/api)
- [Recraft v3 Parameters ↗](https://fal.ai/models/fal-ai/recraft/v3/text-to-image/api)
- [Stable Diffusion v3.5 Parameters ↗](https://fal.ai/models/fal-ai/stable-diffusion-v35-medium/api)
- [Bria 3.2 Parameters ↗](https://fal.ai/models/bria/text-to-image/3.2/api)
## Supported Parameters
Standard OpenAI-compatible parameters that work across all models:
| Parameter | Type | Description | Default |
|-----------|------|-------------|---------|
| `prompt` | string | Text description of desired image | Required |
| `model` | string | Fal AI model to use | Required |
| `n` | integer | Number of images to generate (1-4) | `1` |
| `size` | string | Image dimensions (maps to model-specific format) | Model default |
| `api_key` | string | Your Fal AI API key | Environment variable |
## Getting Started
1. Sign up at [fal.ai](https://fal.ai/)
2. Get your API key from your account settings
3. Set `FAL_AI_API_KEY` environment variable
4. Choose a model from the [Fal AI model gallery](https://fal.ai/models)
5. Start generating images with LiteLLM
## Additional Resources
- [Fal AI Documentation](https://fal.ai/docs)
- [Model Gallery](https://fal.ai/models)
- [API Reference](https://fal.ai/docs/api-reference)
- [Pricing](https://fal.ai/pricing)

View file

@ -1,7 +1,7 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# 🆕 Github
# Github
https://github.com/marketplace/models
:::tip

View file

@ -537,6 +537,7 @@ const sidebars = {
"providers/groq",
"providers/deepseek",
"providers/elevenlabs",
"providers/fal_ai",
"providers/fireworks_ai",
"providers/clarifai",
"providers/compactifai",

View file

@ -342,6 +342,7 @@ def image_generation( # noqa: PLR0915
litellm.LlmProviders.RECRAFT,
litellm.LlmProviders.AIML,
litellm.LlmProviders.GEMINI,
litellm.LlmProviders.FAL_AI,
):
if image_generation_config is None:
raise ValueError(

View file

@ -644,8 +644,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
ctype = content.get("type")
return not (isinstance(ctype, str) and ctype != "text")
def _process_messages(self, messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER) -> List[dict[str, Any]]:
filtered_messages: List[dict[str, Any]] = []
def _process_messages(self, messages: list[Any], max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER) -> List[Dict[str, Any]]:
filtered_messages: List[Dict[str, Any]] = []
for msg in messages:
if not isinstance(msg, dict):
continue

View file

@ -14,7 +14,6 @@ import litellm.types.llms
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_async_httpx_client
from litellm.litellm_core_utils.token_counter import get_image_type
from litellm.types.files import get_file_extension_from_mime_type
from litellm.types.llms.anthropic import *
from litellm.types.llms.bedrock import CachePointBlock

View file

@ -0,0 +1,24 @@
from .cost_calculator import cost_calculator
from .image_generation import (
FalAIBaseConfig,
FalAIBriaConfig,
FalAIFluxProV11UltraConfig,
FalAIImageGenerationConfig,
FalAIImagen4Config,
FalAIRecraftV3Config,
FalAIStableDiffusionConfig,
get_fal_ai_image_generation_config,
)
__all__ = [
"cost_calculator",
"FalAIBaseConfig",
"FalAIImageGenerationConfig",
"FalAIImagen4Config",
"FalAIRecraftV3Config",
"FalAIBriaConfig",
"FalAIFluxProV11UltraConfig",
"FalAIStableDiffusionConfig",
"get_fal_ai_image_generation_config",
]

View file

@ -0,0 +1,26 @@
from typing import Any
import litellm
from litellm.types.utils import ImageResponse
def cost_calculator(
model: str,
image_response: Any,
) -> float:
"""
fal.ai image generation cost calculator
"""
_model_info = litellm.get_model_info(
model=model,
custom_llm_provider=litellm.LlmProviders.FAL_AI.value,
)
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
num_images: int = 0
if isinstance(image_response, ImageResponse):
if image_response.data:
num_images = len(image_response.data)
return output_cost_per_image * num_images
else:
raise ValueError(f"image_response must be of type ImageResponse got type={type(image_response)}")

View file

@ -0,0 +1,49 @@
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from .bria_transformation import FalAIBriaConfig
from .flux_pro_v11_ultra_transformation import FalAIFluxProV11UltraConfig
from .imagen4_transformation import FalAIImagen4Config
from .recraft_v3_transformation import FalAIRecraftV3Config
from .stable_diffusion_transformation import FalAIStableDiffusionConfig
from .transformation import FalAIBaseConfig, FalAIImageGenerationConfig
__all__ = [
"FalAIBaseConfig",
"FalAIImageGenerationConfig",
"FalAIImagen4Config",
"FalAIRecraftV3Config",
"FalAIBriaConfig",
"FalAIFluxProV11UltraConfig",
"FalAIStableDiffusionConfig",
]
def get_fal_ai_image_generation_config(model: str) -> BaseImageGenerationConfig:
"""
Get the appropriate Fal AI image generation configuration based on the model.
Args:
model: The Fal AI model name (e.g., "fal-ai/imagen4/preview", "fal-ai/recraft/v3/text-to-image")
Returns:
The appropriate configuration class for the specified model
"""
model_lower = model.lower()
# Map model names to their corresponding configuration classes
if "imagen4" in model_lower or "imagen-4" in model_lower:
return FalAIImagen4Config()
elif "recraft" in model_lower:
return FalAIRecraftV3Config()
elif "bria" in model_lower:
return FalAIBriaConfig()
elif "flux-pro" in model_lower and "ultra" in model_lower:
return FalAIFluxProV11UltraConfig()
elif "stable-diffusion" in model_lower:
return FalAIStableDiffusionConfig()
# Default to generic Fal AI configuration
return FalAIImageGenerationConfig()

View file

@ -0,0 +1,231 @@
from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageObject, ImageResponse
from .transformation import FalAIBaseConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class FalAIBriaConfig(FalAIBaseConfig):
"""
Configuration for Bria Text-to-Image 3.2 model.
Bria 3.2 is a commercial-grade text-to-image model with prompt enhancement
and multiple aspect ratio options.
Model endpoint: bria/text-to-image/3.2
Documentation: https://fal.ai/models/bria/text-to-image/3.2
"""
IMAGE_GENERATION_ENDPOINT: str = "bria/text-to-image/3.2"
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for Bria 3.2.
"""
return [
"n",
"response_format",
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Bria 3.2 parameters.
Mappings:
- size -> aspect_ratio (1:1, 2:3, 3:2, 3:4, 4:3, 4:5, 5:4, 9:16, 16:9)
- response_format -> ignored (Bria returns URLs)
- n -> ignored (Bria doesn't support multiple images in one call)
"""
supported_params = self.get_supported_openai_params(model)
# Map OpenAI params to Bria params
param_mapping = {
"size": "aspect_ratio",
}
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
# Use mapped parameter name if exists
mapped_key = param_mapping.get(k, k)
mapped_value = non_default_params[k]
# Transform specific parameters
if k == "response_format":
# Bria always returns URLs, so we can ignore this
continue
elif k == "n":
# Bria doesn't support multiple images, ignore
continue
elif k == "size":
# Map OpenAI size format to Bria aspect ratio
mapped_value = self._map_aspect_ratio(mapped_value)
optional_params[mapped_key] = mapped_value
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def _map_aspect_ratio(self, size: str) -> str:
"""
Map OpenAI size format to Bria aspect ratio format.
OpenAI format: "1024x1024", "1792x1024", etc.
Bria format: "1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"
"""
# Map common OpenAI sizes to Bria aspect ratios
size_to_aspect_ratio = {
"1024x1024": "1:1",
"512x512": "1:1",
"1792x1024": "16:9",
"1024x1792": "9:16",
"1024x768": "4:3",
"768x1024": "3:4",
"1280x960": "4:3",
"960x1280": "3:4",
}
if size in size_to_aspect_ratio:
return size_to_aspect_ratio[size]
# Parse custom size format "WIDTHxHEIGHT" and calculate aspect ratio
if "x" in size:
try:
width_str, height_str = size.split("x")
width = int(width_str)
height = int(height_str)
# Calculate aspect ratio and find closest match
ratio = width / height
# Map to closest supported aspect ratio
if 0.95 <= ratio <= 1.05: # Close to 1:1
return "1:1"
elif ratio >= 1.7: # Close to 16:9
return "16:9"
elif ratio <= 0.6: # Close to 9:16
return "9:16"
elif 1.3 <= ratio <= 1.4: # Close to 4:3
return "4:3"
elif 0.7 <= ratio <= 0.8: # Close to 3:4
return "3:4"
elif 1.45 <= ratio <= 1.55: # Close to 3:2
return "3:2"
elif 0.65 <= ratio <= 0.7: # Close to 2:3
return "2:3"
elif 1.2 <= ratio <= 1.3: # Close to 5:4
return "5:4"
elif 0.75 <= ratio <= 0.85: # Close to 4:5
return "4:5"
except (ValueError, AttributeError, ZeroDivisionError):
pass
# Default to 1:1
return "1:1"
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to Bria 3.2 request body.
Required parameters:
- prompt: Prompt for image generation
Optional parameters:
- aspect_ratio: "1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9" (default: "1:1")
- prompt_enhancer: Improve the prompt (default: true)
- sync_mode: Return image directly in response (default: false)
- truncate_prompt: Truncate the prompt (default: true)
- guidance_scale: Guidance scale 1-10 (default: 5)
- num_inference_steps: Inference steps 20-50 (default: 30)
- seed: Random seed for reproducibility (default: 5555)
- negative_prompt: Negative prompt string
"""
bria_request_body = {
"prompt": prompt,
**optional_params,
}
return bria_request_body
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 the Bria 3.2 response to litellm ImageResponse format.
Expected response format:
{
"image": {
"url": "https://...",
"content_type": "image/png",
"file_name": "...",
"file_size": 123456,
"width": 1024,
"height": 1024
}
}
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
# Handle Bria response format - uses "image" (singular) not "images"
image_data = response_data.get("image")
if image_data and isinstance(image_data, dict):
model_response.data.append(
ImageObject(
url=image_data.get("url", None),
b64_json=None, # Bria returns URLs only
)
)
return model_response

View file

@ -0,0 +1,263 @@
from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageObject, ImageResponse
from .transformation import FalAIBaseConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class FalAIFluxProV11UltraConfig(FalAIBaseConfig):
"""
Configuration for Fal AI Flux Pro v1.1-ultra model.
FLUX Pro v1.1-ultra is a high-quality text-to-image model with enhanced detail
and support for image prompts.
Model endpoint: fal-ai/flux-pro/v1.1-ultra
Documentation: https://fal.ai/models/fal-ai/flux-pro/v1.1-ultra
"""
IMAGE_GENERATION_ENDPOINT: str = "fal-ai/flux-pro/v1.1-ultra"
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for Flux Pro v1.1-ultra.
"""
return [
"n",
"response_format",
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Flux Pro v1.1-ultra parameters.
Mappings:
- n -> num_images (1-4, default 1)
- response_format -> output_format (jpeg or png)
- size -> aspect_ratio (21:9, 16:9, 4:3, 3:2, 1:1, 2:3, 3:4, 9:16, 9:21)
"""
supported_params = self.get_supported_openai_params(model)
# Map OpenAI params to Flux Pro v1.1-ultra params
param_mapping = {
"n": "num_images",
"response_format": "output_format",
"size": "aspect_ratio",
}
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
# Use mapped parameter name if exists
mapped_key = param_mapping.get(k, k)
mapped_value = non_default_params[k]
# Transform specific parameters
if k == "response_format":
# Map OpenAI response formats to image formats
if mapped_value in ["b64_json", "url"]:
mapped_value = "jpeg"
elif k == "size":
# Map OpenAI size format to Flux aspect ratio
mapped_value = self._map_aspect_ratio(mapped_value)
optional_params[mapped_key] = mapped_value
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def _map_aspect_ratio(self, size: str) -> str:
"""
Map OpenAI size format to Flux Pro aspect ratio format.
OpenAI format: "1024x1024", "1792x1024", etc.
Flux format: "21:9", "16:9", "4:3", "3:2", "1:1", "2:3", "3:4", "9:16", "9:21"
Default: "16:9"
"""
# Map common OpenAI sizes to Flux aspect ratios
size_to_aspect_ratio = {
"1024x1024": "1:1",
"512x512": "1:1",
"1792x1024": "16:9",
"1024x1792": "9:16",
"1024x768": "4:3",
"768x1024": "3:4",
"1536x1024": "3:2",
"1024x1536": "2:3",
"2048x876": "21:9",
"876x2048": "9:21",
}
if size in size_to_aspect_ratio:
return size_to_aspect_ratio[size]
# Parse custom size format "WIDTHxHEIGHT" and calculate aspect ratio
if "x" in size:
try:
width_str, height_str = size.split("x")
width = int(width_str)
height = int(height_str)
# Calculate aspect ratio and find closest match
ratio = width / height
# Map to closest supported aspect ratio
if 0.95 <= ratio <= 1.05: # Close to 1:1
return "1:1"
elif ratio >= 2.3: # Close to 21:9
return "21:9"
elif 1.7 <= ratio < 2.3: # Close to 16:9
return "16:9"
elif 1.3 <= ratio < 1.7: # Close to 4:3
return "4:3"
elif 1.4 <= ratio < 1.6: # Close to 3:2
return "3:2"
elif 0.6 <= ratio < 0.7: # Close to 3:4
return "3:4"
elif 0.65 <= ratio < 0.75: # Close to 2:3
return "2:3"
elif 0.5 <= ratio < 0.6: # Close to 9:16
return "9:16"
elif ratio < 0.5: # Close to 9:21
return "9:21"
except (ValueError, AttributeError, ZeroDivisionError):
pass
# Default to 16:9
return "16:9"
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to Flux Pro v1.1-ultra request body.
Required parameters:
- prompt: The prompt to generate an image from
Optional parameters:
- num_images: Number of images (1-4, default: 1)
- aspect_ratio: Aspect ratio (default: "16:9")
- raw: Generate less processed images (default: false)
- output_format: "jpeg" or "png" (default: "jpeg")
- image_url: Image URL for image-to-image generation
- sync_mode: Return data URI (default: false)
- safety_tolerance: Safety level "1"-"6" (default: "2")
- enable_safety_checker: Enable safety checker (default: true)
- seed: Random seed for reproducibility
- image_prompt_strength: Strength of image prompt 0-1 (default: 0.1)
- enhance_prompt: Enhance prompt for better results (default: false)
"""
flux_pro_request_body = {
"prompt": prompt,
**optional_params,
}
return flux_pro_request_body
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 the Flux Pro v1.1-ultra response to litellm ImageResponse format.
Expected response format:
{
"images": [
{
"url": "https://...",
"width": 1024,
"height": 768,
"content_type": "image/jpeg"
}
],
"timings": {"inference": 2.5, ...},
"seed": 42,
"has_nsfw_concepts": [false],
"prompt": "original prompt"
}
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
# Handle Flux Pro v1.1-ultra response format
images = response_data.get("images", [])
if isinstance(images, list):
for image_data in images:
if isinstance(image_data, dict):
model_response.data.append(
ImageObject(
url=image_data.get("url", None),
b64_json=None, # Flux Pro returns URLs only
)
)
elif isinstance(image_data, str):
# If images is just a list of URLs
model_response.data.append(
ImageObject(
url=image_data,
b64_json=None,
)
)
# Add additional metadata from Flux Pro response
if hasattr(model_response, "_hidden_params"):
if "seed" in response_data:
model_response._hidden_params["seed"] = response_data["seed"]
if "timings" in response_data:
model_response._hidden_params["timings"] = response_data["timings"]
if "has_nsfw_concepts" in response_data:
model_response._hidden_params["has_nsfw_concepts"] = response_data[
"has_nsfw_concepts"
]
return model_response

View file

@ -0,0 +1,242 @@
from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageObject, ImageResponse
from .transformation import FalAIBaseConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class FalAIImagen4Config(FalAIBaseConfig):
"""
Configuration for Fal AI Imagen4 model.
Google's highest quality image generation model available through Fal AI.
Model variants:
- fal-ai/imagen4/preview (Standard): $0.05 per image
- fal-ai/imagen4/preview/fast (Fast): $0.04 per image
- fal-ai/imagen4/preview/ultra (Ultra): $0.06 per image
Documentation: https://fal.ai/models/fal-ai/imagen4/preview
"""
IMAGE_GENERATION_ENDPOINT: str = "fal-ai/imagen4/preview"
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for Imagen4.
"""
return [
"n",
"response_format",
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Imagen4 parameters.
Mappings:
- n -> num_images (1-4, default 1)
- size -> aspect_ratio (1:1, 16:9, 9:16, 3:4, 4:3)
- response_format -> ignored (Imagen4 returns URLs)
"""
supported_params = self.get_supported_openai_params(model)
# Map OpenAI params to Imagen4 params
param_mapping = {
"n": "num_images",
"size": "aspect_ratio",
}
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
# Use mapped parameter name if exists
mapped_key = param_mapping.get(k, k)
mapped_value = non_default_params[k]
# Transform specific parameters
if k == "response_format":
# Imagen4 always returns URLs, so we can ignore this
continue
elif k == "size":
# Map OpenAI size format to Imagen4 aspect ratio
mapped_value = self._map_aspect_ratio(mapped_value)
optional_params[mapped_key] = mapped_value
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def _map_aspect_ratio(self, size: str) -> str:
"""
Map OpenAI size format to Imagen4 aspect ratio format.
OpenAI format: "1024x1024", "1792x1024", etc.
Imagen4 format: "1:1", "16:9", "9:16", "3:4", "4:3"
Available aspect ratios:
- 1:1 (default)
- 16:9
- 9:16
- 3:4
- 4:3
"""
# Map common OpenAI sizes to Imagen4 aspect ratios
size_to_aspect_ratio = {
"1024x1024": "1:1",
"512x512": "1:1",
"1792x1024": "16:9",
"1024x1792": "9:16",
"1024x768": "4:3",
"768x1024": "3:4",
}
if size in size_to_aspect_ratio:
return size_to_aspect_ratio[size]
# Parse custom size format "WIDTHxHEIGHT" and calculate aspect ratio
if "x" in size:
try:
width_str, height_str = size.split("x")
width = int(width_str)
height = int(height_str)
# Calculate aspect ratio and find closest match
ratio = width / height
# Map to closest supported aspect ratio
if 0.95 <= ratio <= 1.05: # Close to 1:1
return "1:1"
elif ratio >= 1.7: # Close to 16:9
return "16:9"
elif ratio <= 0.6: # Close to 9:16
return "9:16"
elif ratio >= 1.2: # Close to 4:3
return "4:3"
elif ratio <= 0.8: # Close to 3:4
return "3:4"
except (ValueError, AttributeError, ZeroDivisionError):
pass
# Default to 1:1
return "1:1"
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to Imagen4 request body.
Required parameters:
- prompt: The text prompt describing what you want to see
Optional parameters:
- aspect_ratio: "1:1", "16:9", "9:16", "3:4", "4:3" (default: "1:1")
- num_images: Number of images (1-4, default: 1)
- resolution: "1K" or "2K" (default: "1K")
- seed: Random seed for reproducibility
- negative_prompt: Description of what to discourage (default: "")
"""
imagen4_request_body = {
"prompt": prompt,
**optional_params,
}
return imagen4_request_body
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 the Imagen4 response to litellm ImageResponse format.
Expected response format:
{
"images": [
{
"url": "https://...",
"content_type": "image/png",
"file_name": "z9RV14K95DvU.png",
"file_size": 4404019
}
],
"seed": 42
}
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
# Handle Imagen4 response format
images = response_data.get("images", [])
if isinstance(images, list):
for image_data in images:
if isinstance(image_data, dict):
model_response.data.append(
ImageObject(
url=image_data.get("url", None),
b64_json=None, # Imagen4 returns URLs only
)
)
elif isinstance(image_data, str):
# If images is just a list of URLs
model_response.data.append(
ImageObject(
url=image_data,
b64_json=None,
)
)
# Add seed metadata from Imagen4 response
if hasattr(model_response, "_hidden_params"):
if "seed" in response_data:
model_response._hidden_params["seed"] = response_data["seed"]
return model_response

View file

@ -0,0 +1,226 @@
from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageObject, ImageResponse
from .transformation import FalAIBaseConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class FalAIRecraftV3Config(FalAIBaseConfig):
"""
Configuration for Fal AI Recraft v3 Text-to-Image model.
Recraft v3 is a text-to-image model with multiple style options including
realistic images, digital illustrations, and vector illustrations.
Model endpoint: fal-ai/recraft/v3/text-to-image
Documentation: https://fal.ai/models/fal-ai/recraft/v3/text-to-image
"""
IMAGE_GENERATION_ENDPOINT: str = "fal-ai/recraft/v3/text-to-image"
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for Recraft v3.
"""
return [
"n",
"response_format",
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Recraft v3 parameters.
Mappings:
- size -> image_size (can be preset or custom width/height)
- response_format -> ignored (Recraft returns URLs)
- n -> ignored (Recraft doesn't support multiple images)
"""
supported_params = self.get_supported_openai_params(model)
# Map OpenAI params to Recraft v3 params
param_mapping = {
"size": "image_size",
}
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
# Use mapped parameter name if exists
mapped_key = param_mapping.get(k, k)
mapped_value = non_default_params[k]
# Transform specific parameters
if k == "response_format":
# Recraft always returns URLs, so we can ignore this
continue
elif k == "n":
# Recraft doesn't support multiple images, ignore
continue
elif k == "size":
# Map OpenAI size format to Recraft image_size
mapped_value = self._map_image_size(mapped_value)
optional_params[mapped_key] = mapped_value
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def _map_image_size(self, size: str) -> Any:
"""
Map OpenAI size format to Recraft v3 image_size format.
OpenAI format: "1024x1024", "1792x1024", etc.
Recraft format: Can be preset strings or {"width": int, "height": int}
Available presets:
- square_hd (default)
- square
- portrait_4_3
- portrait_16_9
- landscape_4_3
- landscape_16_9
"""
# Map common OpenAI sizes to Recraft presets
size_mapping = {
"1024x1024": "square_hd",
"512x512": "square",
"768x1024": "portrait_4_3",
"576x1024": "portrait_16_9",
"1024x768": "landscape_4_3",
"1024x576": "landscape_16_9",
}
if size in size_mapping:
return size_mapping[size]
# Parse custom size format "WIDTHxHEIGHT"
if "x" in size:
try:
width, height = size.split("x")
return {
"width": int(width),
"height": int(height),
}
except (ValueError, AttributeError):
pass
# Default to square_hd
return "square_hd"
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to Recraft v3 request body.
Required parameters:
- prompt: Text prompt (max 1000 characters)
Optional parameters:
- image_size: Preset or {"width": int, "height": int} (default: "square_hd")
- style: Style preset (default: "realistic_image")
Options: "any", "realistic_image", "digital_illustration", "vector_illustration", etc.
- colors: Array of RGB color objects [{"r": 0-255, "g": 0-255, "b": 0-255}]
- enable_safety_checker: Enable safety checker (default: false)
- style_id: UUID for custom style reference
Note: Vector illustrations cost 2X as much.
"""
recraft_request_body = {
"prompt": prompt,
**optional_params,
}
return recraft_request_body
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 the Recraft v3 response to litellm ImageResponse format.
Expected response format:
{
"images": [
{
"url": "https://...",
"content_type": "image/webp",
"file_name": "...",
"file_size": 123456
}
]
}
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
# Handle Recraft v3 response format
images = response_data.get("images", [])
if isinstance(images, list):
for image_data in images:
if isinstance(image_data, dict):
model_response.data.append(
ImageObject(
url=image_data.get("url", None),
b64_json=None, # Recraft returns URLs only
)
)
elif isinstance(image_data, str):
# If images is just a list of URLs
model_response.data.append(
ImageObject(
url=image_data,
b64_json=None,
)
)
return model_response

View file

@ -0,0 +1,281 @@
from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageObject, ImageResponse
from .transformation import FalAIBaseConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class FalAIStableDiffusionConfig(FalAIBaseConfig):
"""
Configuration for Fal AI Stable Diffusion models.
Supports Stable Diffusion v3.5 variants and other Stable Diffusion models on Fal AI.
Example models:
- fal-ai/stable-diffusion-v35-medium
- fal-ai/stable-diffusion-v35-large
Documentation: https://fal.ai/models/fal-ai/stable-diffusion-v35-medium
"""
IMAGE_GENERATION_ENDPOINT: str = "" # Will be set from model name
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 request.
For Stable Diffusion models, extract the endpoint from the model name.
"""
from litellm.secret_managers.main import get_secret_str
complete_url: str = (
api_base
or get_secret_str("FAL_AI_API_BASE")
or self.DEFAULT_BASE_URL
)
complete_url = complete_url.rstrip("/")
# Extract endpoint from model name
# e.g., "fal-ai/stable-diffusion-v35-medium" or "stable-diffusion-v35-medium"
endpoint = model
if "/" in model and not model.startswith("fal-ai/"):
# If model is like "custom/stable-diffusion-v35-medium", use full path
endpoint = model
elif not model.startswith("fal-ai/"):
# If model is just "stable-diffusion-v35-medium", prepend fal-ai
endpoint = f"fal-ai/{model}"
complete_url = f"{complete_url}/{endpoint}"
return complete_url
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for Stable Diffusion models.
"""
return [
"n",
"response_format",
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Stable Diffusion parameters.
Mappings:
- n -> num_images (1-4, default 1)
- response_format -> output_format (jpeg or png)
- size -> image_size (can be preset or custom width/height)
"""
supported_params = self.get_supported_openai_params(model)
# Map OpenAI params to Stable Diffusion params
param_mapping = {
"n": "num_images",
"response_format": "output_format",
"size": "image_size",
}
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
# Use mapped parameter name if exists
mapped_key = param_mapping.get(k, k)
mapped_value = non_default_params[k]
# Transform specific parameters
if k == "response_format":
# Map OpenAI response formats to image formats
if mapped_value in ["b64_json", "url"]:
mapped_value = "jpeg"
elif k == "size":
# Map OpenAI size format to Stable Diffusion image_size
mapped_value = self._map_image_size(mapped_value)
optional_params[mapped_key] = mapped_value
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def _map_image_size(self, size: str) -> Any:
"""
Map OpenAI size format to Stable Diffusion image_size format.
OpenAI format: "1024x1024", "1792x1024", etc.
Stable Diffusion format: Can be preset strings or {"width": int, "height": int}
Available presets:
- square_hd
- square
- portrait_4_3
- portrait_16_9
- landscape_4_3 (default)
- landscape_16_9
"""
# Map common OpenAI sizes to Stable Diffusion presets
size_mapping = {
"1024x1024": "square_hd",
"512x512": "square",
"768x1024": "portrait_4_3",
"576x1024": "portrait_16_9",
"1024x768": "landscape_4_3",
"1024x576": "landscape_16_9",
}
if size in size_mapping:
return size_mapping[size]
# Parse custom size format "WIDTHxHEIGHT"
if "x" in size:
try:
width, height = size.split("x")
return {
"width": int(width),
"height": int(height),
}
except (ValueError, AttributeError):
pass
# Default to landscape_4_3
return "landscape_4_3"
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to Stable Diffusion request body.
Required parameters:
- prompt: The prompt to generate an image from
Optional parameters:
- num_images: Number of images (1-4, default: 1)
- image_size: Size preset or {"width": int, "height": int} (default: landscape_4_3)
- output_format: "jpeg" or "png" (default: jpeg)
- sync_mode: Wait for image upload before returning (default: false)
- guidance_scale: CFG scale 0-20 (default: 4.5)
- num_inference_steps: Inference steps 1-50 (default: 40)
- seed: Random seed for reproducibility
- negative_prompt: Negative prompt string (default: "")
- enable_safety_checker: Enable safety checker (default: true)
"""
stable_diffusion_request_body = {
"prompt": prompt,
**optional_params,
}
return stable_diffusion_request_body
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 the Stable Diffusion response to litellm ImageResponse format.
Expected response format:
{
"images": [
{
"url": "https://...",
"width": 1024,
"height": 768,
"content_type": "image/jpeg"
}
],
"timings": {"inference": 2.5, ...},
"seed": 42,
"has_nsfw_concepts": [false],
"prompt": "original prompt"
}
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
# Handle Stable Diffusion response format
images = response_data.get("images", [])
if isinstance(images, list):
for image_data in images:
if isinstance(image_data, dict):
model_response.data.append(
ImageObject(
url=image_data.get("url", None),
b64_json=None, # Stable Diffusion returns URLs only
)
)
elif isinstance(image_data, str):
# If images is just a list of URLs
model_response.data.append(
ImageObject(
url=image_data,
b64_json=None,
)
)
# Add additional metadata from Stable Diffusion response
if hasattr(model_response, "_hidden_params"):
if "seed" in response_data:
model_response._hidden_params["seed"] = response_data["seed"]
if "timings" in response_data:
model_response._hidden_params["timings"] = response_data["timings"]
if "has_nsfw_concepts" in response_data:
model_response._hidden_params["has_nsfw_concepts"] = response_data[
"has_nsfw_concepts"
]
return model_response

View file

@ -0,0 +1,176 @@
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.utils import ImageObject, ImageResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class FalAIBaseConfig(BaseImageGenerationConfig):
"""
Base configuration for Fal AI image generation models.
Handles common functionality like URL construction and authentication.
"""
DEFAULT_BASE_URL: str = "https://fal.run"
IMAGE_GENERATION_ENDPOINT: str = ""
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 request
Some providers need `model` in `api_base`
"""
complete_url: str = (
api_base
or get_secret_str("FAL_AI_API_BASE")
or self.DEFAULT_BASE_URL
)
complete_url = complete_url.rstrip("/")
if self.IMAGE_GENERATION_ENDPOINT:
complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}"
return complete_url
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:
final_api_key: Optional[str] = (
api_key or
get_secret_str("FAL_AI_API_KEY")
)
if not final_api_key:
raise ValueError("FAL_AI_API_KEY is not set")
headers["Authorization"] = f"Key {final_api_key}"
return headers
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 the image generation response to the litellm image response
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
if not model_response.data:
model_response.data = []
# Handle fal.ai response format
images = response_data.get("images", [])
if isinstance(images, list):
for image_data in images:
if isinstance(image_data, dict):
model_response.data.append(ImageObject(
url=image_data.get("url", None),
b64_json=image_data.get("b64_json", None),
))
elif isinstance(image_data, str):
# If images is just a list of URLs
model_response.data.append(ImageObject(
url=image_data,
b64_json=None,
))
return model_response
class FalAIImageGenerationConfig(FalAIBaseConfig):
"""
Default Fal AI image generation configuration for generic models.
"""
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for fal.ai image generation
"""
return [
"n",
"response_format",
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_params = self.get_supported_openai_params(model)
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
optional_params[k] = non_default_params[k]
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to the fal.ai image generation request body
"""
fal_ai_image_generation_request_body = {
"prompt": prompt,
**optional_params,
}
return fal_ai_image_generation_request_body

View file

@ -8259,6 +8259,46 @@
"supports_function_calling": true,
"supports_tool_choice": false
},
"fal_ai/bria/text-to-image/3.2": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/flux-pro/v1.1-ultra": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/imagen4/preview": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/recraft/v3/text-to-image": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/stable-diffusion-v35-medium": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"featherless_ai/featherless-ai/Qwerky-72B": {
"litellm_provider": "featherless_ai",
"max_input_tokens": 32768,

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.1 KiB

View file

@ -2554,6 +2554,7 @@ class LlmProviders(str, Enum):
PG_VECTOR = "pg_vector"
HYPERBOLIC = "hyperbolic"
RECRAFT = "recraft"
FAL_AI = "fal_ai"
HEROKU = "heroku"
AIML = "aiml"
COMETAPI = "cometapi"

View file

@ -9,9 +9,9 @@
import ast
import asyncio
import contextvars
import base64
import binascii
import contextvars
import copy
import datetime
import hashlib
@ -387,7 +387,15 @@ def print_verbose(
####### CLIENT ###################
# make it easy to log if completion/embedding runs succeeded or failed + see what happened | Non-Blocking
def load_custom_provider_entrypoints():
found_entry_points = tuple(entry_points().select(group="litellm")) # type: ignore
# Handle both Python 3.9 (returns dict) and Python 3.10+ (returns object with select method)
eps = entry_points()
if hasattr(eps, "select"):
# Python 3.10+
found_entry_points = tuple(eps.select(group="litellm")) # type: ignore
else:
# Python 3.9 and earlier - entry_points() returns a dict
found_entry_points = eps.get("litellm", ()) # type: ignore
for entry_point in found_entry_points:
# types are ignored because of circular dependency issues importing CustomLLM and CustomLLMItem
HandlerClass = entry_point.load()
@ -7658,6 +7666,12 @@ class ProviderConfigManager:
)
return LiteLLMProxyImageGenerationConfig()
elif LlmProviders.FAL_AI == provider:
from litellm.llms.fal_ai.image_generation import (
get_fal_ai_image_generation_config,
)
return get_fal_ai_image_generation_config(model)
return None
@staticmethod

View file

@ -8259,6 +8259,46 @@
"supports_function_calling": true,
"supports_tool_choice": false
},
"fal_ai/bria/text-to-image/3.2": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/flux-pro/v1.1-ultra": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/imagen4/preview": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/recraft/v3/text-to-image": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"fal_ai/fal-ai/stable-diffusion-v35-medium": {
"litellm_provider": "fal_ai",
"mode": "image_generation",
"output_cost_per_image": 0.0398,
"supported_endpoints": [
"/v1/images/generations"
]
},
"featherless_ai/featherless-ai/Qwerky-72B": {
"litellm_provider": "featherless_ai",
"max_input_tokens": 32768,

View file

@ -47,6 +47,7 @@ class BaseImageGenTest(ABC):
async def test_basic_image_generation(self):
"""Test basic image generation"""
try:
litellm._turn_on_debug()
custom_logger = TestCustomLogger()
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [custom_logger]
@ -55,7 +56,7 @@ class BaseImageGenTest(ABC):
response = await litellm.aimage_generation(
**base_image_generation_call_args, prompt="A image of a otter"
)
print(response)
print("FAL AI RESPONSE: ", response)
await asyncio.sleep(1)

View file

@ -168,7 +168,7 @@ def test_get_request_body_stability3():
model = "stability.sd3-large"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["prompt"] == prompt
@ -181,7 +181,7 @@ def test_get_request_body_stability():
model = "stability.stable-diffusion-xl-v1"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["text_prompts"][0]["text"] == prompt
@ -239,7 +239,7 @@ def test_get_request_body_nova_canvas_default():
model = "amazon.nova-canvas-v1"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["taskType"] == "TEXT_IMAGE"
@ -254,7 +254,7 @@ def test_get_request_body_nova_canvas_text_image():
model = "amazon.nova-canvas-v1"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["taskType"] == "TEXT_IMAGE"
@ -273,7 +273,7 @@ def test_get_request_body_nova_canvas_color_guided_generation():
model = "amazon.nova-canvas-v1"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["taskType"] == "COLOR_GUIDED_GENERATION"
@ -434,7 +434,7 @@ def test_get_request_body_nova_canvas_inference_profile_arn():
nova_model = "us.amazon.nova-canvas-v1:0"
result = handler._get_request_body(
model=nova_model, prompt=prompt, optional_params=optional_params
model=nova_model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["taskType"] == "TEXT_IMAGE"
@ -450,7 +450,7 @@ def test_get_request_body_nova_canvas_with_model_id_param():
model = "amazon.nova-canvas-v1"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
# After fix, model_id should not appear in the result
@ -487,7 +487,7 @@ def test_get_request_body_cross_region_inference_profile():
# This should work after the fix - cross-region format should be detected as 'nova'
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["taskType"] == "TEXT_IMAGE"
@ -502,7 +502,7 @@ def test_backward_compatibility_regular_nova_model():
model = "amazon.nova-canvas-v1"
result = handler._get_request_body(
model=model, prompt=prompt, optional_params=optional_params
model=model, bedrock_provider=None, prompt=prompt, optional_params=optional_params
)
assert result["taskType"] == "TEXT_IMAGE"

View file

@ -0,0 +1,68 @@
import asyncio
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import aimage_generation
@pytest.mark.parametrize(
"model",
[
"fal_ai/fal-ai/flux-pro/v1.1-ultra",
"fal_ai/fal-ai/recraft/v3/text-to-image",
"fal_ai/bria/text-to-image/3.2",
"fal_ai/fal-ai/stable-diffusion-v35-medium"
],
)
@pytest.mark.asyncio
async def test_fal_ai_image_generation_basic(model):
"""
Test basic image generation for various Fal AI models.
Tests that each model can:
- Accept a basic text prompt
- Return a valid response with image data
- Handle the response properly through litellm
"""
try:
litellm.set_verbose = True
response = await aimage_generation(
model=model,
prompt="A cute baby sea otter",
)
print(f"\nResponse from {model}:")
print(f" Number of images: {len(response.data)}")
print(f" First image URL: {response.data[0].url if response.data else 'None'}")
# Basic assertions
assert response is not None, f"Response should not be None for {model}"
assert hasattr(response, "data"), f"Response should have data attribute for {model}"
assert len(response.data) > 0, f"Response should have at least one image for {model}"
# Check that we got a URL or b64_json
first_image = response.data[0]
assert (
first_image.url is not None or first_image.b64_json is not None
), f"Image should have either url or b64_json for {model}"
print(f"✓ Test passed for {model}")
except litellm.RateLimitError as e:
pytest.skip(f"Rate limit error for {model}: {str(e)}")
except litellm.ContentPolicyViolationError as e:
pytest.skip(f"Content policy violation for {model}: {str(e)}")
except litellm.InternalServerError as e:
pytest.skip(f"Internal server error for {model}: {str(e)}")
except Exception as e:
if "Your task failed as a result of our safety system" in str(e):
pytest.skip(f"Safety system rejection for {model}")
else:
pytest.fail(f"Test failed for {model}: {str(e)}")

View file

@ -164,6 +164,9 @@ class TestAimlImageGeneration(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "aiml/flux-pro/v1.1"}
class TestFAL_AI_ImageGeneration(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "fal_ai/fal-ai/imagen4/preview"}
class TestGoogleImageGen(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:

View file

@ -10,7 +10,6 @@ import traceback
import openai
import pytest
from pytest_mock import MockerFixture
sys.path.insert(
0, os.path.abspath("../..")
@ -540,13 +539,14 @@ async def test_simple_aembedding():
}
def test_custom_llm_provider_entrypoint(mocker: MockerFixture):
def test_custom_llm_provider_entrypoint():
# This test mocks the use of entry-points in pyproject.toml:
# [project.entry-point.litellm]
# custom_llm = <module>:MyCustomLLM
# another-custom-llm = <module>:AnotherCustomLLM
from litellm.utils import custom_llm_setup
from importlib.metadata import EntryPoints, EntryPoint
class AnotherCustomLLM(CustomLLM):
pass
@ -559,25 +559,22 @@ def test_custom_llm_provider_entrypoint(mocker: MockerFixture):
def load(self):
return providers[self.name]
mocker.patch("importlib.metadata.EntryPoint.load", load)
from importlib.metadata import EntryPoints, EntryPoint
entry_points = EntryPoints([
EntryPoint(group="litellm", name="custom_llm", value="package.module:MyCustomLLM"),
EntryPoint(group="litellm", name="another-custom-llm", value="package.module:AnotherCustomLLM"),
])
mocked = mocker.patch("litellm.utils.entry_points")
mocked.return_value = entry_points
assert litellm.custom_provider_map == []
assert litellm._custom_providers == []
with patch("importlib.metadata.EntryPoint.load", load):
with patch("litellm.utils.entry_points", return_value=entry_points):
assert litellm.custom_provider_map == []
assert litellm._custom_providers == []
custom_llm_setup()
custom_llm_setup()
assert litellm._custom_providers == ['custom_llm', 'another-custom-llm']
assert litellm._custom_providers == ['custom_llm', 'another-custom-llm']
assert litellm.custom_provider_map[0]["provider"] == "custom_llm"
assert isinstance(litellm.custom_provider_map[0]["custom_handler"], CustomLLM)
assert litellm.custom_provider_map[0]["provider"] == "custom_llm"
assert isinstance(litellm.custom_provider_map[0]["custom_handler"], CustomLLM)
assert litellm.custom_provider_map[1]["provider"] == "another-custom-llm"
assert isinstance(litellm.custom_provider_map[1]["custom_handler"], AnotherCustomLLM)
assert litellm.custom_provider_map[1]["provider"] == "another-custom-llm"
assert isinstance(litellm.custom_provider_map[1]["custom_handler"], AnotherCustomLLM)

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.1 KiB

View file

@ -575,6 +575,14 @@ const PROVIDER_CREDENTIAL_FIELDS: Record<Providers, ProviderCredentialField[]> =
placeholder: "http://localhost:7997",
},
],
[Providers.FalAI]: [
{
key: "api_key",
label: "API Key",
type: "password",
required: true,
}
],
};
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ selectedProvider, uploadProps }) => {

View file

@ -14,6 +14,7 @@ export enum Providers {
Deepgram = "Deepgram",
Deepseek = "Deepseek",
ElevenLabs = "ElevenLabs",
FalAI = "Fal AI",
FireworksAI = "Fireworks AI",
Google_AI_Studio = "Google AI Studio",
GradientAI = "GradientAI",
@ -73,6 +74,7 @@ export const provider_map: Record<string, string> = {
Triton: "triton",
Deepgram: "deepgram",
ElevenLabs: "elevenlabs",
FalAI: "fal_ai",
SageMaker: "sagemaker_chat",
Voyage: "voyage",
JinaAI: "jina_ai",
@ -120,6 +122,7 @@ export const providerLogoMap: Record<string, string> = {
[Providers.Triton]: `${asset_logos_folder}nvidia_triton.png`,
[Providers.Deepgram]: `${asset_logos_folder}deepgram.png`,
[Providers.ElevenLabs]: `${asset_logos_folder}elevenlabs.png`,
[Providers.FalAI]: `${asset_logos_folder}fal_ai.jpg`,
[Providers.Voyage]: `${asset_logos_folder}voyage.webp`,
[Providers.JinaAI]: `${asset_logos_folder}jina.png`,
[Providers.VolcEngine]: `${asset_logos_folder}volcengine.png`,
@ -183,6 +186,8 @@ export const getPlaceholder = (selectedProvider: string): string => {
return "volcengine/<any-model-on-volcengine>";
} else if (selectedProvider == Providers.DeepInfra) {
return "deepinfra/<any-model-on-deepinfra>";
} else if (selectedProvider == Providers.FalAI) {
return "fal_ai/fal-ai/flux-pro/v1.1-ultra";
} else {
return "gpt-3.5-turbo";
}