mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(azure_ai): cost embeddings, responses, images, and rerank relays instead of logging zero
This commit is contained in:
parent
c698ddeacb
commit
17b7003592
5 changed files with 261 additions and 24 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue