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:
hx 2026-09-08 11:09:29 +00:00
parent f4d012c9a3
commit c25883d55a
2 changed files with 110 additions and 69 deletions

View file

@ -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,

View file

@ -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
#########################################################