From a039041cc33606dd71e4ffba7041fa06a71c7af2 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 23 Jan 2026 16:12:38 +0530 Subject: [PATCH] Fix: Add image generation streamign --- litellm/images/main.py | 32 ++++++++++++++++++++-- litellm/llms/openai/openai.py | 21 ++++++++++++++ litellm/proxy/image_endpoints/endpoints.py | 26 +++++++++++++++++- 3 files changed, 76 insertions(+), 3 deletions(-) diff --git a/litellm/images/main.py b/litellm/images/main.py index 6c4c502a7b0..78d1ebbd07d 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -85,7 +85,7 @@ def _get_ImageEditRequestUtils() -> "ImageEditRequestUtils": ##### Image Generation ####################### @client -async def aimage_generation(*args, **kwargs) -> ImageResponse: +async def aimage_generation(*args, **kwargs): """ Asynchronously calls the `image_generation` function with the given arguments and keyword arguments. @@ -116,6 +116,13 @@ async def aimage_generation(*args, **kwargs) -> ImageResponse: # Await normally init_response = await loop.run_in_executor(None, func_with_context) + # Check if streaming is enabled + if kwargs.get("stream", False): + # For streaming, return the stream object directly + if asyncio.iscoroutine(init_response): + return await init_response # type: ignore + return init_response + response: Optional[ImageResponse] = None if isinstance(init_response, dict): response = ImageResponse(**init_response) @@ -159,6 +166,8 @@ def image_generation( api_base: Optional[str] = None, api_version: Optional[str] = None, custom_llm_provider=None, + stream: Optional[bool] = False, + partial_images: Optional[bool] = None, *, aimg_generation: Literal[True], **kwargs, @@ -183,6 +192,8 @@ def image_generation( api_base: Optional[str] = None, api_version: Optional[str] = None, custom_llm_provider=None, + stream: Optional[bool] = False, + partial_images: Optional[bool] = None, *, aimg_generation: Literal[False] = False, **kwargs, @@ -207,6 +218,8 @@ def image_generation( # noqa: PLR0915 api_base: Optional[str] = None, api_version: Optional[str] = None, custom_llm_provider=None, + stream: Optional[bool] = False, + partial_images: Optional[bool] = None, **kwargs, ) -> Union[ ImageResponse, @@ -262,6 +275,8 @@ def image_generation( # noqa: PLR0915 "quality", "size", "style", + "stream", + "partial_images", ] litellm_params = all_litellm_params default_params = openai_params + litellm_params @@ -426,6 +441,7 @@ def image_generation( # noqa: PLR0915 logging_obj=litellm_logging_obj, timeout=timeout, client=client, + stream=stream, ) elif custom_llm_provider == "azure_ai": from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo @@ -471,7 +487,7 @@ def image_generation( # noqa: PLR0915 ): # Forward OpenAI organization if present (set by proxy pre-call utils) organization: Optional[str] = kwargs.get("organization", None) - model_response = openai_chat_completions.image_generation( + response = openai_chat_completions.image_generation( model=model, prompt=prompt, timeout=timeout, @@ -483,7 +499,15 @@ def image_generation( # noqa: PLR0915 organization=organization, aimg_generation=aimg_generation, client=client, + stream=stream, + partial_images=partial_images ) + + # If streaming is enabled, return the stream directly + if stream: + return response + + model_response = response elif custom_llm_provider == "bedrock": if model is None: raise Exception("Model needs to be set for bedrock") @@ -755,6 +779,8 @@ def image_edit( # noqa: PLR0915 "size", "style", "async_call", + "stream", + "partial_images", ] litellm_params_list = all_litellm_params default_params = openai_params + litellm_params_list @@ -912,6 +938,7 @@ def image_edit( # noqa: PLR0915 timeout=timeout or DEFAULT_REQUEST_TIMEOUT, _is_async=_is_async, client=kwargs.get("client"), + stream=image_edit_request_params.get("stream", False), ) # Call the handler with _is_async flag instead of directly calling the async handler return base_llm_http_handler.image_edit_handler( @@ -928,6 +955,7 @@ def image_edit( # noqa: PLR0915 timeout=timeout or DEFAULT_REQUEST_TIMEOUT, _is_async=_is_async, client=kwargs.get("client"), + stream=image_edit_request_params.get("stream", False), ) except Exception as e: diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 4d623097478..6dbf2cdc0a1 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1317,6 +1317,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) response = await openai_aclient.images.generate(**data, timeout=timeout) # type: ignore + + # Check if response is a stream + is_stream = data.get("stream", False) + if is_stream: + # Return the stream object directly for streaming responses + return response + stringified_response = response.model_dump() ## LOGGING logging_obj.post_call( @@ -1348,10 +1355,18 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client=None, aimg_generation=None, organization: Optional[str] = None, + stream: Optional[bool] = False, + partial_images: Optional[bool] = None, ) -> ImageResponse: data = {} try: data = {"model": model, "prompt": prompt, **optional_params} + + # Add stream and partial_images if provided + if stream is not None: + data["stream"] = stream + if partial_images is not None: + data["partial_images"] = partial_images max_retries = data.pop("max_retries", 2) if not isinstance(max_retries, int): raise OpenAIError(status_code=422, message="max retries must be an int") @@ -1384,6 +1399,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ## COMPLETION CALL _response = openai_client.images.generate(**data, timeout=timeout) # type: ignore + # Check if response is a stream + is_stream = data.get("stream", False) + if is_stream: + # Return the stream object directly for streaming responses + return _response + response = _response.model_dump() ## LOGGING logging_obj.post_call( diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index 4a2c05f8590..408a92d27bd 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -1,10 +1,11 @@ import asyncio +import json import traceback from typing import List import orjson from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, status -from fastapi.responses import ORJSONResponse +from fastapi.responses import ORJSONResponse, StreamingResponse import litellm from litellm._logging import verbose_proxy_logger @@ -138,6 +139,29 @@ async def image_generation( ) response = await llm_call + # Check if response is a streaming iterator + is_streaming = hasattr(response, '__aiter__') and not isinstance(response, (dict, list, str)) + + if is_streaming: + # Handle streaming response using the same pattern as responses API + from litellm.proxy.proxy_server import select_data_generator + + selected_data_generator = select_data_generator( + response=response, + user_api_key_dict=user_api_key_dict, + request_data=data, + ) + + return StreamingResponse( + selected_data_generator, + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + } + ) + ### ALERTING ### asyncio.create_task( proxy_logging_obj.update_request_status(