[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:
parent
538339e1a8
commit
6a7aab7b84
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -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]:
|
||||
|
||||
@ -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
|
||||
@ -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"
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user