Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/image_endpoints/endpoints.py: 58%
112 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import asyncio
2import io
3from collections.abc import Sequence
4from typing import Final, get_type_hints
6import orjson
7from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, UploadFile, status
8from fastapi.responses import ORJSONResponse
10import litellm
11from litellm.litellm_core_utils.prompt_templates.common_utils import (
12 get_str_from_messages,
13)
14from litellm.proxy._types import *
15from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
16from litellm.proxy.common_request_processing import (
17 ProxyBaseLLMRequestProcessing,
18 log_llm_api_exception,
19 resolve_litellm_call_id,
20)
21from litellm.proxy.common_utils.http_parsing_utils import (
22 coerce_numeric_form_fields,
23 numeric_form_fields,
24)
25from litellm.proxy.common_utils.openai_error_payload import (
26 error_status_code,
27 litellm_call_id_headers,
28 openai_error_param,
29 openai_error_type,
30)
31from litellm.proxy.route_llm_request import route_request
32from litellm.types.images.main import ImageEditRequestParams
33from litellm.types.llms.openai import ChatCompletionUserMessage
35router: Final = APIRouter()
37IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(ImageEditRequestParams))
39IMAGE_ARRAY_FIELD: Final = "image[]"
40MASK_ARRAY_FIELD: Final = "mask[]"
41BRACKETED_FILE_FIELDS: Final = frozenset({IMAGE_ARRAY_FIELD, MASK_ARRAY_FIELD})
44async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO:
45 """
46 Read a FastAPI UploadFile into a BytesIO and set .name so OpenAI SDK
47 infers filename/content-type correctly.
48 """
49 data: Final = await upload.read()
50 buffer: Final = io.BytesIO(data)
51 buffer.name = upload.filename
52 return buffer
55async def batch_to_bytesio(
56 uploads: Sequence[UploadFile] | None,
57) -> list[io.BytesIO] | None:
58 """
59 Convert a sequence of UploadFiles to a list of BytesIO buffers, or None.
60 """
61 if not uploads:
62 return None
63 return [await uploadfile_to_bytesio(u) for u in uploads]
66@router.post(
67 "/v1/images/generations",
68 dependencies=[Depends(user_api_key_auth)],
69 response_class=ORJSONResponse,
70 tags=["images"],
71)
72@router.post(
73 "/images/generations",
74 dependencies=[Depends(user_api_key_auth)],
75 response_class=ORJSONResponse,
76 tags=["images"],
77)
78@router.post(
79 "/openai/deployments/{model:path}/images/generations",
80 dependencies=[Depends(user_api_key_auth)],
81 response_class=ORJSONResponse,
82 tags=["images"],
83) # azure compatible endpoint
84async def image_generation(
85 request: Request,
86 fastapi_response: Response,
87 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
88 model: str | None = None,
89):
90 from litellm.proxy.litellm_pre_call_utils import reject_url_valued_destination
91 from litellm.proxy.proxy_server import (
92 add_litellm_data_to_request,
93 general_settings,
94 llm_router,
95 proxy_config,
96 proxy_logging_obj,
97 user_model,
98 version,
99 )
101 litellm_call_id: Final = resolve_litellm_call_id(request.headers.get("x-litellm-call-id"))
102 data = {"litellm_call_id": litellm_call_id}
103 try:
104 # Use orjson to parse JSON data, orjson speeds up requests significantly
105 body: Final = await request.body()
106 data = orjson.loads(body) | data
108 # Include original request and headers in the data
109 data = await add_litellm_data_to_request(
110 data=data,
111 request=request,
112 general_settings=general_settings,
113 user_api_key_dict=user_api_key_dict,
114 version=version,
115 proxy_config=proxy_config,
116 )
118 if isinstance(model, str):
119 reject_url_valued_destination("model", model)
121 data["model"] = (
122 model
123 or general_settings.get("image_generation_model", None) # server default
124 or user_model # model name passed via cli args
125 or data.get("model", None) # default passed in http request
126 )
127 if user_model:
128 data["model"] = user_model
130 ### MODEL ALIAS MAPPING ###
131 # check if model name in model alias map
132 # get the actual model name
133 if data["model"] in litellm.model_alias_map:
134 data["model"] = litellm.model_alias_map[data["model"]]
136 ### CALL HOOKS ### - modify incoming data / reject request before calling the model
137 prompt_value: Final = data.get("prompt")
138 if prompt_value is not None:
139 # Reformat the image prompt as a chat message so guardrails can process it.
140 user_message: Final[ChatCompletionUserMessage] = {
141 "role": "user",
142 "content": prompt_value,
143 }
144 data["messages"] = [user_message]
145 data = await proxy_logging_obj.pre_call_hook(
146 user_api_key_dict=user_api_key_dict, data=data, call_type="image_generation"
147 )
149 messages: Final = data.get("messages")
150 if isinstance(messages, list) and messages:
151 data["prompt"] = get_str_from_messages(messages)
152 data.pop("messages", None)
154 ## ROUTE TO CORRECT ENDPOINT ##
155 llm_call: Final = await route_request(
156 data=data,
157 route_type="aimage_generation",
158 llm_router=llm_router,
159 user_model=user_model,
160 )
161 response = await llm_call
163 ### ALERTING ###
164 asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
166 ### CALL HOOKS ### - modify outgoing data (guardrails, otel, etc.)
167 response = await proxy_logging_obj.post_call_success_hook(
168 data=data, user_api_key_dict=user_api_key_dict, response=response
169 )
171 ### RESPONSE HEADERS ###
172 hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
173 model_id: Final = hidden_params.get("model_id", None) or ""
174 cache_key: Final = hidden_params.get("cache_key", None) or ""
175 api_base: Final = hidden_params.get("api_base", None) or ""
176 response_cost: Final = hidden_params.get("response_cost", None) or ""
177 response_call_id: Final = hidden_params.get("litellm_call_id", None) or ""
179 fastapi_response.headers.update(
180 ProxyBaseLLMRequestProcessing.get_custom_headers(
181 user_api_key_dict=user_api_key_dict,
182 model_id=model_id,
183 cache_key=cache_key,
184 api_base=api_base,
185 version=version,
186 response_cost=response_cost,
187 model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
188 call_id=response_call_id,
189 request_data=data,
190 hidden_params=hidden_params,
191 )
192 )
194 # Call response headers hook (matches base_process_llm_request behavior)
195 callback_headers: Final = await proxy_logging_obj.post_call_response_headers_hook(
196 data=data,
197 user_api_key_dict=user_api_key_dict,
198 response=response,
199 request_headers=dict(request.headers),
200 )
201 if callback_headers:
202 fastapi_response.headers.update(callback_headers)
204 return response
205 except Exception as e:
206 await proxy_logging_obj.post_call_failure_hook(
207 user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
208 )
209 log_llm_api_exception(e, litellm_call_id)
210 if isinstance(e, HTTPException): 210 ↛ 211line 210 didn't jump to line 211 because the condition on line 210 was never true
211 raise ProxyException(
212 message=getattr(e, "message", str(e)),
213 type=openai_error_type(e, error_status_code(e, status.HTTP_400_BAD_REQUEST)),
214 param=openai_error_param(e),
215 headers=litellm_call_id_headers(litellm_call_id),
216 code=error_status_code(e, status.HTTP_400_BAD_REQUEST),
217 )
218 else:
219 error_msg: Final = f"{e}"
220 raise ProxyException(
221 message=getattr(e, "message", error_msg),
222 type=openai_error_type(e, error_status_code(e, 500)),
223 param=openai_error_param(e),
224 headers=litellm_call_id_headers(litellm_call_id),
225 openai_code=getattr(e, "code", None),
226 code=error_status_code(e, 500),
227 )
230@router.post(
231 "/v1/images/edits",
232 dependencies=[Depends(user_api_key_auth)],
233 tags=["images"],
234)
235@router.post(
236 "/images/edits",
237 dependencies=[Depends(user_api_key_auth)],
238 tags=["images"],
239)
240@router.post(
241 "/openai/deployments/{model:path}/images/edits",
242 dependencies=[Depends(user_api_key_auth)],
243 response_class=ORJSONResponse,
244 tags=["images"],
245) # azure compatible endpoint
246async def image_edit_api(
247 request: Request,
248 fastapi_response: Response,
249 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
250 image: list[UploadFile] | None = File(None),
251 image_array: list[UploadFile] | None = File(None, alias=IMAGE_ARRAY_FIELD),
252 mask: list[UploadFile] | None = File(None),
253 mask_array: list[UploadFile] | None = File(None, alias=MASK_ARRAY_FIELD),
254 model: str | None = None,
255):
256 """
257 Follows the OpenAI Images API spec: https://platform.openai.com/docs/api-reference/images/create
259 ```bash
260 curl -s -D >(grep -i x-request-id >&2) \
261 -o >(jq -r '.data[0].b64_json' | base64 --decode > gift-basket.png) \
262 -X POST "http://localhost:4000/v1/images/edits" \
263 -H "Authorization: Bearer sk-1234" \
264 -F "model=gpt-image-1" \
265 -F "image[]=@soap.png" \
266 -F 'prompt=Create a studio ghibli image of this'
267 ```
268 """
269 if image is not None and image_array is not None: 269 ↛ 270line 269 didn't jump to line 270 because the condition on line 269 was never true
270 raise HTTPException(status_code=422, detail="Cannot specify both 'image' and 'image[]'")
271 if mask is not None and mask_array is not None: 271 ↛ 272line 271 didn't jump to line 272 because the condition on line 271 was never true
272 raise HTTPException(status_code=422, detail="Cannot specify both 'mask' and 'mask[]'")
273 if image is None and image_array is not None: 273 ↛ 274line 273 didn't jump to line 274 because the condition on line 273 was never true
274 image = image_array
275 if mask is None and mask_array is not None: 275 ↛ 276line 275 didn't jump to line 276 because the condition on line 275 was never true
276 mask = mask_array
278 # if image is None:
279 # raise HTTPException(status_code=422, detail="Field required: image")
280 # Note: Image is optional for some models (e.g., Bedrock Stability style-transfer)
281 # The validation will be done at the model level if image is truly required
283 from litellm.proxy.proxy_server import (
284 _read_request_body,
285 general_settings,
286 llm_router,
287 proxy_config,
288 proxy_logging_obj,
289 select_data_generator,
290 user_api_base,
291 user_max_tokens,
292 user_model,
293 user_request_timeout,
294 user_temperature,
295 version,
296 )
298 #########################################################
299 # Read request body and convert UploadFiles to BytesIO
300 #########################################################
301 data: Final = {
302 key: value
303 for key, value in coerce_numeric_form_fields(
304 parsed_body=await _read_request_body(request=request),
305 numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
306 ).items()
307 if key not in BRACKETED_FILE_FIELDS
308 }
309 image_files: Final = await batch_to_bytesio(image)
310 mask_files: Final = await batch_to_bytesio(mask)
311 if image_files: 311 ↛ 312line 311 didn't jump to line 312 because the condition on line 311 was never true
312 data["image"] = image_files
313 if mask_files: 313 ↛ 314line 313 didn't jump to line 314 because the condition on line 313 was never true
314 data["mask"] = mask_files
316 for _field in ("image", "mask"):
317 if _field in data and isinstance(data[_field], str): 317 ↛ 318line 317 didn't jump to line 318 because the condition on line 317 was never true
318 raise HTTPException(
319 status_code=422,
320 detail=f"'{_field}' must be provided as a multipart file upload, not a string.",
321 )
323 # Ensure prompt exists in data (default to None for models that don't require it)
324 if "prompt" not in data: 324 ↛ 327line 324 didn't jump to line 327 because the condition on line 324 was always true
325 data["prompt"] = None
327 data["model"] = (
328 model
329 or general_settings.get("image_generation_model", None) # server default
330 or user_model # model name passed via cli args
331 or data.get("model", None) # default passed in http request
332 )
333 #########################################################
334 # Process request
335 #########################################################
337 processor: Final = ProxyBaseLLMRequestProcessing(data=data)
338 try:
339 return await processor.base_process_llm_request(
340 request=request,
341 fastapi_response=fastapi_response,
342 user_api_key_dict=user_api_key_dict,
343 route_type="aimage_edit",
344 proxy_logging_obj=proxy_logging_obj,
345 llm_router=llm_router,
346 general_settings=general_settings,
347 proxy_config=proxy_config,
348 select_data_generator=select_data_generator,
349 model=None,
350 user_model=user_model,
351 user_temperature=user_temperature,
352 user_request_timeout=user_request_timeout,
353 user_max_tokens=user_max_tokens,
354 user_api_base=user_api_base,
355 version=version,
356 )
357 except Exception as e:
358 raise await processor._handle_llm_api_exception(
359 e=e,
360 user_api_key_dict=user_api_key_dict,
361 proxy_logging_obj=proxy_logging_obj,
362 version=version,
363 )