fix(azure_ai): bill FLUX.2 by the megapixels Azure reports in request_meta

This commit is contained in:
shrey kharbanda 2026-09-25 01:09:07 +00:00
parent 3d7b85bdd0
commit 744ceba2d3
10 changed files with 153 additions and 45 deletions

View file

@ -19,6 +19,7 @@ from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRI
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.llms.azure_ai.image_generation import AzureFoundryFluxImageGenerationConfig
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -1323,7 +1324,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
deployment_name=model,
)
provider_config: Final = get_azure_image_generation_config(data.get("model", "dall-e-2"))
if isinstance(provider_config, AzureFoundryMAIImageGenerationConfig):
if isinstance(
provider_config, (AzureFoundryMAIImageGenerationConfig, AzureFoundryFluxImageGenerationConfig)
):
return provider_config.transform_image_generation_response(
model=data.get("model", "dall-e-2"),
raw_response=httpx_response,

View file

@ -1,5 +1,8 @@
from litellm._logging import verbose_logger
from litellm.llms.azure_ai.image_generation import AzureFoundryMAIImageGenerationConfig
from litellm.llms.azure_ai.image_generation import (
AzureFoundryFluxImageGenerationConfig,
AzureFoundryMAIImageGenerationConfig,
)
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
@ -27,6 +30,8 @@ def get_azure_image_generation_config(model: str) -> BaseImageGenerationConfig:
return AzureDallE3ImageGenerationConfig()
elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model):
return AzureFoundryMAIImageGenerationConfig()
elif AzureFoundryFluxImageGenerationConfig.is_flux2_model(model):
return AzureFoundryFluxImageGenerationConfig()
else:
verbose_logger.debug(
"Using AzureGPTImageGenerationConfig for model: %s. This follows the gpt-image model format.", model

View file

@ -2,8 +2,9 @@ import base64
from collections.abc import Mapping, Sequence
from io import BufferedReader
from types import MappingProxyType
from typing import Any, Final
from typing import TYPE_CHECKING, Any, Final
import httpx
from httpx._types import RequestFiles
import litellm
@ -13,12 +14,17 @@ from litellm.llms.azure_ai.common_utils import (
)
from litellm.llms.azure_ai.image_generation.flux_transformation import (
AzureFoundryFluxImageGenerationConfig,
with_flux2_billed_megapixels,
)
from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.llms.openai import FileTypes
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ImageResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig):
@ -121,6 +127,17 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig):
}
return request_body, []
def transform_image_edit_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
) -> ImageResponse:
return with_flux2_billed_megapixels(
super().transform_image_edit_response(model=model, raw_response=raw_response, logging_obj=logging_obj),
raw_response,
)
def _convert_image_to_base64(self, image: Any) -> str:
"""Convert image file to base64 string"""
if isinstance(image, BufferedReader):

View file

@ -9,7 +9,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
)
from litellm.types.utils import ImageResponse, ModelInfo
_BILLED_PIXELS_PER_INPUT_IMAGE: Final = 1024 * 1024
_PIXELS_PER_MEGAPIXEL: Final = 1024 * 1024
def _input_cost_per_pixel(resolved: ModelInfo) -> float:
@ -55,7 +55,11 @@ def cost_calculator(
if output_cost_per_image:
return output_cost_per_image * num_images
if _input_cost_per_pixel(_model_info):
input_cost_per_pixel: Final = _input_cost_per_pixel(_model_info)
if input_cost_per_pixel and image_response.provider_billed_megapixels is not None:
return input_cost_per_pixel * _PIXELS_PER_MEGAPIXEL * image_response.provider_billed_megapixels
if input_cost_per_pixel:
from litellm.cost_calculator import default_image_cost_calculator
width: Final = optional_params.get("width") if optional_params else None
@ -65,15 +69,13 @@ def cost_calculator(
if type(width) is int and type(height) is int and width > 0 and height > 0
else size or image_response.size
)
output_cost: Final = default_image_cost_calculator(
return default_image_cost_calculator(
model=_model_info.get("key", model),
custom_llm_provider=litellm.LlmProviders.AZURE_AI.value,
size=pixel_size,
n=num_images,
model_info=model_info,
)
input_image_pixels: Final = image_response.input_image_count * _BILLED_PIXELS_PER_INPUT_IMAGE
return output_cost + _input_cost_per_pixel(_model_info) * input_image_pixels
return 0.0
raise ValueError(f"image_response must be of type ImageResponse got type={type(image_response)}")

View file

@ -1,10 +1,19 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from typing import TYPE_CHECKING, Final
import httpx
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from litellm.exceptions import BadRequestError, UnsupportedParamsError
from litellm.llms.openai.image_generation import GPTImageGenerationConfig
from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams
from litellm.types.utils import ImageResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
FLUX2_DROPPED_OPENAI_PARAMS: Final[tuple[OpenAIImageGenerationOptionalParams, ...]] = (
"background",
@ -15,6 +24,28 @@ FLUX2_DROPPED_OPENAI_PARAMS: Final[tuple[OpenAIImageGenerationOptionalParams, ..
)
class _Flux2RequestMeta(BaseModel):
model_config = ConfigDict(frozen=True)
input_mp: float = Field(ge=0)
output_mp: float = Field(ge=0)
class _Flux2ResponseBody(BaseModel):
model_config = ConfigDict(frozen=True)
request_meta: _Flux2RequestMeta
def with_flux2_billed_megapixels(image_response: ImageResponse, raw_response: httpx.Response) -> ImageResponse:
try:
request_meta: Final = _Flux2ResponseBody.model_validate_json(raw_response.content).request_meta
except ValidationError:
return image_response
image_response.set_provider_billed_megapixels(request_meta.input_mp + request_meta.output_mp)
return image_response
class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig):
"""Azure Foundry BFL API configuration for FLUX image generation."""
@ -152,3 +183,32 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig):
}
)
return {**optional_params, **mapped_params} # mutable-ok: inherited config contract returns a dict
def transform_image_generation_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: "LiteLLMLoggingObj",
request_data: dict,
optional_params: dict,
litellm_params: dict,
encoding: "Tokenizer | None",
api_key: str | None = None,
json_mode: bool | None = None,
) -> ImageResponse:
image_response: Final = super().transform_image_generation_response(
model=model,
raw_response=raw_response,
model_response=model_response,
logging_obj=logging_obj,
request_data=request_data,
optional_params=optional_params,
litellm_params=litellm_params,
encoding=encoding,
api_key=api_key,
json_mode=json_mode,
)
return (
with_flux2_billed_megapixels(image_response, raw_response) if self.is_flux2_model(model) else image_response
)

View file

@ -6968,7 +6968,6 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
image_edit_response.set_input_image_count(len(image) if isinstance(image, list) else 1)
return image_edit_response
async def async_image_edit_handler(
@ -7070,7 +7069,6 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
image_edit_response.set_input_image_count(len(image) if isinstance(image, list) else 1)
return image_edit_response
def image_generation_handler(

View file

@ -2578,14 +2578,14 @@ from openai.types.images_response import ImagesResponse as OpenAIImageResponse
class ImageResponse(OpenAIImageResponse, BaseLiteLLMOpenAIResponseObject):
_hidden_params: dict = {}
_input_image_count: int = 0
_provider_billed_megapixels: float | None = None
@property
def input_image_count(self) -> int:
return self._input_image_count
def provider_billed_megapixels(self) -> float | None:
return self._provider_billed_megapixels
def set_input_image_count(self, count: int) -> None:
self._input_image_count = count
def set_provider_billed_megapixels(self, megapixels: float) -> None:
self._provider_billed_megapixels = megapixels
usage: ImageUsage | None = None
"""

View file

@ -197,6 +197,38 @@ def test_flux2_cost_uses_mapped_dimensions_after_response_transformation(dimensi
) == pytest.approx(_flex_rate() * 2048 * 1024 * 2)
@pytest.mark.parametrize(
("request_meta", "expected_megapixels"),
(
({"input_mp": 0.0, "output_mp": 0.39}, 0.39),
({"input_mp": 45.78, "output_mp": 1.0}, 46.78),
({"output_mp": 0.39}, 2048 * 1024 / (1024 * 1024)),
({"input_mp": -1.0, "output_mp": 0.39}, 2048 * 1024 / (1024 * 1024)),
),
)
def test_flux2_generation_bills_azure_reported_megapixels_and_falls_back_to_requested_size(
request_meta: Mapping[str, float], expected_megapixels: float
):
params: Final = {"width": 2048, "height": 1024}
response: Final = get_azure_image_generation_config("FLUX.2-flex").transform_image_generation_response(
model="FLUX.2-flex",
raw_response=httpx.Response(200, json={"data": [{"b64_json": "aW1n"}], "request_meta": request_meta}),
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={"prompt": "A red fox", **params},
optional_params=params,
litellm_params={},
encoding=None,
)
assert litellm.completion_cost(
model="azure_ai/FLUX.2-flex",
completion_response=response,
optional_params=params,
call_type="image_generation",
) == pytest.approx(_flex_rate() * 1024 * 1024 * expected_megapixels)
def test_flux2_flex_cost_accepts_lowercase_model_spelling():
response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n"), ImageObject(b64_json="aW1n")])

View file

@ -4249,28 +4249,6 @@ async def test_async_image_edit_handler_records_upstream_response_headers():
_assert_upstream_headers_recorded(response)
def test_image_edit_handler_records_input_image_count():
client = HTTPHandler(client=httpx.Client(transport=_json_with_upstream_headers({"transformed_by": "sync"})))
response = BaseLLMHTTPHandler().image_edit_handler(
client=client, **{**_image_edit_call_kwargs(), "image": [b"one", b"two", b"three"]}
)
assert response.input_image_count == 3
@pytest.mark.asyncio
async def test_async_image_edit_handler_records_input_image_count():
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"transformed_by": "async"}))
response = await BaseLLMHTTPHandler().async_image_edit_handler(
client=client, **{**_image_edit_call_kwargs(), "image": [b"one", b"two"]}
)
assert response.input_image_count == 2
class _HeaderImageGenerationConfig(BaseImageGenerationConfig):
def get_supported_openai_params(self, model):
return []

View file

@ -144,7 +144,9 @@ def test_flux2_image_edit_rejects_too_many_references(model: str, reference_imag
)
@pytest.mark.parametrize("dimensions", ({"size": "2048x1024"}, {"width": 2048, "height": 1024}, {"width": "2048", "height": "1024"}))
@pytest.mark.parametrize(
"dimensions", ({"size": "2048x1024"}, {"width": 2048, "height": 1024}, {"width": "2048", "height": "1024"})
)
@pytest.mark.usefixtures("local_model_cost_map")
def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[str, int | str]):
def respond(request: httpx.Request) -> httpx.Response:
@ -176,21 +178,32 @@ def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[
)
rate: Final = litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"]
assert response._hidden_params["response_cost"] == pytest.approx(rate * (2048 * 1024 * 2 + 1024 * 1024))
assert response._hidden_params["response_cost"] == pytest.approx(rate * 2048 * 1024 * 2)
@pytest.mark.parametrize("reference_images", (1, 3, 10))
@pytest.mark.parametrize(
("input_mp", "output_mp"),
((0.39, 1.0), (3.0, 1.0), (11.63, 1.0), (2.0, 0.75)),
)
@pytest.mark.usefixtures("local_model_cost_map")
def test_flux2_image_edit_bills_each_reference_image_as_one_megapixel(reference_images: int):
def test_flux2_image_edit_bills_the_megapixels_azure_reports(input_mp: float, output_mp: float):
client: Final = HTTPHandler(
client=httpx.Client(
transport=httpx.MockTransport(lambda request: httpx.Response(200, json={"data": [{"b64_json": "aW1n"}]}))
transport=httpx.MockTransport(
lambda request: httpx.Response(
200,
json={
"data": [{"b64_json": "aW1n"}],
"request_meta": {"cost": 5.0, "input_mp": input_mp, "output_mp": output_mp},
},
)
)
)
)
response: Final = litellm.image_edit(
model="azure_ai/FLUX.2-flex",
image=[b"image"] * reference_images,
image=[b"one", b"two"],
prompt="Blend every reference",
api_key="test-key",
api_base="https://example.services.ai.azure.com",
@ -199,7 +212,7 @@ def test_flux2_image_edit_bills_each_reference_image_as_one_megapixel(reference_
)
rate: Final = litellm.model_cost["azure_ai/FLUX.2-flex"]["input_cost_per_pixel"]
assert response._hidden_params["response_cost"] == pytest.approx(rate * 1024 * 1024 * (1 + reference_images))
assert response._hidden_params["response_cost"] == pytest.approx(rate * 1024 * 1024 * (input_mp + output_mp))
def test_flux2_image_edit_accepts_and_drops_openai_only_parameters():