diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index b9580ba3948..e20efe52bec 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -320,9 +320,11 @@ async def image_edit_api( detail=f"'{_field}' must be provided as a multipart file upload, not a string.", ) - # Ensure prompt exists in data (default to None for models that don't require it) + # Ensure prompt and image exist in data (default to None for models that don't require them) if "prompt" not in data: data["prompt"] = None + if "image" not in data: + data["image"] = None data["model"] = ( model diff --git a/tests/test_litellm/proxy/image_endpoints/test_endpoints.py b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py index ad0901e9eee..f4aebecc11e 100644 --- a/tests/test_litellm/proxy/image_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py @@ -222,6 +222,28 @@ def test_image_edit_multipart_n_that_is_not_a_number_is_left_alone(monkeypatch): assert captured["n"] == "two" +@pytest.mark.parametrize( + "files, form, missing", + [ + ({}, {"model": "stability.stable-style-transfer-v1:0", "prompt": "oil painting"}, "image"), + ( + {"image": ("tree.png", b"\x89PNG\r\n\x1a\n", "image/png")}, + {"model": "stability.stable-image-remove-background-v1:0"}, + "prompt", + ), + ], +) +def test_image_edit_without_an_optional_field_reaches_the_provider_with_it_set_to_none( + monkeypatch, files, form, missing +): + captured: Dict[str, Any] = {} + + response = _image_edit_client(monkeypatch, captured).post("/v1/images/edits", files=files or None, data=form) + + assert response.status_code == 200, response.text + assert missing in captured and captured[missing] is None, captured + + @pytest.mark.asyncio async def test_a_model_the_router_cannot_serve_answers_an_openai_typed_error(monkeypatch: pytest.MonkeyPatch): """A bare HTTPException carries no type or param, so the tail used to ship the @@ -290,7 +312,9 @@ async def test_failure_log_carries_the_callers_litellm_call_id( async def fake_add_litellm_data_to_request(**kwargs: object) -> object: return kwargs["data"] - async def fake_pre_call_hook(*, user_api_key_dict: UserAPIKeyAuth, data: dict[str, object], call_type: str) -> dict[str, object]: + async def fake_pre_call_hook( + *, user_api_key_dict: UserAPIKeyAuth, data: dict[str, object], call_type: str + ) -> dict[str, object]: return data async def fake_post_call_failure_hook(**_: object) -> None: @@ -327,7 +351,9 @@ async def test_failure_log_carries_the_callers_litellm_call_id( ) with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), pytest.raises(ProxyException) as raised: - await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth()) + await endpoints.image_generation( + request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth() + ) assert raised.value.headers["x-litellm-call-id"] == call_id record = next(r for r in caplog.records if "Exception occured" in r.getMessage()) @@ -378,7 +404,9 @@ async def test_failure_before_the_provider_call_bills_the_callers_litellm_call_i ) with pytest.raises(ProxyException) as raised: - await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth()) + await endpoints.image_generation( + request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth() + ) assert raised.value.headers["x-litellm-call-id"] == call_id assert [data["litellm_call_id"] for data in hook_request_data] == [call_id]