Merge remote-tracking branch 'origin' into litellm_org_usage

This commit is contained in:
yuneng-jiang
2025-11-12 18:22:44 -08:00
21 changed files with 945 additions and 124 deletions
@@ -0,0 +1,198 @@
# RunwayML - Image Generation
## Overview
| Property | Details |
|-------|-------|
| Description | RunwayML provides advanced AI-powered image generation with high-quality results |
| Provider Route on LiteLLM | `runwayml/` |
| Supported Operations | [`/images/generations`](#quick-start) |
| Link to Provider Doc | [RunwayML API ↗](https://docs.dev.runwayml.com/) |
LiteLLM supports RunwayML's Gen-4 image generation API, allowing you to generate high-quality images from text prompts.
## Quick Start
```python showLineNumbers title="Basic Image Generation"
from litellm import image_generation
import os
os.environ["RUNWAYML_API_KEY"] = "your-api-key"
response = image_generation(
model="runwayml/gen4_image",
prompt="A serene mountain landscape at sunset",
size="1920x1080"
)
print(response.data[0].url)
```
## Authentication
Set your RunwayML API key:
```python showLineNumbers title="Set API Key"
import os
os.environ["RUNWAYML_API_KEY"] = "your-api-key"
```
## Supported Parameters
| Parameter | Type | Required | Description |
|-----------|------|----------|-------------|
| `model` | string | Yes | Model to use (e.g., `runwayml/gen4_image`) |
| `prompt` | string | Yes | Text description for the image |
| `size` | string | No | Image dimensions (default: `1920x1080`) |
### Supported Sizes
- `1024x1024`
- `1792x1024`
- `1024x1792`
- `1920x1080` (default)
- `1080x1920`
## Async Usage
```python showLineNumbers title="Async Image Generation"
from litellm import aimage_generation
import os
import asyncio
os.environ["RUNWAYML_API_KEY"] = "your-api-key"
async def generate_image():
response = await aimage_generation(
model="runwayml/gen4_image",
prompt="A futuristic city skyline at night",
size="1920x1080"
)
print(response.data[0].url)
asyncio.run(generate_image())
```
## LiteLLM Proxy Usage
Add RunwayML to your proxy configuration:
```yaml showLineNumbers title="config.yaml"
model_list:
- model_name: gen4-image
litellm_params:
model: runwayml/gen4_image
api_key: os.environ/RUNWAYML_API_KEY
```
Start the proxy:
```bash
litellm --config /path/to/config.yaml
```
Generate images through the proxy:
```bash showLineNumbers title="Proxy Request"
curl --location 'http://localhost:4000/v1/images/generations' \
--header 'Content-Type: application/json' \
--header 'x-litellm-api-key: sk-1234' \
--data '{
"model": "runwayml/gen4_image",
"prompt": "A serene mountain landscape at sunset",
"size": "1920x1080"
}'
```
## Supported Models
| Model | Description | Default Size |
|-------|-------------|--------------|
| `runwayml/gen4_image` | High-quality image generation | 1920x1080 |
## Cost Tracking
LiteLLM automatically tracks RunwayML image generation costs:
```python showLineNumbers title="Cost Tracking"
from litellm import image_generation, completion_cost
response = image_generation(
model="runwayml/gen4_image",
prompt="A serene mountain landscape at sunset",
size="1920x1080"
)
cost = completion_cost(completion_response=response)
print(f"Image generation cost: ${cost}")
```
## Supported Features
| Feature | Supported |
|---------|-----------|
| Image Generation | ✅ |
| Cost Tracking | ✅ |
| Logging | ✅ |
| Fallbacks | ✅ |
| Load Balancing | ✅ |
## How It Works
RunwayML uses an asynchronous task-based API pattern. LiteLLM handles the polling and response transformation automatically.
### Complete Flow Diagram
```mermaid
sequenceDiagram
participant Client
box rgb(200, 220, 255) LiteLLM AI Gateway
participant LiteLLM
end
participant RunwayML as RunwayML API
Client->>LiteLLM: POST /images/generations (OpenAI format)
Note over LiteLLM: Transform to RunwayML format
LiteLLM->>RunwayML: POST v1/text_to_image
RunwayML-->>LiteLLM: 200 OK + task ID
Note over LiteLLM: Automatic Polling
loop Every 2 seconds
LiteLLM->>RunwayML: GET v1/tasks/{task_id}
RunwayML-->>LiteLLM: Status: RUNNING
end
LiteLLM->>RunwayML: GET v1/tasks/{task_id}
RunwayML-->>LiteLLM: Status: SUCCEEDED + image URL
Note over LiteLLM: Transform to OpenAI format
LiteLLM-->>Client: Image Response (OpenAI format)
```
### What LiteLLM Does For You
When you call `litellm.image_generation()` or `/v1/images/generations`:
1. **Request Transformation**: Converts OpenAI image generation format → RunwayML format
2. **Submits Task**: Sends transformed request to RunwayML API
3. **Receives Task ID**: Captures the task ID from the initial response
4. **Automatic Polling**:
- Polls the task status endpoint every 2 seconds
- Continues until status is `SUCCEEDED` or `FAILED`
- Default timeout: 10 minutes (configurable via `RUNWAYML_POLLING_TIMEOUT`)
5. **Response Transformation**: Converts RunwayML format → OpenAI format
6. **Returns Result**: Sends unified OpenAI format response to client
**Polling Configuration:**
- Default timeout: 600 seconds (10 minutes)
- Configurable via `RUNWAYML_POLLING_TIMEOUT` environment variable
- Uses sync (`time.sleep()`) or async (`await asyncio.sleep()`) based on call type
:::info
**Typical processing time**: 10-30 seconds depending on image size and complexity
:::
-9
View File
@@ -1,9 +0,0 @@
# Demo App
Here is a demo of the proxy. To log in pass in:
- Username: admin
- Password: sk-1234
[Demo UI](https://demo.litellm.ai/ui)
+86 -86
View File
@@ -28,7 +28,7 @@ fallbacks=[{"gpt-3.5-turbo": ["gpt-4"]}]
```python
from litellm import Router
router = Router(
model_list=[
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
@@ -47,8 +47,8 @@ router = Router(
"rpm": 6
}
}
],
fallbacks=[{"gpt-3.5-turbo": ["gpt-4"]}] # 👈 KEY CHANGE
],
fallbacks=[{"gpt-3.5-turbo": ["gpt-4"]}] # 👈 KEY CHANGE
)
```
@@ -104,9 +104,9 @@ model_list = [{..}, {..}] # defined in Step 1.
router = Router(model_list=model_list, fallbacks=[{"bad-model": ["my-good-model"]}])
response = router.completion(
model="bad-model",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
mock_testing_fallbacks=True,
model="bad-model",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
mock_testing_fallbacks=True,
)
```
@@ -431,32 +431,32 @@ content_policy_fallbacks=[{"claude-2": ["my-fallback-model"]}]
from litellm import Router
router = Router(
model_list=[
{
"model_name": "claude-2",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": Exception("content filtering policy"),
},
},
{
"model_name": "my-fallback-model",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": "This works!",
},
},
],
content_policy_fallbacks=[{"claude-2": ["my-fallback-model"]}], # 👈 KEY CHANGE
# fallbacks=[..], # [OPTIONAL]
# context_window_fallbacks=[..], # [OPTIONAL]
model_list=[
{
"model_name": "claude-2",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": Exception("content filtering policy"),
},
},
{
"model_name": "my-fallback-model",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": "This works!",
},
},
],
content_policy_fallbacks=[{"claude-2": ["my-fallback-model"]}], # 👈 KEY CHANGE
# fallbacks=[..], # [OPTIONAL]
# context_window_fallbacks=[..], # [OPTIONAL]
)
response = router.completion(
model="claude-2",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
model="claude-2",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
)
```
</TabItem>
@@ -466,7 +466,7 @@ In your proxy config.yaml just add this line 👇
```yaml
router_settings:
content_policy_fallbacks=[{"claude-2": ["my-fallback-model"]}]
content_policy_fallbacks=[{"claude-2": ["my-fallback-model"]}]
```
Start proxy
@@ -495,32 +495,32 @@ context_window_fallbacks=[{"claude-2": ["my-fallback-model"]}]
from litellm import Router
router = Router(
model_list=[
{
"model_name": "claude-2",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": Exception("prompt is too long"),
},
},
{
"model_name": "my-fallback-model",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": "This works!",
},
},
],
context_window_fallbacks=[{"claude-2": ["my-fallback-model"]}], # 👈 KEY CHANGE
# fallbacks=[..], # [OPTIONAL]
# content_policy_fallbacks=[..], # [OPTIONAL]
model_list=[
{
"model_name": "claude-2",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": Exception("prompt is too long"),
},
},
{
"model_name": "my-fallback-model",
"litellm_params": {
"model": "claude-2",
"api_key": "",
"mock_response": "This works!",
},
},
],
context_window_fallbacks=[{"claude-2": ["my-fallback-model"]}], # 👈 KEY CHANGE
# fallbacks=[..], # [OPTIONAL]
# content_policy_fallbacks=[..], # [OPTIONAL]
)
response = router.completion(
model="claude-2",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
model="claude-2",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
)
```
</TabItem>
@@ -530,7 +530,7 @@ In your proxy config.yaml just add this line 👇
```yaml
router_settings:
context_window_fallbacks=[{"claude-2": ["my-fallback-model"]}]
context_window_fallbacks=[{"claude-2": ["my-fallback-model"]}]
```
Start proxy
@@ -725,22 +725,22 @@ Filter older instances of a model (e.g. gpt-3.5-turbo) with smaller context wind
```yaml
router_settings:
enable_pre_call_checks: true # 1. Enable pre-call checks
enable_pre_call_checks: true # 1. Enable pre-call checks
model_list:
- model_name: gpt-3.5-turbo
litellm_params:
model: azure/chatgpt-v-2
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2023-07-01-preview"
model_info:
base_model: azure/gpt-4-1106-preview # 2. 👈 (azure-only) SET BASE MODEL
- model_name: gpt-3.5-turbo
litellm_params:
model: gpt-3.5-turbo-1106
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-3.5-turbo
litellm_params:
model: azure/chatgpt-v-2
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2023-07-01-preview"
model_info:
base_model: azure/gpt-4-1106-preview # 2. 👈 (azure-only) SET BASE MODEL
- model_name: gpt-3.5-turbo
litellm_params:
model: gpt-3.5-turbo-1106
api_key: os.environ/OPENAI_API_KEY
```
**2. Start proxy**
@@ -766,8 +766,8 @@ text = "What is the meaning of 42?" * 5000
response = client.chat.completions.create(
model="gpt-3.5-turbo",
messages = [
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
],
)
@@ -782,20 +782,20 @@ Fallback to larger models if current model is too small.
```yaml
router_settings:
enable_pre_call_checks: true # 1. Enable pre-call checks
enable_pre_call_checks: true # 1. Enable pre-call checks
model_list:
- model_name: gpt-3.5-turbo-small
litellm_params:
model: azure/chatgpt-v-2
- model_name: gpt-3.5-turbo-small
litellm_params:
model: azure/chatgpt-v-2
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2023-07-01-preview"
model_info:
base_model: azure/gpt-4-1106-preview # 2. 👈 (azure-only) SET BASE MODEL
- model_name: gpt-3.5-turbo-large
litellm_params:
- model_name: gpt-3.5-turbo-large
litellm_params:
model: gpt-3.5-turbo-1106
api_key: os.environ/OPENAI_API_KEY
@@ -831,8 +831,8 @@ text = "What is the meaning of 42?" * 5000
response = client.chat.completions.create(
model="gpt-3.5-turbo",
messages = [
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
],
)
@@ -849,9 +849,9 @@ Fallback across providers (e.g. from Azure OpenAI to Anthropic) if you hit conte
```yaml
model_list:
- model_name: gpt-3.5-turbo-small
litellm_params:
model: azure/chatgpt-v-2
- model_name: gpt-3.5-turbo-small
litellm_params:
model: azure/chatgpt-v-2
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2023-07-01-preview"
@@ -874,9 +874,9 @@ You can also set default_fallbacks, in case a specific model group is misconfigu
```yaml
model_list:
- model_name: gpt-3.5-turbo-small
litellm_params:
model: azure/chatgpt-v-2
- model_name: gpt-3.5-turbo-small
litellm_params:
model: azure/chatgpt-v-2
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2023-07-01-preview"
@@ -906,7 +906,7 @@ Set 'region_name' of deployment.
```yaml
router_settings:
enable_pre_call_checks: true # 1. Enable pre-call checks
enable_pre_call_checks: true # 1. Enable pre-call checks
model_list:
- model_name: gpt-3.5-turbo
+6 -1
View File
@@ -130,7 +130,11 @@ const sidebars = {
"proxy/release_cycle",
],
},
"proxy/demo",
{
"type": "link",
"label": "Demo LiteLLM Cloud",
"href": "https://www.litellm.ai/cloud"
},
{
type: "category",
label: "Admin UI",
@@ -581,6 +585,7 @@ const sidebars = {
type: "category",
label: "RunwayML",
items: [
"providers/runwayml/images",
"providers/runwayml/videos",
]
},
+1
View File
@@ -86,6 +86,7 @@ MAX_TOKEN_TRIMMING_ATTEMPTS = int(
) # Maximum number of attempts to trim the message
RUNWAYML_DEFAULT_API_VERSION = str(os.getenv("RUNWAYML_DEFAULT_API_VERSION", "2024-11-06"))
RUNWAYML_POLLING_TIMEOUT = int(os.getenv("RUNWAYML_POLLING_TIMEOUT", 600)) # 10 minutes default for image generation
########## Networking constants ##############################################################
_DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client for 1 hour
+1
View File
@@ -343,6 +343,7 @@ def image_generation( # noqa: PLR0915
litellm.LlmProviders.AIML,
litellm.LlmProviders.GEMINI,
litellm.LlmProviders.FAL_AI,
litellm.LlmProviders.RUNWAYML,
):
if image_generation_config is None:
raise ValueError(
@@ -735,6 +735,15 @@ class CostCalculatorUtils:
model=model,
image_response=completion_response,
)
elif custom_llm_provider == litellm.LlmProviders.RUNWAYML.value:
from litellm.llms.runwayml.cost_calculator import (
cost_calculator as runwayml_image_cost_calculator,
)
return runwayml_image_cost_calculator(
model=model,
image_response=completion_response,
)
else:
return default_image_cost_calculator(
model=model,
-2
View File
@@ -1,2 +0,0 @@
# RunwayML integration for LiteLLM
+6
View File
@@ -0,0 +1,6 @@
# RunwayML integration for LiteLLM
from .cost_calculator import cost_calculator
from .videos.transformation import RunwayMLVideoConfig
__all__ = ["RunwayMLVideoConfig", "cost_calculator"]
+31
View File
@@ -0,0 +1,31 @@
from typing import Any
import litellm
from litellm.types.utils import ImageResponse
def cost_calculator(
model: str,
image_response: Any,
) -> float:
"""
RunwayML image generation cost calculator.
RunwayML charges per image generated, not per pixel.
Pricing is stored in model_prices_and_context_window.json with output_cost_per_image.
"""
_model_info = litellm.get_model_info(
model=model,
custom_llm_provider=litellm.LlmProviders.RUNWAYML.value,
)
output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0
num_images: int = 0
if isinstance(image_response, ImageResponse):
if image_response.data:
num_images = len(image_response.data)
return output_cost_per_image * num_images
else:
raise ValueError(
f"image_response must be of type ImageResponse, got type={type(image_response)}"
)
@@ -0,0 +1,13 @@
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from .transformation import RunwayMLImageGenerationConfig
__all__ = [
"RunwayMLImageGenerationConfig",
]
def get_runwayml_image_generation_config(model: str) -> BaseImageGenerationConfig:
return RunwayMLImageGenerationConfig()
@@ -0,0 +1,513 @@
import asyncio
import time
from typing import TYPE_CHECKING, Any, Dict, List, Optional
import httpx
from litellm._logging import verbose_logger
from litellm.constants import (
RUNWAYML_DEFAULT_API_VERSION,
RUNWAYML_POLLING_TIMEOUT,
)
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIImageGenerationOptionalParams,
)
from litellm.types.utils import ImageObject, ImageResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class RunwayMLImageGenerationConfig(BaseImageGenerationConfig):
"""
Configuration for RunwayML image generation models.
"""
DEFAULT_BASE_URL: str = "https://api.dev.runwayml.com"
IMAGE_GENERATION_ENDPOINT: str = "v1/text_to_image"
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""
Get the complete url for the request
Some providers need `model` in `api_base`
"""
complete_url: str = (
api_base
or get_secret_str("RUNWAYML_API_BASE")
or self.DEFAULT_BASE_URL
)
complete_url = complete_url.rstrip("/")
if self.IMAGE_GENERATION_ENDPOINT:
complete_url = f"{complete_url}/{self.IMAGE_GENERATION_ENDPOINT}"
return complete_url
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
final_api_key: Optional[str] = (
api_key or
get_secret_str("RUNWAYML_API_SECRET") or
get_secret_str("RUNWAYML_API_KEY")
)
if not final_api_key:
raise ValueError("RUNWAYML_API_SECRET or RUNWAYML_API_KEY is not set")
headers["Authorization"] = f"Bearer {final_api_key}"
headers["X-Runway-Version"] = RUNWAYML_DEFAULT_API_VERSION
return headers
@staticmethod
def _transform_runwayml_response_to_openai(
response_data: Dict[str, Any],
model_response: ImageResponse,
) -> ImageResponse:
"""
Transform RunwayML response format to OpenAI ImageResponse format.
RunwayML response format (after polling):
{
"id": "task_123...",
"status": "SUCCEEDED",
"output": ["https://cloudfront.net/.../image.png"],
"completedAt": "2025-11-13T..."
}
OpenAI ImageResponse format:
{
"data": [
{
"url": "https://cloudfront.net/.../image.png",
"b64_json": null
}
]
}
Args:
response_data: JSON response from RunwayML (after polling completes)
model_response: ImageResponse object to populate
Returns:
Populated ImageResponse in OpenAI format
"""
if not model_response.data:
model_response.data = []
# Handle RunwayML response format
# Response contains task.output with image URL(s)
output = response_data.get("output", [])
if isinstance(output, list):
for image_item in output:
if isinstance(image_item, str):
# If output is a list of URL strings
model_response.data.append(ImageObject(
url=image_item,
b64_json=None,
))
elif isinstance(image_item, dict):
# If output contains dict with url/b64_json
model_response.data.append(ImageObject(
url=image_item.get("url", None),
b64_json=image_item.get("b64_json", None),
))
return model_response
@staticmethod
def _check_timeout(start_time: float, timeout_secs: float) -> None:
"""
Check if operation has timed out.
Args:
start_time: Start time of the operation
timeout_secs: Timeout duration in seconds
Raises:
TimeoutError: If operation has exceeded timeout
"""
if time.time() - start_time > timeout_secs:
raise TimeoutError(
f"RunwayML task polling timed out after {timeout_secs} seconds"
)
@staticmethod
def _check_task_status(response_data: Dict[str, Any]) -> str:
"""
Check RunwayML task status from response.
RunwayML statuses: PENDING, RUNNING, SUCCEEDED, FAILED, CANCELLED, THROTTLED
Args:
response_data: JSON response from RunwayML task endpoint
Returns:
Normalized status string: "running", "succeeded", or raises on failure
Raises:
ValueError: If task failed or status is unknown
"""
status = response_data.get("status", "").upper()
verbose_logger.debug(f"RunwayML task status: {status}")
if status == "SUCCEEDED":
return "succeeded"
elif status == "FAILED":
failure_reason = response_data.get("failure", "Unknown error")
failure_code = response_data.get("failureCode", "unknown")
raise ValueError(
f"RunwayML image generation failed: {failure_reason} (code: {failure_code})"
)
elif status == "CANCELLED":
raise ValueError("RunwayML image generation was cancelled")
elif status in ["PENDING", "RUNNING", "THROTTLED"]:
return "running"
else:
raise ValueError(f"Unknown RunwayML task status: {status}")
def _poll_task_sync(
self,
task_id: str,
api_base: str,
headers: Dict[str, str],
timeout_secs: float = 600,
) -> httpx.Response:
"""
Poll RunwayML task until completion (sync).
RunwayML POST returns immediately with a task that has status PENDING/RUNNING.
We need to poll GET /v1/tasks/{task_id} until status is SUCCEEDED or FAILED.
Args:
task_id: The task ID to poll
api_base: Base URL for RunwayML API
headers: Request headers (including auth)
timeout_secs: Total timeout in seconds (default: 600s = 10 minutes)
Returns:
Final response with completed task
"""
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
client = _get_httpx_client()
start_time = time.time()
# Build task status URL
api_base = api_base.rstrip("/")
task_url = f"{api_base}/v1/tasks/{task_id}"
verbose_logger.debug(f"Polling RunwayML task: {task_url}")
while True:
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
# Poll the task status
response = client.get(url=task_url, headers=headers)
response.raise_for_status()
response_data = response.json()
# Check task status
status = self._check_task_status(response_data=response_data)
if status == "succeeded":
return response
elif status == "running":
# Wait before polling again (RunwayML recommends 1-2 second intervals)
time.sleep(2)
async def _poll_task_async(
self,
task_id: str,
api_base: str,
headers: Dict[str, str],
timeout_secs: float = 600,
) -> httpx.Response:
"""
Poll RunwayML task until completion (async).
Args:
task_id: The task ID to poll
api_base: Base URL for RunwayML API
headers: Request headers (including auth)
timeout_secs: Total timeout in seconds (default: 600s = 10 minutes)
Returns:
Final response with completed task
"""
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
client = get_async_httpx_client(llm_provider=litellm.LlmProviders.RUNWAYML)
start_time = time.time()
# Build task status URL
api_base = api_base.rstrip("/")
task_url = f"{api_base}/v1/tasks/{task_id}"
verbose_logger.debug(f"Polling RunwayML task (async): {task_url}")
while True:
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
# Poll the task status
response = await client.get(url=task_url, headers=headers)
response.raise_for_status()
response_data = response.json()
# Check task status
status = self._check_task_status(response_data=response_data)
if status == "succeeded":
return response
elif status == "running":
# Wait before polling again (RunwayML recommends 1-2 second intervals)
await asyncio.sleep(2)
def transform_image_generation_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ImageResponse:
"""
Transform the image generation response to the litellm image response.
RunwayML returns a task immediately with status PENDING/RUNNING.
We need to poll the task until it completes (status SUCCEEDED).
Initial response:
{
"id": "task_123...",
"status": "PENDING" | "RUNNING",
"createdAt": "2025-11-13T..."
}
After polling:
{
"id": "task_123...",
"status": "SUCCEEDED",
"output": ["https://cloudfront.net/.../image.png"],
"completedAt": "2025-11-13T..."
}
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
verbose_logger.debug(
"RunwayML starting polling..."
)
# Get task ID
task_id = response_data.get("id")
if not task_id:
raise ValueError("RunwayML response missing task ID")
# Get headers for polling (need auth)
poll_headers = {
"Authorization": raw_response.request.headers.get("Authorization", ""),
"X-Runway-Version": raw_response.request.headers.get("X-Runway-Version", RUNWAYML_DEFAULT_API_VERSION),
}
# Poll until task completes
raw_response = self._poll_task_sync(
task_id=task_id,
api_base=self.DEFAULT_BASE_URL,
headers=poll_headers,
timeout_secs=RUNWAYML_POLLING_TIMEOUT,
)
# Update response_data with polled result
response_data = raw_response.json()
verbose_logger.debug("RunwayML polling complete, transforming to OpenAI format")
# Transform RunwayML response to OpenAI format
return self._transform_runwayml_response_to_openai(
response_data=response_data,
model_response=model_response,
)
async def async_transform_image_generation_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ImageResponse:
"""
Async transform the image generation response to the litellm image response.
RunwayML returns a task immediately with status PENDING/RUNNING.
We need to poll the task until it completes (status SUCCEEDED) using async polling.
"""
try:
response_data = raw_response.json()
except Exception as e:
raise self.get_error_class(
error_message=f"Error transforming image generation response: {e}",
status_code=raw_response.status_code,
headers=raw_response.headers,
)
verbose_logger.debug(
"RunwayML starting polling (async)..."
)
# Get task ID
task_id = response_data.get("id")
if not task_id:
raise ValueError("RunwayML response missing task ID")
# Get headers for polling (need auth)
poll_headers = {
"Authorization": raw_response.request.headers.get("Authorization", ""),
"X-Runway-Version": raw_response.request.headers.get("X-Runway-Version", RUNWAYML_DEFAULT_API_VERSION),
}
# Poll until task completes (async)
raw_response = await self._poll_task_async(
task_id=task_id,
api_base=self.DEFAULT_BASE_URL,
headers=poll_headers,
timeout_secs=RUNWAYML_POLLING_TIMEOUT,
)
# Update response_data with polled result
response_data = raw_response.json()
verbose_logger.debug("RunwayML polling complete (async), transforming to OpenAI format")
# Transform RunwayML response to OpenAI format
return self._transform_runwayml_response_to_openai(
response_data=response_data,
model_response=model_response,
)
def get_supported_openai_params(
self, model: str
) -> List[OpenAIImageGenerationOptionalParams]:
"""
Get supported OpenAI parameters for RunwayML image generation
"""
return [
"size",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_params = self.get_supported_openai_params(model)
# Map OpenAI 'size' parameter to RunwayML 'ratio' parameter
if "size" in non_default_params:
size = non_default_params["size"]
# Map common OpenAI sizes to RunwayML ratios
size_to_ratio_map = {
"1024x1024": "1024:1024",
"1792x1024": "1792:1024",
"1024x1792": "1024:1792",
"1920x1080": "1920:1080",
"1080x1920": "1080:1920",
}
optional_params["ratio"] = size_to_ratio_map.get(size, "1920:1080")
for k in non_default_params.keys():
if k not in optional_params.keys():
if k in supported_params:
optional_params[k] = non_default_params[k]
elif drop_params:
pass
else:
raise ValueError(
f"Parameter {k} is not supported for model {model}. Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters."
)
return optional_params
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the image generation request to the RunwayML image generation request body
RunwayML expects:
- model: The model to use (e.g., 'gen4_image')
- promptText: The text prompt
- ratio: The aspect ratio (e.g., '1920:1080', '1080:1920', '1024:1024')
"""
runwayml_request_body = {
"model": model or "gen4_image",
"promptText": prompt,
}
# Add any RunwayML-specific parameters
if "ratio" in optional_params:
runwayml_request_body["ratio"] = optional_params["ratio"]
else:
# Set default ratio if not provided
runwayml_request_body["ratio"] = "1920:1080"
# Add any other optional parameters
for k, v in optional_params.items():
if k not in runwayml_request_body and k not in ["size"]:
runwayml_request_body[k] = v
return runwayml_request_body
+7 -1
View File
@@ -7635,6 +7635,12 @@ class ProviderConfigManager:
)
return get_fal_ai_image_generation_config(model)
elif LlmProviders.RUNWAYML == provider:
from litellm.llms.runwayml.image_generation import (
get_runwayml_image_generation_config,
)
return get_runwayml_image_generation_config(model)
return None
@staticmethod
@@ -7661,7 +7667,7 @@ class ProviderConfigManager:
return VertexAIVideoConfig()
elif LlmProviders.RUNWAYML == provider:
from litellm.llms.runway.videos.transformation import RunwayMLVideoConfig
from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig
return RunwayMLVideoConfig()
return None
+1 -1
View File
@@ -1366,7 +1366,7 @@
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"image_generations": true,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
@@ -175,6 +175,10 @@ class TestGoogleImageGen(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "gemini/imagen-4.0-generate-001"}
class TestRunwaymlImageGeneration(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "runwayml/gen4_image"}
class TestAzureOpenAIDalle3(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
@@ -1935,3 +1935,43 @@ async def test_asearch_with_fallbacks_helper_missing_search_provider():
original_generic_function=mock_original_function,
query="test query"
)
def test_get_first_default_fallback():
"""Test _get_first_default_fallback method"""
# Test with default fallback ("*")
model_list = [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "gpt-3.5-turbo", "api_key": "fake-key"},
}
]
router = Router(
model_list=model_list,
fallbacks=[{"*": ["gpt-3.5-turbo"]}]
)
result = router._get_first_default_fallback()
assert result == "gpt-3.5-turbo"
# Test with no fallbacks
router_no_fallbacks = Router(model_list=model_list)
result = router_no_fallbacks._get_first_default_fallback()
assert result is None
# Test with fallbacks but no default
router_no_default = Router(
model_list=model_list,
fallbacks=[{"gpt-4": ["gpt-3.5-turbo"]}]
)
result = router_no_default._get_first_default_fallback()
assert result is None
# Test with empty default list
router_empty_list = Router(
model_list=model_list,
fallbacks=[{"*": []}]
)
result = router_empty_list._get_first_default_fallback()
assert result is None
@@ -6,7 +6,7 @@ from unittest.mock import Mock
import httpx
import pytest
from litellm.llms.runway.videos.transformation import RunwayMLVideoConfig
from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.main import VideoObject
@@ -83,7 +83,7 @@ describe("SSOModals", () => {
fireEvent.change(emailInput, { target: { value: "test@example.com" } });
// Fill in an invalid URL
const urlInput = getByLabelText("PROXY BASE URL");
const urlInput = getByLabelText("Proxy Base URL");
fireEvent.change(urlInput, { target: { value: "invalid-url" } });
// Submit the form
@@ -137,7 +137,7 @@ describe("SSOModals", () => {
fireEvent.change(emailInput, { target: { value: "test@example.com" } });
// Fill in a URL with trailing slash
const urlInput = getByLabelText("PROXY BASE URL") as HTMLInputElement;
const urlInput = getByLabelText("Proxy Base URL") as HTMLInputElement;
fireEvent.change(urlInput, { target: { value: "https://example.com/" } });
// Submit the form
@@ -171,7 +171,7 @@ describe("SSOModals", () => {
const { getByLabelText } = render(<TestWrapper />);
const urlInput = getByLabelText("PROXY BASE URL") as HTMLInputElement;
const urlInput = getByLabelText("Proxy Base URL") as HTMLInputElement;
// Simulate user typing "https://"
fireEvent.change(urlInput, { target: { value: "h" } });
@@ -237,7 +237,7 @@ describe("SSOModals", () => {
fireEvent.change(emailInput, { target: { value: "test@example.com" } });
// Fill in an incomplete URL like "http:"
const urlInput = getByLabelText("PROXY BASE URL");
const urlInput = getByLabelText("Proxy Base URL");
fireEvent.change(urlInput, { target: { value: "http:" } });
// Submit the form
@@ -43,8 +43,8 @@ const ssoProviderConfigs: Record<string, SSOProviderConfig> = {
google_client_secret: "GOOGLE_CLIENT_SECRET",
},
fields: [
{ label: "GOOGLE CLIENT ID", name: "google_client_id" },
{ label: "GOOGLE CLIENT SECRET", name: "google_client_secret" },
{ label: "Google Client ID", name: "google_client_id" },
{ label: "Google Client Secret", name: "google_client_secret" },
],
},
microsoft: {
@@ -54,9 +54,9 @@ const ssoProviderConfigs: Record<string, SSOProviderConfig> = {
microsoft_tenant: "MICROSOFT_TENANT",
},
fields: [
{ label: "MICROSOFT CLIENT ID", name: "microsoft_client_id" },
{ label: "MICROSOFT CLIENT SECRET", name: "microsoft_client_secret" },
{ label: "MICROSOFT TENANT", name: "microsoft_tenant" },
{ label: "Microsoft Client ID", name: "microsoft_client_id" },
{ label: "Microsoft Client Secret", name: "microsoft_client_secret" },
{ label: "Microsoft Tenant", name: "microsoft_tenant" },
],
},
okta: {
@@ -68,18 +68,18 @@ const ssoProviderConfigs: Record<string, SSOProviderConfig> = {
generic_userinfo_endpoint: "GENERIC_USERINFO_ENDPOINT",
},
fields: [
{ label: "GENERIC CLIENT ID", name: "generic_client_id" },
{ label: "GENERIC CLIENT SECRET", name: "generic_client_secret" },
{ label: "Generic Client ID", name: "generic_client_id" },
{ label: "Generic Client Secret", name: "generic_client_secret" },
{
label: "AUTHORIZATION ENDPOINT",
label: "Authorization Endpoint",
name: "generic_authorization_endpoint",
placeholder: "https://your-okta-domain/authorize",
placeholder: "https://your-domain/authorize",
},
{ label: "TOKEN ENDPOINT", name: "generic_token_endpoint", placeholder: "https://your-okta-domain/token" },
{ label: "Token Endpoint", name: "generic_token_endpoint", placeholder: "https://your-domain/token" },
{
label: "USERINFO ENDPOINT",
label: "Userinfo Endpoint",
name: "generic_userinfo_endpoint",
placeholder: "https://your-okta-domain/userinfo",
placeholder: "https://your-domain/userinfo",
},
],
},
@@ -92,11 +92,11 @@ const ssoProviderConfigs: Record<string, SSOProviderConfig> = {
generic_userinfo_endpoint: "GENERIC_USERINFO_ENDPOINT",
},
fields: [
{ label: "GENERIC CLIENT ID", name: "generic_client_id" },
{ label: "GENERIC CLIENT SECRET", name: "generic_client_secret" },
{ label: "AUTHORIZATION ENDPOINT", name: "generic_authorization_endpoint" },
{ label: "TOKEN ENDPOINT", name: "generic_token_endpoint" },
{ label: "USERINFO ENDPOINT", name: "generic_userinfo_endpoint" },
{ label: "Generic Client ID", name: "generic_client_id" },
{ label: "Generic Client Secret", name: "generic_client_secret" },
{ label: "Authorization Endpoint", name: "generic_authorization_endpoint" },
{ label: "Token Endpoint", name: "generic_token_endpoint" },
{ label: "Userinfo Endpoint", name: "generic_userinfo_endpoint" },
],
},
};
@@ -282,7 +282,12 @@ const SSOModals: React.FC<SSOModalsProps> = ({
style={{ height: 24, width: 24, marginRight: 12, objectFit: "contain" }}
/>
)}
<span>{value.charAt(0).toUpperCase() + value.slice(1)} SSO</span>
<span>
{value.toLowerCase() === "okta"
? "Okta / Auth0"
: value.charAt(0).toUpperCase() + value.slice(1)}{" "}
SSO
</span>
</div>
</Select.Option>
))}
@@ -307,7 +312,7 @@ const SSOModals: React.FC<SSOModalsProps> = ({
<TextInput />
</Form.Item>
<Form.Item
label="PROXY BASE URL"
label="Proxy Base URL"
name="proxy_base_url"
normalize={(value) => value?.trim()}
rules={[