diff --git a/docs/my-website/docs/providers/fal_ai.md b/docs/my-website/docs/providers/fal_ai.md
new file mode 100644
index 00000000000..d42182b57a1
--- /dev/null
+++ b/docs/my-website/docs/providers/fal_ai.md
@@ -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
+
+
+
+
+```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)
+```
+
+
+
+
+
+```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)
+```
+
+
+
+
+
+```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)
+```
+
+
+
+
+
+```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())
+```
+
+
+
+
+
+```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}")
+```
+
+
+
+
+### 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
+
+
+
+
+```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)
+```
+
+
+
+
+
+```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)
+```
+
+
+
+
+
+```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"
+}'
+```
+
+
+
+
+
+
+## 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)
+
diff --git a/docs/my-website/docs/providers/github.md b/docs/my-website/docs/providers/github.md
index b9e525ef5c1..51220166140 100644
--- a/docs/my-website/docs/providers/github.md
+++ b/docs/my-website/docs/providers/github.md
@@ -1,7 +1,7 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
-# 🆕 Github
+# Github
https://github.com/marketplace/models
:::tip
diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js
index f83f454aa91..3b993e4620c 100644
--- a/docs/my-website/sidebars.js
+++ b/docs/my-website/sidebars.js
@@ -537,6 +537,7 @@ const sidebars = {
"providers/groq",
"providers/deepseek",
"providers/elevenlabs",
+ "providers/fal_ai",
"providers/fireworks_ai",
"providers/clarifai",
"providers/compactifai",
diff --git a/litellm/images/main.py b/litellm/images/main.py
index 63603411fb1..5be5f993814 100644
--- a/litellm/images/main.py
+++ b/litellm/images/main.py
@@ -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(
diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py
index 637446798b7..1df16a28177 100644
--- a/litellm/integrations/custom_logger.py
+++ b/litellm/integrations/custom_logger.py
@@ -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
diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py
index 337298688e0..18bd9fc1684 100644
--- a/litellm/litellm_core_utils/prompt_templates/factory.py
+++ b/litellm/litellm_core_utils/prompt_templates/factory.py
@@ -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
diff --git a/litellm/llms/fal_ai/__init__.py b/litellm/llms/fal_ai/__init__.py
new file mode 100644
index 00000000000..492197951e9
--- /dev/null
+++ b/litellm/llms/fal_ai/__init__.py
@@ -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",
+]
+
diff --git a/litellm/llms/fal_ai/cost_calculator.py b/litellm/llms/fal_ai/cost_calculator.py
new file mode 100644
index 00000000000..b7caae3834f
--- /dev/null
+++ b/litellm/llms/fal_ai/cost_calculator.py
@@ -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)}")
+
diff --git a/litellm/llms/fal_ai/image_generation/__init__.py b/litellm/llms/fal_ai/image_generation/__init__.py
new file mode 100644
index 00000000000..74d3b434b87
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/__init__.py
@@ -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()
+
diff --git a/litellm/llms/fal_ai/image_generation/bria_transformation.py b/litellm/llms/fal_ai/image_generation/bria_transformation.py
new file mode 100644
index 00000000000..cb5aa6b761d
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/bria_transformation.py
@@ -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
+
diff --git a/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py b/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py
new file mode 100644
index 00000000000..664f11d40dc
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/flux_pro_v11_ultra_transformation.py
@@ -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
+
diff --git a/litellm/llms/fal_ai/image_generation/imagen4_transformation.py b/litellm/llms/fal_ai/image_generation/imagen4_transformation.py
new file mode 100644
index 00000000000..f38ced65313
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/imagen4_transformation.py
@@ -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
+
diff --git a/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py
new file mode 100644
index 00000000000..572a8a0f1c3
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/recraft_v3_transformation.py
@@ -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
+
diff --git a/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py
new file mode 100644
index 00000000000..10e2c6b4161
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/stable_diffusion_transformation.py
@@ -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
+
diff --git a/litellm/llms/fal_ai/image_generation/transformation.py b/litellm/llms/fal_ai/image_generation/transformation.py
new file mode 100644
index 00000000000..04b7b167523
--- /dev/null
+++ b/litellm/llms/fal_ai/image_generation/transformation.py
@@ -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
+
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index 0623bc8a904..489cd8b8b56 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -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,
diff --git a/litellm/proxy/_experimental/out/assets/logos/fal_ai.jpg b/litellm/proxy/_experimental/out/assets/logos/fal_ai.jpg
new file mode 100644
index 00000000000..5de52c9188b
Binary files /dev/null and b/litellm/proxy/_experimental/out/assets/logos/fal_ai.jpg differ
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 964047fc57b..b6ed7c38b9b 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -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"
diff --git a/litellm/utils.py b/litellm/utils.py
index 0b9bc436c9e..065e7a172d7 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -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
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index 0623bc8a904..489cd8b8b56 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -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,
diff --git a/tests/image_gen_tests/base_image_generation_test.py b/tests/image_gen_tests/base_image_generation_test.py
index 6e8470525d6..e18094dceb0 100644
--- a/tests/image_gen_tests/base_image_generation_test.py
+++ b/tests/image_gen_tests/base_image_generation_test.py
@@ -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)
diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py
index a2bafa85c57..99bfccb1a00 100644
--- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py
+++ b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py
@@ -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"
diff --git a/tests/image_gen_tests/test_fal_ai_image_generation.py b/tests/image_gen_tests/test_fal_ai_image_generation.py
new file mode 100644
index 00000000000..949606a58ac
--- /dev/null
+++ b/tests/image_gen_tests/test_fal_ai_image_generation.py
@@ -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)}")
+
diff --git a/tests/image_gen_tests/test_image_generation.py b/tests/image_gen_tests/test_image_generation.py
index 0f980854c08..4fc8a18a9d3 100644
--- a/tests/image_gen_tests/test_image_generation.py
+++ b/tests/image_gen_tests/test_image_generation.py
@@ -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:
diff --git a/tests/local_testing/test_custom_llm.py b/tests/local_testing/test_custom_llm.py
index 16d20984d5b..0f1af3b56e1 100644
--- a/tests/local_testing/test_custom_llm.py
+++ b/tests/local_testing/test_custom_llm.py
@@ -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 = :MyCustomLLM
# another-custom-llm = :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)
diff --git a/ui/litellm-dashboard/public/assets/logos/fal_ai.jpg b/ui/litellm-dashboard/public/assets/logos/fal_ai.jpg
new file mode 100644
index 00000000000..5de52c9188b
Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/fal_ai.jpg differ
diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx
index e42f1bfba6a..c7686fb0b0c 100644
--- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx
+++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx
@@ -575,6 +575,14 @@ const PROVIDER_CREDENTIAL_FIELDS: Record =
placeholder: "http://localhost:7997",
},
],
+ [Providers.FalAI]: [
+ {
+ key: "api_key",
+ label: "API Key",
+ type: "password",
+ required: true,
+ }
+ ],
};
const ProviderSpecificFields: React.FC = ({ selectedProvider, uploadProps }) => {
diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx
index fb09b9fa82f..fb6162f304f 100644
--- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx
+++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx
@@ -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 = {
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 = {
[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/";
} else if (selectedProvider == Providers.DeepInfra) {
return "deepinfra/";
+ } else if (selectedProvider == Providers.FalAI) {
+ return "fal_ai/fal-ai/flux-pro/v1.1-ultra";
} else {
return "gpt-3.5-turbo";
}