diff --git a/litellm/integrations/otel/__init__.py b/litellm/integrations/otel/__init__.py index 9c1205bb277..d7627d4d63d 100644 --- a/litellm/integrations/otel/__init__.py +++ b/litellm/integrations/otel/__init__.py @@ -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", diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index 79487e69ac4..5e3401cd62c 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -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. diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index aba9cc80240..4e4ed4b7513 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -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, ) diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index ada2822ba66..1647e0a5bd1 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -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()) diff --git a/litellm/litellm_core_utils/logging_utils.py b/litellm/litellm_core_utils/logging_utils.py index a17415f3ab8..91c8ba36b26 100644 --- a/litellm/litellm_core_utils/logging_utils.py +++ b/litellm/litellm_core_utils/logging_utils.py @@ -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. diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index c8f94b575ad..980b27cda55 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -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, diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 4fc6655ca54..ee0efb88a38 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -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, diff --git a/litellm/main.py b/litellm/main.py index 52785e7a393..2cf53833c5a 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index 19d0cfc0b18..2a66d5ee139 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -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 diff --git a/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py b/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py index 06871edb773..55ef74abd7b 100644 --- a/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py +++ b/tests/test_litellm/llms/openai/image_generation/test_openai_image_generation_extra_headers.py @@ -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 ): diff --git a/tests/test_litellm/test_non_chat_routes_open_llm_spans.py b/tests/test_litellm/test_non_chat_routes_open_llm_spans.py new file mode 100644 index 00000000000..d62959ccd43 --- /dev/null +++ b/tests/test_litellm/test_non_chat_routes_open_llm_spans.py @@ -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/"]