mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(proxy): accept image-edit URL and data-URI form fields
FastAPI File() rejected string image/mask values before the handler ran, so Grok Imagine edits that pass an http(s) URL or data:image URI never reached xAI. Parse those fields from the form instead and keep multipart uploads working Co-authored-by: HX <HXhlx@users.noreply.github.com>
This commit is contained in:
parent
f4d012c9a3
commit
c25883d55a
2 changed files with 110 additions and 69 deletions
|
|
@ -39,6 +39,8 @@ _SIZE_TO_ASPECT_RATIO: Final = {
|
|||
"1920x1080": "16:9",
|
||||
"1080x1920": "9:16",
|
||||
}
|
||||
|
||||
|
||||
def _duration_from_seconds(seconds: object) -> int:
|
||||
try:
|
||||
return int(seconds) if seconds is not None else 6
|
||||
|
|
@ -110,9 +112,7 @@ class XAIVideoConfig(BaseVideoConfig):
|
|||
from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth
|
||||
|
||||
params: Final = (
|
||||
litellm_params.model_dump()
|
||||
if isinstance(litellm_params, GenericLiteLLMParams)
|
||||
else (litellm_params or {})
|
||||
litellm_params.model_dump() if isinstance(litellm_params, GenericLiteLLMParams) else (litellm_params or {})
|
||||
)
|
||||
if should_use_xai_oauth(params) and not XAIModelInfo.get_api_key(api_key):
|
||||
return XAIOAuthAuthenticator().get_api_base().rstrip("/")
|
||||
|
|
@ -243,9 +243,7 @@ class XAIVideoConfig(BaseVideoConfig):
|
|||
progress=0,
|
||||
)
|
||||
if custom_llm_provider:
|
||||
video_obj.id = encode_video_id_with_provider(
|
||||
video_obj.id, custom_llm_provider, model
|
||||
)
|
||||
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model)
|
||||
video_obj.usage = usage if isinstance(usage, dict) else {}
|
||||
video_obj._hidden_params["video_url"] = None
|
||||
return video_obj
|
||||
|
|
@ -294,9 +292,7 @@ class XAIVideoConfig(BaseVideoConfig):
|
|||
)
|
||||
video_obj._hidden_params["video_url"] = video_url
|
||||
if custom_llm_provider and video_obj.id and video_obj.id != "unknown":
|
||||
video_obj.id = encode_video_id_with_provider(
|
||||
video_obj.id, custom_llm_provider, response_data.get("model")
|
||||
)
|
||||
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, response_data.get("model"))
|
||||
return video_obj
|
||||
|
||||
def transform_video_content_request(
|
||||
|
|
@ -321,9 +317,7 @@ class XAIVideoConfig(BaseVideoConfig):
|
|||
url: Final = video_meta.get("url") if isinstance(video_meta, dict) else None
|
||||
if isinstance(url, str) and url:
|
||||
return url
|
||||
raise ValueError(
|
||||
f"xAI video not ready for download (status={payload.get('status')}): {payload}"
|
||||
)
|
||||
raise ValueError(f"xAI video not ready for download (status={payload.get('status')}): {payload}")
|
||||
|
||||
def transform_video_content_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -5,8 +5,9 @@ from collections.abc import Sequence
|
|||
from typing import Final, get_type_hints
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, UploadFile, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from fastapi.responses import ORJSONResponse
|
||||
from starlette.datastructures import UploadFile
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -17,6 +18,7 @@ from litellm.proxy._types import *
|
|||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_is_form_content_type,
|
||||
coerce_numeric_form_fields,
|
||||
numeric_form_fields,
|
||||
)
|
||||
|
|
@ -27,6 +29,11 @@ from litellm.types.llms.openai import ChatCompletionUserMessage
|
|||
router: Final = APIRouter()
|
||||
|
||||
IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(ImageEditRequestParams))
|
||||
_IMAGE_REFERENCE_PREFIXES: Final = ("http://", "https://", "data:image/")
|
||||
_IMAGE_EDIT_FILE_FIELDS: Final = (
|
||||
("image", "image[]"),
|
||||
("mask", "mask[]"),
|
||||
)
|
||||
|
||||
|
||||
async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO:
|
||||
|
|
@ -51,6 +58,93 @@ async def batch_to_bytesio(
|
|||
return [await uploadfile_to_bytesio(u) for u in uploads]
|
||||
|
||||
|
||||
def _is_image_reference_string(value: str) -> bool:
|
||||
return value.startswith(_IMAGE_REFERENCE_PREFIXES)
|
||||
|
||||
|
||||
def _invalid_image_field_error(field: str) -> HTTPException:
|
||||
return HTTPException(
|
||||
status_code=422,
|
||||
detail=f"'{field}' must be a multipart file, http(s) URL, or data:image URI.",
|
||||
)
|
||||
|
||||
|
||||
def _form_field_values(form: object, name: str) -> tuple[object, ...]:
|
||||
getlist: Final = getattr(form, "getlist", None)
|
||||
if not callable(getlist):
|
||||
return ()
|
||||
return tuple(getlist(name))
|
||||
|
||||
|
||||
async def _coerce_image_part(value: object, field: str) -> io.BytesIO | str:
|
||||
if isinstance(value, UploadFile):
|
||||
return await uploadfile_to_bytesio(value)
|
||||
if isinstance(value, str) and _is_image_reference_string(value):
|
||||
return value
|
||||
raise _invalid_image_field_error(field)
|
||||
|
||||
|
||||
async def _normalize_image_values(values: tuple[object, ...], field: str) -> object | None:
|
||||
if not values:
|
||||
return None
|
||||
coerced: Final = tuple([await _coerce_image_part(value, field) for value in values])
|
||||
if len(coerced) == 1 and isinstance(coerced[0], str):
|
||||
return coerced[0]
|
||||
return list(coerced)
|
||||
|
||||
|
||||
def _json_image_values(data: dict[str, object], field: str) -> tuple[object, ...]:
|
||||
if field not in data:
|
||||
return ()
|
||||
raw: Final = data[field]
|
||||
if isinstance(raw, list):
|
||||
return tuple(raw)
|
||||
return (raw,)
|
||||
|
||||
|
||||
async def _normalized_image_edit_fields(
|
||||
values_by_field: dict[str, tuple[object, ...]],
|
||||
) -> dict[str, object]:
|
||||
image: Final = await _normalize_image_values(values_by_field["image"], "image")
|
||||
mask: Final = await _normalize_image_values(values_by_field["mask"], "mask")
|
||||
return {
|
||||
**({"image": image} if image is not None else {}),
|
||||
**({"mask": mask} if mask is not None else {}),
|
||||
}
|
||||
|
||||
|
||||
async def _image_edit_assets_from_request(
|
||||
request: Request,
|
||||
data: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
form: Final = await request.form() if _is_form_content_type(request.headers.get("content-type", "")) else None
|
||||
if form is None:
|
||||
return {
|
||||
**{key: value for key, value in data.items() if key not in {"image[]", "mask[]"}},
|
||||
**await _normalized_image_edit_fields(
|
||||
{field: _json_image_values(data, field) for field, _alias in _IMAGE_EDIT_FILE_FIELDS}
|
||||
),
|
||||
}
|
||||
|
||||
form_values: Final = {
|
||||
name: _form_field_values(form, name) for field, alias in _IMAGE_EDIT_FILE_FIELDS for name in (field, alias)
|
||||
}
|
||||
conflicts: Final = tuple(
|
||||
field for field, alias in _IMAGE_EDIT_FILE_FIELDS if form_values[field] and form_values[alias]
|
||||
)
|
||||
if conflicts:
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail=f"Cannot specify both '{conflicts[0]}' and '{conflicts[0]}[]'",
|
||||
)
|
||||
return {
|
||||
**{key: value for key, value in data.items() if key not in {"image[]", "mask[]"}},
|
||||
**await _normalized_image_edit_fields(
|
||||
{field: form_values[field] or form_values[alias] for field, alias in _IMAGE_EDIT_FILE_FIELDS}
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/images/generations",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -235,10 +329,6 @@ async def image_edit_api(
|
|||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
image: list[UploadFile] | None = File(None),
|
||||
image_array: list[UploadFile] | None = File(None, alias="image[]"),
|
||||
mask: list[UploadFile] | None = File(None),
|
||||
mask_array: list[UploadFile] | None = File(None, alias="mask[]"),
|
||||
model: str | None = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -254,20 +344,6 @@ async def image_edit_api(
|
|||
-F 'prompt=Create a studio ghibli image of this'
|
||||
```
|
||||
"""
|
||||
if image is not None and image_array is not None:
|
||||
raise HTTPException(status_code=422, detail="Cannot specify both 'image' and 'image[]'")
|
||||
if mask is not None and mask_array is not None:
|
||||
raise HTTPException(status_code=422, detail="Cannot specify both 'mask' and 'mask[]'")
|
||||
if image is None and image_array is not None:
|
||||
image = image_array
|
||||
if mask is None and mask_array is not None:
|
||||
mask = mask_array
|
||||
|
||||
# if image is None:
|
||||
# raise HTTPException(status_code=422, detail="Field required: image")
|
||||
# Note: Image is optional for some models (e.g., Bedrock Stability style-transfer)
|
||||
# The validation will be done at the model level if image is truly required
|
||||
|
||||
from litellm.proxy.proxy_server import (
|
||||
_read_request_body,
|
||||
general_settings,
|
||||
|
|
@ -283,49 +359,20 @@ async def image_edit_api(
|
|||
version,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Read request body and convert UploadFiles to BytesIO
|
||||
#########################################################
|
||||
data: Final = dict(
|
||||
parsed_body: Final = dict(
|
||||
coerce_numeric_form_fields(
|
||||
parsed_body=await _read_request_body(request=request),
|
||||
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
|
||||
)
|
||||
)
|
||||
image_files: Final = await batch_to_bytesio(image)
|
||||
mask_files: Final = await batch_to_bytesio(mask)
|
||||
if image_files:
|
||||
data["image"] = image_files
|
||||
if mask_files:
|
||||
data["mask"] = mask_files
|
||||
|
||||
invalid_image_fields: Final = tuple(
|
||||
field
|
||||
for field in ("image", "mask")
|
||||
if field in data
|
||||
and isinstance(data[field], str)
|
||||
and not (
|
||||
data[field].startswith("http://")
|
||||
or data[field].startswith("https://")
|
||||
or data[field].startswith("data:image/")
|
||||
)
|
||||
)
|
||||
if invalid_image_fields:
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail=f"'{invalid_image_fields[0]}' must be a multipart file, http(s) URL, or data:image URI.",
|
||||
)
|
||||
|
||||
# Ensure prompt exists in data (default to None for models that don't require it)
|
||||
if "prompt" not in data:
|
||||
data["prompt"] = None
|
||||
|
||||
data["model"] = (
|
||||
model
|
||||
or general_settings.get("image_generation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
)
|
||||
with_assets: Final = await _image_edit_assets_from_request(request, parsed_body)
|
||||
data: Final = {
|
||||
**with_assets,
|
||||
**({} if "prompt" in with_assets else {"prompt": None}),
|
||||
"model": (
|
||||
model or general_settings.get("image_generation_model", None) or user_model or with_assets.get("model")
|
||||
),
|
||||
}
|
||||
#########################################################
|
||||
# Process request
|
||||
#########################################################
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue