From 17b70035924b7bb2e97065a7f3512c540e105d66 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 7 Sep 2026 18:11:17 -0700 Subject: [PATCH] fix(azure_ai): cost embeddings, responses, images, and rerank relays instead of logging zero --- .../llms/azure/passthrough/transformation.py | 45 ++++++- .../azure_ai/passthrough/transformation.py | 32 ++++- .../base_llm/passthrough/transformation.py | 20 ++- .../test_azure_passthrough_transformation.py | 120 +++++++++++++++--- ...est_azure_ai_passthrough_transformation.py | 68 +++++++++- 5 files changed, 261 insertions(+), 24 deletions(-) diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index 062c60d7231..d7078d94a0b 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -1,4 +1,5 @@ -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass from typing import TYPE_CHECKING, Final, Optional import httpx @@ -9,12 +10,14 @@ from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.base_llm.passthrough.transformation import ( BasePassthroughConfig, + relayed_json_object, replace_path_segment, strip_leading_model_segment, ) from litellm.secret_managers.main import get_secret_str -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import CallTypes, EmbeddingResponse, ImageResponse if TYPE_CHECKING: from httpx import URL @@ -38,6 +41,40 @@ def _relayed_messages(litellm_logging_obj: Logging) -> Sequence[Mapping[str, obj return details.request_data.messages if details.request_data else None +@dataclass(frozen=True, slots=True) +class OpenAIRelayShape: + path_suffix: str + call_type: CallTypes + parse: Callable[[Mapping[str, object]], EmbeddingResponse | ImageResponse | ResponsesAPIResponse] + + +OPENAI_RELAY_SHAPES: Final = ( + OpenAIRelayShape("/embeddings", CallTypes.aembedding, EmbeddingResponse.model_validate), + OpenAIRelayShape("/responses", CallTypes.aresponses, ResponsesAPIResponse.model_validate), + OpenAIRelayShape("/images/generations", CallTypes.aimage_generation, ImageResponse.model_validate), +) + + +def logged_openai_response( + httpx_response: Response, logging_obj: Logging, endpoint: str +) -> EmbeddingResponse | ImageResponse | ResponsesAPIResponse | None: + relayed_path: Final = f"/{endpoint.strip('/')}" + shape: Final = next( + (candidate for candidate in OPENAI_RELAY_SHAPES if relayed_path.endswith(candidate.path_suffix)), None + ) + body: Final = relayed_json_object(httpx_response) if shape else None + if shape is None or body is None: + return None + try: + parsed: Final = shape.parse(body) + except ValidationError: + return None + logging_obj.call_type = ( + shape.call_type.value + ) # rebind-ok: routes cost calculation to the relayed shape's pricing path + return parsed + + class AzurePassthroughConfig(BasePassthroughConfig): def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: return bool(request_data.get("stream")) @@ -114,13 +151,13 @@ class AzurePassthroughConfig(BasePassthroughConfig): request_data: dict, logging_obj: Logging, endpoint: str, - ) -> Optional["CostResponseTypes"]: + ) -> Optional["CostResponseTypes | ResponsesAPIResponse"]: from litellm import encoding from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.types.utils import ModelResponse if "chat/completions" not in endpoint: - return None + return logged_openai_response(httpx_response, logging_obj, endpoint) openai_chat_config: Final = OpenAIGPTConfig() diff --git a/litellm/llms/azure_ai/passthrough/transformation.py b/litellm/llms/azure_ai/passthrough/transformation.py index 4cb93b4f8f9..c49e10566ed 100644 --- a/litellm/llms/azure_ai/passthrough/transformation.py +++ b/litellm/llms/azure_ai/passthrough/transformation.py @@ -14,8 +14,13 @@ from litellm.llms.azure_ai.common_utils import ( get_azure_ai_auth_headers, ) from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config -from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, strip_leading_model_segment +from litellm.llms.base_llm.passthrough.transformation import ( + BasePassthroughConfig, + relayed_json_object, + strip_leading_model_segment, +) from litellm.types.llms.openai import AllMessageValues +from litellm.types.rerank import RerankResponse from litellm.types.utils import CallTypes, StandardPassThroughResponseObject if TYPE_CHECKING: @@ -23,6 +28,7 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse + from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import CostResponseTypes @@ -77,6 +83,18 @@ def relayed_body(httpx_response: Response) -> str | dict: return body if isinstance(body, dict) else httpx_response.text +def logged_rerank_response(httpx_response: Response, logging_obj: Logging, endpoint: str) -> RerankResponse | None: + body: Final = relayed_json_object(httpx_response) if f"/{endpoint.strip('/')}".endswith("/rerank") else None + if body is None: + return None + try: + rerank_response: Final = RerankResponse.model_validate(body) + except ValidationError: + return None + logging_obj.call_type = CallTypes.arerank.value # rebind-ok: routes cost calculation to the per-query rerank path + return rerank_response + + class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig): def __init__(self, ocr_config_for: Callable[[str], BaseOCRConfig | None] = get_azure_ai_ocr_config) -> None: super().__init__() @@ -130,7 +148,14 @@ class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig): request_data: Mapping[str, object], logging_obj: Logging, endpoint: str, - ) -> CostResponseTypes | OCRResponse | StandardPassThroughResponseObject | None: + ) -> ( + CostResponseTypes + | OCRResponse + | RerankResponse + | ResponsesAPIResponse + | StandardPassThroughResponseObject + | None + ): from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig chat_result: Final = AzurePassthroughConfig().logging_non_streaming_response( # pyright: ignore[reportUnknownMemberType] # the Azure config still types request_data as a bare dict @@ -146,6 +171,9 @@ class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig): ocr_result: Final = self.logged_ocr_response(model, httpx_response, logging_obj, endpoint) if ocr_result is not None: return ocr_result + rerank_result: Final = logged_rerank_response(httpx_response, logging_obj, endpoint) + if rerank_result is not None: + return rerank_result return StandardPassThroughResponseObject(response=relayed_body(httpx_response)) def logged_ocr_response( diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index 2c21df01b5b..975b74137db 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -3,18 +3,25 @@ from abc import abstractmethod from collections.abc import Mapping from typing import TYPE_CHECKING, Final, Optional, Union +from pydantic import TypeAdapter, ValidationError + from ..base_utils import BaseLLMModelInfo if TYPE_CHECKING: from httpx import URL, Headers, Response from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.llms.openai import ResponsesAPIResponse + from litellm.types.rerank import RerankResponse from litellm.types.utils import CostResponseTypes, StandardPassThroughResponseObject from ..chat.transformation import BaseLLMException from ..ocr.transformation import OCRResponse +RELAYED_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object]) + + def strip_leading_model_segment(endpoint: str, model_names: tuple[str, ...]) -> str: path: Final = endpoint.lstrip("/") for model_name in model_names: @@ -32,6 +39,15 @@ def replace_path_segment(endpoint: str, segment: str, replacement: str) -> str: return bounded_segment.sub(lambda _: replacement, endpoint) +def relayed_json_object(httpx_response: "Response") -> Mapping[str, object] | None: + if httpx_response.status_code != 200: + return None + try: + return RELAYED_JSON_OBJECT.validate_python(httpx_response.json()) + except (ValueError, ValidationError): + return None + + class BasePassthroughConfig(BaseLLMModelInfo): @abstractmethod def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: @@ -123,7 +139,9 @@ class BasePassthroughConfig(BaseLLMModelInfo): request_data: dict, logging_obj: "LiteLLMLoggingObj", endpoint: str, - ) -> Optional["CostResponseTypes | OCRResponse | StandardPassThroughResponseObject"]: + ) -> Optional[ + "CostResponseTypes | OCRResponse | RerankResponse | ResponsesAPIResponse | StandardPassThroughResponseObject" + ]: pass def handle_logging_collected_chunks( diff --git a/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py b/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py index 262a281aeb3..2723a13c3fc 100644 --- a/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py +++ b/tests/test_litellm/llms/azure/passthrough/test_azure_passthrough_transformation.py @@ -1,15 +1,16 @@ import json +from datetime import datetime from unittest.mock import MagicMock import httpx import pytest import litellm - - +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.token_counter import high_detail_image_token_upper_bound from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig -from litellm.types.utils import ModelResponse +from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.utils import EmbeddingResponse, ModelResponse def _azure_chat_completion_body(): @@ -77,25 +78,112 @@ def test_azure_passthrough_logging_non_streaming_response_chat_completions(): assert result.usage.total_tokens == 18 -def test_azure_passthrough_logging_non_streaming_response_unknown_endpoint_returns_none(): - """ - Endpoints other than chat/completions (responses, messages, images) fall - through to None — matches base-class behavior and Bedrock's "unknown - endpoint" handling. Not a regression; just scoping. - """ - config = AzurePassthroughConfig() - logging_obj = MagicMock() - - result = config.logging_non_streaming_response( - model="gpt-4.1-mini", +def _relay_logging_obj(model: str) -> Logging: + logging_obj = Logging( + model=model, + messages=[], + stream=False, + call_type="allm_passthrough_route", + start_time=datetime.now(), + litellm_call_id="call-1", + function_id="fn-1", + ) + logging_obj.update_environment_variables( + model=model, + litellm_params={"api_base": "https://my-resource.openai.azure.com", "custom_llm_provider": "azure"}, + optional_params={}, custom_llm_provider="azure", - httpx_response=_make_httpx_response(_azure_chat_completion_body()), + ) + return logging_obj + + +def _relay_logging_result(model: str, endpoint: str, body, status_code: int = 200): + logging_obj = _relay_logging_obj(model) + response = httpx.Response( + status_code=status_code, + headers={"content-type": "application/json"}, + content=json.dumps(body).encode("utf-8"), + request=httpx.Request("POST", f"https://my-resource.openai.azure.com/{endpoint}?api-version=2025-04-01-preview"), + ) + result = AzurePassthroughConfig().logging_non_streaming_response( + model=model, + custom_llm_provider="azure", + httpx_response=response, request_data={}, logging_obj=logging_obj, - endpoint="openai/responses", + endpoint=endpoint, + ) + return result, logging_obj + + +EMBEDDINGS_BODY = { + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1000, "total_tokens": 1000}, +} + +RESPONSES_BODY = { + "id": "resp_1", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4.1-mini", + "output": [ + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + "usage": {"input_tokens": 1000, "output_tokens": 100, "total_tokens": 1100}, +} + + +def test_azure_passthrough_embeddings_relay_is_costed_per_input_token(): + result, logging_obj = _relay_logging_result( + "text-embedding-3-small", "openai/deployments/text-embedding-3-small/embeddings", EMBEDDINGS_BODY + ) + per_token = litellm.get_model_info("azure/text-embedding-3-small")["input_cost_per_token"] + + assert isinstance(result, EmbeddingResponse) + assert logging_obj.call_type == "aembedding" + assert per_token > 0 + assert logging_obj._response_cost_calculator(result=result) == pytest.approx(1000 * per_token) + + +def test_azure_passthrough_responses_relay_is_costed_per_token(): + result, logging_obj = _relay_logging_result("gpt-4.1-mini", "openai/responses", RESPONSES_BODY) + info = litellm.get_model_info("azure/gpt-4.1-mini") + + assert isinstance(result, ResponsesAPIResponse) + assert logging_obj.call_type == "aresponses" + assert logging_obj._response_cost_calculator(result=result) == pytest.approx( + 1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"] + ) + + +def test_azure_passthrough_failed_embeddings_relay_is_not_costed(): + result, logging_obj = _relay_logging_result( + "text-embedding-3-small", + "openai/deployments/text-embedding-3-small/embeddings", + {"error": {"code": "429", "message": "rate limited"}}, + status_code=429, ) assert result is None + assert logging_obj.call_type == "allm_passthrough_route" + + +def test_azure_passthrough_logging_non_streaming_response_unknown_endpoint_returns_none(): + result, logging_obj = _relay_logging_result( + "gpt-4o-mini-tts", "openai/deployments/gpt-4o-mini-tts/audio/speech", {"audio": "..."} + ) + + assert result is None + assert logging_obj.call_type == "allm_passthrough_route" def _sse_line(payload: dict) -> str: diff --git a/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py b/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py index 644316c2e37..788aced50f8 100644 --- a/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py +++ b/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py @@ -9,7 +9,8 @@ import litellm from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.azure_ai.passthrough.transformation import AzureAIPassthroughConfig from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.types.utils import LlmProviders, ModelResponse +from litellm.types.rerank import RerankResponse +from litellm.types.utils import EmbeddingResponse, ImageResponse, LlmProviders, ModelResponse from litellm.utils import ProviderConfigManager FOUNDRY_BASE = "https://my-resource.services.ai.azure.com" @@ -370,6 +371,71 @@ def test_unparseable_ocr_body_falls_back_to_the_passthrough_object(): assert logging_obj.call_type == "allm_passthrough_route" +EMBEDDINGS_BODY = { + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "model": "embed-v-4-0", + "usage": {"prompt_tokens": 1200, "total_tokens": 1200}, +} + +RERANK_BODY = { + "id": "rerank-1", + "results": [{"index": 1, "relevance_score": 0.9}, {"index": 0, "relevance_score": 0.2}], + "meta": {"api_version": {"version": "2"}, "billed_units": {"search_units": 2}}, +} + +IMAGE_BODY = {"created": 1, "data": [{"b64_json": "AAAA"}]} + + +def test_foundry_embeddings_relay_is_costed_per_input_token(): + result, logging_obj = _relay_logging_result( + AzureAIPassthroughConfig(), "embed-v-4-0", "models/embeddings", EMBEDDINGS_BODY + ) + per_token = litellm.get_model_info("azure_ai/embed-v-4-0")["input_cost_per_token"] + + assert isinstance(result, EmbeddingResponse) + assert logging_obj.call_type == "aembedding" + assert per_token > 0 + assert logging_obj._response_cost_calculator(result=result) == pytest.approx(1200 * per_token) + + +def test_cohere_rerank_relay_is_costed_per_search_unit(): + result, logging_obj = _relay_logging_result( + AzureAIPassthroughConfig(), "cohere-rerank-v4.0-fast", "providers/cohere/v2/rerank", RERANK_BODY + ) + per_query = litellm.get_model_info("azure_ai/cohere-rerank-v4.0-fast")["input_cost_per_query"] + + assert isinstance(result, RerankResponse) + assert logging_obj.call_type == "arerank" + assert per_query > 0 + assert logging_obj._response_cost_calculator(result=result) == pytest.approx(2 * per_query) + + +def test_image_generation_relay_is_costed_per_image(): + result, logging_obj = _relay_logging_result( + AzureAIPassthroughConfig(), "FLUX.2-pro", "openai/deployments/FLUX.2-pro/images/generations", IMAGE_BODY + ) + per_image = litellm.get_model_info("azure_ai/FLUX.2-pro")["output_cost_per_image"] + + assert isinstance(result, ImageResponse) + assert logging_obj.call_type == "aimage_generation" + assert per_image > 0 + assert logging_obj._response_cost_calculator(result=result) == pytest.approx(per_image) + + +def test_rejected_rerank_relay_keeps_the_passthrough_object_and_call_type(): + result, logging_obj = _relay_logging_result( + AzureAIPassthroughConfig(), + "cohere-rerank-v4.0-fast", + "providers/cohere/v2/rerank", + {"message": "invalid request"}, + status_code=400, + ) + + assert result == {"response": {"message": "invalid request"}} + assert logging_obj.call_type == "allm_passthrough_route" + + def test_streaming_chat_completion_chunks_are_costed_like_azure(): head = {"id": "chatcmpl-1", "object": "chat.completion.chunk", "created": 1, "model": "gpt-5.4-mini"} chunks = [