diff --git a/docs/my-website/docs/image_edits.md b/docs/my-website/docs/image_edits.md index a843833454..4d72aa5bdf 100644 --- a/docs/my-website/docs/image_edits.md +++ b/docs/my-website/docs/image_edits.md @@ -16,7 +16,7 @@ LiteLLM provides image editing functionality that maps to OpenAI's `/images/edit | Supported operations | Create image edits | Single and multiple images supported | | Supported LiteLLM SDK Versions | 1.63.8+ | Gemini support requires 1.79.3+ | | Supported LiteLLM Proxy Versions | 1.71.1+ | Gemini support requires 1.79.3+ | -| Supported LLM providers | **OpenAI**, **Gemini (Google AI Studio)**, **Vertex AI**, **Stability AI**, **AWS Bedrock (Stability)** | Gemini supports the new `gemini-2.5-flash-image` family. Vertex AI supports both Gemini and Imagen models. Stability AI and Bedrock Stability support various image editing operations. | +| Supported LLM providers | **OpenAI**, **Gemini (Google AI Studio)**, **Vertex AI**, **Stability AI**, **AWS Bedrock (Stability)**, **Black Forest Labs** | Gemini supports the new `gemini-2.5-flash-image` family. Vertex AI supports both Gemini and Imagen models. Stability AI and Bedrock Stability support various image editing operations. Black Forest Labs supports FLUX Kontext models. | #### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/) @@ -199,6 +199,63 @@ for idx, image_obj in enumerate(response.data): + + +#### Basic Image Edit +```python showLineNumbers title="Black Forest Labs Image Edit" +import os +import litellm + +os.environ["BFL_API_KEY"] = "your-api-key" + +response = litellm.image_edit( + model="black_forest_labs/flux-kontext-pro", + image=open("original_image.png", "rb"), + prompt="Add a green leaf to the scene", +) + +print(response.data[0].url) +``` + +#### Inpainting with Mask +```python showLineNumbers title="Black Forest Labs Inpainting" +import os +import litellm + +os.environ["BFL_API_KEY"] = "your-api-key" + +# Use flux-pro-1.0-fill for inpainting +response = litellm.image_edit( + model="black_forest_labs/flux-pro-1.0-fill", + image=open("original_image.png", "rb"), + mask=open("mask_image.png", "rb"), + prompt="Replace with a garden", +) + +print(response.data[0].url) +``` + +#### Outpainting (Expand) +```python showLineNumbers title="Black Forest Labs Outpainting" +import os +import litellm + +os.environ["BFL_API_KEY"] = "your-api-key" + +# Use flux-pro-1.0-expand to extend image borders +response = litellm.image_edit( + model="black_forest_labs/flux-pro-1.0-expand", + image=open("original_image.png", "rb"), + prompt="Continue the scene with mountains", + top=256, + bottom=256, +) + +print(response.data[0].url) +``` + + + #### Basic Image Edit (Gemini) @@ -351,6 +408,35 @@ curl -X POST "http://0.0.0.0:4000/v1/images/edits" \ + + +1. Add Black Forest Labs image edit models to your `config.yaml`: +```yaml showLineNumbers title="Black Forest Labs Proxy Configuration" +model_list: + - model_name: bfl-kontext-pro + litellm_params: + model: black_forest_labs/flux-kontext-pro + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_edit +``` + +2. Start the LiteLLM proxy server: +```bash showLineNumbers title="Start LiteLLM Proxy Server" +litellm --config /path/to/config.yaml +``` + +3. Make an image edit request: +```bash showLineNumbers title="Black Forest Labs Proxy Image Edit" +curl -X POST "http://0.0.0.0:4000/v1/images/edits" \ + -H "Authorization: Bearer " \ + -F "model=bfl-kontext-pro" \ + -F "image=@original_image.png" \ + -F "prompt=Add a sunset in the background" +``` + + + 1. Add Vertex AI image edit models to your `config.yaml`: diff --git a/docs/my-website/docs/image_generation.md b/docs/my-website/docs/image_generation.md index 7f27f48f91..9002927d5f 100644 --- a/docs/my-website/docs/image_generation.md +++ b/docs/my-website/docs/image_generation.md @@ -15,7 +15,7 @@ import TabItem from '@theme/TabItem'; | Fallbacks | ✅ | Works between supported models | | Loadbalancing | ✅ | Works between supported models | | Guardrails | ✅ | Applies to input prompts (non-streaming only) | -| Supported Providers | OpenAI, Azure, Google AI Studio, Vertex AI, AWS Bedrock, Recraft, OpenRouter, Xinference, Nscale | | +| Supported Providers | OpenAI, Azure, Google AI Studio, Vertex AI, AWS Bedrock, Black Forest Labs, Recraft, OpenRouter, Xinference, Nscale | | ## Quick Start diff --git a/docs/my-website/docs/providers/black_forest_labs.md b/docs/my-website/docs/providers/black_forest_labs.md new file mode 100644 index 0000000000..7074fa1f13 --- /dev/null +++ b/docs/my-website/docs/providers/black_forest_labs.md @@ -0,0 +1,291 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Black Forest Labs Image Generation + +Black Forest Labs provides state-of-the-art text-to-image generation using their FLUX models. + +## Overview + +| Property | Details | +|----------|---------| +| Description | Black Forest Labs FLUX models for high-quality text-to-image generation | +| Provider Route on LiteLLM | `black_forest_labs/` | +| Provider Doc | [Black Forest Labs API ↗](https://docs.bfl.ai/) | +| Supported Operations | [`/images/generations`](#image-generation) | + +## Setup + +### API Key + +```python showLineNumbers +import os + +# Set your Black Forest Labs API key +os.environ["BFL_API_KEY"] = "your-api-key-here" +``` + +Get your API key from [Black Forest Labs](https://blackforestlabs.ai/). + +## Supported Models + +| Model Name | Description | Price | +|------------|-------------|-------| +| `black_forest_labs/flux-pro-1.1` | Fast & reliable standard generation | $0.04/image | +| `black_forest_labs/flux-pro-1.1-ultra` | Ultra high-resolution (up to 4MP) | $0.06/image | +| `black_forest_labs/flux-dev` | Development/open-source variant | $0.025/image | +| `black_forest_labs/flux-pro` | Original pro model | $0.05/image | + +## Image Generation + +### Usage - LiteLLM Python SDK + + + + +```python showLineNumbers title="Basic Image Generation" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Generate an image +response = litellm.image_generation( + model="black_forest_labs/flux-pro-1.1", + prompt="A beautiful sunset over the ocean with sailing boats", +) + +# BFL returns URLs +print(response.data[0].url) +``` + + + + + +```python showLineNumbers title="Async Image Generation" +import os +import asyncio +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +async def generate_image(): + response = await litellm.aimage_generation( + model="black_forest_labs/flux-pro-1.1", + prompt="A futuristic city skyline at night", + ) + print(response.data[0].url) + +# Run the async function +asyncio.run(generate_image()) +``` + + + + + +```python showLineNumbers title="Image Generation with Custom Size" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Generate with specific dimensions +response = litellm.image_generation( + model="black_forest_labs/flux-pro-1.1", + prompt="A majestic mountain landscape", + size="1792x1024", # Maps to width/height +) + +print(response.data[0].url) +``` + + + + + +```python showLineNumbers title="Ultra High Resolution with flux-pro-1.1-ultra" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Generate ultra high-resolution image +response = litellm.image_generation( + model="black_forest_labs/flux-pro-1.1-ultra", + prompt="Detailed portrait of a fantasy character", + size="2048x2048", # Up to 4MP supported + quality="hd", # Maps to raw=True for natural look +) + +print(response.data[0].url) +``` + + + + + +```python showLineNumbers title="Advanced Image Generation with BFL Parameters" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Generate with BFL-specific parameters +response = litellm.image_generation( + model="black_forest_labs/flux-pro-1.1", + prompt="A cute orange cat sitting on a windowsill", + seed=42, # For reproducible results + output_format="png", # png or jpeg + safety_tolerance=2, # 0-6, higher = more permissive + prompt_upsampling=True, # Enhance prompt for better results +) + +print(response.data[0].url) +``` + + + + +### Usage - LiteLLM Proxy Server + +#### 1. Configure your config.yaml + +```yaml showLineNumbers title="Black Forest Labs Image Generation Configuration" +model_list: + - model_name: flux-pro + litellm_params: + model: black_forest_labs/flux-pro-1.1 + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_generation + + - model_name: flux-ultra + litellm_params: + model: black_forest_labs/flux-pro-1.1-ultra + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_generation + + - model_name: flux-dev + litellm_params: + model: black_forest_labs/flux-dev + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_generation + +general_settings: + master_key: sk-1234 +``` + +#### 2. Start LiteLLM Proxy Server + +```bash showLineNumbers title="Start LiteLLM Proxy Server" +litellm --config /path/to/config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +#### 3. Make image generation requests + + + + +```python showLineNumbers title="Black Forest Labs via Proxy - OpenAI SDK" +from openai import OpenAI + +# Initialize client with your proxy URL +client = OpenAI( + base_url="http://localhost:4000", + api_key="sk-1234" +) + +# Generate image with FLUX Pro +response = client.images.generate( + model="flux-pro", + prompt="A beautiful garden with colorful flowers", + size="1024x1024", +) + +print(response.data[0].url) +``` + + + + + +```bash showLineNumbers title="Black Forest Labs via Proxy - cURL" +curl -X POST 'http://localhost:4000/v1/images/generations' \ + -H 'Content-Type: application/json' \ + -H 'Authorization: Bearer sk-1234' \ + -d '{ + "model": "flux-pro", + "prompt": "A beautiful garden with colorful flowers", + "size": "1024x1024" + }' +``` + + + + +## Supported Parameters + +### OpenAI-Compatible Parameters + +| Parameter | Type | Description | Mapping | +|-----------|------|-------------|---------| +| `prompt` | string | Text description of the image to generate | Direct | +| `model` | string | The FLUX model to use | Direct | +| `size` | string | Image dimensions (e.g., `1024x1024`) | Maps to `width` and `height` | +| `n` | integer | Number of images (ultra model only, up to 4) | Maps to `num_images` | +| `quality` | string | `hd` for natural look | Maps to `raw=True` for ultra | +| `response_format` | string | `url` or `b64_json` | Direct | + +### Black Forest Labs Specific Parameters + +| Parameter | Type | Description | Default | +|-----------|------|-------------|---------| +| `width` | integer | Image width (256-1920, multiples of 16) | 1024 | +| `height` | integer | Image height (256-1920, multiples of 16) | 1024 | +| `aspect_ratio` | string | Alternative to width/height (e.g., `16:9`, `1:1`) | - | +| `seed` | integer | Seed for reproducible results | Random | +| `output_format` | string | Output format: `png` or `jpeg` | `png` | +| `safety_tolerance` | integer | Safety filter tolerance (0-6, higher = more permissive) | 2 | +| `prompt_upsampling` | boolean | Enhance prompt for better results | `false` | + +### Ultra Model Specific Parameters + +| Parameter | Type | Description | Default | +|-----------|------|-------------|---------| +| `raw` | boolean | Raw mode for more natural, less synthetic look | `false` | +| `num_images` | integer | Number of images to generate (1-4) | 1 | + +## How It Works + +Black Forest Labs uses a polling-based API: + +1. **Submit Request**: LiteLLM sends your prompt to BFL +2. **Get Task ID**: BFL returns a task ID and polling URL +3. **Poll for Result**: LiteLLM automatically polls until the image is ready +4. **Return Result**: The generated image URL is returned + +This polling is handled automatically by LiteLLM - you just call `image_generation()` and get the result. + +## Getting Started + +1. Create an account at [Black Forest Labs](https://blackforestlabs.ai/) +2. Get your API key from the dashboard +3. Set your `BFL_API_KEY` environment variable +4. Use `litellm.image_generation()` with any supported model + +## Additional Resources + +- [Black Forest Labs Documentation](https://docs.bfl.ai/) +- [Black Forest Labs Image Editing](./black_forest_labs_img_edit.md) - For editing existing images +- [FLUX Model Information](https://blackforestlabs.ai/) diff --git a/docs/my-website/docs/providers/black_forest_labs_img_edit.md b/docs/my-website/docs/providers/black_forest_labs_img_edit.md new file mode 100644 index 0000000000..592ad0f9ef --- /dev/null +++ b/docs/my-website/docs/providers/black_forest_labs_img_edit.md @@ -0,0 +1,301 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Black Forest Labs Image Editing + +Black Forest Labs provides powerful image editing capabilities using their FLUX models to modify existing images based on text descriptions. + +## Overview + +| Property | Details | +|----------|---------| +| Description | Black Forest Labs Image Editing uses FLUX Kontext and other models to modify, inpaint, and expand images based on text prompts. | +| Provider Route on LiteLLM | `black_forest_labs/` | +| Provider Doc | [Black Forest Labs API ↗](https://docs.bfl.ai/) | +| Supported Operations | [`/images/edits`](#image-editing) | + +## Setup + +### API Key + +```python showLineNumbers +import os + +# Set your Black Forest Labs API key +os.environ["BFL_API_KEY"] = "your-api-key-here" +``` + +Get your API key from [Black Forest Labs](https://blackforestlabs.ai/). + +## Supported Models + +| Model Name | Description | Use Case | +|------------|-------------|----------| +| `black_forest_labs/flux-kontext-pro` | FLUX Kontext Pro - General image editing with prompts | General editing, style transfer | +| `black_forest_labs/flux-kontext-max` | FLUX Kontext Max - Premium quality editing | High-quality edits | +| `black_forest_labs/flux-pro-1.0-fill` | FLUX Pro Fill - Inpainting with mask | Remove/replace objects | +| `black_forest_labs/flux-pro-1.0-expand` | FLUX Pro Expand - Outpainting | Expand image borders | + +## Image Editing + +### Usage - LiteLLM Python SDK + + + + +```python showLineNumbers title="Basic Image Editing" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Edit an image with a prompt +response = litellm.image_edit( + model="black_forest_labs/flux-kontext-pro", + image=open("path/to/your/image.png", "rb"), + prompt="Add a green leaf to the scene", +) + +# BFL returns URLs +print(response.data[0].url) +``` + + + + + +```python showLineNumbers title="Async Image Editing" +import os +import asyncio +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +async def edit_image(): + response = await litellm.aimage_edit( + model="black_forest_labs/flux-kontext-pro", + image=open("path/to/your/image.png", "rb"), + prompt="Make this image look like a watercolor painting", + ) + print(response.data[0].url) + +# Run the async function +asyncio.run(edit_image()) +``` + + + + + +```python showLineNumbers title="Inpainting with Mask" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Use flux-pro-1.0-fill for inpainting +response = litellm.image_edit( + model="black_forest_labs/flux-pro-1.0-fill", + image=open("path/to/your/image.png", "rb"), + mask=open("path/to/mask.png", "rb"), # White areas will be edited + prompt="Replace with a beautiful garden", + steps=50, # BFL-specific parameter + guidance=30, # BFL-specific parameter +) + +print(response.data[0].url) +``` + + + + + +```python showLineNumbers title="Outpainting - Expand Image Borders" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Use flux-pro-1.0-expand to extend image borders +response = litellm.image_edit( + model="black_forest_labs/flux-pro-1.0-expand", + image=open("path/to/your/image.png", "rb"), + prompt="Continue the scene with a mountain landscape", + top=256, # Expand 256 pixels at top + bottom=256, # Expand 256 pixels at bottom + left=128, # Expand 128 pixels at left + right=128, # Expand 128 pixels at right +) + +print(response.data[0].url) +``` + + + + + +```python showLineNumbers title="Advanced Image Editing with BFL Parameters" +import os +import litellm + +# Set your API key +os.environ["BFL_API_KEY"] = "your-api-key-here" + +# Edit image with BFL-specific parameters +response = litellm.image_edit( + model="black_forest_labs/flux-kontext-pro", + image=open("path/to/your/image.png", "rb"), + prompt="Transform into cyberpunk style with neon lights", + seed=42, # For reproducible results + output_format="png", # png or jpeg + safety_tolerance=2, # 0-6, higher = more permissive + aspect_ratio="16:9", # Output aspect ratio +) + +print(response.data[0].url) +``` + + + + +### Usage - LiteLLM Proxy Server + +#### 1. Configure your config.yaml + +```yaml showLineNumbers title="Black Forest Labs Image Editing Configuration" +model_list: + - model_name: bfl-kontext-pro + litellm_params: + model: black_forest_labs/flux-kontext-pro + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_edit + + - model_name: bfl-kontext-max + litellm_params: + model: black_forest_labs/flux-kontext-max + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_edit + + - model_name: bfl-fill + litellm_params: + model: black_forest_labs/flux-pro-1.0-fill + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_edit + + - model_name: bfl-expand + litellm_params: + model: black_forest_labs/flux-pro-1.0-expand + api_key: os.environ/BFL_API_KEY + model_info: + mode: image_edit + +general_settings: + master_key: sk-1234 +``` + +#### 2. Start LiteLLM Proxy Server + +```bash showLineNumbers title="Start LiteLLM Proxy Server" +litellm --config /path/to/config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +#### 3. Make image editing requests + + + + +```python showLineNumbers title="Black Forest Labs via Proxy - OpenAI SDK" +from openai import OpenAI + +# Initialize client with your proxy URL +client = OpenAI( + base_url="http://localhost:4000", + api_key="sk-1234" +) + +# Edit image with FLUX Kontext Pro +response = client.images.edit( + model="bfl-kontext-pro", + image=open("path/to/your/image.png", "rb"), + prompt="Add magical sparkles and fairy dust", +) + +print(response.data[0].url) +``` + + + + + +```bash showLineNumbers title="Black Forest Labs via Proxy - cURL" +curl --location 'http://localhost:4000/v1/images/edits' \ +--header 'Authorization: Bearer sk-1234' \ +--form 'model="bfl-kontext-pro"' \ +--form 'prompt="Add a sunset in the background"' \ +--form 'image=@"path/to/your/image.png"' +``` + + + + +## Supported Parameters + +### OpenAI-Compatible Parameters + +| Parameter | Type | Description | Default | +|-----------|------|-------------|---------| +| `image` | file | The image file to edit | Required | +| `prompt` | string | Text description of the desired changes | Required | +| `model` | string | The FLUX model to use | Required | +| `mask` | file | Mask image for inpainting (flux-pro-1.0-fill) | Optional | +| `n` | integer | Number of images (BFL returns 1 per request) | `1` | +| `size` | string | Maps to aspect_ratio | Optional | +| `response_format` | string | `url` or `b64_json` | `url` | + +### Black Forest Labs Specific Parameters + +| Parameter | Type | Description | Default | Models | +|-----------|------|-------------|---------|--------| +| `seed` | integer | Seed for reproducible results | Random | All | +| `output_format` | string | Output format: `png` or `jpeg` | `png` | All | +| `safety_tolerance` | integer | Safety filter tolerance (0-6) | 2 | All | +| `aspect_ratio` | string | Output aspect ratio (e.g., `16:9`, `1:1`) | Original | Kontext models | +| `steps` | integer | Number of inference steps | Model default | Fill | +| `guidance` | float | Guidance scale | Model default | Fill | +| `grow_mask` | integer | Pixels to grow mask | 0 | Fill | +| `top` | integer | Pixels to expand at top | 0 | Expand | +| `bottom` | integer | Pixels to expand at bottom | 0 | Expand | +| `left` | integer | Pixels to expand at left | 0 | Expand | +| `right` | integer | Pixels to expand at right | 0 | Expand | + +## How It Works + +Black Forest Labs uses a polling-based API: + +1. **Submit Request**: LiteLLM sends your image and prompt to BFL +2. **Get Task ID**: BFL returns a task ID and polling URL +3. **Poll for Result**: LiteLLM automatically polls until the image is ready +4. **Return Result**: The generated image URL is returned + +This polling is handled automatically by LiteLLM - you just call `image_edit()` and get the result. + +## Getting Started + +1. Create an account at [Black Forest Labs](https://blackforestlabs.ai/) +2. Get your API key from the dashboard +3. Set your `BFL_API_KEY` environment variable +4. Use `litellm.image_edit()` with any supported model + +## Additional Resources + +- [Black Forest Labs Documentation](https://docs.bfl.ai/) +- [FLUX Model Information](https://blackforestlabs.ai/) diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index f7487d24b1..a2a10dab0e 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -805,6 +805,8 @@ const sidebars = { "providers/anyscale", "providers/apertis", "providers/baseten", + "providers/black_forest_labs", + "providers/black_forest_labs_img_edit", "providers/bytez", "providers/cerebras", "providers/chutes", diff --git a/litellm/__init__.py b/litellm/__init__.py index 84b8e47c46..1af1f68b5b 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -575,6 +575,7 @@ v0_models: Set = set() morph_models: Set = set() lambda_ai_models: Set = set() hyperbolic_models: Set = set() +black_forest_labs_models: Set = set() recraft_models: Set = set() cometapi_models: Set = set() oci_models: Set = set() @@ -821,6 +822,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None): lambda_ai_models.add(key) elif value.get("litellm_provider") == "hyperbolic": hyperbolic_models.add(key) + elif value.get("litellm_provider") == "black_forest_labs": + black_forest_labs_models.add(key) elif value.get("litellm_provider") == "recraft": recraft_models.add(key) elif value.get("litellm_provider") == "cometapi": @@ -952,6 +955,7 @@ model_list = list( | v0_models | morph_models | lambda_ai_models + | black_forest_labs_models | recraft_models | cometapi_models | oci_models @@ -1049,6 +1053,7 @@ models_by_provider: dict = { "morph": morph_models, "lambda_ai": lambda_ai_models, "hyperbolic": hyperbolic_models, + "black_forest_labs": black_forest_labs_models, "recraft": recraft_models, "cometapi": cometapi_models, "oci": oci_models, diff --git a/litellm/images/main.py b/litellm/images/main.py index eb6aa0c209..553aa26da9 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -50,6 +50,10 @@ from litellm.main import ( openai_image_variations, ) +# BFL handlers +from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit +from litellm.llms.black_forest_labs.image_generation.handler import bfl_image_generation + ########################################### from litellm.secret_managers.main import get_secret_str from litellm.types.images.main import ImageEditOptionalRequestParams @@ -404,7 +408,7 @@ def image_generation( # noqa: PLR0915 litellm.LlmProviders.STABILITY, litellm.LlmProviders.RUNWAYML, litellm.LlmProviders.VERTEX_AI, - litellm.LlmProviders.OPENROUTER + litellm.LlmProviders.OPENROUTER, ): if image_generation_config is None: raise ValueError( @@ -427,6 +431,22 @@ def image_generation( # noqa: PLR0915 timeout=timeout, client=client, ) + elif custom_llm_provider == "black_forest_labs": + # Route to BFL-specific handler (polling required) + if model is None: + raise Exception("Model needs to be set for black_forest_labs") + return bfl_image_generation.image_generation( + model=model, + prompt=prompt, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params_dict, + logging_obj=litellm_logging_obj, + timeout=timeout, + extra_headers=extra_headers, + client=client, + aimg_generation=aimg_generation, + ) elif custom_llm_provider == "azure_ai": from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo @@ -920,6 +940,23 @@ def image_edit( # noqa: PLR0915 _is_async=_is_async, client=kwargs.get("client"), ) + elif custom_llm_provider == "black_forest_labs": + # Route to BFL-specific handler (polling required) + if model is None: + raise Exception("Model needs to be set for black_forest_labs") + image_edit_request_params.update(non_default_params) + return bfl_image_edit.image_edit( + model=model, + image=images, + prompt=prompt, + image_edit_optional_request_params=image_edit_request_params, + litellm_params=litellm_params, + logging_obj=litellm_logging_obj, + timeout=timeout or DEFAULT_REQUEST_TIMEOUT, + extra_headers=extra_headers, + client=kwargs.get("client"), + aimage_edit=_is_async, + ) # Call the handler with _is_async flag instead of directly calling the async handler return base_llm_http_handler.image_edit_handler( model=model, diff --git a/litellm/llms/black_forest_labs/__init__.py b/litellm/llms/black_forest_labs/__init__.py new file mode 100644 index 0000000000..7a78638c8c --- /dev/null +++ b/litellm/llms/black_forest_labs/__init__.py @@ -0,0 +1,21 @@ +from .common_utils import ( + DEFAULT_API_BASE, + DEFAULT_MAX_POLLING_TIME, + DEFAULT_POLLING_INTERVAL, + IMAGE_EDIT_MODELS, + IMAGE_GENERATION_MODELS, + BlackForestLabsError, +) +from .image_edit import BlackForestLabsImageEditConfig +from .image_generation import BlackForestLabsImageGenerationConfig + +__all__ = [ + "BlackForestLabsError", + "BlackForestLabsImageEditConfig", + "BlackForestLabsImageGenerationConfig", + "DEFAULT_API_BASE", + "DEFAULT_MAX_POLLING_TIME", + "DEFAULT_POLLING_INTERVAL", + "IMAGE_EDIT_MODELS", + "IMAGE_GENERATION_MODELS", +] diff --git a/litellm/llms/black_forest_labs/common_utils.py b/litellm/llms/black_forest_labs/common_utils.py new file mode 100644 index 0000000000..507ef17c50 --- /dev/null +++ b/litellm/llms/black_forest_labs/common_utils.py @@ -0,0 +1,42 @@ +""" +Black Forest Labs Common Utilities + +Common utilities, constants, and error handling for Black Forest Labs API. +""" + +from typing import Dict + +from litellm.llms.base_llm.chat.transformation import BaseLLMException + + +class BlackForestLabsError(BaseLLMException): + """Exception class for Black Forest Labs API errors.""" + + pass + + +# API Constants +DEFAULT_API_BASE = "https://api.bfl.ai" + +# Polling configuration +DEFAULT_POLLING_INTERVAL = 1.5 # seconds +DEFAULT_MAX_POLLING_TIME = 300 # 5 minutes + +# Model to endpoint mapping for image edit +IMAGE_EDIT_MODELS: Dict[str, str] = { + "flux-kontext-pro": "/v1/flux-kontext-pro", + "flux-kontext-max": "/v1/flux-kontext-max", + "flux-pro-1.0-fill": "/v1/flux-pro-1.0-fill", + "flux-pro-1.0-expand": "/v1/flux-pro-1.0-expand", +} + +# Model to endpoint mapping for image generation +IMAGE_GENERATION_MODELS: Dict[str, str] = { + "flux-pro-1.1": "/v1/flux-pro-1.1", + "flux-pro-1.1-ultra": "/v1/flux-pro-1.1-ultra", + "flux-dev": "/v1/flux-dev", + "flux-pro": "/v1/flux-pro", + # Kontext models support both text-to-image and image editing + "flux-kontext-pro": "/v1/flux-kontext-pro", + "flux-kontext-max": "/v1/flux-kontext-max", +} diff --git a/litellm/llms/black_forest_labs/image_edit/__init__.py b/litellm/llms/black_forest_labs/image_edit/__init__.py new file mode 100644 index 0000000000..73af716e06 --- /dev/null +++ b/litellm/llms/black_forest_labs/image_edit/__init__.py @@ -0,0 +1,8 @@ +from .handler import BlackForestLabsImageEdit, bfl_image_edit +from .transformation import BlackForestLabsImageEditConfig + +__all__ = [ + "BlackForestLabsImageEditConfig", + "BlackForestLabsImageEdit", + "bfl_image_edit", +] diff --git a/litellm/llms/black_forest_labs/image_edit/handler.py b/litellm/llms/black_forest_labs/image_edit/handler.py new file mode 100644 index 0000000000..44a102ec48 --- /dev/null +++ b/litellm/llms/black_forest_labs/image_edit/handler.py @@ -0,0 +1,454 @@ +""" +Black Forest Labs Image Edit Handler + +Handles image edit requests for Black Forest Labs models. +BFL uses an async polling pattern - the initial request returns a task ID, +then we poll until the result is ready. +""" + +import asyncio +import time +from typing import Any, Dict, List, Optional, Union + +import httpx + +import litellm +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import FileTypes, ImageResponse + +from ..common_utils import ( + DEFAULT_MAX_POLLING_TIME, + DEFAULT_POLLING_INTERVAL, + BlackForestLabsError, +) +from .transformation import BlackForestLabsImageEditConfig + + +class BlackForestLabsImageEdit: + """ + Black Forest Labs Image Edit handler. + + Handles the HTTP requests and polling logic, delegating data transformation + to the BlackForestLabsImageEditConfig class. + """ + + def __init__(self): + self.config = BlackForestLabsImageEditConfig() + + def image_edit( + self, + model: str, + image: Union[FileTypes, List[FileTypes]], + prompt: Optional[str], + image_edit_optional_request_params: Dict, + litellm_params: Union[GenericLiteLLMParams, Dict], + logging_obj: LiteLLMLoggingObj, + timeout: Optional[Union[float, httpx.Timeout]], + extra_headers: Optional[Dict[str, Any]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + aimage_edit: bool = False, + ) -> Union[ImageResponse, Any]: + """ + Main entry point for image edit requests. + + Args: + model: The model to use (e.g., "black_forest_labs/flux-kontext-pro") + image: The image(s) to edit + prompt: The edit instruction + image_edit_optional_request_params: Optional parameters for the request + litellm_params: LiteLLM parameters including api_key, api_base + logging_obj: Logging object + timeout: Request timeout + extra_headers: Additional headers + client: HTTP client to use + aimage_edit: If True, return async coroutine + + Returns: + ImageResponse or coroutine if aimage_edit=True + """ + # Handle litellm_params as dict or object + if isinstance(litellm_params, dict): + api_key = litellm_params.get("api_key") + api_base = litellm_params.get("api_base") + litellm_params_dict = litellm_params + else: + api_key = litellm_params.api_key + api_base = litellm_params.api_base + litellm_params_dict = dict(litellm_params) + + if aimage_edit: + return self.async_image_edit( + model=model, + image=image, + prompt=prompt, + image_edit_optional_request_params=image_edit_optional_request_params, + litellm_params=litellm_params, + logging_obj=logging_obj, + timeout=timeout, + extra_headers=extra_headers, + client=client if isinstance(client, AsyncHTTPHandler) else None, + ) + + # Sync version + if client is None or not isinstance(client, HTTPHandler): + sync_client = _get_httpx_client() + else: + sync_client = client + + # Validate environment and get headers + headers = self.config.validate_environment( + api_key=api_key, + headers=image_edit_optional_request_params.get("extra_headers", {}) or {}, + model=model, + ) + if extra_headers: + headers.update(extra_headers) + + # Get complete URL + complete_url = self.config.get_complete_url( + model=model, + api_base=api_base, + litellm_params=litellm_params_dict, + ) + + # Transform request + # Handle image list vs single image + if isinstance(image, list): + if not image: + raise BlackForestLabsError(status_code=400, message="No image provided") + image_input = image[0] + else: + image_input = image + data, _ = self.config.transform_image_edit_request( + model=model, + prompt=prompt or "", + image=image_input, + image_edit_optional_request_params=image_edit_optional_request_params, + litellm_params=litellm_params_dict, + headers=headers, + ) + + # Logging + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": complete_url, + "headers": headers, + }, + ) + + # Make initial request + try: + response = sync_client.post( + url=complete_url, + headers=headers, + json=data, + timeout=timeout, + ) + except Exception as e: + raise BlackForestLabsError( + status_code=500, + message=f"Request failed: {str(e)}", + ) + + # Poll for result + final_response = self._poll_for_result_sync( + initial_response=response, + headers=headers, + sync_client=sync_client, + ) + + # Transform response + return self.config.transform_image_edit_response( + model=model, + raw_response=final_response, + logging_obj=logging_obj, + ) + + async def async_image_edit( + self, + model: str, + image: Union[FileTypes, List[FileTypes]], + prompt: Optional[str], + image_edit_optional_request_params: Dict, + litellm_params: Union[GenericLiteLLMParams, Dict], + logging_obj: LiteLLMLoggingObj, + timeout: Optional[Union[float, httpx.Timeout]], + extra_headers: Optional[Dict[str, Any]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> ImageResponse: + """ + Async version of image edit. + """ + # Handle litellm_params as dict or object + if isinstance(litellm_params, dict): + api_key = litellm_params.get("api_key") + api_base = litellm_params.get("api_base") + litellm_params_dict = litellm_params + else: + api_key = litellm_params.api_key + api_base = litellm_params.api_base + litellm_params_dict = dict(litellm_params) + + if client is None: + async_client = get_async_httpx_client( + llm_provider=litellm.LlmProviders.BLACK_FOREST_LABS, + ) + else: + async_client = client + + # Validate environment and get headers + headers = self.config.validate_environment( + api_key=api_key, + headers=image_edit_optional_request_params.get("extra_headers", {}) or {}, + model=model, + ) + if extra_headers: + headers.update(extra_headers) + + # Get complete URL + complete_url = self.config.get_complete_url( + model=model, + api_base=api_base, + litellm_params=litellm_params_dict, + ) + + # Transform request + if isinstance(image, list): + if not image: + raise BlackForestLabsError(status_code=400, message="No image provided") + image_input = image[0] + else: + image_input = image + data, _ = self.config.transform_image_edit_request( + model=model, + prompt=prompt or "", + image=image_input, + image_edit_optional_request_params=image_edit_optional_request_params, + litellm_params=litellm_params_dict, + headers=headers, + ) + + # Logging + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": complete_url, + "headers": headers, + }, + ) + + # Make initial request + try: + response = await async_client.post( + url=complete_url, + headers=headers, + json=data, + timeout=timeout, + ) + except Exception as e: + raise BlackForestLabsError( + status_code=500, + message=f"Request failed: {str(e)}", + ) + + # Poll for result + final_response = await self._poll_for_result_async( + initial_response=response, + headers=headers, + async_client=async_client, + ) + + # Transform response + return self.config.transform_image_edit_response( + model=model, + raw_response=final_response, + logging_obj=logging_obj, + ) + + def _poll_for_result_sync( + self, + initial_response: httpx.Response, + headers: dict, + sync_client: HTTPHandler, + max_wait: float = DEFAULT_MAX_POLLING_TIME, + interval: float = DEFAULT_POLLING_INTERVAL, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> httpx.Response: + """ + Poll BFL API until result is ready (sync version). + + Args: + initial_response: The initial response containing polling_url + headers: Headers to use for polling (must include x-key) + sync_client: HTTP client + max_wait: Maximum time to wait in seconds + interval: Polling interval in seconds + timeout: Timeout for each individual polling request + + Returns: + Final response with completed result + """ + # Validate initial response status code + if initial_response.status_code >= 400: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL initial request failed: {initial_response.text}", + ) + + # Parse initial response to get polling URL + try: + response_data = initial_response.json() + except Exception as e: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"Error parsing initial response: {e}", + ) + + # Check for immediate errors + if "errors" in response_data: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL error: {response_data['errors']}", + ) + + polling_url = response_data.get("polling_url") + if not polling_url: + raise BlackForestLabsError( + status_code=500, + message="No polling_url in BFL response", + ) + + # Get just the auth header for polling + polling_headers = {"x-key": headers.get("x-key", "")} + + start_time = time.time() + verbose_logger.debug(f"BFL starting sync polling at {polling_url}") + + while time.time() - start_time < max_wait: + response = sync_client.get( + url=polling_url, + headers=polling_headers, + ) + + if response.status_code != 200: + raise BlackForestLabsError( + status_code=response.status_code, + message=f"Polling failed: {response.text}", + ) + + data = response.json() + status = data.get("status") + + verbose_logger.debug(f"BFL poll status: {status}") + + if status == "Ready": + return response + elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + raise BlackForestLabsError( + status_code=400, + message=f"Image generation failed: {status}", + ) + + time.sleep(interval) + + raise BlackForestLabsError( + status_code=408, + message=f"Polling timed out after {max_wait} seconds", + ) + + async def _poll_for_result_async( + self, + initial_response: httpx.Response, + headers: dict, + async_client: AsyncHTTPHandler, + max_wait: float = DEFAULT_MAX_POLLING_TIME, + interval: float = DEFAULT_POLLING_INTERVAL, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> httpx.Response: + """ + Poll BFL API until result is ready (async version). + """ + # Validate initial response status code + if initial_response.status_code >= 400: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL initial request failed: {initial_response.text}", + ) + + # Parse initial response to get polling URL + try: + response_data = initial_response.json() + except Exception as e: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"Error parsing initial response: {e}", + ) + + # Check for immediate errors + if "errors" in response_data: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL error: {response_data['errors']}", + ) + + polling_url = response_data.get("polling_url") + if not polling_url: + raise BlackForestLabsError( + status_code=500, + message="No polling_url in BFL response", + ) + + # Get just the auth header for polling + polling_headers = {"x-key": headers.get("x-key", "")} + + start_time = time.time() + verbose_logger.debug(f"BFL starting async polling at {polling_url}") + + while time.time() - start_time < max_wait: + response = await async_client.get( + url=polling_url, + headers=polling_headers, + ) + + if response.status_code != 200: + raise BlackForestLabsError( + status_code=response.status_code, + message=f"Polling failed: {response.text}", + ) + + data = response.json() + status = data.get("status") + + verbose_logger.debug(f"BFL poll status: {status}") + + if status == "Ready": + return response + elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + raise BlackForestLabsError( + status_code=400, + message=f"Image generation failed: {status}", + ) + + await asyncio.sleep(interval) + + raise BlackForestLabsError( + status_code=408, + message=f"Polling timed out after {max_wait} seconds", + ) + + +# Singleton instance for use in images/main.py +bfl_image_edit = BlackForestLabsImageEdit() diff --git a/litellm/llms/black_forest_labs/image_edit/transformation.py b/litellm/llms/black_forest_labs/image_edit/transformation.py new file mode 100644 index 0000000000..78898345bf --- /dev/null +++ b/litellm/llms/black_forest_labs/image_edit/transformation.py @@ -0,0 +1,308 @@ +""" +Black Forest Labs Image Edit Configuration + +Handles transformation between OpenAI-compatible format and Black Forest Labs API format +for image editing endpoints (flux-kontext-pro, flux-kontext-max, etc.). + +API Reference: https://docs.bfl.ai/ +""" + +import base64 +import time +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union + +import httpx +from httpx._types import RequestFiles + +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import FileTypes, ImageObject, ImageResponse + +from ..common_utils import ( + DEFAULT_API_BASE, + IMAGE_EDIT_MODELS, + BlackForestLabsError, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class BlackForestLabsImageEditConfig(BaseImageEditConfig): + """ + Configuration for Black Forest Labs image editing. + + Supports: + - flux-kontext-pro: General image editing with prompts + - flux-kontext-max: Premium quality editing + - flux-pro-1.0-fill: Inpainting with mask + - flux-pro-1.0-expand: Outpainting (expand image borders) + + Note: HTTP requests and polling are handled by the handler (handler.py). + This class only handles data transformation. + """ + + def get_supported_openai_params(self, model: str) -> List[str]: + """ + Return list of OpenAI params supported by Black Forest Labs. + + Note: BFL uses different parameter names, these are mapped in map_openai_params. + """ + return [ + "mask", + "seed", + "output_format", + "safety_tolerance", + "prompt_upsampling", + "aspect_ratio", + "steps", + "guidance", + "grow_mask", + "top", + "bottom", + "left", + "right", + ] + + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """ + Map OpenAI parameters to Black Forest Labs parameters. + + BFL-specific params are passed through directly. + """ + optional_params: Dict[str, Any] = {} + + # Pass through BFL-specific params + bfl_params = [ + "seed", + "output_format", + "safety_tolerance", + "prompt_upsampling", + # Kontext-specific + "aspect_ratio", + # Fill/Inpaint-specific + "steps", + "guidance", + "grow_mask", + # Expand-specific + "top", + "bottom", + "left", + "right", + ] + + # Convert TypedDict to regular dict for access + params_dict = dict(image_edit_optional_params) + + for param in bfl_params: + if param in params_dict: + value = params_dict[param] + if value is not None: + optional_params[param] = value + + # Set default output format + if "output_format" not in optional_params: + optional_params["output_format"] = "png" + + return optional_params + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + ) -> dict: + """ + Validate environment and set up headers for Black Forest Labs. + + BFL uses x-key header for authentication. + """ + final_api_key: Optional[str] = ( + api_key + or get_secret_str("BFL_API_KEY") + or get_secret_str("BLACK_FOREST_LABS_API_KEY") + ) + + if not final_api_key: + raise BlackForestLabsError( + status_code=401, + message="BFL_API_KEY is not set. Please set it via environment variable or pass api_key parameter.", + ) + + headers["x-key"] = final_api_key + headers["Content-Type"] = "application/json" + headers["Accept"] = "application/json" + + return headers + + def use_multipart_form_data(self) -> bool: + """ + BFL uses JSON requests, not multipart/form-data. + """ + return False + + def _get_model_endpoint(self, model: str) -> str: + """ + Get the API endpoint for a given model. + """ + # Remove provider prefix if present (e.g., "black_forest_labs/flux-kontext-pro") + model_name = model.lower() + if "/" in model_name: + model_name = model_name.split("/")[-1] + + # Check if model is in our mapping + if model_name in IMAGE_EDIT_MODELS: + return IMAGE_EDIT_MODELS[model_name] + + raise ValueError( + f"Unknown BFL image edit model: {model_name}. " + f"Supported models: {list(IMAGE_EDIT_MODELS.keys())}" + ) + + def get_complete_url( + self, + model: str, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Get the complete URL for the Black Forest Labs API request. + """ + base_url: str = ( + api_base + or get_secret_str("BFL_API_BASE") + or DEFAULT_API_BASE + ) + base_url = base_url.rstrip("/") + + endpoint = self._get_model_endpoint(model) + return f"{base_url}{endpoint}" + + def _read_image_bytes(self, image: Any) -> bytes: + """Read image bytes from various input types.""" + if isinstance(image, bytes): + return image + elif isinstance(image, list): + # If it's a list, take the first image + return self._read_image_bytes(image[0]) + elif isinstance(image, str): + if image.startswith(("http://", "https://")): + # Download image from URL + response = httpx.get(image, timeout=60.0) + response.raise_for_status() + return response.content + else: + # Assume it's a file path + with open(image, "rb") as f: + return f.read() + elif hasattr(image, "read"): + # File-like object + pos = getattr(image, "tell", lambda: 0)() + if hasattr(image, "seek"): + image.seek(0) + data = image.read() + if hasattr(image, "seek"): + image.seek(pos) + return data + else: + raise ValueError( + f"Unsupported image type: {type(image)}. " + "Expected bytes, str (URL or file path), or file-like object." + ) + + def transform_image_edit_request( + self, + model: str, + prompt: str, + image: FileTypes, + image_edit_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[Dict, RequestFiles]: + """ + Transform OpenAI-style request to Black Forest Labs request format. + + BFL uses JSON body with base64-encoded images, not multipart/form-data. + """ + # Read and encode image + image_bytes = self._read_image_bytes(image) + b64_image = base64.b64encode(image_bytes).decode("utf-8") + + # Build request body + request_body: Dict[str, Any] = { + "prompt": prompt, + "input_image": b64_image, + } + + # Add optional params (only BFL-recognized parameters) + bfl_request_params = [ + "seed", "output_format", "safety_tolerance", "prompt_upsampling", + "aspect_ratio", "steps", "guidance", "grow_mask", + "top", "bottom", "left", "right", + ] + for key, value in image_edit_optional_request_params.items(): + if key in bfl_request_params and value is not None: + request_body[key] = value + + # Handle mask if provided (for inpainting) + if "mask" in image_edit_optional_request_params: + mask = image_edit_optional_request_params["mask"] + mask_bytes = self._read_image_bytes(mask) + request_body["mask"] = base64.b64encode(mask_bytes).decode("utf-8") + + # BFL uses JSON, not multipart - return empty files + return request_body, [] + + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ImageResponse: + """ + Transform Black Forest Labs response to OpenAI-compatible ImageResponse. + + This is called with the FINAL polled response (after handler does polling). + The response contains: {"status": "Ready", "result": {"sample": "https://..."}} + """ + try: + response_data = raw_response.json() + except Exception as e: + raise BlackForestLabsError( + status_code=raw_response.status_code, + message=f"Error parsing BFL response: {e}", + ) + + # Get image URL from result + image_url = response_data.get("result", {}).get("sample") + if not image_url: + raise BlackForestLabsError( + status_code=500, + message="No image URL in BFL result", + ) + + # Build ImageResponse + return ImageResponse( + created=int(time.time()), + data=[ImageObject(url=image_url)], + ) + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BlackForestLabsError: + """Return the appropriate error class for Black Forest Labs.""" + return BlackForestLabsError( + status_code=status_code, + message=error_message, + ) diff --git a/litellm/llms/black_forest_labs/image_generation/__init__.py b/litellm/llms/black_forest_labs/image_generation/__init__.py new file mode 100644 index 0000000000..2ccee2069e --- /dev/null +++ b/litellm/llms/black_forest_labs/image_generation/__init__.py @@ -0,0 +1,12 @@ +from .handler import BlackForestLabsImageGeneration, bfl_image_generation +from .transformation import ( + BlackForestLabsImageGenerationConfig, + get_black_forest_labs_image_generation_config, +) + +__all__ = [ + "BlackForestLabsImageGenerationConfig", + "get_black_forest_labs_image_generation_config", + "BlackForestLabsImageGeneration", + "bfl_image_generation", +] diff --git a/litellm/llms/black_forest_labs/image_generation/handler.py b/litellm/llms/black_forest_labs/image_generation/handler.py new file mode 100644 index 0000000000..99dc2feca3 --- /dev/null +++ b/litellm/llms/black_forest_labs/image_generation/handler.py @@ -0,0 +1,440 @@ +""" +Black Forest Labs Image Generation Handler + +Handles image generation requests for Black Forest Labs models. +BFL uses an async polling pattern - the initial request returns a task ID, +then we poll until the result is ready. +""" + +import asyncio +import time +from typing import Any, Dict, Optional, Union + +import httpx + +import litellm +from litellm._logging import verbose_logger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import ImageResponse + +from ..common_utils import ( + DEFAULT_MAX_POLLING_TIME, + DEFAULT_POLLING_INTERVAL, + BlackForestLabsError, +) +from .transformation import BlackForestLabsImageGenerationConfig + + +class BlackForestLabsImageGeneration: + """ + Black Forest Labs Image Generation handler. + + Handles the HTTP requests and polling logic, delegating data transformation + to the BlackForestLabsImageGenerationConfig class. + """ + + def __init__(self): + self.config = BlackForestLabsImageGenerationConfig() + + def image_generation( + self, + model: str, + prompt: str, + model_response: ImageResponse, + optional_params: Dict, + litellm_params: Union[GenericLiteLLMParams, Dict], + logging_obj: LiteLLMLoggingObj, + timeout: Optional[Union[float, httpx.Timeout]], + extra_headers: Optional[Dict[str, Any]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + aimg_generation: bool = False, + ) -> Union[ImageResponse, Any]: + """ + Main entry point for image generation requests. + + Args: + model: The model to use (e.g., "black_forest_labs/flux-pro-1.1") + prompt: The text prompt for image generation + model_response: ImageResponse object to populate + optional_params: Optional parameters for the request + litellm_params: LiteLLM parameters including api_key, api_base + logging_obj: Logging object + timeout: Request timeout + extra_headers: Additional headers + client: HTTP client to use + aimg_generation: If True, return async coroutine + + Returns: + ImageResponse or coroutine if aimg_generation=True + """ + # Handle litellm_params as dict or object + if isinstance(litellm_params, dict): + api_key = litellm_params.get("api_key") + api_base = litellm_params.get("api_base") + litellm_params_dict = litellm_params + else: + api_key = litellm_params.api_key + api_base = litellm_params.api_base + litellm_params_dict = dict(litellm_params) + + if aimg_generation: + return self.async_image_generation( + model=model, + prompt=prompt, + model_response=model_response, + optional_params=optional_params, + litellm_params=litellm_params, + logging_obj=logging_obj, + timeout=timeout, + extra_headers=extra_headers, + client=client if isinstance(client, AsyncHTTPHandler) else None, + ) + + # Sync version + if client is None or not isinstance(client, HTTPHandler): + sync_client = _get_httpx_client() + else: + sync_client = client + + # Validate environment and get headers + headers = self.config.validate_environment( + api_key=api_key, + headers={}, + model=model, + messages=[], + optional_params=optional_params, + litellm_params=litellm_params_dict, + ) + if extra_headers: + headers.update(extra_headers) + + # Get complete URL + complete_url = self.config.get_complete_url( + api_base=api_base, + api_key=api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params_dict, + ) + + # Transform request + data = self.config.transform_image_generation_request( + model=model, + prompt=prompt, + optional_params=optional_params, + litellm_params=litellm_params_dict, + headers=headers, + ) + + # Logging + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": complete_url, + "headers": headers, + }, + ) + + # Make initial request + try: + response = sync_client.post( + url=complete_url, + headers=headers, + json=data, + timeout=timeout, + ) + except Exception as e: + raise BlackForestLabsError( + status_code=500, + message=f"Request failed: {str(e)}", + ) + + # Poll for result + final_response = self._poll_for_result_sync( + initial_response=response, + headers=headers, + sync_client=sync_client, + ) + + # Transform response + return self.config.transform_image_generation_response( + model=model, + raw_response=final_response, + model_response=model_response, + logging_obj=logging_obj, + ) + + async def async_image_generation( + self, + model: str, + prompt: str, + model_response: ImageResponse, + optional_params: Dict, + litellm_params: Union[GenericLiteLLMParams, Dict], + logging_obj: LiteLLMLoggingObj, + timeout: Optional[Union[float, httpx.Timeout]], + extra_headers: Optional[Dict[str, Any]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> ImageResponse: + """ + Async version of image generation. + """ + # Handle litellm_params as dict or object + if isinstance(litellm_params, dict): + api_key = litellm_params.get("api_key") + api_base = litellm_params.get("api_base") + litellm_params_dict = litellm_params + else: + api_key = litellm_params.api_key + api_base = litellm_params.api_base + litellm_params_dict = dict(litellm_params) + + if client is None: + async_client = get_async_httpx_client( + llm_provider=litellm.LlmProviders.BLACK_FOREST_LABS, + ) + else: + async_client = client + + # Validate environment and get headers + headers = self.config.validate_environment( + api_key=api_key, + headers={}, + model=model, + messages=[], + optional_params=optional_params, + litellm_params=litellm_params_dict, + ) + if extra_headers: + headers.update(extra_headers) + + # Get complete URL + complete_url = self.config.get_complete_url( + api_base=api_base, + api_key=api_key, + model=model, + optional_params=optional_params, + litellm_params=litellm_params_dict, + ) + + # Transform request + data = self.config.transform_image_generation_request( + model=model, + prompt=prompt, + optional_params=optional_params, + litellm_params=litellm_params_dict, + headers=headers, + ) + + # Logging + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": complete_url, + "headers": headers, + }, + ) + + # Make initial request + try: + response = await async_client.post( + url=complete_url, + headers=headers, + json=data, + timeout=timeout, + ) + except Exception as e: + raise BlackForestLabsError( + status_code=500, + message=f"Request failed: {str(e)}", + ) + + # Poll for result + final_response = await self._poll_for_result_async( + initial_response=response, + headers=headers, + async_client=async_client, + ) + + # Transform response + return self.config.transform_image_generation_response( + model=model, + raw_response=final_response, + model_response=model_response, + logging_obj=logging_obj, + ) + + def _poll_for_result_sync( + self, + initial_response: httpx.Response, + headers: dict, + sync_client: HTTPHandler, + max_wait: float = DEFAULT_MAX_POLLING_TIME, + interval: float = DEFAULT_POLLING_INTERVAL, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> httpx.Response: + """ + Poll BFL API until result is ready (sync version). + """ + # Validate initial response status code + if initial_response.status_code >= 400: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL initial request failed: {initial_response.text}", + ) + + # Parse initial response to get polling URL + try: + response_data = initial_response.json() + except Exception as e: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"Error parsing initial response: {e}", + ) + + # Check for immediate errors + if "errors" in response_data: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL error: {response_data['errors']}", + ) + + polling_url = response_data.get("polling_url") + if not polling_url: + raise BlackForestLabsError( + status_code=500, + message="No polling_url in BFL response", + ) + + # Get just the auth header for polling + polling_headers = {"x-key": headers.get("x-key", "")} + + start_time = time.time() + verbose_logger.debug(f"BFL starting sync polling at {polling_url}") + + while time.time() - start_time < max_wait: + response = sync_client.get( + url=polling_url, + headers=polling_headers, + ) + + if response.status_code != 200: + raise BlackForestLabsError( + status_code=response.status_code, + message=f"Polling failed: {response.text}", + ) + + data = response.json() + status = data.get("status") + + verbose_logger.debug(f"BFL poll status: {status}") + + if status == "Ready": + return response + elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + raise BlackForestLabsError( + status_code=400, + message=f"Image generation failed: {status}", + ) + + time.sleep(interval) + + raise BlackForestLabsError( + status_code=408, + message=f"Polling timed out after {max_wait} seconds", + ) + + async def _poll_for_result_async( + self, + initial_response: httpx.Response, + headers: dict, + async_client: AsyncHTTPHandler, + max_wait: float = DEFAULT_MAX_POLLING_TIME, + interval: float = DEFAULT_POLLING_INTERVAL, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> httpx.Response: + """ + Poll BFL API until result is ready (async version). + """ + # Validate initial response status code + if initial_response.status_code >= 400: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL initial request failed: {initial_response.text}", + ) + + # Parse initial response to get polling URL + try: + response_data = initial_response.json() + except Exception as e: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"Error parsing initial response: {e}", + ) + + # Check for immediate errors + if "errors" in response_data: + raise BlackForestLabsError( + status_code=initial_response.status_code, + message=f"BFL error: {response_data['errors']}", + ) + + polling_url = response_data.get("polling_url") + if not polling_url: + raise BlackForestLabsError( + status_code=500, + message="No polling_url in BFL response", + ) + + # Get just the auth header for polling + polling_headers = {"x-key": headers.get("x-key", "")} + + start_time = time.time() + verbose_logger.debug(f"BFL starting async polling at {polling_url}") + + while time.time() - start_time < max_wait: + response = await async_client.get( + url=polling_url, + headers=polling_headers, + ) + + if response.status_code != 200: + raise BlackForestLabsError( + status_code=response.status_code, + message=f"Polling failed: {response.text}", + ) + + data = response.json() + status = data.get("status") + + verbose_logger.debug(f"BFL poll status: {status}") + + if status == "Ready": + return response + elif status in ["Error", "Failed", "Content Moderated", "Request Moderated"]: + raise BlackForestLabsError( + status_code=400, + message=f"Image generation failed: {status}", + ) + + await asyncio.sleep(interval) + + raise BlackForestLabsError( + status_code=408, + message=f"Polling timed out after {max_wait} seconds", + ) + + +# Singleton instance for use in images/main.py +bfl_image_generation = BlackForestLabsImageGeneration() diff --git a/litellm/llms/black_forest_labs/image_generation/transformation.py b/litellm/llms/black_forest_labs/image_generation/transformation.py new file mode 100644 index 0000000000..fd664b3ea7 --- /dev/null +++ b/litellm/llms/black_forest_labs/image_generation/transformation.py @@ -0,0 +1,324 @@ +""" +Black Forest Labs Image Generation Configuration + +Handles transformation between OpenAI-compatible format and Black Forest Labs API format +for image generation endpoints (flux-pro-1.1, flux-pro-1.1-ultra, flux-dev, flux-pro). + +API Reference: https://docs.bfl.ai/ +""" + +import time +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union + +import httpx + +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIImageGenerationOptionalParams, +) +from litellm.types.utils import ImageObject, ImageResponse + +from ..common_utils import ( + DEFAULT_API_BASE, + IMAGE_GENERATION_MODELS, + BlackForestLabsError, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class BlackForestLabsImageGenerationConfig(BaseImageGenerationConfig): + """ + Configuration for Black Forest Labs image generation (text-to-image). + + Supports: + - flux-pro-1.1: Fast & reliable standard generation + - flux-pro-1.1-ultra: Ultra high-resolution (up to 4MP) + - flux-dev: Development/open-source variant + - flux-pro: Original pro model + + Note: HTTP requests and polling are handled by the handler (handler.py). + This class only handles data transformation. + """ + + def get_supported_openai_params( + self, model: str + ) -> List[OpenAIImageGenerationOptionalParams]: + """ + Return list of OpenAI params supported by Black Forest Labs. + + Note: BFL uses different parameter names, these are mapped in map_openai_params. + """ + return [ + "n", # Number of images (BFL returns 1 per request, but ultra supports up to 4) + "size", # Maps to width/height or aspect_ratio + "quality", # Maps to raw mode for ultra + "seed", + "output_format", + "safety_tolerance", + "prompt_upsampling", + "raw", + "num_images", + "image_url", + "image_prompt_strength", + "aspect_ratio", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + """ + Map OpenAI parameters to Black Forest Labs parameters. + + BFL-specific params are passed through directly. + """ + supported_params = self.get_supported_openai_params(model) + + for k, v in non_default_params.items(): + if k in optional_params: + continue + + if k in supported_params: + # Map OpenAI 'size' to BFL width/height + if k == "size" and v: + self._map_size_param(v, optional_params) + elif k == "n": + if "ultra" in model.lower(): + optional_params["num_images"] = v + # non-ultra: silently skip (n=1 is BFL default) + elif k == "quality": + if v == "hd" and "ultra" in model.lower(): + optional_params["raw"] = True + # other quality values have no BFL mapping + else: + optional_params[k] = v + elif not drop_params: + raise ValueError( + f"Parameter {k} is not supported for model {model}. " + f"Supported parameters are {supported_params}. " + f"Set drop_params=True to drop unsupported parameters." + ) + + return optional_params + + def _map_size_param(self, size: str, optional_params: dict) -> None: + """Map OpenAI size parameter to BFL width/height.""" + # Common size mappings + size_mapping = { + "1024x1024": (1024, 1024), + "1792x1024": (1792, 1024), + "1024x1792": (1024, 1792), + "512x512": (512, 512), + "256x256": (256, 256), + } + + if size in size_mapping: + width, height = size_mapping[size] + optional_params["width"] = width + optional_params["height"] = height + elif "x" in size: + # Parse custom size + try: + width, height = map(int, size.lower().split("x")) + optional_params["width"] = width + optional_params["height"] = height + except ValueError: + raise ValueError( + f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')." + ) + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + """ + Validate environment and set up headers for Black Forest Labs. + + BFL uses x-key header for authentication. + """ + final_api_key: Optional[str] = ( + api_key + or get_secret_str("BFL_API_KEY") + or get_secret_str("BLACK_FOREST_LABS_API_KEY") + ) + + if not final_api_key: + raise BlackForestLabsError( + status_code=401, + message="BFL_API_KEY is not set. Please set it via environment variable or pass api_key parameter.", + ) + + headers["x-key"] = final_api_key + headers["Content-Type"] = "application/json" + headers["Accept"] = "application/json" + + return headers + + def _get_model_endpoint(self, model: str) -> str: + """ + Get the API endpoint for a given model. + """ + # Remove provider prefix if present (e.g., "black_forest_labs/flux-pro-1.1") + model_name = model.lower() + if "/" in model_name: + model_name = model_name.split("/")[-1] + + # Check if model is in our mapping + if model_name in IMAGE_GENERATION_MODELS: + return IMAGE_GENERATION_MODELS[model_name] + + raise ValueError( + f"Unknown BFL image generation model: {model_name}. " + f"Supported models: {list(IMAGE_GENERATION_MODELS.keys())}" + ) + + 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 Black Forest Labs API request. + """ + base_url: str = ( + api_base or get_secret_str("BFL_API_BASE") or DEFAULT_API_BASE + ) + base_url = base_url.rstrip("/") + + endpoint = self._get_model_endpoint(model) + return f"{base_url}{endpoint}" + + def transform_image_generation_request( + self, + model: str, + prompt: str, + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transform OpenAI-style request to Black Forest Labs request format. + + https://docs.bfl.ai/flux_models/flux_1_1_pro + """ + # Build request body with prompt + request_body: Dict[str, Any] = { + "prompt": prompt, + } + + # BFL-specific params that can be passed through + bfl_params = [ + "width", + "height", + "aspect_ratio", + "seed", + "output_format", + "safety_tolerance", + "prompt_upsampling", + # Ultra-specific + "raw", + "num_images", + "image_url", + "image_prompt_strength", + ] + + for param in bfl_params: + if param in optional_params and optional_params[param] is not None: + request_body[param] = optional_params[param] + + # Set default output format if not specified + if "output_format" not in request_body: + request_body["output_format"] = "png" + + return request_body + + def transform_image_generation_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ImageResponse, + logging_obj: LiteLLMLoggingObj, + **kwargs, + ) -> ImageResponse: + """ + Transform Black Forest Labs response to OpenAI-compatible ImageResponse. + + This is called with the FINAL polled response (after handler does polling). + The response contains: {"status": "Ready", "result": {"sample": "https://..."}} + """ + try: + response_data = raw_response.json() + except Exception as e: + raise BlackForestLabsError( + status_code=raw_response.status_code, + message=f"Error parsing BFL response: {e}", + ) + + result = response_data.get("result", {}) + + if not model_response.data: + model_response.data = [] + + # Handle single image (sample) or multiple images + if isinstance(result, dict) and "sample" in result: + model_response.data.append(ImageObject(url=result["sample"])) + elif isinstance(result, list): + # Multiple images returned + for img in result: + if isinstance(img, str): + model_response.data.append(ImageObject(url=img)) + elif isinstance(img, dict) and "url" in img: + model_response.data.append(ImageObject(url=img["url"])) + + if not model_response.data: + raise BlackForestLabsError( + status_code=500, + message="No image URL in BFL result", + ) + + model_response.created = int(time.time()) + return model_response + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BlackForestLabsError: + """Return the appropriate error class for Black Forest Labs.""" + return BlackForestLabsError( + status_code=status_code, + message=error_message, + ) + + +def get_black_forest_labs_image_generation_config( + model: str, +) -> BlackForestLabsImageGenerationConfig: + """ + Get the appropriate image generation config for a Black Forest Labs model. + + Currently returns a single config class, but can be extended + for model-specific configurations if needed. + """ + return BlackForestLabsImageGenerationConfig() diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 0f5ccaf17a..a330eaa7c3 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3106,6 +3106,7 @@ class LlmProviders(str, Enum): GEMINI = "gemini" AI21 = "ai21" BASETEN = "baseten" + BLACK_FOREST_LABS = "black_forest_labs" AZURE = "azure" AZURE_TEXT = "azure_text" AZURE_AI = "azure_ai" diff --git a/litellm/utils.py b/litellm/utils.py index caab91c1c7..a4b4e26aec 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8692,6 +8692,12 @@ class ProviderConfigManager: ) return get_runwayml_image_generation_config(model) + elif LlmProviders.BLACK_FOREST_LABS == provider: + from litellm.llms.black_forest_labs.image_generation import ( + get_black_forest_labs_image_generation_config, + ) + + return get_black_forest_labs_image_generation_config(model) elif LlmProviders.VERTEX_AI == provider: from litellm.llms.vertex_ai.image_generation import ( get_vertex_ai_image_generation_config, @@ -8777,6 +8783,12 @@ class ProviderConfigManager: ) return RecraftImageEditConfig() + elif LlmProviders.BLACK_FOREST_LABS == provider: + from litellm.llms.black_forest_labs.image_edit.transformation import ( + BlackForestLabsImageEditConfig, + ) + + return BlackForestLabsImageEditConfig() elif LlmProviders.AZURE_AI == provider: from litellm.llms.azure_ai.image_edit import get_azure_ai_image_edit_config diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b92e272797..5043f9684b 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -7786,6 +7786,80 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "black_forest_labs/flux-kontext-pro": { + "litellm_provider": "black_forest_labs", + "mode": "image_edit", + "output_cost_per_image": 0.04, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/edits", + "/v1/images/generations" + ] + }, + "black_forest_labs/flux-kontext-max": { + "litellm_provider": "black_forest_labs", + "mode": "image_edit", + "output_cost_per_image": 0.08, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/edits", + "/v1/images/generations" + ] + }, + "black_forest_labs/flux-pro-1.0-fill": { + "litellm_provider": "black_forest_labs", + "mode": "image_edit", + "output_cost_per_image": 0.05, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/edits" + ] + }, + "black_forest_labs/flux-pro-1.0-expand": { + "litellm_provider": "black_forest_labs", + "mode": "image_edit", + "output_cost_per_image": 0.05, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/edits" + ] + }, + "black_forest_labs/flux-pro-1.1": { + "litellm_provider": "black_forest_labs", + "mode": "image_generation", + "output_cost_per_image": 0.04, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, + "black_forest_labs/flux-pro-1.1-ultra": { + "litellm_provider": "black_forest_labs", + "mode": "image_generation", + "output_cost_per_image": 0.06, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, + "black_forest_labs/flux-dev": { + "litellm_provider": "black_forest_labs", + "mode": "image_generation", + "output_cost_per_image": 0.025, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, + "black_forest_labs/flux-pro": { + "litellm_provider": "black_forest_labs", + "mode": "image_generation", + "output_cost_per_image": 0.05, + "source": "https://bfl.ai/pricing", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "cerebras/llama-3.3-70b": { "input_cost_per_token": 8.5e-07, "litellm_provider": "cerebras", diff --git a/tests/test_litellm/llms/black_forest_labs/__init__.py b/tests/test_litellm/llms/black_forest_labs/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/test_litellm/llms/black_forest_labs/image_edit/__init__.py b/tests/test_litellm/llms/black_forest_labs/image_edit/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/test_litellm/llms/black_forest_labs/image_edit/test_bfl_image_edit_transformation.py b/tests/test_litellm/llms/black_forest_labs/image_edit/test_bfl_image_edit_transformation.py new file mode 100644 index 0000000000..7709734e5e --- /dev/null +++ b/tests/test_litellm/llms/black_forest_labs/image_edit/test_bfl_image_edit_transformation.py @@ -0,0 +1,304 @@ +""" +Unit tests for Black Forest Labs image edit transformation functionality. + +Note: Polling tests are now in test_bfl_image_edit_handler.py +since polling logic was moved to the handler. +""" + +import base64 +import json +import os +import sys +import time +from io import BytesIO +from typing import Dict, List +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.black_forest_labs.image_edit.transformation import ( + BlackForestLabsImageEditConfig, +) +from litellm.llms.black_forest_labs.common_utils import BlackForestLabsError +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import ImageObject, ImageResponse + + +class TestBlackForestLabsImageEditTransformation: + """ + Unit tests for Black Forest Labs image edit transformation functionality. + """ + + def setup_method(self): + """Set up test fixtures before each test method.""" + self.config = BlackForestLabsImageEditConfig() + self.model = "flux-kontext-pro" + self.logging_obj = MagicMock() + self.prompt = "Add a red hat to the person in the image" + + def test_get_supported_openai_params(self): + """Test that supported OpenAI params are returned correctly.""" + params = self.config.get_supported_openai_params(self.model) + + # BFL image edit supports BFL-specific params passed through directly + assert isinstance(params, list) + assert len(params) > 0 + assert "seed" in params + assert "output_format" in params + assert "safety_tolerance" in params + + def test_map_openai_params_basic(self): + """Test mapping of OpenAI params to BFL params.""" + optional_params = ImageEditOptionalRequestParams() + + result = self.config.map_openai_params( + image_edit_optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + # Should have default output_format + assert result.get("output_format") == "png" + + def test_map_openai_params_with_bfl_specific(self): + """Test that BFL-specific params are passed through.""" + # BFL-specific params are passed as dict keys + optional_params: ImageEditOptionalRequestParams = { + "seed": 42, + "safety_tolerance": 2, + "aspect_ratio": "16:9", + } + + result = self.config.map_openai_params( + image_edit_optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result.get("seed") == 42 + assert result.get("safety_tolerance") == 2 + assert result.get("aspect_ratio") == "16:9" + assert result.get("output_format") == "png" + + def test_validate_environment_with_api_key(self): + """Test environment validation with provided API key.""" + headers = {} + + result = self.config.validate_environment( + headers=headers, + model=self.model, + api_key="test-api-key", + ) + + assert result["x-key"] == "test-api-key" + assert result["Content-Type"] == "application/json" + assert result["Accept"] == "application/json" + + def test_validate_environment_missing_api_key(self): + """Test that missing API key raises error.""" + headers = {} + + with patch("litellm.llms.black_forest_labs.image_edit.transformation.get_secret_str") as mock_get_secret: + mock_get_secret.return_value = None + + with pytest.raises(BlackForestLabsError) as exc_info: + self.config.validate_environment( + headers=headers, + model=self.model, + api_key=None, + ) + + assert exc_info.value.status_code == 401 + assert "BFL_API_KEY is not set" in exc_info.value.message + + def test_get_model_endpoint_kontext_pro(self): + """Test endpoint resolution for flux-kontext-pro.""" + endpoint = self.config._get_model_endpoint("flux-kontext-pro") + assert endpoint == "/v1/flux-kontext-pro" + + def test_get_model_endpoint_kontext_max(self): + """Test endpoint resolution for flux-kontext-max.""" + endpoint = self.config._get_model_endpoint("flux-kontext-max") + assert endpoint == "/v1/flux-kontext-max" + + def test_get_model_endpoint_with_provider_prefix(self): + """Test endpoint resolution with provider prefix.""" + endpoint = self.config._get_model_endpoint("black_forest_labs/flux-kontext-pro") + assert endpoint == "/v1/flux-kontext-pro" + + def test_get_model_endpoint_fill(self): + """Test endpoint resolution for flux-pro-1.0-fill.""" + endpoint = self.config._get_model_endpoint("flux-pro-1.0-fill") + assert endpoint == "/v1/flux-pro-1.0-fill" + + def test_get_complete_url(self): + """Test complete URL generation.""" + url = self.config.get_complete_url( + model="flux-kontext-pro", + api_base=None, + litellm_params={}, + ) + + assert url == "https://api.bfl.ai/v1/flux-kontext-pro" + + def test_get_complete_url_custom_base(self): + """Test complete URL generation with custom base.""" + url = self.config.get_complete_url( + model="flux-kontext-pro", + api_base="https://custom.api.com/", + litellm_params={}, + ) + + assert url == "https://custom.api.com/v1/flux-kontext-pro" + + def test_transform_image_edit_request(self): + """Test request transformation to BFL format.""" + image_data = b"fake_image_data" + image = BytesIO(image_data) + + image_edit_optional_params = { + "seed": 123, + "output_format": "jpeg", + } + + litellm_params = GenericLiteLLMParams() + headers = {} + + data, files = self.config.transform_image_edit_request( + model=self.model, + prompt=self.prompt, + image=image, + image_edit_optional_request_params=image_edit_optional_params, + litellm_params=litellm_params, + headers=headers, + ) + + # Check that data contains the expected parameters + assert data["prompt"] == self.prompt + assert "input_image" in data + # Verify base64 encoding + decoded = base64.b64decode(data["input_image"]) + assert decoded == image_data + assert data["seed"] == 123 + assert data["output_format"] == "jpeg" + + # BFL uses JSON, not multipart - files should be empty + assert files == [] + + def test_transform_image_edit_request_with_mask(self): + """Test request transformation with mask for inpainting.""" + image_data = b"fake_image_data" + mask_data = b"fake_mask_data" + image = BytesIO(image_data) + + image_edit_optional_params = { + "mask": BytesIO(mask_data), + "output_format": "png", + } + + litellm_params = GenericLiteLLMParams() + headers = {} + + data, files = self.config.transform_image_edit_request( + model="flux-pro-1.0-fill", + prompt=self.prompt, + image=image, + image_edit_optional_request_params=image_edit_optional_params, + litellm_params=litellm_params, + headers=headers, + ) + + # Check mask is base64 encoded + assert "mask" in data + decoded_mask = base64.b64decode(data["mask"]) + assert decoded_mask == mask_data + + def test_read_image_bytes_from_bytes(self): + """Test reading image bytes from bytes input.""" + image_data = b"test_image_bytes" + result = self.config._read_image_bytes(image_data) + assert result == image_data + + def test_read_image_bytes_from_file_like(self): + """Test reading image bytes from file-like object.""" + image_data = b"test_image_bytes" + image = BytesIO(image_data) + result = self.config._read_image_bytes(image) + assert result == image_data + + def test_read_image_bytes_from_list(self): + """Test reading image bytes from list (takes first).""" + image_data = b"test_image_bytes" + images = [BytesIO(image_data), BytesIO(b"other")] + result = self.config._read_image_bytes(images) + assert result == image_data + + def test_transform_image_edit_response_success(self): + """Test response transformation with final polled response.""" + # The response is now the FINAL polled response from handler + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "status": "Ready", + "result": {"sample": "https://example.com/edited_image.png"}, + } + mock_response.status_code = 200 + + result = self.config.transform_image_edit_response( + model=self.model, + raw_response=mock_response, + logging_obj=self.logging_obj, + ) + + assert len(result.data) == 1 + assert result.data[0].url == "https://example.com/edited_image.png" + + def test_transform_image_edit_response_no_image_url(self): + """Test response transformation when no image URL is present.""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "status": "Ready", + "result": {}, + } + mock_response.status_code = 200 + + with pytest.raises(BlackForestLabsError, match="No image URL"): + self.config.transform_image_edit_response( + model=self.model, + raw_response=mock_response, + logging_obj=self.logging_obj, + ) + + def test_transform_image_edit_response_json_parse_error(self): + """Test response transformation with JSON parse error.""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.side_effect = json.JSONDecodeError("error", "doc", 0) + mock_response.status_code = 200 + + with pytest.raises(BlackForestLabsError, match="Error parsing"): + self.config.transform_image_edit_response( + model=self.model, + raw_response=mock_response, + logging_obj=self.logging_obj, + ) + + def test_get_error_class(self): + """Test that get_error_class returns BlackForestLabsError.""" + error = self.config.get_error_class( + error_message="Test error", + status_code=400, + headers={}, + ) + + assert isinstance(error, BlackForestLabsError) + assert error.status_code == 400 + assert "Test error" in str(error.message) + + def test_use_multipart_form_data_returns_false(self): + """Test that use_multipart_form_data returns False for BFL.""" + assert self.config.use_multipart_form_data() is False diff --git a/tests/test_litellm/llms/black_forest_labs/image_generation/__init__.py b/tests/test_litellm/llms/black_forest_labs/image_generation/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/test_litellm/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py b/tests/test_litellm/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py new file mode 100644 index 0000000000..a839983f8e --- /dev/null +++ b/tests/test_litellm/llms/black_forest_labs/image_generation/test_bfl_image_generation_transformation.py @@ -0,0 +1,350 @@ +""" +Unit tests for Black Forest Labs image generation transformation functionality. + +Note: Polling tests are now in test_bfl_image_generation_handler.py +since polling logic was moved to the handler. +""" + +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.black_forest_labs.image_generation.transformation import ( + BlackForestLabsImageGenerationConfig, + get_black_forest_labs_image_generation_config, +) +from litellm.llms.black_forest_labs.common_utils import BlackForestLabsError +from litellm.types.utils import ImageObject, ImageResponse + + +class TestBlackForestLabsImageGenerationTransformation: + """ + Unit tests for Black Forest Labs image generation transformation functionality. + """ + + def setup_method(self): + """Set up test fixtures before each test method.""" + self.config = BlackForestLabsImageGenerationConfig() + self.model = "flux-pro-1.1" + self.logging_obj = MagicMock() + self.prompt = "A beautiful sunset over the ocean" + + def test_get_supported_openai_params(self): + """Test that supported OpenAI params are returned correctly.""" + params = self.config.get_supported_openai_params(self.model) + + assert "n" in params + assert "size" in params + assert "quality" in params + + def test_map_openai_params_basic(self): + """Test mapping of OpenAI params to BFL params.""" + non_default_params = {} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params, optional_params, self.model, drop_params=False + ) + + # Empty input should return empty output + assert result == {} + + def test_map_openai_params_size_mapping(self): + """Test that OpenAI size is mapped to BFL width/height.""" + non_default_params = {"size": "1024x1024"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params, optional_params, self.model, drop_params=False + ) + + assert result["width"] == 1024 + assert result["height"] == 1024 + + def test_map_openai_params_size_custom(self): + """Test custom size parsing.""" + non_default_params = {"size": "800x600"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params, optional_params, self.model, drop_params=False + ) + + assert result["width"] == 800 + assert result["height"] == 600 + + def test_map_openai_params_n_for_ultra(self): + """Test that n is mapped to num_images for ultra model.""" + non_default_params = {"n": 4} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params, optional_params, "flux-pro-1.1-ultra", drop_params=False + ) + + assert result["num_images"] == 4 + + def test_map_openai_params_quality_hd_for_ultra(self): + """Test that 'hd' quality maps to raw=True for ultra model.""" + non_default_params = {"quality": "hd"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params, optional_params, "flux-pro-1.1-ultra", drop_params=False + ) + + assert result["raw"] is True + + def test_map_openai_params_unsupported_raises(self): + """Test that unsupported params raise ValueError when drop_params=False.""" + non_default_params = {"unsupported_param": "value"} + optional_params = {} + + with pytest.raises(ValueError, match="not supported"): + self.config.map_openai_params( + non_default_params, optional_params, self.model, drop_params=False + ) + + def test_map_openai_params_unsupported_dropped(self): + """Test that unsupported params are dropped when drop_params=True.""" + non_default_params = {"unsupported_param": "value"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params, optional_params, self.model, drop_params=True + ) + + assert "unsupported_param" not in result + + def test_validate_environment_with_api_key(self): + """Test that validate_environment sets headers correctly.""" + headers = {} + + result = self.config.validate_environment( + headers=headers, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + api_key="test_api_key", + ) + + assert result["x-key"] == "test_api_key" + assert result["Content-Type"] == "application/json" + + def test_validate_environment_missing_api_key(self): + """Test that validate_environment raises error when API key is missing.""" + headers = {} + + with patch( + "litellm.llms.black_forest_labs.image_generation.transformation.get_secret_str", + return_value=None, + ): + with pytest.raises(BlackForestLabsError, match="BFL_API_KEY"): + self.config.validate_environment( + headers=headers, + model=self.model, + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + def test_get_model_endpoint_flux_pro_1_1(self): + """Test endpoint for flux-pro-1.1 model.""" + endpoint = self.config._get_model_endpoint("flux-pro-1.1") + assert endpoint == "/v1/flux-pro-1.1" + + def test_get_model_endpoint_flux_pro_1_1_ultra(self): + """Test endpoint for flux-pro-1.1-ultra model.""" + endpoint = self.config._get_model_endpoint("flux-pro-1.1-ultra") + assert endpoint == "/v1/flux-pro-1.1-ultra" + + def test_get_model_endpoint_flux_dev(self): + """Test endpoint for flux-dev model.""" + endpoint = self.config._get_model_endpoint("flux-dev") + assert endpoint == "/v1/flux-dev" + + def test_get_model_endpoint_flux_pro(self): + """Test endpoint for flux-pro model.""" + endpoint = self.config._get_model_endpoint("flux-pro") + assert endpoint == "/v1/flux-pro" + + def test_get_model_endpoint_flux_kontext_pro(self): + """Test endpoint for flux-kontext-pro model (supports both generation and editing).""" + endpoint = self.config._get_model_endpoint("flux-kontext-pro") + assert endpoint == "/v1/flux-kontext-pro" + + def test_get_model_endpoint_flux_kontext_max(self): + """Test endpoint for flux-kontext-max model (supports both generation and editing).""" + endpoint = self.config._get_model_endpoint("flux-kontext-max") + assert endpoint == "/v1/flux-kontext-max" + + def test_get_model_endpoint_unknown_raises(self): + """Test that unknown models raise ValueError.""" + with pytest.raises(ValueError, match="Unknown BFL image generation model"): + self.config._get_model_endpoint("unknown-model") + + def test_get_model_endpoint_with_provider_prefix(self): + """Test that provider prefix is stripped from model name.""" + endpoint = self.config._get_model_endpoint("black_forest_labs/flux-pro-1.1") + assert endpoint == "/v1/flux-pro-1.1" + + def test_get_complete_url(self): + """Test URL construction with default base.""" + url = self.config.get_complete_url( + api_base=None, + api_key=None, + model="flux-pro-1.1", + optional_params={}, + litellm_params={}, + ) + + assert "https://api.bfl.ai/v1/flux-pro-1.1" == url + + def test_get_complete_url_custom_base(self): + """Test URL construction with custom base.""" + url = self.config.get_complete_url( + api_base="https://custom.api.com", + api_key=None, + model="flux-pro-1.1", + optional_params={}, + litellm_params={}, + ) + + assert "https://custom.api.com/v1/flux-pro-1.1" == url + + def test_transform_image_generation_request(self): + """Test request body transformation.""" + request = self.config.transform_image_generation_request( + model=self.model, + prompt=self.prompt, + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["prompt"] == self.prompt + assert request["output_format"] == "png" + + def test_transform_image_generation_request_custom_format(self): + """Test request body with custom output format.""" + request = self.config.transform_image_generation_request( + model=self.model, + prompt=self.prompt, + optional_params={"output_format": "jpeg"}, + litellm_params={}, + headers={}, + ) + + assert request["output_format"] == "jpeg" + + def test_transform_image_generation_request_ultra_params(self): + """Test request body with ultra-specific params.""" + request = self.config.transform_image_generation_request( + model="flux-pro-1.1-ultra", + prompt=self.prompt, + optional_params={ + "raw": True, + "num_images": 2, + "aspect_ratio": "16:9", + }, + litellm_params={}, + headers={}, + ) + + assert request["raw"] is True + assert request["num_images"] == 2 + assert request["aspect_ratio"] == "16:9" + + def test_transform_image_generation_response_success(self): + """Test response transformation with final polled response.""" + # The response is now the FINAL polled response from handler + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "status": "Ready", + "result": {"sample": "https://example.com/image.png"}, + } + mock_response.status_code = 200 + + model_response = ImageResponse(created=0, data=[]) + + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + ) + + assert len(result.data) == 1 + assert result.data[0].url == "https://example.com/image.png" + + def test_transform_image_generation_response_multiple_images(self): + """Test response transformation with multiple images.""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "status": "Ready", + "result": [ + "https://example.com/image1.png", + "https://example.com/image2.png", + ], + } + mock_response.status_code = 200 + + model_response = ImageResponse(created=0, data=[]) + + result = self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + ) + + assert len(result.data) == 2 + assert result.data[0].url == "https://example.com/image1.png" + assert result.data[1].url == "https://example.com/image2.png" + + def test_transform_image_generation_response_no_image(self): + """Test response transformation when no image URL is present.""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "status": "Ready", + "result": {}, + } + mock_response.status_code = 200 + + model_response = ImageResponse(created=0, data=[]) + + with pytest.raises(BlackForestLabsError, match="No image URL"): + self.config.transform_image_generation_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=self.logging_obj, + ) + + def test_get_error_class(self): + """Test that get_error_class returns BlackForestLabsError.""" + error = self.config.get_error_class( + error_message="Test error", + status_code=400, + headers={}, + ) + + assert isinstance(error, BlackForestLabsError) + assert error.status_code == 400 + assert "Test error" in str(error.message) + + def test_get_black_forest_labs_image_generation_config(self): + """Test the factory function.""" + config = get_black_forest_labs_image_generation_config("flux-pro-1.1") + + assert isinstance(config, BlackForestLabsImageGenerationConfig)