fix(datadog): drain cost-management queue + opt-in FinOps tag allowlist (#28487)

* fix(datadog): drain cost-management queue + opt-in FinOps tag allowlist

* fix(datadog): guard non-dict callback_specific_params + log empty aggregation

* fix(datadog): block user-controlled tags from overwriting reserved cost-attribution dimensions

* fix(datadog): cast metadata to dict[str, Any] to satisfy mypy
This commit is contained in:
michelligabriele 2026-05-28 21:04:04 +02:00 • committed by GitHub
parent 69afcd09d0
commit 928f09f8a4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 343 additions and 45 deletions

View file

@ -2,10 +2,17 @@ import asyncio
import os
import time
from datetime import datetime
from typing import Dict, List, Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple, cast
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.datadog.datadog_handler import (
get_datadog_env,
get_datadog_hostname,
get_datadog_pod_name,
get_datadog_service,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -15,9 +22,30 @@ from litellm.types.integrations.datadog_cost_management import (
)
from litellm.types.utils import StandardLoggingPayload
# Reserved tag keys whose values come from trusted sources (infra env, LiteLLM
# core payload fields, or proxy-controlled auth metadata). User-supplied
# request_tags / metadata cannot overwrite these, even when the key is
# allowlisted via cost_tag_keys, because that would let an authenticated caller
# spoof cost attribution (e.g. request_tags=["team:victim-team"]).
_RESERVED_TAG_KEYS: frozenset = frozenset(
{
"env",
"service",
"host",
"pod_name",
"provider",
"model",
"model_id",
"team",
"user",
"model_group",
}
)
class DatadogCostManagementLogger(CustomBatchLogger):
def __init__(self, **kwargs):
def __init__(self, cost_tag_keys: Optional[List[str]] = None, **kwargs):
self.cost_tag_keys: List[str] = list(cost_tag_keys) if cost_tag_keys else []
self.dd_api_key = os.getenv("DD_API_KEY")
self.dd_app_key = os.getenv("DD_APP_KEY")
self.dd_site = os.getenv("DD_SITE", "datadoghq.com")
@ -68,20 +96,21 @@ class DatadogCostManagementLogger(CustomBatchLogger):
if not self.log_queue:
return
batch_to_send = self.log_queue[:]
self.log_queue = []
try:
# Aggregate costs from the batch
aggregated_entries = self._aggregate_costs(self.log_queue)
aggregated_entries = self._aggregate_costs(batch_to_send)
if not aggregated_entries:
verbose_logger.debug(
"Datadog Cost Management: batch produced no aggregable entries; "
"dropping %d log(s) from queue.",
len(batch_to_send),
)
return
# Send to Datadog
await self._upload_to_datadog(aggregated_entries)
# Clear queue only on success (or if we decide to drop on failure)
# CustomBatchLogger clears queue in flush_queue, so we just process here
except Exception as e:
self.log_queue = batch_to_send + self.log_queue
verbose_logger.exception(
f"Datadog Cost Management: Error in async_send_batch: {str(e)}"
)
@ -151,45 +180,81 @@ class DatadogCostManagementLogger(CustomBatchLogger):
return list(aggregator.values())
def _extract_tags(self, log: StandardLoggingPayload) -> Dict[str, str]:
from litellm.integrations.datadog.datadog_handler import (
get_datadog_env,
get_datadog_hostname,
get_datadog_pod_name,
get_datadog_service,
)
tags = {
tags: Dict[str, str] = {
"env": get_datadog_env(),
"service": get_datadog_service(),
"host": get_datadog_hostname(),
"pod_name": get_datadog_pod_name(),
}
# Add metadata as tags
metadata = log.get("metadata", {})
if metadata:
# Add user info
# Add user info
if metadata.get("user_api_key_alias"):
tags["user"] = str(metadata["user_api_key_alias"])
# Always-on canonical FOCUS dimensions from top-level payload fields.
# Non-sensitive and required for Datadog Custom Costs per-model attribution.
self._add_tag(tags, "provider", log.get("custom_llm_provider"))
self._add_tag(tags, "model", log.get("model"))
self._add_tag(tags, "model_id", log.get("model_id"))
# Add Team Tag
team_tag = (
metadata.get("user_api_key_team_alias")
or metadata.get("team_alias") # type: ignore
or metadata.get("user_api_key_team_id")
or metadata.get("team_id") # type: ignore
)
# cast because StandardLoggingMetadata is a TypedDict; we iterate it
# as a generic mapping below.
metadata: Dict[str, Any] = cast(Dict[str, Any], log.get("metadata") or {})
if team_tag:
tags["team"] = str(team_tag)
# model_group is not in StandardLoggingMetadata TypedDict, so we need to access it via dict.get()
model_group = metadata.get("model_group") # type: ignore[misc]
if model_group:
tags["model_group"] = str(model_group)
# Backwards-compat: team/user/model_group preserved regardless of allowlist.
if metadata.get("user_api_key_alias"):
tags["user"] = str(metadata["user_api_key_alias"])
team_tag = (
metadata.get("user_api_key_team_alias")
or metadata.get("team_alias")
or metadata.get("user_api_key_team_id")
or metadata.get("team_id")
)
if team_tag:
tags["team"] = str(team_tag)
if metadata.get("model_group"):
tags["model_group"] = str(metadata["model_group"])
# Allowlist-gated: request_tags (split on `:`) and arbitrary metadata.*.
# Reserved keys are hard-blocked here regardless of allowlist membership —
# see _RESERVED_TAG_KEYS for the rationale.
if self.cost_tag_keys:
allow = set(self.cost_tag_keys)
for rt in log.get("request_tags") or []:
if not isinstance(rt, str) or ":" not in rt:
continue
k, _, v = rt.partition(":")
if k in allow and v:
self._set_custom_tag(tags, k, v)
for k, v in metadata.items():
if k in allow and v is not None and not isinstance(v, (dict, list)):
self._set_custom_tag(tags, k, str(v))
for nested_key in ("spend_logs_metadata", "requester_metadata"):
nested = metadata.get(nested_key)
if isinstance(nested, dict):
for k, v in nested.items():
if (
k in allow
and v is not None
and not isinstance(v, (dict, list))
):
self._set_custom_tag(tags, k, str(v))
return tags
@staticmethod
def _set_custom_tag(tags: Dict[str, str], key: str, value: str) -> None:
if key in _RESERVED_TAG_KEYS:
verbose_logger.debug(
"Datadog Cost Management: dropping user-supplied tag %r=%r — "
"key is reserved for trusted cost attribution.",
key,
value,
)
return
tags[key] = value
@staticmethod
def _add_tag(tags: Dict[str, str], key: str, value: Any) -> None:
if value:
tags[key] = str(value)
async def _upload_to_datadog(self, payload: List[Dict]):
if not self.dd_api_key or not self.dd_app_key:
return
@ -201,8 +266,6 @@ class DatadogCostManagementLogger(CustomBatchLogger):
}
# The API endpoint expects a list of objects directly in the body (file content behavior)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
data_json = safe_dumps(payload)
response = await self.async_client.put(

View file

@ -317,7 +317,15 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
DatadogCostManagementLogger,
)
datadog_cost_management_obj = DatadogCostManagementLogger()
init_params = {}
if (
"datadog_cost_management" in callback_specific_params
and isinstance(
callback_specific_params["datadog_cost_management"], dict
)
):
init_params = callback_specific_params["datadog_cost_management"]
datadog_cost_management_obj = DatadogCostManagementLogger(**init_params)
imported_list.append(datadog_cost_management_obj)
elif isinstance(callback, CustomLogger):
imported_list.append(callback)

View file

@ -1,4 +1,4 @@
from typing import Dict, Optional, TypedDict
from typing import Dict, List, Optional, TypedDict
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
@ -9,7 +9,7 @@ class DatadogCostManagementInitParams(StandardCustomLoggerInitParams):
Init params for Datadog Cost Management
"""
datadog_cost_management_params: Optional[Dict] = None
cost_tag_keys: Optional[List[str]] = None
class DatadogFOCUSCostEntry(TypedDict):

View file

@ -3,7 +3,7 @@ import time
from unittest.mock import AsyncMock
import pytest
from httpx import Response
from httpx import Request, Response
from litellm.integrations.datadog.datadog_cost_management import (
DatadogCostManagementLogger,
@ -167,3 +167,230 @@ async def test_async_send_batch(clean_env):
content = json.loads(call_args[1]["content"])
assert content[0]["ProviderName"] == "openai"
assert content[0]["BilledCost"] == 0.01
_PUT_REQUEST = Request("PUT", "https://api.test.datadoghq.com/api/v2/cost/custom_costs")
@pytest.mark.asyncio
async def test_async_send_batch_clears_queue_on_success(clean_env):
"""Bug 1 regression: log_queue must be empty after a successful upload."""
logger = DatadogCostManagementLogger()
logger.async_client = AsyncMock()
logger.async_client.put.return_value = Response(
202, json={"status": "ok"}, request=_PUT_REQUEST
)
logger.log_queue = [
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
]
await logger.async_send_batch()
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_async_send_batch_preserves_events_added_during_upload(clean_env):
"""Events appended while the upload is in flight survive (land on the cleared queue)."""
logger = DatadogCostManagementLogger()
later_event = StandardLoggingPayload(
custom_llm_provider="anthropic",
model="claude-3",
response_cost=0.02,
startTime=time.time(),
)
async def slow_put(*args, **kwargs):
logger.log_queue.append(later_event)
return Response(202, json={"status": "ok"}, request=_PUT_REQUEST)
logger.async_client = AsyncMock()
logger.async_client.put.side_effect = slow_put
logger.log_queue = [
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
]
await logger.async_send_batch()
assert logger.log_queue == [later_event]
@pytest.mark.asyncio
async def test_async_send_batch_requeues_on_upload_failure(clean_env):
"""Failed upload requeues the original batch (no data loss)."""
logger = DatadogCostManagementLogger()
logger.async_client = AsyncMock()
logger.async_client.put.side_effect = Exception("boom")
original = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
logger.log_queue = [original]
await logger.async_send_batch()
assert logger.log_queue == [original]
@pytest.mark.asyncio
async def test_extract_tags_emits_canonical_focus_dimensions(clean_env):
"""provider, model, model_id always emitted regardless of cost_tag_keys."""
logger = DatadogCostManagementLogger()
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4o",
model_id="router-id-123",
response_cost=0.01,
startTime=time.time(),
)
tags = logger._extract_tags(log)
assert tags["provider"] == "openai"
assert tags["model"] == "gpt-4o"
assert tags["model_id"] == "router-id-123"
@pytest.mark.asyncio
async def test_extract_tags_allowlist_filters_request_tags(clean_env):
"""Only request_tags whose key is in cost_tag_keys reach the Tags dict."""
logger = DatadogCostManagementLogger(cost_tag_keys=["capability", "tier"])
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
request_tags=["capability:chat", "tier:gold", "secret:disallowed"],
)
tags = logger._extract_tags(log)
assert tags["capability"] == "chat"
assert tags["tier"] == "gold"
assert "secret" not in tags
@pytest.mark.asyncio
async def test_extract_tags_allowlist_filters_metadata(clean_env):
"""Only metadata keys in cost_tag_keys flow through; others (and dict/list values) are dropped."""
logger = DatadogCostManagementLogger(cost_tag_keys=["capability", "owner"])
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
metadata={
"capability": "chat",
"owner": "team-x",
"secret_field": "sensitive",
"nested_obj": {"a": 1},
},
)
tags = logger._extract_tags(log)
assert tags["capability"] == "chat"
assert tags["owner"] == "team-x"
assert "secret_field" not in tags
assert "nested_obj" not in tags
@pytest.mark.asyncio
async def test_extract_tags_empty_allowlist_default(clean_env):
"""With no cost_tag_keys, request_tags and arbitrary metadata.* do NOT leak into Tags."""
logger = DatadogCostManagementLogger()
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
request_tags=["capability:chat"],
metadata={"capability": "chat", "user_api_key_alias": "alice"},
)
tags = logger._extract_tags(log)
assert "capability" not in tags
# Backwards-compat keys still flow:
assert tags["user"] == "alice"
@pytest.mark.asyncio
async def test_extract_tags_nested_metadata_allowlisted(clean_env):
"""spend_logs_metadata and requester_metadata get spread one level under the allowlist."""
logger = DatadogCostManagementLogger(cost_tag_keys=["env", "platform"])
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
metadata={
"spend_logs_metadata": {"platform": "web", "ignored": "x"},
"requester_metadata": {"env": "prod"},
},
)
tags = logger._extract_tags(log)
assert tags["platform"] == "web"
# "env" is a reserved trusted dimension — requester_metadata.env must NOT
# overwrite the value sourced from get_datadog_env().
assert tags["env"] != "prod"
assert "ignored" not in tags
@pytest.mark.asyncio
async def test_extract_tags_allowlist_cannot_override_reserved_dimensions(clean_env):
"""
Reserved tag keys (env, service, host, pod_name, provider, model, model_id,
team, user, model_group) must not be overwritten by user-controlled
request_tags or metadata, even when listed in cost_tag_keys.
"""
reserved = [
"env",
"service",
"host",
"pod_name",
"provider",
"model",
"model_id",
"team",
"user",
"model_group",
]
logger = DatadogCostManagementLogger(cost_tag_keys=reserved)
metadata_attack = {k: f"attacker-meta-{k}" for k in reserved}
metadata_attack["user_api_key_alias"] = "trusted-user"
metadata_attack["user_api_key_team_alias"] = "trusted-team"
metadata_attack["model_group"] = "trusted-group"
metadata_attack["spend_logs_metadata"] = {
k: f"attacker-spend-{k}" for k in reserved
}
metadata_attack["requester_metadata"] = {k: f"attacker-req-{k}" for k in reserved}
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
model_id="router-id-123",
response_cost=0.01,
startTime=time.time(),
request_tags=[f"{k}:attacker-rt-{k}" for k in reserved],
metadata=metadata_attack,
)
tags = logger._extract_tags(log)
# Canonical FOCUS dims keep their trusted (top-level payload) values.
assert tags["provider"] == "openai"
assert tags["model"] == "gpt-4"
assert tags["model_id"] == "router-id-123"
# Backwards-compat trusted dims keep their proxy-controlled metadata values.
assert tags["user"] == "trusted-user"
assert tags["team"] == "trusted-team"
assert tags["model_group"] == "trusted-group"
# No reserved key carries an attacker-supplied prefix from any path.
for k in reserved:
assert not tags[k].startswith("attacker-"), (
f"reserved key {k!r} was overwritten by user-controlled input: "
f"{tags[k]!r}"
)