fix(otel): emit LLM Call spans for speech, image, moderation, ocr and transcription (#37752)

* fix(otel): emit LLM Call spans for speech, image, moderation, ocr and transcription

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(otel): log the image request before caller headers are merged in

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(otel): map non-chat routes to standard genai operations

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(otel): stop caller image headers aliasing the logged request body

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(otel): keep resolved api_base in async moderation pre_call

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(otel): log resolved client endpoint for speech pre_call

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* chore(otel): justify mutable request payloads in speech and image pre_call

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(otel): keep caller headers out of the logged speech request body

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-08-22 11:11:21 -07:00 • committed by GitHub
parent 076eebb520
commit 28887f12c5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 473 additions and 10 deletions

View file

@ -49,6 +49,7 @@ from litellm.integrations.otel.model.semconv import (
Error,
GenAI,
GenAIOperation,
GenAIOutputType,
GenAIProvider,
JsonRpc,
LiteLLM,
@ -60,6 +61,7 @@ from litellm.integrations.otel.model.semconv import (
RpcSystem,
Server,
resolve_operation,
resolve_output_type,
resolve_provider,
)
from litellm.integrations.otel.model.spans import (
@ -84,6 +86,7 @@ __all__ = [
"Error",
"GenAI",
"GenAIOperation",
"GenAIOutputType",
"GenAIProvider",
"GuardrailSpanData",
"JsonRpc",
@ -116,6 +119,7 @@ __all__ = [
"is_otel_v2_enabled",
"promoted_baggage",
"resolve_operation",
"resolve_output_type",
"resolve_provider",
"span_role_for_service",
"validate_registry",

View file

@ -42,6 +42,7 @@ class GenAIMapper:
_LLM_CALL_ATTRS: dict[str, Callable[[LLMCallSpanData], AttrValue | None]] = {
GenAI.OPERATION_NAME: lambda d: d.operation.value,
GenAI.PROVIDER_NAME: lambda d: d.provider or None,
GenAI.OUTPUT_TYPE: lambda d: d.output_type.value if d.output_type else None,
GenAI.REQUEST_MODEL: lambda d: d.request_model or None,
GenAI.REQUEST_TEMPERATURE: lambda d: d.request_params.temperature,
GenAI.REQUEST_TOP_P: lambda d: d.request_params.top_p,
@ -65,6 +66,7 @@ class GenAIMapper:
Server.ADDRESS: lambda d: d.server.address if d.server else None,
Server.PORT: lambda d: d.server.port if d.server else None,
LiteLLM.CALL_ID: lambda d: d.identity.call_id or None,
LiteLLM.CALL_TYPE: lambda d: d.call_type,
# The provider/underlying model is only known once routing has picked a
# deployment, so it can't ride identity Baggage (seeded at auth, before
# routing) onto the boundary-born LLM span — stamp it directly here.

View file

@ -15,8 +15,10 @@ from litellm.integrations.otel.model.metadata import (
)
from litellm.integrations.otel.model.semconv import (
GenAIOperation,
GenAIOutputType,
MCPMethod,
resolve_operation,
resolve_output_type,
resolve_provider,
)
from litellm.integrations.otel.model.utils import (
@ -310,6 +312,11 @@ class LLMCallSpanData:
choices_out: tuple[Mapping[str, object], ...] = ()
system_fingerprint: str | None = None
time_to_first_chunk_seconds: float | None = None
# The requested output modality, set only on the routes that pin one (image
# generation, speech, transcription, OCR), and the litellm route itself, which
# keeps routes the convention folds into one operation distinguishable.
output_type: GenAIOutputType | None = None
call_type: str | None = None
@classmethod
def from_standard_logging_payload(
@ -334,8 +341,9 @@ class LLMCallSpanData:
# otherwise the content-bearing mappers receive empty sequences and emit
# no prompt/response text.
finish_reasons: Final = _finish_reasons(choices_out)
call_type: Final = as_str(payload.get("call_type"))
return cls(
operation=resolve_operation(as_str(payload.get("call_type"))),
operation=resolve_operation(call_type),
provider=resolve_provider(as_str(payload.get("custom_llm_provider"))),
request_model=context.request_model,
response_model=context.response_model,
@ -358,6 +366,8 @@ class LLMCallSpanData:
choices_out=choices_out if capture_content else (),
system_fingerprint=as_str(response.get("system_fingerprint")),
time_to_first_chunk_seconds=time_to_first_chunk_seconds,
output_type=resolve_output_type(call_type),
call_type=call_type or None,
)

View file

@ -3,7 +3,9 @@ Keys follow the OpenTelemetry GenAI semantic conventions (experimental). Anythin
without a semconv equivalent lives under the ``litellm.*`` vendor namespace.
"""
from collections.abc import Mapping
from enum import Enum
from types import MappingProxyType
from typing import Final
from litellm._logging import verbose_logger
@ -30,6 +32,21 @@ class GenAIOperation(str, Enum):
EXECUTE_TOOL = "execute_tool" # MCP tool-call spans
LITELLM_VECTOR_STORE_MANAGEMENT = "litellm.vector_store_management"
LITELLM_VECTOR_STORE_FILE_MANAGEMENT = "litellm.vector_store_file_management"
LITELLM_MODERATION = "litellm.moderation"
class GenAIOutputType(str, Enum):
"""Values for ``gen_ai.output.type``, the modality the client asked for.
It is what separates the inference routes that share ``generate_content``:
image generation requests ``image``, speech requests ``speech``, and
transcription and OCR both request ``text``.
"""
TEXT = "text"
JSON = "json"
IMAGE = "image"
SPEECH = "speech"
class GenAIProvider(str, Enum):
@ -258,6 +275,11 @@ class LiteLLM:
"""Vendor-extension keys (no semconv equivalent). Always ``litellm.*``."""
CALL_ID: Final = "litellm.call_id"
# The litellm route that produced the call. Needed because the convention maps
# several routes onto one operation: transcription and OCR are both
# ``generate_content`` with a ``text`` output type, so this is the only thing
# that tells them apart.
CALL_TYPE: Final = "litellm.call_type"
COST_PREFIX: Final = "litellm.cost."
METADATA_PREFIX: Final = "litellm.metadata."
TEAM_ID: Final = "litellm.team.id"
@ -352,6 +374,16 @@ _OPERATION_BY_CALL_TYPE: Final[dict[str, GenAIOperation]] = {
"aembedding": GenAIOperation.EMBEDDINGS,
"responses": GenAIOperation.CHAT,
"aresponses": GenAIOperation.CHAT,
"image_generation": GenAIOperation.GENERATE_CONTENT,
"aimage_generation": GenAIOperation.GENERATE_CONTENT,
"moderation": GenAIOperation.LITELLM_MODERATION,
"amoderation": GenAIOperation.LITELLM_MODERATION,
"ocr": GenAIOperation.GENERATE_CONTENT,
"aocr": GenAIOperation.GENERATE_CONTENT,
"speech": GenAIOperation.GENERATE_CONTENT,
"aspeech": GenAIOperation.GENERATE_CONTENT,
"transcription": GenAIOperation.GENERATE_CONTENT,
"atranscription": GenAIOperation.GENERATE_CONTENT,
"call_mcp_tool": GenAIOperation.EXECUTE_TOOL,
"vector_store_search": GenAIOperation.RETRIEVAL,
"avector_store_search": GenAIOperation.RETRIEVAL,
@ -385,6 +417,23 @@ _OPERATION_BY_CALL_TYPE: Final[dict[str, GenAIOperation]] = {
}
# litellm ``call_type`` -> ``gen_ai.output.type``. Only the call types whose route
# fixes the requested modality are listed; the attribute is conditionally required
# on a request that asks for an output format, so anything else is left unstamped.
_OUTPUT_TYPE_BY_CALL_TYPE: Final[Mapping[str, GenAIOutputType]] = MappingProxyType(
{
"image_generation": GenAIOutputType.IMAGE,
"aimage_generation": GenAIOutputType.IMAGE,
"speech": GenAIOutputType.SPEECH,
"aspeech": GenAIOutputType.SPEECH,
"transcription": GenAIOutputType.TEXT,
"atranscription": GenAIOutputType.TEXT,
"ocr": GenAIOutputType.TEXT,
"aocr": GenAIOutputType.TEXT,
}
)
def resolve_provider(custom_llm_provider: str | None) -> str:
"""Map a litellm provider string to a ``gen_ai.provider.name`` value.
@ -416,3 +465,11 @@ def resolve_operation(call_type: str | None) -> GenAIOperation:
GenAIOperation.CHAT.value,
)
return GenAIOperation.CHAT
def resolve_output_type(call_type: str | None) -> GenAIOutputType | None:
"""Map a litellm ``call_type`` to a ``gen_ai.output.type`` value, or ``None``
for a route that doesn't pin the output modality."""
if not call_type:
return None
return _OUTPUT_TYPE_BY_CALL_TYPE.get(call_type.lower())

View file

@ -3,6 +3,7 @@ import functools
import inspect
import re
import time
from collections.abc import Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final
@ -268,6 +269,16 @@ def _set_duration_in_model_call_details(
verbose_logger.warning("Error setting `llm_api_duration_ms`: %s", e)
def speech_request_body(model: str, voice: str, optional_params: Mapping[str, object]) -> Mapping[str, object]:
"""Speech request body for telemetry, without the caller headers the provider SDKs
take as request kwargs rather than body fields."""
return { # mutable-ok: loggers isinstance-check the request body as a dict
"model": model,
"voice": voice,
**{key: value for key, value in optional_params.items() if key != "extra_headers"},
}
def track_llm_api_timing():
"""
Decorator to track LLM API call timing for both sync and async functions.

View file

@ -17,7 +17,7 @@ from openai import (
import litellm
from litellm.constants import AZURE_OPERATION_POLLING_TIMEOUT, DEFAULT_MAX_RETRIES
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
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.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -1352,6 +1352,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
organization: str | None,
max_retries: int,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
azure_ad_token: str | None = None,
azure_ad_token_provider: Callable | None = None,
aspeech: bool | None = None,
@ -1373,6 +1374,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
azure_ad_token_provider=azure_ad_token_provider,
max_retries=max_retries,
timeout=timeout,
logging_obj=logging_obj,
client=client,
litellm_params=litellm_params,
)
@ -1387,6 +1389,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
litellm_params=litellm_params,
)
logging_obj.pre_call(
input=input,
api_key=api_key,
additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict
"complete_input_dict": speech_request_body(model, voice, optional_params),
"api_base": str(azure_client.base_url),
},
)
response: Final = azure_client.audio.speech.create(
model=model,
voice=voice,
@ -1408,6 +1419,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
azure_ad_token_provider: Callable | None,
max_retries: int,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
client=None,
litellm_params: dict | None = None,
) -> HttpxBinaryResponseContent:
@ -1421,6 +1433,15 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
litellm_params=litellm_params,
)
logging_obj.pre_call(
input=input,
api_key=api_key,
additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict
"complete_input_dict": speech_request_body(model, voice, optional_params),
"api_base": str(azure_client.base_url),
},
)
azure_response: Final = await azure_client.audio.speech.create(
model=model,
voice=voice,

View file

@ -22,7 +22,7 @@ from litellm._logging import verbose_logger
from litellm.constants import DEFAULT_MAX_RETRIES
from litellm.files.types import FileContentStreamingResult
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.chat.invoke_handler import MockResponseIterator
@ -1365,9 +1365,21 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=client,
)
if headers:
data["extra_headers"] = headers
response = await openai_aclient.images.generate(**data, timeout=timeout)
logging_obj.pre_call(
input=prompt,
api_key=openai_aclient.api_key,
additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict
"headers": {"Authorization": f"Bearer {openai_aclient.api_key}"}, # mutable-ok: logged header map
"api_base": str(openai_aclient.base_url),
"acompletion": True,
"complete_input_dict": data,
},
)
request_data: Final = ( # mutable-ok: the OpenAI SDK takes the request body as a dict
{**data, "extra_headers": headers} if headers else data
)
response = await openai_aclient.images.generate(**request_data, timeout=timeout)
stringified_response: Final = response.model_dump()
## LOGGING
logging_obj.post_call(
@ -1450,9 +1462,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
## COMPLETION CALL
if headers:
data["extra_headers"] = headers
_response: Final = openai_client.images.generate(**data, timeout=timeout)
request_data: Final = ( # mutable-ok: the OpenAI SDK takes the request body as a dict
{**data, "extra_headers": headers} if headers else data
)
_response: Final = openai_client.images.generate(**request_data, timeout=timeout)
response: Final = _response.model_dump()
## LOGGING
@ -1501,6 +1514,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
project: str | None,
max_retries: int,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
aspeech: bool | None = None,
client=None,
shared_session: Optional["ClientSession"] = None,
@ -1517,6 +1531,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
project=project,
max_retries=max_retries,
timeout=timeout,
logging_obj=logging_obj,
client=client,
shared_session=shared_session,
)
@ -1531,7 +1546,17 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
shared_session=shared_session,
)
response: Final = cast(OpenAI, openai_client).audio.speech.create(
sync_client: Final = cast(OpenAI, openai_client)
logging_obj.pre_call(
input=input,
api_key=api_key,
additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict
"complete_input_dict": speech_request_body(model, voice, optional_params),
"api_base": str(sync_client.base_url),
},
)
response: Final = sync_client.audio.speech.create(
model=model,
voice=voice,
input=input,
@ -1551,6 +1576,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
project: str | None,
max_retries: int,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
client=None,
shared_session: Optional["ClientSession"] = None,
) -> HttpxBinaryResponseContent:
@ -1567,6 +1593,15 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
),
)
logging_obj.pre_call(
input=input,
api_key=api_key,
additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict
"complete_input_dict": speech_request_body(model, voice, optional_params),
"api_base": str(openai_client.base_url),
},
)
response: Final = await openai_client.audio.speech.create(
model=model,
voice=voice,

View file

@ -7537,6 +7537,15 @@ async def amoderation(
},
custom_llm_provider=custom_llm_provider,
)
moderation_request: Final = {"input": input, "model": model} # mutable-ok: logged as the raw request body
litellm_logging_obj.pre_call(
input=input,
api_key=api_key,
additional_args={ # mutable-ok: loggers isinstance-check this payload as a dict
"complete_input_dict": moderation_request,
"api_base": str(_openai_client.base_url),
},
)
if model is not None:
response = await _openai_client.moderations.create(input=input, model=model)
@ -8042,6 +8051,7 @@ def speech(
project=project,
max_retries=max_retries,
timeout=timeout,
logging_obj=logging_obj,
client=client, # pass AsyncOpenAI, OpenAI client
aspeech=aspeech,
shared_session=shared_session,
@ -8120,6 +8130,7 @@ def speech(
organization=organization,
max_retries=max_retries,
timeout=timeout,
logging_obj=logging_obj,
client=client, # pass AsyncOpenAI, OpenAI client
aspeech=aspeech,
litellm_params=litellm_params_dict,

View file

@ -4,6 +4,7 @@ and the typed StandardLoggingPayload adapter. These need no OTel SDK."""
import logging
import re
from pathlib import Path
from typing import Final
import pytest
@ -14,6 +15,7 @@ from litellm.integrations.otel import (
Error,
GenAI,
GenAIOperation,
GenAIOutputType,
HTTP,
LiteLLM,
OpenTelemetryV2Config,
@ -21,8 +23,10 @@ from litellm.integrations.otel import (
is_otel_v2_enabled,
promoted_baggage,
resolve_operation,
resolve_output_type,
resolve_provider,
)
from litellm.integrations.otel.mappers.genai import GenAIMapper
from litellm.integrations.otel.model import spans as spans_mod
from litellm.integrations.otel.model.payloads import (
LLMCallSpanData,
@ -264,6 +268,74 @@ def test_vector_store_file_management_is_not_chat(call_type):
assert resolve_operation(call_type).value == "litellm.vector_store_file_management"
_NON_CHAT_ROUTES: Final = (
("image_generation", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.IMAGE),
("speech", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.SPEECH),
("transcription", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.TEXT),
("ocr", GenAIOperation.GENERATE_CONTENT, GenAIOutputType.TEXT),
("moderation", GenAIOperation.LITELLM_MODERATION, None),
)
@pytest.mark.parametrize(
("call_type", "operation", "output_type"),
[
(f"{prefix}{call_type}", operation, output_type)
for call_type, operation, output_type in _NON_CHAT_ROUTES
for prefix in ("", "a")
],
)
def test_non_chat_inference_routes_follow_genai_semconv(call_type, operation, output_type):
"""Image generation, speech, transcription and OCR all produce content, so the
convention names them ``generate_content`` and separates them by the requested
output modality rather than by an invented operation. Moderation classifies
instead of generating and the convention names nothing for it, so it keeps a
vendor value. Either way the spans must not land in the chat series a dashboard
reads."""
assert resolve_operation(call_type) is operation
assert resolve_output_type(call_type) is output_type
@pytest.mark.parametrize(
("call_type", "operation", "output_type"),
[(f"a{call_type}", operation, output_type) for call_type, operation, output_type in _NON_CHAT_ROUTES],
)
def test_non_chat_route_spans_carry_semconv_name_and_modality(call_type, operation, output_type):
"""The emitted span, not just the mapping table: name is
``{gen_ai.operation.name} {gen_ai.request.model}``, the modality rides
``gen_ai.output.type``, and the route stays recoverable from
``litellm.call_type`` now that several routes share one operation."""
data = LLMCallSpanData.from_standard_logging_payload(
_sample_payload(call_type=call_type, model="some-model", custom_llm_provider="openai")
)
attrs = GenAIMapper().map(data)
assert spans_mod.llm_call_span_name(data) == f"{operation.value} some-model"
assert attrs[GenAI.OPERATION_NAME] == operation.value
assert attrs[GenAI.PROVIDER_NAME] == "openai"
assert attrs[GenAI.REQUEST_MODEL] == "some-model"
assert attrs[LiteLLM.CALL_TYPE] == call_type
assert attrs.get(GenAI.OUTPUT_TYPE) == (output_type.value if output_type else None)
def test_non_chat_route_error_span_keeps_error_attributes():
"""Modality mapping must not cost the failure signal: a failed non-chat call
still carries the error type alongside the standardized operation."""
data = LLMCallSpanData.from_standard_logging_payload(
_sample_payload(
call_type="aspeech",
model="tts-1",
status="failure",
error_information={"error_class": "BadRequestError"},
)
)
attrs = GenAIMapper().map(data)
assert attrs[GenAI.OPERATION_NAME] == GenAIOperation.GENERATE_CONTENT.value
assert attrs[GenAI.OUTPUT_TYPE] == GenAIOutputType.SPEECH.value
assert attrs[Error.TYPE] == "BadRequestError"
def test_vendor_operation_values_are_namespaced():
"""A vendor value must stay under the ``litellm.`` prefix: an unprefixed invented
name could collide with a value the convention adds later, silently changing what

View file

@ -148,6 +148,58 @@ class TestImageGenerationExtraHeaders:
_, kwargs = mock_openai_client.images.generate.call_args
assert "extra_headers" not in kwargs
@pytest.mark.parametrize("is_async", [False, True])
@pytest.mark.asyncio
async def test_caller_headers_never_reach_the_logged_request_body(
self, openai_chat_completions, mock_logging_obj, is_async
):
"""The body handed to pre_call is also what telemetry reads at close time, so
merging caller headers into that same dict would publish a customer's auth
header as a span attribute. The upstream call still gets them."""
mock_image_data = MagicMock()
mock_image_data.model_dump.return_value = {
"created": 1700000000,
"data": [{"url": "https://example.com/image.png"}],
}
mock_openai_client = MagicMock()
mock_openai_client.api_key = "test-key"
mock_openai_client._base_url._uri_reference = "https://api.openai.com"
test_headers = {"cf-aig-authorization": "Bearer custom-token"}
if is_async:
mock_openai_client.images.generate = AsyncMock(return_value=mock_image_data)
await openai_chat_completions.aimage_generation(
prompt="A white cat",
data={"model": "dall-e-3", "prompt": "A white cat"},
model_response=MagicMock(),
timeout=60.0,
logging_obj=mock_logging_obj,
api_key="test-key",
headers=test_headers,
client=mock_openai_client,
)
else:
mock_openai_client.images.generate.return_value = mock_image_data
openai_chat_completions.image_generation(
model="dall-e-3",
prompt="A white cat",
timeout=60.0,
optional_params={},
logging_obj=mock_logging_obj,
api_key="test-key",
headers=test_headers,
client=mock_openai_client,
)
logged_body = mock_logging_obj.pre_call.call_args[1]["additional_args"][
"complete_input_dict"
]
assert "extra_headers" not in logged_body
_, kwargs = mock_openai_client.images.generate.call_args
assert kwargs.get("extra_headers") == test_headers
def test_sync_image_generation_forwards_headers_to_async(
self, openai_chat_completions, mock_logging_obj
):

View file

@ -0,0 +1,188 @@
"""Regression tests: every route that issues an upstream call must fire the
``pre_call`` input hook.
Tracing integrations open their LLM-call span there (``OpenTelemetryV2`` keys the
span off ``log_pre_api_call`` and treats "no pre_call" as "the request never
reached a provider"), so a handler that skips it leaves the call with no LLM-call
span in the trace at all. Speech, async image generation and moderation each used
to skip it.
"""
import asyncio
from typing import Any, Final
import httpx
import pytest
from openai import AsyncAzureOpenAI, AsyncOpenAI
import litellm
from litellm.integrations.custom_logger import CustomLogger
class _PreCallRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.call_types: list[str] = [] # mutable-ok: test recorder of hook calls
self.api_bases: list[str] = [] # mutable-ok: test recorder of hook calls
self.request_bodies: list[Any] = [] # mutable-ok: test recorder of hook calls
def log_pre_api_call(self, model, messages, kwargs) -> None:
self.call_types.append(str(kwargs.get("call_type")))
self.api_bases.append(str(kwargs.get("litellm_params", {}).get("api_base")))
self.request_bodies.append(kwargs.get("additional_args", {}).get("complete_input_dict"))
class _FakeSpeech:
def __init__(self) -> None:
self.calls: list[dict[str, Any]] = [] # mutable-ok: test recorder of SDK calls
async def create(self, **kwargs: Any) -> Any:
self.calls.append(kwargs)
request: Final = httpx.Request("POST", "https://api.openai.com/v1/audio/speech")
return type(
"_Speech",
(),
{"response": httpx.Response(200, content=b"audio-bytes", request=request)},
)()
class _FakeImages:
async def generate(self, **kwargs: Any) -> Any:
return type(
"_Images",
(),
{
"model_dump": lambda self: {
"created": 1,
"data": [{"url": "https://example.com/img.png"}],
}
},
)()
class _FakeModerations:
async def create(self, **kwargs: Any) -> Any:
return type(
"_Moderations",
(),
{
"model_dump": lambda self: {
"id": "modr-1",
"model": "omni-moderation-latest",
"results": [
{
"flagged": False,
"categories": {},
"category_scores": {},
"category_applied_input_types": {},
}
],
}
},
)()
class _FakeAsyncOpenAI(AsyncOpenAI):
"""Stands in for the injected client: a real ``AsyncOpenAI`` (``amoderation``
type-checks it) whose resource namespaces answer without a network call."""
def __init__(self, base_url: str = "https://api.openai.com/v1") -> None:
super().__init__(api_key="sk-test", base_url=base_url)
self.speech = _FakeSpeech()
self.audio = type("_Audio", (), {"speech": self.speech})()
self.images = _FakeImages()
self.moderations = _FakeModerations()
class _FakeAsyncAzureOpenAI(AsyncAzureOpenAI):
"""Same idea for the Azure entrypoint, which resolves no default endpoint of
its own when ``AZURE_API_BASE`` is unset."""
def __init__(self) -> None:
super().__init__(
api_key="sk-test",
api_version="2024-02-01",
azure_endpoint="https://unit-test.openai.azure.com",
)
self.speech = _FakeSpeech()
self.audio = type("_Audio", (), {"speech": self.speech})()
@pytest.fixture
def recorder(monkeypatch):
recorder: Final = _PreCallRecorder()
monkeypatch.setattr(litellm, "callbacks", [recorder])
monkeypatch.setattr(litellm, "success_callback", [])
return recorder
def test_async_speech_opens_an_llm_span(recorder):
asyncio.run(
litellm.aspeech(
model="openai/tts-1",
input="hello",
voice="alloy",
client=_FakeAsyncOpenAI(),
)
)
assert recorder.call_types == ["aspeech"]
def test_azure_async_speech_opens_an_llm_span_without_api_base(recorder, monkeypatch):
"""Azure resolves no default endpoint, so a missing ``api_base`` used to reach
``_get_masked_api_base`` as ``None``; the ``TypeError`` was swallowed and the
whole callback dispatch was skipped."""
monkeypatch.delenv("AZURE_API_BASE", raising=False)
asyncio.run(
litellm.aspeech(
model="azure/tts-deployment",
input="hello",
voice="alloy",
client=_FakeAsyncAzureOpenAI(),
)
)
assert recorder.call_types == ["aspeech"]
assert recorder.api_bases == ["https://unit-test.openai.azure.com/openai/"]
def test_azure_async_speech_keeps_caller_headers_out_of_the_logged_body(recorder):
"""The Azure entrypoint carries caller headers in ``optional_params``, so they reach
the provider as a request kwarg; telemetry reads the logged body, which must stay
free of them."""
headers: Final = {"authorization": "Bearer caller-secret"}
client: Final = _FakeAsyncAzureOpenAI()
asyncio.run(
litellm.aspeech(
model="azure/tts-deployment",
input="hello",
voice="alloy",
extra_headers=headers,
client=client,
)
)
assert recorder.call_types == ["aspeech"]
assert "extra_headers" not in recorder.request_bodies[0]
assert client.speech.calls[0]["extra_headers"] == headers
def test_async_image_generation_opens_an_llm_span(recorder):
asyncio.run(
litellm.aimage_generation(
model="openai/dall-e-3",
prompt="a cat",
client=_FakeAsyncOpenAI(),
)
)
assert recorder.call_types == ["aimage_generation"]
def test_async_moderation_opens_an_llm_span(recorder):
asyncio.run(
litellm.amoderation(
model="omni-moderation-latest",
input="hello",
client=_FakeAsyncOpenAI(base_url="https://gateway.example/v1"),
)
)
assert recorder.call_types == ["amoderation"]
assert recorder.api_bases == ["https://gateway.example/v1/"]