fix(images): forward scalar-array edit params as repeated multipart fields

Flatten dict-backed multipart bodies so a scalar list becomes one field with a
tuple value, which httpx emits as a repeated part per element, instead of
collapsing to the last element under dict.update. Nested objects still flatten
to key[subkey] like the OpenAI SDK, and the file-tuple video path is untouched.
This commit is contained in:
mateo-berri 2026-08-24 12:48:44 -07:00
parent 4eb09ad56e
commit a703378915
2 changed files with 68 additions and 7 deletions

View file

@ -27,20 +27,46 @@ def _flatten_form_field(key: str, value: object) -> tuple[tuple[str, str], ...]:
return ((key, serialized),)
def flatten_form_field_values(*sources: Mapping[str, object] | None) -> tuple[tuple[str, str], ...]:
def _is_form_scalar(value: object) -> bool:
return value is not None and not isinstance(value, (Mapping, list, tuple))
def _flatten_form_data_field(key: str, value: object) -> tuple[tuple[str, str | tuple[str, ...]], ...]:
if isinstance(value, Mapping):
return tuple(
item
for subkey, subvalue in value.items()
for item in _flatten_form_data_field(f"{key}[{subkey}]", subvalue)
)
if isinstance(value, (list, tuple)):
if all(_is_form_scalar(entry) for entry in value):
serialized_fields: Final = tuple(field for entry in value if (field := _form_field_value(entry)))
return ((key, serialized_fields),) if serialized_fields else ()
return tuple(item for entry in value for item in _flatten_form_data_field(f"{key}[]", entry))
if value is None:
return ()
serialized: Final = _form_field_value(value)
if not serialized:
return ()
return ((key, serialized),)
def flatten_form_field_values(*sources: Mapping[str, object] | None) -> tuple[tuple[str, str | tuple[str, ...]], ...]:
"""
Flatten JSON-shaped bodies into primitive ``(name, value)`` form fields the way the
OpenAI SDK serializes multipart bodies, applying ``sources`` in order so a later source
wins on a key collision under ``dict.update``. Lets provider params reach a multipart
request without handing the httpx encoder a nested value it rejects with
``Invalid type for value``.
Flatten JSON-shaped bodies into ``(name, value)`` form fields for a ``dict``-backed
multipart body, applying ``sources`` in order so a later source wins on a key collision
under ``dict.update``. Nested objects become ``key[subkey]`` fields the way the OpenAI SDK
serializes them, so provider params reach a multipart request without handing the httpx
encoder a nested value it rejects with ``Invalid type for value``. A scalar list becomes a
single field carrying a tuple value, which httpx emits as one repeated part per element, so
every element survives instead of collapsing to the last under ``dict.update``.
"""
return tuple(
pair
for source in sources
if source is not None
for top_key, top_value in source.items()
for pair in _flatten_form_field(top_key, top_value)
for pair in _flatten_form_data_field(top_key, top_value)
)

View file

@ -1,9 +1,24 @@
import httpx
from litellm.litellm_core_utils.llm_request_utils import (
flatten_form_field_values,
serialize_multipart_form_fields,
)
def _multipart_field_names(data: dict) -> list[str]:
request = httpx.Request(
"POST",
"http://backend/v1/images/edits",
data=data,
files=[("image[]", ("in.png", b"stub", "image/png"))],
)
request.read()
body = request.content.decode("utf-8", "replace")
prefix = 'Content-Disposition: form-data; name="'
return [line[len(prefix) : line.index('"', len(prefix))] for line in body.splitlines() if line.startswith(prefix)]
def test_serialize_multipart_form_fields_flattens_like_the_openai_sdk():
fields = serialize_multipart_form_fields(
{
@ -62,3 +77,23 @@ def test_flatten_form_field_values_later_source_wins_on_collision():
("seed", "2"),
)
assert dict(flatten_form_field_values({"seed": 1}, {"seed": 2}))["seed"] == "2"
def test_flatten_form_field_values_keeps_scalar_lists_as_repeated_fields():
assert flatten_form_field_values(
{"loras": ["a", "b", "c"], "generation_config": {"tags": [1, 2]}, "seed": 42}
) == (
("loras", ("a", "b", "c")),
("generation_config[tags]", ("1", "2")),
("seed", "42"),
)
def test_flatten_form_field_values_scalar_list_survives_update_into_multipart():
request_params: dict = {"model": "my-edit-model"}
request_params.update(flatten_form_field_values({"loras": ["style_a", "style_b"]}))
names = _multipart_field_names(request_params)
assert names.count("loras") == 2
assert names.count("model") == 1