mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Added streaming support for response api streaming image generation (#15269)
This commit is contained in:
parent
faeb7484db
commit
c0d0424eb8
5 changed files with 220 additions and 0 deletions
|
|
@ -37,6 +37,29 @@ for event in response:
|
|||
print(event)
|
||||
```
|
||||
|
||||
#### Image Generation with Streaming
|
||||
```python showLineNumbers title="OpenAI Streaming Image Generation"
|
||||
import litellm
|
||||
import base64
|
||||
|
||||
# Streaming image generation with partial images
|
||||
stream = litellm.responses(
|
||||
model="gpt-4.1", # Use an actual image generation model
|
||||
input="Generate a gorgeous image of a river made of white owl feathers",
|
||||
stream=True,
|
||||
tools=[{"type": "image_generation", "partial_images": 2}],
|
||||
|
||||
)
|
||||
|
||||
for event in stream:
|
||||
if event.type == "response.image_generation_call.partial_image":
|
||||
idx = event.partial_image_index
|
||||
image_base64 = event.partial_image_b64
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
with open(f"river{idx}.png", "wb") as f:
|
||||
f.write(image_bytes)
|
||||
```
|
||||
|
||||
#### GET a Response
|
||||
```python showLineNumbers title="Get Response by ID"
|
||||
import litellm
|
||||
|
|
@ -150,6 +173,33 @@ for event in response:
|
|||
print(event)
|
||||
```
|
||||
|
||||
#### Image Generation with Streaming
|
||||
```python showLineNumbers title="OpenAI Proxy Streaming Image Generation"
|
||||
from openai import OpenAI
|
||||
import base64
|
||||
|
||||
# Initialize client with your proxy URL
|
||||
client = OpenAI(api_key="sk-1234", base_url="http://localhost:4000")
|
||||
|
||||
stream = client.responses.create(
|
||||
model="gpt-4.1",
|
||||
input="Draw a gorgeous image of a river made of white owl feathers, snaking its way through a serene winter landscape",
|
||||
stream=True,
|
||||
tools=[{"type": "image_generation", "partial_images": 2}],
|
||||
)
|
||||
|
||||
|
||||
for event in stream:
|
||||
print(f"event: {event}")
|
||||
if event.type == "response.image_generation_call.partial_image":
|
||||
idx = event.partial_image_index
|
||||
image_base64 = event.partial_image_b64
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
with open(f"river{idx}.png", "wb") as f:
|
||||
f.write(image_bytes)
|
||||
|
||||
```
|
||||
|
||||
#### GET a Response
|
||||
```python showLineNumbers title="Get Response by ID with OpenAI SDK"
|
||||
from openai import OpenAI
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ Requests to /chat/completions may be bridged here automatically when the provide
|
|||
| Logging | ✅ | Works across all integrations |
|
||||
| End-user Tracking | ✅ | |
|
||||
| Streaming | ✅ | |
|
||||
| Image Generation Streaming | ✅ | Progressive image generation with partial images (1-3) |
|
||||
| Fallbacks | ✅ | Works between supported models |
|
||||
| Loadbalancing | ✅ | Works between supported models |
|
||||
| Supported operations | Create a response, Get a response, Delete a response | |
|
||||
|
|
@ -56,6 +57,29 @@ for event in response:
|
|||
print(event)
|
||||
```
|
||||
|
||||
#### Image Generation with Streaming
|
||||
```python showLineNumbers title="OpenAI Streaming Image Generation"
|
||||
import litellm
|
||||
import base64
|
||||
|
||||
# Streaming image generation with partial images
|
||||
stream = litellm.responses(
|
||||
model="gpt-4.1", # Use an actual image generation model
|
||||
input="Generate a gorgeous image of a river made of white owl feathers",
|
||||
stream=True,
|
||||
tools=[{"type": "image_generation", "partial_images": 2}],
|
||||
|
||||
)
|
||||
|
||||
for event in stream:
|
||||
if event.type == "response.image_generation_call.partial_image":
|
||||
idx = event.partial_image_index
|
||||
image_base64 = event.partial_image_b64
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
with open(f"river{idx}.png", "wb") as f:
|
||||
f.write(image_bytes)
|
||||
```
|
||||
|
||||
#### GET a Response
|
||||
```python showLineNumbers title="Get Response by ID"
|
||||
import litellm
|
||||
|
|
@ -380,6 +404,32 @@ for event in response:
|
|||
print(event)
|
||||
```
|
||||
|
||||
#### Image Generation with Streaming
|
||||
```python showLineNumbers title="OpenAI Proxy Streaming Image Generation"
|
||||
from openai import OpenAI
|
||||
import base64
|
||||
|
||||
client = OpenAI(api_key="sk-1234", base_url="http://localhost:4000")
|
||||
|
||||
stream = client.responses.create(
|
||||
model="gpt-4.1",
|
||||
input="Draw a gorgeous image of a river made of white owl feathers, snaking its way through a serene winter landscape",
|
||||
stream=True,
|
||||
tools=[{"type": "image_generation", "partial_images": 2}],
|
||||
)
|
||||
|
||||
|
||||
for event in stream:
|
||||
print(f"event: {event}")
|
||||
if event.type == "response.image_generation_call.partial_image":
|
||||
idx = event.partial_image_index
|
||||
image_base64 = event.partial_image_b64
|
||||
image_bytes = base64.b64decode(image_base64)
|
||||
with open(f"river{idx}.png", "wb") as f:
|
||||
f.write(image_bytes)
|
||||
|
||||
```
|
||||
|
||||
#### GET a Response
|
||||
```python showLineNumbers title="Get Response by ID with OpenAI SDK"
|
||||
from openai import OpenAI
|
||||
|
|
|
|||
|
|
@ -271,6 +271,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
ResponsesAPIStreamEvents.MCP_CALL_ARGUMENTS_DONE: MCPCallArgumentsDoneEvent,
|
||||
ResponsesAPIStreamEvents.MCP_CALL_COMPLETED: MCPCallCompletedEvent,
|
||||
ResponsesAPIStreamEvents.MCP_CALL_FAILED: MCPCallFailedEvent,
|
||||
ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE: ImageGenerationPartialImageEvent,
|
||||
ResponsesAPIStreamEvents.ERROR: ErrorEvent,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -992,6 +992,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
|
|||
prompt_cache_key: Optional[str]
|
||||
stream_options: Optional[dict]
|
||||
top_logprobs: Optional[int]
|
||||
partial_images: Optional[int] # Number of partial images to generate (1-3) for streaming image generation
|
||||
|
||||
|
||||
class ResponsesAPIRequestParams(ResponsesAPIOptionalRequestParams, total=False):
|
||||
|
|
@ -1129,6 +1130,9 @@ class ResponsesAPIStreamEvents(str, Enum):
|
|||
MCP_CALL_COMPLETED = "response.mcp_call.completed"
|
||||
MCP_CALL_FAILED = "response.mcp_call.failed"
|
||||
|
||||
# Image generation events
|
||||
IMAGE_GENERATION_PARTIAL_IMAGE = "image_generation.partial_image"
|
||||
|
||||
# Error event
|
||||
ERROR = "error"
|
||||
|
||||
|
|
@ -1352,6 +1356,12 @@ class MCPCallFailedEvent(BaseLiteLLMOpenAIResponseObject):
|
|||
output_index: int
|
||||
|
||||
|
||||
class ImageGenerationPartialImageEvent(BaseLiteLLMOpenAIResponseObject):
|
||||
type: Literal[ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE]
|
||||
partial_image_index: int
|
||||
b64_json: str
|
||||
|
||||
|
||||
class ErrorEvent(BaseLiteLLMOpenAIResponseObject):
|
||||
type: Literal[ResponsesAPIStreamEvents.ERROR]
|
||||
code: Optional[str]
|
||||
|
|
@ -1400,6 +1410,7 @@ ResponsesAPIStreamingResponse = Annotated[
|
|||
MCPCallArgumentsDoneEvent,
|
||||
MCPCallCompletedEvent,
|
||||
MCPCallFailedEvent,
|
||||
ImageGenerationPartialImageEvent,
|
||||
ErrorEvent,
|
||||
GenericEvent,
|
||||
],
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import litellm
|
|||
from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.types.llms.openai import (
|
||||
ImageGenerationPartialImageEvent,
|
||||
OutputTextDeltaEvent,
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIRequestParams,
|
||||
|
|
@ -282,6 +283,113 @@ class TestOpenAIResponsesAPIConfig:
|
|||
assert isinstance(result, GenericEvent)
|
||||
assert result.type == "test"
|
||||
|
||||
def test_get_event_model_class_image_generation_partial_image(self):
|
||||
"""Test that get_event_model_class returns ImageGenerationPartialImageEvent for image generation events"""
|
||||
event_type = ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE
|
||||
result = self.config.get_event_model_class(event_type)
|
||||
assert result == ImageGenerationPartialImageEvent
|
||||
|
||||
def test_transform_streaming_response_image_generation_partial_image(self):
|
||||
"""Test streaming response transformation for image generation partial image events"""
|
||||
# Test with a partial image event - simulating OpenAI's streaming image generation
|
||||
chunk = {
|
||||
"type": "image_generation.partial_image",
|
||||
"partial_image_index": 0,
|
||||
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==", # 1x1 red pixel PNG
|
||||
}
|
||||
|
||||
result = self.config.transform_streaming_response(
|
||||
model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj
|
||||
)
|
||||
|
||||
# Verify the result is the correct event type
|
||||
assert isinstance(result, ImageGenerationPartialImageEvent)
|
||||
assert result.type == ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE
|
||||
assert result.partial_image_index == 0
|
||||
assert result.b64_json == chunk["b64_json"]
|
||||
assert len(result.b64_json) > 0 # Verify we have image data
|
||||
|
||||
def test_transform_streaming_response_multiple_partial_images(self):
|
||||
"""Test streaming response with multiple partial images (simulating progressive image generation)"""
|
||||
# Test with multiple partial images (as would happen with partial_images=2 or 3)
|
||||
test_cases = [
|
||||
{
|
||||
"type": "image_generation.partial_image",
|
||||
"partial_image_index": 0,
|
||||
"b64_json": "base64data_partial_0",
|
||||
},
|
||||
{
|
||||
"type": "image_generation.partial_image",
|
||||
"partial_image_index": 1,
|
||||
"b64_json": "base64data_partial_1",
|
||||
},
|
||||
{
|
||||
"type": "image_generation.partial_image",
|
||||
"partial_image_index": 2,
|
||||
"b64_json": "base64data_partial_2",
|
||||
},
|
||||
]
|
||||
|
||||
for idx, chunk in enumerate(test_cases):
|
||||
result = self.config.transform_streaming_response(
|
||||
model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj
|
||||
)
|
||||
|
||||
assert isinstance(result, ImageGenerationPartialImageEvent)
|
||||
assert result.type == ResponsesAPIStreamEvents.IMAGE_GENERATION_PARTIAL_IMAGE
|
||||
assert result.partial_image_index == idx
|
||||
assert result.b64_json == chunk["b64_json"]
|
||||
|
||||
def test_transform_responses_api_request_with_partial_images_param(self):
|
||||
"""Test request transformation with partial_images parameter for streaming image generation"""
|
||||
input_text = "Generate a beautiful landscape"
|
||||
optional_params = {
|
||||
"temperature": 0.7,
|
||||
"stream": True,
|
||||
"partial_images": 2, # Request 2 partial images during generation
|
||||
}
|
||||
|
||||
result = self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input=input_text,
|
||||
response_api_optional_request_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Validate the result includes partial_images parameter
|
||||
expected_fields = {
|
||||
"model": self.model,
|
||||
"input": input_text,
|
||||
"temperature": 0.7,
|
||||
"stream": True,
|
||||
"partial_images": 2,
|
||||
}
|
||||
|
||||
self.validate_responses_api_request_params(result, expected_fields)
|
||||
|
||||
def test_partial_images_parameter_validation(self):
|
||||
"""Test that partial_images parameter accepts valid values (1-3)"""
|
||||
input_text = "Generate an image"
|
||||
|
||||
# Test with different valid partial_images values
|
||||
for partial_images_value in [1, 2, 3]:
|
||||
optional_params = {
|
||||
"stream": True,
|
||||
"partial_images": partial_images_value,
|
||||
}
|
||||
|
||||
result = self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input=input_text,
|
||||
response_api_optional_request_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result["partial_images"] == partial_images_value
|
||||
assert result["stream"] is True
|
||||
|
||||
|
||||
class TestAzureResponsesAPIConfig:
|
||||
def setup_method(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue