diff --git a/litellm/__init__.py b/litellm/__init__.py index 96c1552c36..dd09810d46 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -132,22 +132,22 @@ prometheus_initialize_budget_metrics: Optional[bool] = False require_auth_for_metrics_endpoint: Optional[bool] = False argilla_batch_size: Optional[int] = None datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload -gcs_pub_sub_use_v1: Optional[bool] = ( - False # if you want to use v1 gcs pubsub logged payload -) -generic_api_use_v1: Optional[bool] = ( - False # if you want to use v1 generic api logged payload -) +gcs_pub_sub_use_v1: Optional[ + bool +] = False # if you want to use v1 gcs pubsub logged payload +generic_api_use_v1: Optional[ + bool +] = False # if you want to use v1 generic api logged payload argilla_transformation_object: Optional[Dict[str, Any]] = None -_async_input_callback: List[Union[str, Callable, CustomLogger]] = ( - [] -) # internal variable - async custom callbacks are routed here. -_async_success_callback: List[Union[str, Callable, CustomLogger]] = ( - [] -) # internal variable - async custom callbacks are routed here. -_async_failure_callback: List[Union[str, Callable, CustomLogger]] = ( - [] -) # internal variable - async custom callbacks are routed here. +_async_input_callback: List[ + Union[str, Callable, CustomLogger] +] = [] # internal variable - async custom callbacks are routed here. +_async_success_callback: List[ + Union[str, Callable, CustomLogger] +] = [] # internal variable - async custom callbacks are routed here. +_async_failure_callback: List[ + Union[str, Callable, CustomLogger] +] = [] # internal variable - async custom callbacks are routed here. pre_call_rules: List[Callable] = [] post_call_rules: List[Callable] = [] turn_off_message_logging: Optional[bool] = False @@ -155,18 +155,18 @@ log_raw_request_response: bool = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False filter_invalid_headers: Optional[bool] = False -add_user_information_to_llm_headers: Optional[bool] = ( - None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers -) +add_user_information_to_llm_headers: Optional[ + bool +] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers store_audit_logs = False # Enterprise feature, allow users to see audit logs ### end of callbacks ############# -email: Optional[str] = ( - None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -) -token: Optional[str] = ( - None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -) +email: Optional[ + str +] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +token: Optional[ + str +] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 telemetry = True max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False)) @@ -247,24 +247,20 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None enable_caching_on_provider_specific_optional_params: bool = ( False # feature-flag for caching on optional params - e.g. 'top_k' ) -caching: bool = ( - False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -) -caching_with_models: bool = ( - False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -) -cache: Optional[Cache] = ( - None # cache object <- use this - https://docs.litellm.ai/docs/caching -) +caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +cache: Optional[ + Cache +] = None # cache object <- use this - https://docs.litellm.ai/docs/caching default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None default_redis_batch_cache_expiry: Optional[float] = None model_alias_map: Dict[str, str] = {} model_group_alias_map: Dict[str, str] = {} max_budget: float = 0.0 # set the max budget across all providers -budget_duration: Optional[str] = ( - None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). -) +budget_duration: Optional[ + str +] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). default_soft_budget: float = ( DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0 ) @@ -273,15 +269,11 @@ forward_traceparent_to_llm_provider: bool = False _current_cost = 0.0 # private variable, used if max budget is set error_logs: Dict = {} -add_function_to_prompt: bool = ( - False # if function calling not supported by api, append function call details to system prompt -) +add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt client_session: Optional[httpx.Client] = None aclient_session: Optional[httpx.AsyncClient] = None model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks' -model_cost_map_url: str = ( - "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json" -) +model_cost_map_url: str = "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json" suppress_debug_info = False dynamodb_table_name: Optional[str] = None s3_callback_params: Optional[Dict] = None @@ -304,9 +296,7 @@ disable_end_user_cost_tracking_prometheus_only: Optional[bool] = None custom_prometheus_metadata_labels: List[str] = [] #### REQUEST PRIORITIZATION #### priority_reservation: Optional[Dict[str, float]] = None -force_ipv4: bool = ( - False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. -) +force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. module_level_aclient = AsyncHTTPHandler( timeout=request_timeout, client_alias="module level aclient" ) @@ -320,13 +310,13 @@ fallbacks: Optional[List] = None context_window_fallbacks: Optional[List] = None content_policy_fallbacks: Optional[List] = None allowed_fails: int = 3 -num_retries_per_request: Optional[int] = ( - None # for the request overall (incl. fallbacks + model retries) -) +num_retries_per_request: Optional[ + int +] = None # for the request overall (incl. fallbacks + model retries) ####### SECRET MANAGERS ##################### -secret_manager_client: Optional[Any] = ( - None # list of instantiated key management clients - e.g. azure kv, infisical, etc. -) +secret_manager_client: Optional[ + Any +] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc. _google_kms_resource_name: Optional[str] = None _key_management_system: Optional[KeyManagementSystem] = None _key_management_settings: KeyManagementSettings = KeyManagementSettings() @@ -448,6 +438,7 @@ snowflake_models: List = [] llama_models: List = [] nscale_models: List = [] + def is_bedrock_pricing_only_model(key: str) -> bool: """ Excludes keys with the pattern 'bedrock//'. These are in the model_prices_and_context_window.json file for pricing purposes only. @@ -1091,6 +1082,7 @@ from .proxy.proxy_cli import run_server from .router import Router from .assistants.main import * from .batches.main import * +from .images.main import * from .batch_completion.main import * # type: ignore from .rerank_api.main import * from .llms.anthropic.experimental_pass_through.messages.handler import * @@ -1117,10 +1109,10 @@ from .types.llms.custom_llm import CustomLLMItem from .types.utils import GenericStreamingChunk custom_provider_map: List[CustomLLMItem] = [] -_custom_providers: List[str] = ( - [] -) # internal helper util, used to track names of custom providers -disable_hf_tokenizer_download: Optional[bool] = ( - None # disable huggingface tokenizer download. Defaults to openai clk100 -) +_custom_providers: List[ + str +] = [] # internal helper util, used to track names of custom providers +disable_hf_tokenizer_download: Optional[ + bool +] = None # disable huggingface tokenizer download. Defaults to openai clk100 global_disable_no_log_param: bool = False diff --git a/litellm/constants.py b/litellm/constants.py index e224c0dd69..37d44f64c5 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -153,12 +153,14 @@ FIREWORKS_AI_80_B = int(os.getenv("FIREWORKS_AI_80_B", 80)) #### Logging callback constants #### REDACTED_BY_LITELM_STRING = "REDACTED_BY_LITELM" +############### LLM Provider Constants ############### ### ANTHROPIC CONSTANTS ### ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES = { "low": 1, "medium": 5, "high": 10, } +DEFAULT_IMAGE_ENDPOINT_MODEL = "dall-e-2" LITELLM_CHAT_PROVIDERS = [ "openai", diff --git a/litellm/images/main.py b/litellm/images/main.py new file mode 100644 index 0000000000..cf62b3a365 --- /dev/null +++ b/litellm/images/main.py @@ -0,0 +1,736 @@ +import asyncio +import contextvars +from functools import partial +from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast + +import httpx + +import litellm +from litellm import Logging, client, exception_type, get_litellm_params +from litellm.constants import DEFAULT_IMAGE_ENDPOINT_MODEL +from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT +from litellm.exceptions import LiteLLMUnknownProvider +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.mock_functions import mock_image_generation +from litellm.llms.base_llm import BaseImageEditConfig, BaseImageGenerationConfig +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.custom_llm import CustomLLM + +#################### Initialize provider clients #################### +from litellm.main import ( + azure_chat_completions, + base_llm_aiohttp_handler, + base_llm_http_handler, + bedrock_image_generation, + openai_chat_completions, + openai_image_variations, + vertex_image_generation, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.llms.openai import ImageGenerationRequestQuality +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import ( + LITELLM_IMAGE_VARIATION_PROVIDERS, + FileTypes, + LlmProviders, + all_litellm_params, +) +from litellm.utils import ( + ImageResponse, + ProviderConfigManager, + get_llm_provider, + get_optional_params_image_gen, +) + +from .utils import ImageEditRequestUtils + + +##### Image Generation ####################### +@client +async def aimage_generation(*args, **kwargs) -> ImageResponse: + """ + Asynchronously calls the `image_generation` function with the given arguments and keyword arguments. + + Parameters: + - `args` (tuple): Positional arguments to be passed to the `image_generation` function. + - `kwargs` (dict): Keyword arguments to be passed to the `image_generation` function. + + Returns: + - `response` (Any): The response returned by the `image_generation` function. + """ + loop = asyncio.get_event_loop() + model = args[0] if len(args) > 0 else kwargs["model"] + ### PASS ARGS TO Image Generation ### + kwargs["aimg_generation"] = True + custom_llm_provider = None + try: + # Use a partial function to pass your keyword arguments + func = partial(image_generation, *args, **kwargs) + + # Add the context to the function + ctx = contextvars.copy_context() + func_with_context = partial(ctx.run, func) + + _, custom_llm_provider, _, _ = get_llm_provider( + model=model, api_base=kwargs.get("api_base", None) + ) + + # Await normally + init_response = await loop.run_in_executor(None, func_with_context) + if isinstance(init_response, dict) or isinstance( + init_response, ImageResponse + ): ## CACHING SCENARIO + if isinstance(init_response, dict): + init_response = ImageResponse(**init_response) + response = init_response + elif asyncio.iscoroutine(init_response): + response = await init_response # type: ignore + else: + # Call the synchronous function using run_in_executor + response = await loop.run_in_executor(None, func_with_context) + return response + except Exception as e: + custom_llm_provider = custom_llm_provider or "openai" + raise exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=args, + extra_kwargs=kwargs, + ) + + +@client +def image_generation( # noqa: PLR0915 + prompt: str, + model: Optional[str] = None, + n: Optional[int] = None, + quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, + response_format: Optional[str] = None, + size: Optional[str] = None, + style: Optional[str] = None, + user: Optional[str] = None, + timeout=600, # default to 10 minutes + api_key: Optional[str] = None, + api_base: Optional[str] = None, + api_version: Optional[str] = None, + custom_llm_provider=None, + **kwargs, +) -> ImageResponse: + """ + Maps the https://api.openai.com/v1/images/generations endpoint. + + Currently supports just Azure + OpenAI. + """ + try: + args = locals() + aimg_generation = kwargs.get("aimg_generation", False) + litellm_call_id = kwargs.get("litellm_call_id", None) + logger_fn = kwargs.get("logger_fn", None) + mock_response: Optional[str] = kwargs.get("mock_response", None) # type: ignore + proxy_server_request = kwargs.get("proxy_server_request", None) + azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None) + model_info = kwargs.get("model_info", None) + metadata = kwargs.get("metadata", {}) + litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore + client = kwargs.get("client", None) + extra_headers = kwargs.get("extra_headers", None) + headers: dict = kwargs.get("headers", None) or {} + base_model = kwargs.get("base_model", None) + if extra_headers is not None: + headers.update(extra_headers) + model_response: ImageResponse = litellm.utils.ImageResponse() + dynamic_api_key: Optional[str] = None + if model is not None or custom_llm_provider is not None: + model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( + model=model, # type: ignore + custom_llm_provider=custom_llm_provider, + api_base=api_base, + ) + else: + model = "dall-e-2" + custom_llm_provider = "openai" # default to dall-e-2 on openai + model_response._hidden_params["model"] = model + openai_params = [ + "user", + "request_timeout", + "api_base", + "api_version", + "api_key", + "deployment_id", + "organization", + "base_url", + "default_headers", + "timeout", + "max_retries", + "n", + "quality", + "size", + "style", + ] + litellm_params = all_litellm_params + default_params = openai_params + litellm_params + non_default_params = { + k: v for k, v in kwargs.items() if k not in default_params + } # model-specific params - pass them straight to the model/provider + + image_generation_config: Optional[BaseImageGenerationConfig] = None + if ( + custom_llm_provider is not None + and custom_llm_provider in LlmProviders._member_map_.values() + ): + image_generation_config = ( + ProviderConfigManager.get_provider_image_generation_config( + model=base_model or model, + provider=LlmProviders(custom_llm_provider), + ) + ) + + optional_params = get_optional_params_image_gen( + model=base_model or model, + n=n, + quality=quality, + response_format=response_format, + size=size, + style=style, + user=user, + custom_llm_provider=custom_llm_provider, + provider_config=image_generation_config, + **non_default_params, + ) + + litellm_params_dict = get_litellm_params(**kwargs) + + logging: Logging = litellm_logging_obj + logging.update_environment_variables( + model=model, + user=user, + optional_params=optional_params, + litellm_params={ + "timeout": timeout, + "azure": False, + "litellm_call_id": litellm_call_id, + "logger_fn": logger_fn, + "proxy_server_request": proxy_server_request, + "model_info": model_info, + "metadata": metadata, + "preset_cache_key": None, + "stream_response": {}, + }, + custom_llm_provider=custom_llm_provider, + ) + if "custom_llm_provider" not in logging.model_call_details: + logging.model_call_details["custom_llm_provider"] = custom_llm_provider + if mock_response is not None: + return mock_image_generation(model=model, mock_response=mock_response) + + if custom_llm_provider == "azure": + # azure configs + api_type = get_secret_str("AZURE_API_TYPE") or "azure" + + api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") + + api_version = ( + api_version + or litellm.api_version + or get_secret_str("AZURE_API_VERSION") + ) + + api_key = ( + api_key + or litellm.api_key + or litellm.azure_key + or get_secret_str("AZURE_OPENAI_API_KEY") + or get_secret_str("AZURE_API_KEY") + ) + + azure_ad_token = optional_params.pop( + "azure_ad_token", None + ) or get_secret_str("AZURE_AD_TOKEN") + + default_headers = { + "Content-Type": "application/json;", + "api-key": api_key, + } + for k, v in default_headers.items(): + if k not in headers: + headers[k] = v + + model_response = azure_chat_completions.image_generation( + model=model, + prompt=prompt, + timeout=timeout, + api_key=api_key, + api_base=api_base, + azure_ad_token=azure_ad_token, + azure_ad_token_provider=azure_ad_token_provider, + logging_obj=litellm_logging_obj, + optional_params=optional_params, + model_response=model_response, + api_version=api_version, + aimg_generation=aimg_generation, + client=client, + headers=headers, + litellm_params=litellm_params_dict, + ) + elif ( + custom_llm_provider == "openai" + or custom_llm_provider in litellm.openai_compatible_providers + ): + model_response = openai_chat_completions.image_generation( + model=model, + prompt=prompt, + timeout=timeout, + api_key=api_key or dynamic_api_key, + api_base=api_base, + logging_obj=litellm_logging_obj, + optional_params=optional_params, + model_response=model_response, + aimg_generation=aimg_generation, + client=client, + ) + elif custom_llm_provider == "bedrock": + if model is None: + raise Exception("Model needs to be set for bedrock") + model_response = bedrock_image_generation.image_generation( # type: ignore + model=model, + prompt=prompt, + timeout=timeout, + logging_obj=litellm_logging_obj, + optional_params=optional_params, + model_response=model_response, + aimg_generation=aimg_generation, + client=client, + ) + elif custom_llm_provider == "vertex_ai": + vertex_ai_project = ( + optional_params.pop("vertex_project", None) + or optional_params.pop("vertex_ai_project", None) + or litellm.vertex_project + or get_secret_str("VERTEXAI_PROJECT") + ) + vertex_ai_location = ( + optional_params.pop("vertex_location", None) + or optional_params.pop("vertex_ai_location", None) + or litellm.vertex_location + or get_secret_str("VERTEXAI_LOCATION") + ) + vertex_credentials = ( + optional_params.pop("vertex_credentials", None) + or optional_params.pop("vertex_ai_credentials", None) + or get_secret_str("VERTEXAI_CREDENTIALS") + ) + + api_base = ( + api_base + or litellm.api_base + or get_secret_str("VERTEXAI_API_BASE") + or get_secret_str("VERTEX_API_BASE") + ) + + model_response = vertex_image_generation.image_generation( + model=model, + prompt=prompt, + timeout=timeout, + logging_obj=litellm_logging_obj, + optional_params=optional_params, + model_response=model_response, + vertex_project=vertex_ai_project, + vertex_location=vertex_ai_location, + vertex_credentials=vertex_credentials, + aimg_generation=aimg_generation, + api_base=api_base, + client=client, + ) + elif ( + custom_llm_provider in litellm._custom_providers + ): # Assume custom LLM provider + # Get the Custom Handler + custom_handler: Optional[CustomLLM] = None + for item in litellm.custom_provider_map: + if item["provider"] == custom_llm_provider: + custom_handler = item["custom_handler"] + + if custom_handler is None: + raise LiteLLMUnknownProvider( + model=model, custom_llm_provider=custom_llm_provider + ) + + ## ROUTE LLM CALL ## + if aimg_generation is True: + async_custom_client: Optional[AsyncHTTPHandler] = None + if client is not None and isinstance(client, AsyncHTTPHandler): + async_custom_client = client + + ## CALL FUNCTION + model_response = custom_handler.aimage_generation( # type: ignore + model=model, + prompt=prompt, + api_key=api_key, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + logging_obj=litellm_logging_obj, + timeout=timeout, + client=async_custom_client, + ) + else: + custom_client: Optional[HTTPHandler] = None + if client is not None and isinstance(client, HTTPHandler): + custom_client = client + + ## CALL FUNCTION + model_response = custom_handler.image_generation( + model=model, + prompt=prompt, + api_key=api_key, + api_base=api_base, + model_response=model_response, + optional_params=optional_params, + logging_obj=litellm_logging_obj, + timeout=timeout, + client=custom_client, + ) + + return model_response + except Exception as e: + ## Map to OpenAI Exception + raise exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=locals(), + extra_kwargs=kwargs, + ) + + +@client +async def aimage_variation(*args, **kwargs) -> ImageResponse: + """ + Asynchronously calls the `image_variation` function with the given arguments and keyword arguments. + + Parameters: + - `args` (tuple): Positional arguments to be passed to the `image_variation` function. + - `kwargs` (dict): Keyword arguments to be passed to the `image_variation` function. + + Returns: + - `response` (Any): The response returned by the `image_variation` function. + """ + loop = asyncio.get_event_loop() + model = kwargs.get("model", None) + custom_llm_provider = kwargs.get("custom_llm_provider", None) + ### PASS ARGS TO Image Generation ### + kwargs["async_call"] = True + try: + # Use a partial function to pass your keyword arguments + func = partial(image_variation, *args, **kwargs) + + # Add the context to the function + ctx = contextvars.copy_context() + func_with_context = partial(ctx.run, func) + + if custom_llm_provider is None and model is not None: + _, custom_llm_provider, _, _ = get_llm_provider( + model=model, api_base=kwargs.get("api_base", None) + ) + + # Await normally + init_response = await loop.run_in_executor(None, func_with_context) + if isinstance(init_response, dict) or isinstance( + init_response, ImageResponse + ): ## CACHING SCENARIO + if isinstance(init_response, dict): + init_response = ImageResponse(**init_response) + response = init_response + elif asyncio.iscoroutine(init_response): + response = await init_response # type: ignore + else: + # Call the synchronous function using run_in_executor + response = await loop.run_in_executor(None, func_with_context) + return response + except Exception as e: + custom_llm_provider = custom_llm_provider or "openai" + raise exception_type( + model=model, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=args, + extra_kwargs=kwargs, + ) + + +@client +def image_variation( + image: FileTypes, + model: str = "dall-e-2", # set to dall-e-2 by default - like OpenAI. + n: int = 1, + response_format: Literal["url", "b64_json"] = "url", + size: Optional[str] = None, + user: Optional[str] = None, + **kwargs, +) -> ImageResponse: + # get non-default params + client = kwargs.get("client", None) + # get logging object + litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj")) + + # get the litellm params + litellm_params = get_litellm_params(**kwargs) + # get the custom llm provider + model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( + model=model, + custom_llm_provider=litellm_params.get("custom_llm_provider", None), + api_base=litellm_params.get("api_base", None), + api_key=litellm_params.get("api_key", None), + ) + + # route to the correct provider w/ the params + try: + llm_provider = LlmProviders(custom_llm_provider) + image_variation_provider = LITELLM_IMAGE_VARIATION_PROVIDERS(llm_provider) + except ValueError: + raise ValueError( + f"Invalid image variation provider: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" + ) + model_response = ImageResponse() + + response: Optional[ImageResponse] = None + + provider_config = ProviderConfigManager.get_provider_model_info( + model=model or "", # openai defaults to dall-e-2 + provider=llm_provider, + ) + + if provider_config is None: + raise ValueError( + f"image variation provider has no known model info config - required for getting api keys, etc.: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" + ) + + api_key = provider_config.get_api_key(litellm_params.get("api_key", None)) + api_base = provider_config.get_api_base(litellm_params.get("api_base", None)) + + if image_variation_provider == LITELLM_IMAGE_VARIATION_PROVIDERS.OPENAI: + if api_key is None: + raise ValueError("API key is required for OpenAI image variations") + if api_base is None: + raise ValueError("API base is required for OpenAI image variations") + + response = openai_image_variations.image_variations( + model_response=model_response, + api_key=api_key, + api_base=api_base, + model=model, + image=image, + timeout=litellm_params.get("timeout", None), + custom_llm_provider=custom_llm_provider, + logging_obj=litellm_logging_obj, + optional_params={}, + litellm_params=litellm_params, + ) + elif image_variation_provider == LITELLM_IMAGE_VARIATION_PROVIDERS.TOPAZ: + if api_key is None: + raise ValueError("API key is required for Topaz image variations") + if api_base is None: + raise ValueError("API base is required for Topaz image variations") + + response = base_llm_aiohttp_handler.image_variations( + model_response=model_response, + api_key=api_key, + api_base=api_base, + model=model, + image=image, + timeout=litellm_params.get("timeout", None) or DEFAULT_REQUEST_TIMEOUT, + custom_llm_provider=custom_llm_provider, + logging_obj=litellm_logging_obj, + optional_params={}, + litellm_params=litellm_params, + client=client, + ) + + # return the response + if response is None: + raise ValueError( + f"Invalid image variation provider: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" + ) + return response + + +@client +def image_edit( + image: FileTypes, + prompt: str, + model: Optional[str] = None, + mask: Optional[str] = None, + n: Optional[int] = None, + quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, + response_format: Optional[str] = None, + size: Optional[str] = None, + user: Optional[str] = None, + # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. + # The extra values given here take precedence over values defined on the client or passed to this method. + extra_headers: Optional[Dict[str, Any]] = None, + extra_query: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + # LiteLLM specific params, + custom_llm_provider: Optional[str] = None, + **kwargs, +) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse]]: + """ + Maps the image edit functionality, similar to OpenAI's images/edits endpoint. + """ + local_vars = locals() + try: + litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + _is_async = kwargs.pop("adelete_responses", False) is True + + # get llm provider logic + litellm_params = GenericLiteLLMParams(**kwargs) + model, custom_llm_provider, _, _ = get_llm_provider( + model=model or DEFAULT_IMAGE_ENDPOINT_MODEL, + custom_llm_provider=custom_llm_provider, + ) + + # get provider config + image_edit_provider_config: Optional[ + BaseImageEditConfig + ] = ProviderConfigManager.get_provider_image_edit_config( + model=model, + provider=litellm.LlmProviders(custom_llm_provider), + ) + + if image_edit_provider_config is None: + raise ValueError(f"image edit is not supported for {custom_llm_provider}") + + local_vars.update(kwargs) + # Get ImageEditOptionalRequestParams with only valid parameters + image_edit_optional_params: ImageEditOptionalRequestParams = ( + ImageEditRequestUtils.get_requested_image_edit_optional_param(local_vars) + ) + + # Get optional parameters for the responses API + image_edit_request_params: Dict = ( + ImageEditRequestUtils.get_optional_params_image_edit( + model=model, + image_edit_provider_config=image_edit_provider_config, + image_edit_optional_params=image_edit_optional_params, + ) + ) + + # Pre Call logging + litellm_logging_obj.update_environment_variables( + model=model, + user=user, + optional_params=dict(image_edit_request_params), + litellm_params={ + "litellm_call_id": litellm_call_id, + **image_edit_request_params, + }, + custom_llm_provider=custom_llm_provider, + ) + + # Call the handler with _is_async flag instead of directly calling the async handler + return base_llm_http_handler.image_edit_handler( + model=model, + image=image, + prompt=prompt, + image_edit_provider_config=image_edit_provider_config, + image_edit_optional_request_params=image_edit_request_params, + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + logging_obj=litellm_logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout or DEFAULT_REQUEST_TIMEOUT, + _is_async=_is_async, + client=kwargs.get("client"), + ) + + except Exception as e: + raise litellm.exception_type( + model=None, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +async def aimage_edit( + image: FileTypes, + model: str, + prompt: str, + mask: Optional[str] = None, + n: Optional[int] = None, + quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, + response_format: Optional[str] = None, + size: Optional[str] = None, + user: Optional[str] = None, + # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. + # The extra values given here take precedence over values defined on the client or passed to this method. + extra_headers: Optional[Dict[str, Any]] = None, + extra_query: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + # LiteLLM specific params, + custom_llm_provider: Optional[str] = None, + **kwargs, +) -> ImageResponse: + """ + Asynchronously calls the `image_edit` function with the given arguments and keyword arguments. + + Parameters: + - `args` (tuple): Positional arguments to be passed to the `image_edit` function. + - `kwargs` (dict): Keyword arguments to be passed to the `image_edit` function. + + Returns: + - `response` (Any): The response returned by the `image_edit` function. + """ + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["async_call"] = True + + # get custom llm provider so we can use this for mapping exceptions + if custom_llm_provider is None: + _, custom_llm_provider, _, _ = litellm.get_llm_provider( + model=model, api_base=local_vars.get("base_url", None) + ) + + func = partial( + image_edit, + image=image, + prompt=prompt, + mask=mask, + model=model, + n=n, + quality=quality, + response_format=response_format, + size=size, + user=user, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + **kwargs, + ) + + ctx = contextvars.copy_context() + func_with_context = partial(ctx.run, func) + init_response = await loop.run_in_executor(None, func_with_context) + + if asyncio.iscoroutine(init_response): + response = await init_response + else: + response = init_response + + return response + except Exception as e: + raise litellm.exception_type( + model=None, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) diff --git a/litellm/images/utils.py b/litellm/images/utils.py new file mode 100644 index 0000000000..4bf338605a --- /dev/null +++ b/litellm/images/utils.py @@ -0,0 +1,71 @@ +from typing import Any, Dict, cast, get_type_hints + +import litellm +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.types.images.main import ImageEditOptionalRequestParams + + +class ImageEditRequestUtils: + @staticmethod + def get_optional_params_image_edit( + model: str, + image_edit_provider_config: BaseImageEditConfig, + image_edit_optional_params: ImageEditOptionalRequestParams, + ) -> Dict: + """ + Get optional parameters for the image edit API. + + Args: + params: Dictionary of all parameters + model: The model name + image_edit_provider_config: The provider configuration for image edit API + + Returns: + A dictionary of supported parameters for the image edit API + """ + # Remove None values and internal parameters + + # Get supported parameters for the model + supported_params = image_edit_provider_config.get_supported_openai_params(model) + + # Check for unsupported parameters + unsupported_params = [ + param + for param in image_edit_optional_params + if param not in supported_params + ] + + if unsupported_params: + raise litellm.UnsupportedParamsError( + model=model, + message=f"The following parameters are not supported for model {model}: {', '.join(unsupported_params)}", + ) + + # Map parameters to provider-specific format + mapped_params = image_edit_provider_config.map_openai_params( + image_edit_optional_params=image_edit_optional_params, + model=model, + drop_params=litellm.drop_params, + ) + + return mapped_params + + @staticmethod + def get_requested_image_edit_optional_param( + params: Dict[str, Any], + ) -> ImageEditOptionalRequestParams: + """ + Filter parameters to only include those defined in ImageEditOptionalRequestParams. + + Args: + params: Dictionary of parameters to filter + + Returns: + ImageEditOptionalRequestParams instance with only the valid parameters + """ + valid_keys = get_type_hints(ImageEditOptionalRequestParams).keys() + filtered_params = { + k: v for k, v in params.items() if k in valid_keys and v is not None + } + + return cast(ImageEditOptionalRequestParams, filtered_params) diff --git a/litellm/llms/base_llm/__init__.py b/litellm/llms/base_llm/__init__.py index cd682a0dbe..187c985fd6 100644 --- a/litellm/llms/base_llm/__init__.py +++ b/litellm/llms/base_llm/__init__.py @@ -2,6 +2,7 @@ from .anthropic_messages.transformation import BaseAnthropicMessagesConfig from .audio_transcription.transformation import BaseAudioTranscriptionConfig from .chat.transformation import BaseConfig from .embedding.transformation import BaseEmbeddingConfig +from .image_edit.transformation import BaseImageEditConfig from .image_generation.transformation import BaseImageGenerationConfig __all__ = [ @@ -10,4 +11,5 @@ __all__ = [ "BaseAudioTranscriptionConfig", "BaseAnthropicMessagesConfig", "BaseEmbeddingConfig", + "BaseImageEditConfig", ] diff --git a/litellm/llms/base_llm/image_edit/transformation.py b/litellm/llms/base_llm/image_edit/transformation.py new file mode 100644 index 0000000000..d471f496af --- /dev/null +++ b/litellm/llms/base_llm/image_edit/transformation.py @@ -0,0 +1,120 @@ +import types +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple + +import httpx +from httpx._types import RequestFiles + +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.responses.main import * +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import FileTypes + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + from litellm.utils import ImageResponse as _ImageResponse + + from ..chat.transformation import BaseLLMException as _BaseLLMException + + LiteLLMLoggingObj = _LiteLLMLoggingObj + BaseLLMException = _BaseLLMException + ImageResponse = _ImageResponse +else: + LiteLLMLoggingObj = Any + BaseLLMException = Any + ImageResponse = Any + + +class BaseImageEditConfig(ABC): + def __init__(self): + pass + + @classmethod + def get_config(cls): + return { + k: v + for k, v in cls.__dict__.items() + if not k.startswith("__") + and not k.startswith("_abc") + and not isinstance( + v, + ( + types.FunctionType, + types.BuiltinFunctionType, + classmethod, + staticmethod, + ), + ) + and v is not None + } + + @abstractmethod + def get_supported_openai_params(self, model: str) -> list: + pass + + @abstractmethod + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + pass + + @abstractmethod + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + ) -> dict: + return {} + + @abstractmethod + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + OPTIONAL + + Get the complete url for the request + + Some providers need `model` in `api_base` + """ + if api_base is None: + raise ValueError("api_base is required") + return api_base + + @abstractmethod + 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]: + pass + + @abstractmethod + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ImageResponse: + pass + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + from ..chat.transformation import BaseLLMException + + raise BaseLLMException( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/llms/custom_httpx/aiohttp_handler.py b/litellm/llms/custom_httpx/aiohttp_handler.py index 13141fc19a..5a1d420865 100644 --- a/litellm/llms/custom_httpx/aiohttp_handler.py +++ b/litellm/llms/custom_httpx/aiohttp_handler.py @@ -102,7 +102,7 @@ class BaseLLMAIOHTTPHandler: api_base: str, headers: dict, data: dict, - timeout: Union[float, httpx.Timeout], + timeout: Optional[Union[float, httpx.Timeout]], litellm_params: dict, stream: bool = False, files: Optional[dict] = None, diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index f99e04ab9d..d026810c44 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, Callable, List, Mapping, Optional, Union import httpx from httpx import USE_CLIENT_DEFAULT, AsyncHTTPTransport, HTTPTransport +from httpx._types import RequestFiles import litellm from litellm.constants import _DEFAULT_TTL_FOR_HTTPX_CLIENTS @@ -199,6 +200,7 @@ class AsyncHTTPHandler: timeout: Optional[Union[float, httpx.Timeout]] = None, stream: bool = False, logging_obj: Optional[LiteLLMLoggingObject] = None, + files: Optional[RequestFiles] = None, ): start_time = time.time() try: @@ -206,7 +208,14 @@ class AsyncHTTPHandler: timeout = self.timeout req = self.client.build_request( - "POST", url, data=data, json=json, params=params, headers=headers, timeout=timeout # type: ignore + "POST", + url, + data=data, # type: ignore + json=json, + params=params, + headers=headers, + timeout=timeout, + files=files, ) response = await self.client.send(req, stream=stream) response.raise_for_status() @@ -533,7 +542,7 @@ class HTTPHandler: headers: Optional[dict] = None, stream: bool = False, timeout: Optional[Union[float, httpx.Timeout]] = None, - files: Optional[dict] = None, + files: Optional[Union[dict, RequestFiles]] = None, content: Any = None, logging_obj: Optional[LiteLLMLoggingObject] = None, ): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 4cf89accfc..cde532e7b0 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -30,6 +30,7 @@ from litellm.llms.base_llm.base_model_iterator import MockResponseIterator from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.base_llm.files.transformation import BaseFilesConfig +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig @@ -58,7 +59,12 @@ from litellm.types.rerank import OptionalRerankParams, RerankResponse from litellm.types.responses.main import DeleteResponseResult from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import EmbeddingResponse, FileTypes, TranscriptionResponse -from litellm.utils import CustomStreamWrapper, ModelResponse, ProviderConfigManager +from litellm.utils import ( + CustomStreamWrapper, + ImageResponse, + ModelResponse, + ProviderConfigManager, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -2018,7 +2024,9 @@ class BaseLLMHTTPHandler: def _handle_error( self, e: Exception, - provider_config: Union[BaseConfig, BaseRerankConfig, BaseResponsesAPIConfig], + provider_config: Union[ + BaseConfig, BaseRerankConfig, BaseResponsesAPIConfig, BaseImageEditConfig + ], ): status_code = getattr(e, "status_code", 500) error_headers = getattr(e, "headers", None) @@ -2098,3 +2106,191 @@ class BaseLLMHTTPHandler: raise Exception( f"Unexpected error while closing WebSocket: {close_error}" ) + + def image_edit_handler( + self, + model: str, + image: Any, + prompt: str, + image_edit_provider_config: BaseImageEditConfig, + image_edit_optional_request_params: Dict, + custom_llm_provider: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + timeout: Union[float, httpx.Timeout], + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + _is_async: bool = False, + fake_stream: bool = False, + litellm_metadata: Optional[Dict[str, Any]] = None, + ) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]: + """ + + Handles image edit requests. + When _is_async=True, returns a coroutine instead of making the call directly. + """ + if _is_async: + # Return the async coroutine if called with _is_async=True + return self.async_image_edit_handler( + model=model, + image=image, + prompt=prompt, + image_edit_provider_config=image_edit_provider_config, + image_edit_optional_request_params=image_edit_optional_request_params, + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + client=client if isinstance(client, AsyncHTTPHandler) else None, + fake_stream=fake_stream, + litellm_metadata=litellm_metadata, + ) + + if client is None or not isinstance(client, HTTPHandler): + sync_httpx_client = _get_httpx_client( + params={"ssl_verify": litellm_params.get("ssl_verify", None)} + ) + else: + sync_httpx_client = client + + headers = image_edit_provider_config.validate_environment( + api_key=litellm_params.api_key, + headers=image_edit_optional_request_params.get("extra_headers", {}) or {}, + model=model, + ) + + if extra_headers: + headers.update(extra_headers) + + api_base = image_edit_provider_config.get_complete_url( + api_base=litellm_params.api_base, + litellm_params=dict(litellm_params), + ) + + data, files = image_edit_provider_config.transform_image_edit_request( + model=model, + image=image, + prompt=prompt, + image_edit_optional_request_params=image_edit_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) + + ## LOGGING + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": headers, + }, + ) + + try: + response = sync_httpx_client.post( + url=api_base, + headers=headers, + data=data, + files=files, + timeout=timeout, + ) + + except Exception as e: + raise self._handle_error( + e=e, + provider_config=image_edit_provider_config, + ) + + return image_edit_provider_config.transform_image_edit_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) + + async def async_image_edit_handler( + self, + model: str, + image: FileTypes, + prompt: str, + image_edit_provider_config: BaseImageEditConfig, + image_edit_optional_request_params: Dict, + custom_llm_provider: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + timeout: Union[float, httpx.Timeout], + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + fake_stream: bool = False, + litellm_metadata: Optional[Dict[str, Any]] = None, + ) -> ImageResponse: + """ + Async version of the image edit handler. + Uses async HTTP client to make requests. + """ + if client is None or not isinstance(client, AsyncHTTPHandler): + async_httpx_client = get_async_httpx_client( + llm_provider=litellm.LlmProviders(custom_llm_provider), + params={"ssl_verify": litellm_params.get("ssl_verify", None)}, + ) + else: + async_httpx_client = client + + headers = image_edit_provider_config.validate_environment( + api_key=litellm_params.api_key, + headers=image_edit_optional_request_params.get("extra_headers", {}) or {}, + model=model, + ) + + if extra_headers: + headers.update(extra_headers) + + api_base = image_edit_provider_config.get_complete_url( + api_base=litellm_params.api_base, + litellm_params=dict(litellm_params), + ) + + data, files = image_edit_provider_config.transform_image_edit_request( + model=model, + image=image, + prompt=prompt, + image_edit_optional_request_params=image_edit_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) + + ## LOGGING + logging_obj.pre_call( + input=prompt, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": headers, + }, + ) + + try: + response = await async_httpx_client.post( + url=api_base, + headers=headers, + data=data, + files=files, + timeout=timeout, + ) + + except Exception as e: + raise self._handle_error( + e=e, + provider_config=image_edit_provider_config, + ) + + return image_edit_provider_config.transform_image_edit_response( + model=model, + raw_response=response, + logging_obj=logging_obj, + ) diff --git a/litellm/llms/openai/image_edit/transformation.py b/litellm/llms/openai/image_edit/transformation.py new file mode 100644 index 0000000000..bcec4aa029 --- /dev/null +++ b/litellm/llms/openai/image_edit/transformation.py @@ -0,0 +1,147 @@ +from io import BufferedReader +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast + +import httpx +from httpx._types import RequestFiles + +import litellm +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, + ImageEditRequestParams, +) +from litellm.types.llms.openai import FileTypes +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import ImageResponse + +from ..common_utils import OpenAIError + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class OpenAIImageEditConfig(BaseImageEditConfig): + def get_supported_openai_params(self, model: str) -> list: + """ + All OpenAI Image Edits params are supported + """ + return [ + "image", + "prompt", + "background", + "mask", + "model", + "n", + "quality", + "response_format", + "size", + "user", + "extra_headers", + "extra_query", + "extra_body", + "timeout", + ] + + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """No mapping applied since inputs are in OpenAI spec already""" + return dict(image_edit_optional_params) + + 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]: + """ + No transform applied since inputs are in OpenAI spec already + + This handles buffered readers as images to be sent as multipart/form-data for OpenAI + """ + request = ImageEditRequestParams( + model=model, + image=image, + prompt=prompt, + **image_edit_optional_request_params, + ) + request_dict = cast(Dict, request) + + ######################################################### + # Separate images as `files` and send other parameters as `data` + ######################################################### + _images = request_dict.get("image") or [] + data_without_images = {k: v for k, v in request_dict.items() if k != "image"} + files_list: List[Tuple[str, Any]] = [] + for _image in _images: + if isinstance(_image, BufferedReader): + files_list.append(("image[]", (_image.name, _image, "image/png"))) + else: + files_list.append(("image[]", (_image, "image/png"))) + return data_without_images, files_list + + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ImageResponse: + """No transform applied since outputs are in OpenAI spec already""" + try: + raw_response_json = raw_response.json() + except Exception: + raise OpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + return ImageResponse(**raw_response_json) + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + ) -> dict: + api_key = ( + api_key + or litellm.api_key + or litellm.openai_key + or get_secret_str("OPENAI_API_KEY") + ) + headers.update( + { + "Authorization": f"Bearer {api_key}", + } + ) + return headers + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Get the endpoint for OpenAI responses API + """ + api_base = ( + api_base + or litellm.api_base + or get_secret_str("OPENAI_BASE_URL") + or get_secret_str("OPENAI_API_BASE") + or "https://api.openai.com/v1" + ) + + # Remove trailing slashes + api_base = api_base.rstrip("/") + + return f"{api_base}/images/edits" diff --git a/litellm/llms/openai/image_variations/handler.py b/litellm/llms/openai/image_variations/handler.py index f738115a29..8b96fb6ef7 100644 --- a/litellm/llms/openai/image_variations/handler.py +++ b/litellm/llms/openai/image_variations/handler.py @@ -50,7 +50,7 @@ class OpenAIImageVariationsHandler: data: dict, headers: dict, model: Optional[str], - timeout: float, + timeout: Optional[float], max_retries: int, logging_obj: LiteLLMLoggingObj, model_response: ImageResponse, @@ -123,7 +123,7 @@ class OpenAIImageVariationsHandler: api_base: str, model: Optional[str], image: FileTypes, - timeout: float, + timeout: Optional[float], custom_llm_provider: str, logging_obj: LiteLLMLoggingObj, optional_params: dict, diff --git a/litellm/main.py b/litellm/main.py index 68589d7127..7cae5acd97 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -183,12 +183,10 @@ from .types.llms.openai import ( ChatCompletionPredictionContentParam, ChatCompletionUserMessage, HttpxBinaryResponseContent, - ImageGenerationRequestQuality, OpenAIModerationResponse, OpenAIWebSearchOptions, ) from .types.utils import ( - LITELLM_IMAGE_VARIATION_PROVIDERS, AdapterCompletionStreamWrapper, ChatCompletionMessageToolCall, CompletionTokensDetails, @@ -204,7 +202,6 @@ encoding = tiktoken.get_encoding("cl100k_base") from litellm.utils import ( Choices, EmbeddingResponse, - ImageResponse, Message, ModelResponse, TextChoices, @@ -4578,516 +4575,6 @@ async def amoderation( ) -##### Image Generation ####################### -@client -async def aimage_generation(*args, **kwargs) -> ImageResponse: - """ - Asynchronously calls the `image_generation` function with the given arguments and keyword arguments. - - Parameters: - - `args` (tuple): Positional arguments to be passed to the `image_generation` function. - - `kwargs` (dict): Keyword arguments to be passed to the `image_generation` function. - - Returns: - - `response` (Any): The response returned by the `image_generation` function. - """ - loop = asyncio.get_event_loop() - model = args[0] if len(args) > 0 else kwargs["model"] - ### PASS ARGS TO Image Generation ### - kwargs["aimg_generation"] = True - custom_llm_provider = None - try: - # Use a partial function to pass your keyword arguments - func = partial(image_generation, *args, **kwargs) - - # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - - _, custom_llm_provider, _, _ = get_llm_provider( - model=model, api_base=kwargs.get("api_base", None) - ) - - # Await normally - init_response = await loop.run_in_executor(None, func_with_context) - if isinstance(init_response, dict) or isinstance( - init_response, ImageResponse - ): ## CACHING SCENARIO - if isinstance(init_response, dict): - init_response = ImageResponse(**init_response) - response = init_response - elif asyncio.iscoroutine(init_response): - response = await init_response # type: ignore - else: - # Call the synchronous function using run_in_executor - response = await loop.run_in_executor(None, func_with_context) - return response - except Exception as e: - custom_llm_provider = custom_llm_provider or "openai" - raise exception_type( - model=model, - custom_llm_provider=custom_llm_provider, - original_exception=e, - completion_kwargs=args, - extra_kwargs=kwargs, - ) - - -@client -def image_generation( # noqa: PLR0915 - prompt: str, - model: Optional[str] = None, - n: Optional[int] = None, - quality: Optional[Union[str, ImageGenerationRequestQuality]] = None, - response_format: Optional[str] = None, - size: Optional[str] = None, - style: Optional[str] = None, - user: Optional[str] = None, - timeout=600, # default to 10 minutes - api_key: Optional[str] = None, - api_base: Optional[str] = None, - api_version: Optional[str] = None, - custom_llm_provider=None, - **kwargs, -) -> ImageResponse: - """ - Maps the https://api.openai.com/v1/images/generations endpoint. - - Currently supports just Azure + OpenAI. - """ - try: - args = locals() - aimg_generation = kwargs.get("aimg_generation", False) - litellm_call_id = kwargs.get("litellm_call_id", None) - logger_fn = kwargs.get("logger_fn", None) - mock_response: Optional[str] = kwargs.get("mock_response", None) # type: ignore - proxy_server_request = kwargs.get("proxy_server_request", None) - azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None) - model_info = kwargs.get("model_info", None) - metadata = kwargs.get("metadata", {}) - litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore - client = kwargs.get("client", None) - extra_headers = kwargs.get("extra_headers", None) - headers: dict = kwargs.get("headers", None) or {} - base_model = kwargs.get("base_model", None) - if extra_headers is not None: - headers.update(extra_headers) - model_response: ImageResponse = litellm.utils.ImageResponse() - dynamic_api_key: Optional[str] = None - if model is not None or custom_llm_provider is not None: - model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( - model=model, # type: ignore - custom_llm_provider=custom_llm_provider, - api_base=api_base, - ) - else: - model = "dall-e-2" - custom_llm_provider = "openai" # default to dall-e-2 on openai - model_response._hidden_params["model"] = model - openai_params = [ - "user", - "request_timeout", - "api_base", - "api_version", - "api_key", - "deployment_id", - "organization", - "base_url", - "default_headers", - "timeout", - "max_retries", - "n", - "quality", - "size", - "style", - ] - litellm_params = all_litellm_params - default_params = openai_params + litellm_params - non_default_params = { - k: v for k, v in kwargs.items() if k not in default_params - } # model-specific params - pass them straight to the model/provider - - image_generation_config: Optional[BaseImageGenerationConfig] = None - if ( - custom_llm_provider is not None - and custom_llm_provider in LlmProviders._member_map_.values() - ): - image_generation_config = ( - ProviderConfigManager.get_provider_image_generation_config( - model=base_model or model, - provider=LlmProviders(custom_llm_provider), - ) - ) - - optional_params = get_optional_params_image_gen( - model=base_model or model, - n=n, - quality=quality, - response_format=response_format, - size=size, - style=style, - user=user, - custom_llm_provider=custom_llm_provider, - provider_config=image_generation_config, - **non_default_params, - ) - - litellm_params_dict = get_litellm_params(**kwargs) - - logging: Logging = litellm_logging_obj - logging.update_environment_variables( - model=model, - user=user, - optional_params=optional_params, - litellm_params={ - "timeout": timeout, - "azure": False, - "litellm_call_id": litellm_call_id, - "logger_fn": logger_fn, - "proxy_server_request": proxy_server_request, - "model_info": model_info, - "metadata": metadata, - "preset_cache_key": None, - "stream_response": {}, - }, - custom_llm_provider=custom_llm_provider, - ) - if "custom_llm_provider" not in logging.model_call_details: - logging.model_call_details["custom_llm_provider"] = custom_llm_provider - if mock_response is not None: - return mock_image_generation(model=model, mock_response=mock_response) - - if custom_llm_provider == "azure": - # azure configs - api_type = get_secret_str("AZURE_API_TYPE") or "azure" - - api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") - - api_version = ( - api_version - or litellm.api_version - or get_secret_str("AZURE_API_VERSION") - ) - - api_key = ( - api_key - or litellm.api_key - or litellm.azure_key - or get_secret_str("AZURE_OPENAI_API_KEY") - or get_secret_str("AZURE_API_KEY") - ) - - azure_ad_token = optional_params.pop( - "azure_ad_token", None - ) or get_secret_str("AZURE_AD_TOKEN") - - default_headers = { - "Content-Type": "application/json;", - "api-key": api_key, - } - for k, v in default_headers.items(): - if k not in headers: - headers[k] = v - - model_response = azure_chat_completions.image_generation( - model=model, - prompt=prompt, - timeout=timeout, - api_key=api_key, - api_base=api_base, - azure_ad_token=azure_ad_token, - azure_ad_token_provider=azure_ad_token_provider, - logging_obj=litellm_logging_obj, - optional_params=optional_params, - model_response=model_response, - api_version=api_version, - aimg_generation=aimg_generation, - client=client, - headers=headers, - litellm_params=litellm_params_dict, - ) - elif ( - custom_llm_provider == "openai" - or custom_llm_provider in litellm.openai_compatible_providers - ): - model_response = openai_chat_completions.image_generation( - model=model, - prompt=prompt, - timeout=timeout, - api_key=api_key or dynamic_api_key, - api_base=api_base, - logging_obj=litellm_logging_obj, - optional_params=optional_params, - model_response=model_response, - aimg_generation=aimg_generation, - client=client, - ) - elif custom_llm_provider == "bedrock": - if model is None: - raise Exception("Model needs to be set for bedrock") - model_response = bedrock_image_generation.image_generation( # type: ignore - model=model, - prompt=prompt, - timeout=timeout, - logging_obj=litellm_logging_obj, - optional_params=optional_params, - model_response=model_response, - aimg_generation=aimg_generation, - client=client, - ) - elif custom_llm_provider == "vertex_ai": - vertex_ai_project = ( - optional_params.pop("vertex_project", None) - or optional_params.pop("vertex_ai_project", None) - or litellm.vertex_project - or get_secret_str("VERTEXAI_PROJECT") - ) - vertex_ai_location = ( - optional_params.pop("vertex_location", None) - or optional_params.pop("vertex_ai_location", None) - or litellm.vertex_location - or get_secret_str("VERTEXAI_LOCATION") - ) - vertex_credentials = ( - optional_params.pop("vertex_credentials", None) - or optional_params.pop("vertex_ai_credentials", None) - or get_secret_str("VERTEXAI_CREDENTIALS") - ) - - api_base = ( - api_base - or litellm.api_base - or get_secret_str("VERTEXAI_API_BASE") - or get_secret_str("VERTEX_API_BASE") - ) - - model_response = vertex_image_generation.image_generation( - model=model, - prompt=prompt, - timeout=timeout, - logging_obj=litellm_logging_obj, - optional_params=optional_params, - model_response=model_response, - vertex_project=vertex_ai_project, - vertex_location=vertex_ai_location, - vertex_credentials=vertex_credentials, - aimg_generation=aimg_generation, - api_base=api_base, - client=client, - ) - elif ( - custom_llm_provider in litellm._custom_providers - ): # Assume custom LLM provider - # Get the Custom Handler - custom_handler: Optional[CustomLLM] = None - for item in litellm.custom_provider_map: - if item["provider"] == custom_llm_provider: - custom_handler = item["custom_handler"] - - if custom_handler is None: - raise LiteLLMUnknownProvider( - model=model, custom_llm_provider=custom_llm_provider - ) - - ## ROUTE LLM CALL ## - if aimg_generation is True: - async_custom_client: Optional[AsyncHTTPHandler] = None - if client is not None and isinstance(client, AsyncHTTPHandler): - async_custom_client = client - - ## CALL FUNCTION - model_response = custom_handler.aimage_generation( # type: ignore - model=model, - prompt=prompt, - api_key=api_key, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - logging_obj=litellm_logging_obj, - timeout=timeout, - client=async_custom_client, - ) - else: - custom_client: Optional[HTTPHandler] = None - if client is not None and isinstance(client, HTTPHandler): - custom_client = client - - ## CALL FUNCTION - model_response = custom_handler.image_generation( - model=model, - prompt=prompt, - api_key=api_key, - api_base=api_base, - model_response=model_response, - optional_params=optional_params, - logging_obj=litellm_logging_obj, - timeout=timeout, - client=custom_client, - ) - - return model_response - except Exception as e: - ## Map to OpenAI Exception - raise exception_type( - model=model, - custom_llm_provider=custom_llm_provider, - original_exception=e, - completion_kwargs=locals(), - extra_kwargs=kwargs, - ) - - -@client -async def aimage_variation(*args, **kwargs) -> ImageResponse: - """ - Asynchronously calls the `image_variation` function with the given arguments and keyword arguments. - - Parameters: - - `args` (tuple): Positional arguments to be passed to the `image_variation` function. - - `kwargs` (dict): Keyword arguments to be passed to the `image_variation` function. - - Returns: - - `response` (Any): The response returned by the `image_variation` function. - """ - loop = asyncio.get_event_loop() - model = kwargs.get("model", None) - custom_llm_provider = kwargs.get("custom_llm_provider", None) - ### PASS ARGS TO Image Generation ### - kwargs["async_call"] = True - try: - # Use a partial function to pass your keyword arguments - func = partial(image_variation, *args, **kwargs) - - # Add the context to the function - ctx = contextvars.copy_context() - func_with_context = partial(ctx.run, func) - - if custom_llm_provider is None and model is not None: - _, custom_llm_provider, _, _ = get_llm_provider( - model=model, api_base=kwargs.get("api_base", None) - ) - - # Await normally - init_response = await loop.run_in_executor(None, func_with_context) - if isinstance(init_response, dict) or isinstance( - init_response, ImageResponse - ): ## CACHING SCENARIO - if isinstance(init_response, dict): - init_response = ImageResponse(**init_response) - response = init_response - elif asyncio.iscoroutine(init_response): - response = await init_response # type: ignore - else: - # Call the synchronous function using run_in_executor - response = await loop.run_in_executor(None, func_with_context) - return response - except Exception as e: - custom_llm_provider = custom_llm_provider or "openai" - raise exception_type( - model=model, - custom_llm_provider=custom_llm_provider, - original_exception=e, - completion_kwargs=args, - extra_kwargs=kwargs, - ) - - -@client -def image_variation( - image: FileTypes, - model: str = "dall-e-2", # set to dall-e-2 by default - like OpenAI. - n: int = 1, - response_format: Literal["url", "b64_json"] = "url", - size: Optional[str] = None, - user: Optional[str] = None, - **kwargs, -) -> ImageResponse: - # get non-default params - client = kwargs.get("client", None) - # get logging object - litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj")) - - # get the litellm params - litellm_params = get_litellm_params(**kwargs) - # get the custom llm provider - model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( - model=model, - custom_llm_provider=litellm_params.get("custom_llm_provider", None), - api_base=litellm_params.get("api_base", None), - api_key=litellm_params.get("api_key", None), - ) - - # route to the correct provider w/ the params - try: - llm_provider = LlmProviders(custom_llm_provider) - image_variation_provider = LITELLM_IMAGE_VARIATION_PROVIDERS(llm_provider) - except ValueError: - raise ValueError( - f"Invalid image variation provider: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" - ) - model_response = ImageResponse() - - response: Optional[ImageResponse] = None - - provider_config = ProviderConfigManager.get_provider_model_info( - model=model or "", # openai defaults to dall-e-2 - provider=llm_provider, - ) - - if provider_config is None: - raise ValueError( - f"image variation provider has no known model info config - required for getting api keys, etc.: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" - ) - - api_key = provider_config.get_api_key(litellm_params.get("api_key", None)) - api_base = provider_config.get_api_base(litellm_params.get("api_base", None)) - - if image_variation_provider == LITELLM_IMAGE_VARIATION_PROVIDERS.OPENAI: - if api_key is None: - raise ValueError("API key is required for OpenAI image variations") - if api_base is None: - raise ValueError("API base is required for OpenAI image variations") - - response = openai_image_variations.image_variations( - model_response=model_response, - api_key=api_key, - api_base=api_base, - model=model, - image=image, - timeout=litellm_params.get("timeout", None), - custom_llm_provider=custom_llm_provider, - logging_obj=litellm_logging_obj, - optional_params={}, - litellm_params=litellm_params, - ) - elif image_variation_provider == LITELLM_IMAGE_VARIATION_PROVIDERS.TOPAZ: - if api_key is None: - raise ValueError("API key is required for Topaz image variations") - if api_base is None: - raise ValueError("API base is required for Topaz image variations") - - response = base_llm_aiohttp_handler.image_variations( - model_response=model_response, - api_key=api_key, - api_base=api_base, - model=model, - image=image, - timeout=litellm_params.get("timeout", None), - custom_llm_provider=custom_llm_provider, - logging_obj=litellm_logging_obj, - optional_params={}, - litellm_params=litellm_params, - client=client, - ) - - # return the response - if response is None: - raise ValueError( - f"Invalid image variation provider: {custom_llm_provider}. Supported providers are: {LITELLM_IMAGE_VARIATION_PROVIDERS}" - ) - return response - - ##### Transcription ####################### diff --git a/litellm/types/images/main.py b/litellm/types/images/main.py new file mode 100644 index 0000000000..1e7466f4e8 --- /dev/null +++ b/litellm/types/images/main.py @@ -0,0 +1,31 @@ +from typing import Any, Dict, List, Literal, Optional, TypedDict, Union + +from litellm.types.utils import FileTypes + + +class ImageEditOptionalRequestParams(TypedDict, total=False): + """ + TypedDict for Optional parameters supported by OpenAI's image edit API. + + Params here: https://platform.openai.com/docs/api-reference/images/createEdit + """ + + background: Optional[Literal["transparent", "opaque", "auto"]] + mask: Optional[str] + n: Optional[int] + quality: Optional[Literal["high", "medium", "low", "standard", "auto"]] + response_format: Optional[Literal["url", "b64_json"]] + size: Optional[str] + user: Optional[str] + + +class ImageEditRequestParams(ImageEditOptionalRequestParams, total=False): + """ + TypedDict for request parameters supported by OpenAI's image edit API. + + Params here: https://platform.openai.com/docs/api-reference/images/createEdit + """ + + image: FileTypes + prompt: str + model: Optional[str] diff --git a/litellm/utils.py b/litellm/utils.py index 0dd4c8cd1e..773196077d 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -229,6 +229,7 @@ from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.llms.base_llm.files.transformation import BaseFilesConfig +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) @@ -1967,6 +1968,7 @@ def supports_prompt_caching( key="supports_prompt_caching", ) + def supports_computer_use( model: str, custom_llm_provider: Optional[str] = None ) -> bool: @@ -6681,6 +6683,19 @@ class ProviderConfigManager: return GeminiRealtimeConfig() return None + @staticmethod + def get_provider_image_edit_config( + model: str, + provider: LlmProviders, + ) -> Optional[BaseImageEditConfig]: + if LlmProviders.OPENAI == provider: + from litellm.llms.openai.image_edit.transformation import ( + OpenAIImageEditConfig, + ) + + return OpenAIImageEditConfig() + return None + def get_end_user_id_for_cost_tracking( litellm_params: dict, diff --git a/tests/image_gen_tests/ishaan_github.png b/tests/image_gen_tests/ishaan_github.png new file mode 100644 index 0000000000..2e6c0dadc9 Binary files /dev/null and b/tests/image_gen_tests/ishaan_github.png differ diff --git a/tests/image_gen_tests/litellm_site.png b/tests/image_gen_tests/litellm_site.png new file mode 100644 index 0000000000..786df87b93 Binary files /dev/null and b/tests/image_gen_tests/litellm_site.png differ diff --git a/tests/image_gen_tests/test_image_edit.png b/tests/image_gen_tests/test_image_edit.png new file mode 100644 index 0000000000..c63c01cc7c Binary files /dev/null and b/tests/image_gen_tests/test_image_edit.png differ diff --git a/tests/image_gen_tests/test_image_edits.py b/tests/image_gen_tests/test_image_edits.py new file mode 100644 index 0000000000..67c529a9bf --- /dev/null +++ b/tests/image_gen_tests/test_image_edits.py @@ -0,0 +1,57 @@ + + +import logging +import os +import sys +import traceback +import pytest +import base64 + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.utils import ImageResponse +# Get the current directory of the file being run +pwd = os.path.dirname(os.path.realpath(__file__)) + +TEST_IMAGES = [ + open(os.path.join(pwd, "ishaan_github.png"), "rb"), + open(os.path.join(pwd, "litellm_site.png"), "rb"), +] + +@pytest.mark.parametrize("sync_mode", [True]) +@pytest.mark.asyncio +async def test_openai_image_edit_litellm_sdk(sync_mode): + from litellm import image_edit, aimage_edit + litellm._turn_on_debug() + + prompt = """ + Create a studio ghibli style image that combines all the reference images. Make sure the person looks like a CTO. + """ + + if sync_mode: + result = image_edit( + prompt=prompt, + model="gpt-image-1", + image=TEST_IMAGES, + ) + else: + result = await aimage_edit( + prompt=prompt, + model="gpt-image-1", + image=TEST_IMAGES, + ) + print("result from image edit", result) + + # Validate the response meets expected schema + ImageResponse.model_validate(result) + + if isinstance(result, ImageResponse): + image_base64 = result.data[0].b64_json + image_bytes = base64.b64decode(image_base64) + + # Save the image to a file + with open("test_image_edit.png", "wb") as f: + f.write(image_bytes)