mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(azure_ai): bill FLUX.2 by the megapixels Azure reports in request_meta
This commit is contained in:
parent
3d7b85bdd0
commit
744ceba2d3
10 changed files with 153 additions and 45 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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")])
|
||||
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue