fix(azure_ai): cost embeddings, responses, images, and rerank relays instead of logging zero

This commit is contained in:
mateo-berri 2026-09-07 18:11:17 -07:00
parent c698ddeacb
commit 17b7003592
5 changed files with 261 additions and 24 deletions

View file

@ -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()

View file

@ -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(

View file

@ -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(

View file

@ -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:

View file

@ -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 = [