[Bug Fix] StandardLoggingPayload on cache_hits should track custom llm provider + DD LLM Obs span type (#12652)

* bug fix - ensure custom llm provider is tracked on cache hit

* fix config.yaml

* test_cache_hit_includes_custom_llm_provider

* fix _get_datadog_span_kind

* test_datadog_span_kind_mapping

* fix ruff check

* test_datadog_span_kind_mapping
This commit is contained in:
Ishaan Jaff 2025-07-16 15:43:15 -07:00 committed by GitHub
parent 538339e1a8
commit 6a7aab7b84
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 231 additions and 16 deletions

View File

@ -141,7 +141,7 @@ class LLMCachingHandler:
verbose_logger.debug("Cache Hit!")
cache_hit = True
end_time = datetime.datetime.now()
model, _, _, _ = litellm.get_llm_provider(
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model,
custom_llm_provider=kwargs.get("custom_llm_provider", None),
api_base=kwargs.get("api_base", None),
@ -153,6 +153,7 @@ class LLMCachingHandler:
kwargs=kwargs,
cached_result=cached_result,
is_async=True,
custom_llm_provider=custom_llm_provider,
)
call_type = original_function.__name__
@ -882,6 +883,7 @@ class LLMCachingHandler:
cached_result: Any,
is_async: bool,
is_embedding: bool = False,
custom_llm_provider: Optional[str] = None,
):
"""
Helper function to update the LiteLLMLoggingObj environment variables.
@ -893,6 +895,7 @@ class LLMCachingHandler:
cached_result (Any): The cached result to log.
is_async (bool): Whether the call is asynchronous or not.
is_embedding (bool): Whether the call is for embeddings or not.
custom_llm_provider (Optional[str]): The custom llm provider being used.
Returns:
None
@ -905,6 +908,7 @@ class LLMCachingHandler:
"model_info": kwargs.get("model_info", {}),
"proxy_server_request": kwargs.get("proxy_server_request", None),
"stream_response": kwargs.get("stream_response", {}),
"custom_llm_provider": custom_llm_provider,
}
if litellm.cache is not None:
@ -928,6 +932,7 @@ class LLMCachingHandler:
original_response=str(cached_result),
additional_args=None,
stream=kwargs.get("stream", False),
custom_llm_provider=custom_llm_provider,
)

View File

@ -11,7 +11,7 @@ import json
import os
import uuid
from datetime import datetime
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, List, Literal, Optional, Union
import httpx
@ -27,7 +27,7 @@ from litellm.llms.custom_httpx.http_handler import (
httpxSpecialProvider,
)
from litellm.types.integrations.datadog_llm_obs import *
from litellm.types.utils import StandardLoggingPayload
from litellm.types.utils import CallTypes, StandardLoggingPayload
class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
@ -149,7 +149,7 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
output_meta = OutputMeta(messages=self._get_response_messages(response_obj))
meta = Meta(
kind="llm",
kind=self._get_datadog_span_kind(standard_logging_payload.get("call_type")),
input=input_meta,
output=output_meta,
metadata=self._get_dd_llm_obs_payload_metadata(standard_logging_payload),
@ -208,6 +208,105 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
return [response_obj["choices"][0]["message"].json()]
return []
def _get_datadog_span_kind(self, call_type: Optional[str]) -> Literal["llm", "tool", "task", "embedding", "retrieval"]:
"""
Map liteLLM call_type to appropriate DataDog LLM Observability span kind.
Available DataDog span kinds: "llm", "tool", "task", "embedding", "retrieval"
"""
if call_type is None:
return "llm"
# Embedding operations
if call_type in [CallTypes.embedding.value, CallTypes.aembedding.value]:
return "embedding"
# LLM completion operations
if call_type in [
CallTypes.completion.value,
CallTypes.acompletion.value,
CallTypes.text_completion.value,
CallTypes.atext_completion.value,
CallTypes.generate_content.value,
CallTypes.agenerate_content.value,
CallTypes.generate_content_stream.value,
CallTypes.agenerate_content_stream.value,
CallTypes.anthropic_messages.value
]:
return "llm"
# Tool operations
if call_type in [CallTypes.call_mcp_tool.value]:
return "tool"
# Retrieval operations
if call_type in [
CallTypes.get_assistants.value,
CallTypes.aget_assistants.value,
CallTypes.get_thread.value,
CallTypes.aget_thread.value,
CallTypes.get_messages.value,
CallTypes.aget_messages.value,
CallTypes.afile_retrieve.value,
CallTypes.file_retrieve.value,
CallTypes.afile_list.value,
CallTypes.file_list.value,
CallTypes.afile_content.value,
CallTypes.file_content.value,
CallTypes.retrieve_batch.value,
CallTypes.aretrieve_batch.value,
CallTypes.retrieve_fine_tuning_job.value,
CallTypes.aretrieve_fine_tuning_job.value,
CallTypes.responses.value,
CallTypes.aresponses.value,
CallTypes.alist_input_items.value
]:
return "retrieval"
# Task operations (batch, fine-tuning, file operations, etc.)
if call_type in [
CallTypes.create_batch.value,
CallTypes.acreate_batch.value,
CallTypes.create_fine_tuning_job.value,
CallTypes.acreate_fine_tuning_job.value,
CallTypes.cancel_fine_tuning_job.value,
CallTypes.acancel_fine_tuning_job.value,
CallTypes.list_fine_tuning_jobs.value,
CallTypes.alist_fine_tuning_jobs.value,
CallTypes.create_assistants.value,
CallTypes.acreate_assistants.value,
CallTypes.delete_assistant.value,
CallTypes.adelete_assistant.value,
CallTypes.create_thread.value,
CallTypes.acreate_thread.value,
CallTypes.add_message.value,
CallTypes.a_add_message.value,
CallTypes.run_thread.value,
CallTypes.arun_thread.value,
CallTypes.run_thread_stream.value,
CallTypes.arun_thread_stream.value,
CallTypes.file_delete.value,
CallTypes.afile_delete.value,
CallTypes.create_file.value,
CallTypes.acreate_file.value,
CallTypes.image_generation.value,
CallTypes.aimage_generation.value,
CallTypes.image_edit.value,
CallTypes.aimage_edit.value,
CallTypes.moderation.value,
CallTypes.amoderation.value,
CallTypes.transcription.value,
CallTypes.atranscription.value,
CallTypes.speech.value,
CallTypes.aspeech.value,
CallTypes.rerank.value,
CallTypes.arerank.value
]:
return "task"
# Default fallback for unknown or passthrough operations
return "llm"
def _ensure_string_content(
self, messages: Optional[Union[str, List[Any], Dict[Any, Any]]]
) -> List[Any]:

View File

@ -7,16 +7,6 @@ model_list:
model: openai/*
litellm_settings:
callbacks: ["s3://litellm-proxy/custom_ui_sso_hook.custom_ui_sso_sign_in_handler"]
- model_name: gemini/*
litellm_params:
model: gemini/*
litellm_settings:
callbacks: ["datadog_llm_observability"]
mcp_servers:
# HTTP Streamable Server
deepwiki_mcp:
url: "https://mcp.deepwiki.com/mcp"
cache: true

View File

@ -174,3 +174,41 @@ class TestDataDogLLMObsLogger:
assert time_to_first_token == 2.0 # 1002.0 - 1000.0 = 2.0 seconds
def test_datadog_span_kind_mapping(self, mock_env_vars):
"""Test that call_type values are correctly mapped to DataDog span kinds"""
from litellm.types.utils import CallTypes
with patch('litellm.integrations.datadog.datadog_llm_obs.get_async_httpx_client'), \
patch('asyncio.create_task'):
logger = DataDogLLMObsLogger()
# Test embedding operations
assert logger._get_datadog_span_kind(CallTypes.embedding.value) == "embedding"
assert logger._get_datadog_span_kind(CallTypes.aembedding.value) == "embedding"
# Test LLM completion operations
assert logger._get_datadog_span_kind(CallTypes.completion.value) == "llm"
assert logger._get_datadog_span_kind(CallTypes.acompletion.value) == "llm"
assert logger._get_datadog_span_kind(CallTypes.text_completion.value) == "llm"
assert logger._get_datadog_span_kind(CallTypes.generate_content.value) == "llm"
assert logger._get_datadog_span_kind(CallTypes.anthropic_messages.value) == "llm"
# Test tool operations
assert logger._get_datadog_span_kind(CallTypes.call_mcp_tool.value) == "tool"
# Test retrieval operations
assert logger._get_datadog_span_kind(CallTypes.get_assistants.value) == "retrieval"
assert logger._get_datadog_span_kind(CallTypes.file_retrieve.value) == "retrieval"
assert logger._get_datadog_span_kind(CallTypes.retrieve_batch.value) == "retrieval"
# Test task operations
assert logger._get_datadog_span_kind(CallTypes.create_batch.value) == "task"
assert logger._get_datadog_span_kind(CallTypes.image_generation.value) == "task"
assert logger._get_datadog_span_kind(CallTypes.moderation.value) == "task"
assert logger._get_datadog_span_kind(CallTypes.transcription.value) == "task"
# Test default fallback
assert logger._get_datadog_span_kind("unknown_call_type") == "llm"
assert logger._get_datadog_span_kind(None) == "llm"

View File

@ -1,3 +1,4 @@
import asyncio
import datetime
import json
import os
@ -22,15 +23,29 @@ import litellm
from litellm._logging import (
ALL_LOGGERS,
_initialize_loggers_with_handler,
_turn_on_json,
verbose_logger,
verbose_proxy_logger,
verbose_router_logger,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload
class CacheHitCustomLogger(CustomLogger):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.logged_standard_logging_payloads: List[StandardLoggingPayload] = []
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
standard_logging_payload = kwargs.get("standard_logging_object", None)
if standard_logging_payload:
self.logged_standard_logging_payloads.append(standard_logging_payload)
def test_json_mode_emits_one_record_per_logger(capfd):
# Turn on JSON logging
litellm._logging._turn_on_json()
_turn_on_json()
# Make sure our loggers will emit INFO-level records
for lg in (verbose_logger, verbose_router_logger, verbose_proxy_logger):
lg.setLevel(logging.INFO)
@ -69,3 +84,71 @@ def test_initialize_loggers_with_handler_sets_propagate_false():
assert (
logger.propagate is False
), f"Logger {logger.name} has propagate set to {logger.propagate}, expected False"
@pytest.mark.asyncio
async def test_cache_hit_includes_custom_llm_provider():
"""
Test that when there's a cache hit, the standard logging payload includes the custom_llm_provider
"""
# Set up caching and custom logger
litellm.cache = litellm.Cache()
test_custom_logger = CacheHitCustomLogger()
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
litellm.callbacks = [test_custom_logger]
try:
# First call - should be a cache miss
response1 = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "test cache hit message"}],
mock_response="test response",
caching=True,
)
# Wait for logging to complete
await asyncio.sleep(0.5)
# Second identical call - should be a cache hit
response2 = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "test cache hit message"}],
mock_response="test response",
caching=True,
)
# Wait for logging to complete
await asyncio.sleep(0.5)
# Verify we have logged events
assert len(test_custom_logger.logged_standard_logging_payloads) >= 2, \
f"Expected at least 2 logged events, got {len(test_custom_logger.logged_standard_logging_payloads)}"
# Find the cache hit event (should be the second call)
cache_hit_payload = None
for payload in test_custom_logger.logged_standard_logging_payloads:
if payload.get("cache_hit") is True:
cache_hit_payload = payload
break
# Verify cache hit event was found
assert cache_hit_payload is not None, "No cache hit event found in logged payloads"
# Verify custom_llm_provider is included in the cache hit payload
assert "custom_llm_provider" in cache_hit_payload, \
"custom_llm_provider missing from cache hit standard logging payload"
# Verify custom_llm_provider has a valid value (should be "openai" for gpt-3.5-turbo)
custom_llm_provider = cache_hit_payload["custom_llm_provider"]
assert custom_llm_provider is not None and custom_llm_provider != "", \
f"custom_llm_provider should not be None or empty, got: {custom_llm_provider}"
print(
f"Cache hit standard logging payload with custom_llm_provider: {custom_llm_provider}",
json.dumps(cache_hit_payload, indent=2),
)
finally:
# Clean up
litellm.callbacks = original_callbacks
litellm.cache = None