Addressing feedback.

- A few more cases were found where the dictionary access might not return the correct value.
- Handling cases where `traceparent` is not lower cased
This commit is contained in:
Josh Bonczkowski 2026-03-13 16:09:43 -04:00
parent acb4f624dd
commit 1688fd496a
2 changed files with 22 additions and 5 deletions

View file

@ -253,7 +253,10 @@ class NewRelicLogger(CustomLogger):
litellm_params = kwargs.get("litellm_params") or {}
metadata = litellm_params.get("metadata") or {}
headers = metadata.get("headers") or {}
traceparent = headers.get("traceparent", None)
# Normalize header key lookup to be case-insensitive per W3C spec
traceparent = next(
(v for k, v in headers.items() if k.lower() == "traceparent"), None
)
trace_id = None
span_id = None
@ -300,7 +303,7 @@ class NewRelicLogger(CustomLogger):
def _get_vendor(self, kwargs: Dict) -> str:
"""Extract vendor/provider from kwargs."""
litellm_params = kwargs.get("litellm_params", {}) or {}
return litellm_params.get("custom_llm_provider", "litellm")
return litellm_params.get("custom_llm_provider") or "litellm"
def _get_model_names(
self, kwargs: Dict, response_obj: ModelResponse
@ -335,7 +338,7 @@ class NewRelicLogger(CustomLogger):
"""
choices = response_obj.get("choices") or []
if choices and len(choices) > 0:
return choices[0].get("finish_reason", "unknown")
return choices[0].get("finish_reason") or "unknown"
return "unknown"
def _to_epoch_ms(self, t: Any) -> float:
@ -441,7 +444,7 @@ class NewRelicLogger(CustomLogger):
request_messages = kwargs.get("messages") or []
for msg in request_messages:
message_data = {
"role": msg.get("role", "user"),
"role": msg.get("role") or "user",
"sequence": sequence,
"response.model": response_model,
"vendor": vendor,
@ -465,7 +468,7 @@ class NewRelicLogger(CustomLogger):
message = choice.get("message", None)
if message:
message_data = {
"role": message.get("role", "assistant"),
"role": message.get("role") or "assistant",
"sequence": sequence,
"response.model": response_model,
"vendor": vendor,

View file

@ -220,6 +220,16 @@ class TestGetTraceContext:
assert trace_id is not None
assert len(trace_id) == 32
def test_extracts_trace_id_from_mixed_case_traceparent_header(self):
# Callers passing headers directly may not normalise case; per W3C spec
# header names are case-insensitive, so "Traceparent" must work too.
kwargs = make_kwargs()
kwargs["litellm_params"]["metadata"]["headers"] = {
"Traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00"
}
trace_id, span_id = self.logger._get_trace_context(kwargs)
assert trace_id == "4bf92f3577b34da6a3ce929d0e0e4736"
# ---------------------------------------------------------------------------
# 4. _extract_message_content edge cases
@ -414,6 +424,10 @@ class TestGetFinishReason:
def test_returns_unknown_when_choices_missing(self):
assert self.logger._get_finish_reason({}) == "unknown"
def test_returns_unknown_when_finish_reason_explicitly_none(self):
response = {"choices": [{"finish_reason": None}]}
assert self.logger._get_finish_reason(response) == "unknown"
class TestToEpochMs:
def setup_method(self):