diff --git a/litellm/images/main.py b/litellm/images/main.py index 1f722eb752a..4794785b866 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -51,6 +51,7 @@ from litellm.types.llms.openai import ImageGenerationRequestQuality from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( LITELLM_IMAGE_VARIATION_PROVIDERS, + CustomPricingLiteLLMParams, LlmProviders, all_litellm_params, ) @@ -883,6 +884,7 @@ def image_edit( optional_params=dict(image_edit_request_params), litellm_params={ **image_edit_request_params, + **litellm_params.model_dump(include=set(CustomPricingLiteLLMParams.model_fields), exclude_none=True), "litellm_call_id": litellm_call_id, "model_info": model_info, }, diff --git a/tests/unit/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py b/tests/unit/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py index 50bbeefb520..f5083094c65 100644 --- a/tests/unit/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py +++ b/tests/unit/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py @@ -6,6 +6,7 @@ import pathlib import struct import tempfile from collections.abc import Callable, Iterator, Mapping +from datetime import datetime from typing import Final import httpx @@ -564,6 +565,52 @@ async def test_flux2_router_image_edit_bills_the_deployment_rates(monkeypatch: p ) +@pytest.mark.parametrize( + ("deployment_prices", "expected_cost"), + ( + ({"output_cost_per_image": 0.5}, 0.5), + ( + {"input_cost_per_pixel": 1e-07, "input_cost_per_reference_pixel": 2e-07}, + 1e-07 * 2 * 1024 * 1024 + 2e-07 * 2 * 1024 * 1024, + ), + ), + ids=("flat", "per-pixel"), +) +async def test_flux2_router_image_edit_bills_the_deployment_rates_with_a_logger_built_before_routing( + monkeypatch: pytest.MonkeyPatch, deployment_prices: Mapping[str, float], expected_cost: float +): + mock_client: Final = AsyncHTTPHandler() + mock_client.client = httpx.AsyncClient(transport=httpx.MockTransport(_edit_ok)) + monkeypatch.setattr(llm_http_handler_module, "get_async_httpx_client", lambda **_kwargs: mock_client) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "flux2-pro-deployment", + "litellm_params": { + "model": "azure_ai/flux.2-pro", + "api_base": "https://example.services.ai.azure.com", + "api_key": "test-key", + **deployment_prices, + }, + } + ] + ) + request: Final = { + "model": "flux2-pro-deployment", + "prompt": "Make it a watercolor", + "image": [_png(1024, 1280)], + "size": "1024x1280", + "litellm_call_id": "proxy-call-id", + } + logging_obj, routed_request = litellm.utils.function_setup( + original_function="aimage_edit", rules_obj=litellm.utils.Rules(), start_time=datetime.now(), **request + ) + + response: Final = await router.aimage_edit(**routed_request, litellm_logging_obj=logging_obj) + + assert response._hidden_params["response_cost"] == pytest.approx(expected_cost) + + def _edit_returning(image: bytes) -> Callable[[httpx.Request], httpx.Response]: def respond(request: httpx.Request) -> httpx.Response: return httpx.Response(200, json={"data": [{"b64_json": base64.b64encode(image).decode()}]})