fix(integrations): keep batch line items out of the built-in metering sinks

With store_batch_line_items_in_callbacks on, per-line batch events carry a
real response_cost and reached six metering sinks that lacked the
is_batch_line_item_event guard, so a completed batch metered aggregate +
per-line (roughly 2x true spend): prometheus (litellm_spend_metric and
request counters), openmeter, lago, datadog cost management (FOCUS
BilledCost), newrelic metrics, and the OTel v1/v2 gen_ai.usage.cost
metrics. Skip line-item events in each sink's success/failure metering
entry points, matching the guards already on the Postgres and ClickHouse
spend sinks. OTel spans for line items are still emitted; only their
cost/token metrics are skipped. Tests drive real sink instances on the
real dispatch lists through a completed 2-line batch and assert each
sink meters exactly the aggregate.
This commit is contained in:
Yucheng He 2026-10-05 12:26:15 -07:00
parent 8ec50a22f6
commit a35d4ee464
8 changed files with 419 additions and 2 deletions

View file

@ -14,6 +14,7 @@ from litellm.integrations.datadog.datadog_handler import (
get_datadog_service,
normalize_datadog_tag_value,
)
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -72,6 +73,10 @@ class DatadogCostManagementLogger(CustomBatchLogger):
super().__init__(**kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is billed by the aggregate aretrieve_batch event; a
# per-line FOCUS BilledCost row would double-count cloud spend.
if is_batch_line_item_event(kwargs):
return
try:
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)

View file

@ -11,6 +11,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
@ -121,6 +122,9 @@ class LagoLogger(CustomLogger):
return returned_val
def log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is billed by the aggregate aretrieve_batch event.
if is_batch_line_item_event(kwargs):
return
_url = os.getenv("LAGO_API_BASE")
assert _url is not None and isinstance(_url, str), (
f"LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={_url}"
@ -153,6 +157,8 @@ class LagoLogger(CustomLogger):
raise e
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
try:
verbose_logger.debug("ENTERS LAGO CALLBACK")
_url = os.getenv("LAGO_API_BASE")

View file

@ -34,6 +34,7 @@ from httpx import HTTPStatusError, Response
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -301,12 +302,17 @@ class NewRelicMetricsLogger(CustomBatchLogger):
await self._final_drain()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
# A batch line item is metered by the aggregate aretrieve_batch event.
if is_batch_line_item_event(kwargs):
return
try:
await self._log_async_event(standard_logging_object=kwargs.get("standard_logging_object", None))
except Exception as e: # noqa: BLE001 # logging must never break the request path
verbose_logger.exception("New Relic Metrics Layer Error - %s\n%s", e, traceback.format_exc())
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
if is_batch_line_item_event(kwargs):
return
try:
await self._log_async_event(standard_logging_object=kwargs.get("standard_logging_object", None))
except Exception as e: # noqa: BLE001 # logging must never break the request path

View file

@ -9,6 +9,7 @@ import httpx
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
@ -97,6 +98,9 @@ class OpenMeterLogger(CustomLogger):
}
def log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is billed by the aggregate aretrieve_batch event.
if is_batch_line_item_event(kwargs):
return
_url = os.getenv("OPENMETER_API_ENDPOINT", "https://openmeter.cloud")
if _url.endswith("/"):
_url += "api/v1/events"
@ -123,6 +127,8 @@ class OpenMeterLogger(CustomLogger):
raise e
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
_url = os.getenv("OPENMETER_API_ENDPOINT", "https://openmeter.cloud")
if _url.endswith("/"):
_url += "api/v1/events"

View file

@ -28,6 +28,7 @@ from litellm.integrations.otel.model.db_endpoint import db_span_attributes
from litellm.integrations.otel.model.metadata import flatten_metadata
from litellm.integrations.otel.model.semconv import LiteLLM, Metric
from litellm.integrations.otel.plumbing.otlp_tls import resolve_otlp_http_tls
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
@ -1655,6 +1656,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return True
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
# A batch line item's tokens and cost are metered by the aggregate
# aretrieve_batch event; per-line samples here would double-count them.
# Spans for line items are still emitted by _handle_success.
if is_batch_line_item_event(kwargs):
return
duration_s: Final = (end_time - start_time).total_seconds()
params: Final = kwargs.get("litellm_params") or {}
provider: Final = _provider_label(params.get("custom_llm_provider"))

View file

@ -33,6 +33,7 @@ from litellm.integrations.otel.model.semconv import (
resolve_provider,
)
from litellm.integrations.otel.model.utils import to_seconds
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -226,13 +227,17 @@ class GenAIMetricRecorder:
usage_is_replayed: Final = is_unbilled_non_inference_call_from_params(
kwargs.get("call_type"), kwargs.get("litellm_params"), response_obj
)
# A batch line item's tokens and cost are metered by the aggregate
# aretrieve_batch event; per-line samples here would double-count them.
# Duration samples for line items are still recorded.
batch_line_item: Final = is_batch_line_item_event(kwargs)
self._metrics.operation_duration.record(duration_s, attributes=common_attrs)
if not usage_is_replayed:
if not usage_is_replayed and not batch_line_item:
self._record_token_usage(response_obj, common_attrs)
cost: Final = kwargs.get("response_cost")
if cost:
if cost and not batch_line_item:
self._metrics.token_cost.record(cost, attributes=common_attrs)
self._record_time_to_first_token(kwargs, common_attrs)

View file

@ -35,6 +35,7 @@ from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker i
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
is_batch_line_item_event,
)
from litellm.litellm_core_utils.service_tier_utils import (
get_service_tier_from_standard_logging_payload,
@ -1345,6 +1346,10 @@ class PrometheusLogger(CustomLogger):
self._track_end_user_metric_series(counter, metric_name, _labels)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is metered by the aggregate aretrieve_batch event; a
# per-line sample here would double-count spend and requests.
if is_batch_line_item_event(kwargs):
return
# Define prometheus client
verbose_logger.debug(
"prometheus Logging - Enters success logging function (kwargs keys: %s)",
@ -2365,6 +2370,8 @@ class PrometheusLogger(CustomLogger):
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
verbose_logger.debug(
"prometheus Logging - Enters failure logging function (kwargs keys: %s)",
list(kwargs.keys()) if isinstance(kwargs, dict) else type(kwargs).__name__,

View file

@ -0,0 +1,376 @@
"""
Guards that keep batch line-item callback events out of the built-in metering sinks.
With ``litellm.store_batch_line_items_in_callbacks`` on, a completed batch emits
one child callback per JSONL line on top of the aggregate ``aretrieve_batch``
event. Every line carries a real ``response_cost``, so any billing/metering sink
without the ``is_batch_line_item_event`` guard meters aggregate + per-line and
reports roughly 2x the true spend.
These tests put REAL sink instances on the REAL dispatch lists, drive the REAL
``Logging.async_success_handler`` for a completed 2-line batch (1 success line +
1 error line, aggregate cost $1.50, per-line cost $0.03), and assert each sink
meters exactly the aggregate. OTel spans for line items must still be emitted:
the guard belongs on cost/token metrics, not on tracing.
"""
import io
import json
import os
import time
import uuid
from contextlib import redirect_stdout
from datetime import datetime
from types import SimpleNamespace
from typing import Any, Final
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import LiteLLMBatch, Usage
AGGREGATE_COST: Final[float] = 1.5
INPUT_JSONL: Final[bytes] = b"\n".join(
[
json.dumps(
{
"custom_id": "a",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi a"}]},
}
).encode(),
json.dumps(
{
"custom_id": "b",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi b"}]},
}
).encode(),
]
)
OUTPUT_JSONL: Final[bytes] = json.dumps(
{
"custom_id": "a",
"response": {
"status_code": 200,
"body": {
"id": "chatcmpl-line-1",
"model": "gpt-4o",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
},
},
}
).encode()
ERROR_JSONL: Final[bytes] = json.dumps(
{
"custom_id": "b",
"response": {"status_code": 400, "body": {"error": {"message": "boom"}}},
"error": {"message": "boom"},
}
).encode()
_FILE_BYTES: Final[dict[str, bytes]] = {
"input-file-1": INPUT_JSONL,
"output-file-1": OUTPUT_JSONL,
"error-file-1": ERROR_JSONL,
}
def _file_content(file_id: str, **_kwargs: Any) -> SimpleNamespace:
return SimpleNamespace(content=_FILE_BYTES[file_id])
def _batch() -> LiteLLMBatch:
return LiteLLMBatch(
id=f"batch_{uuid.uuid4().hex[:8]}",
object="batch",
endpoint="/v1/chat/completions",
input_file_id="input-file-1",
output_file_id="output-file-1",
error_file_id="error-file-1",
status="completed",
completion_window="24h",
created_at=1,
)
def _parent_logging() -> Logging:
logging_obj = Logging(
model="gpt-4o",
messages=[{"role": "user", "content": "<retrieve_batch>"}],
stream=False,
call_type="aretrieve_batch",
start_time=datetime.now(),
litellm_call_id=str(uuid.uuid4()),
function_id=str(uuid.uuid4()),
)
logging_obj.update_environment_variables(
litellm_params={
"metadata": {
"model_info": {"id": "dep-1"},
"model_group": "gpt-4o",
"user_api_key_user_id": "user-77",
"user_api_key_team_id": "team-7",
"user_api_key_team_alias": "team-seven",
"user_api_key_alias": "key-alias-1",
}
},
optional_params={},
custom_llm_provider="openai",
)
return logging_obj
class RecordingHTTP:
"""Stands in for a sink's HTTP egress object only; the sink logic is real."""
def __init__(self) -> None:
self.posts: list[dict[str, Any]] = []
async def post(self, url: str, data: Any = None, content: Any = None, **_kw: Any) -> SimpleNamespace:
self.posts.append({"url": url, "body": data if data is not None else content})
return SimpleNamespace(status_code=200, text="ok", raise_for_status=lambda: None)
async def put(self, url: str, content: Any = None, **_kw: Any) -> SimpleNamespace:
return SimpleNamespace(status_code=202, text="ok", raise_for_status=lambda: None)
async def _log_completed_batch(monkeypatch: pytest.MonkeyPatch, loggers: list) -> None:
batch_line_item_claim_cache.in_memory_cache.flush_cache()
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
saved_success = list(litellm._async_success_callback)
saved_failure = list(litellm._async_failure_callback)
litellm._async_success_callback = list(loggers)
litellm._async_failure_callback = list(loggers)
buf = io.StringIO()
try:
with (
patch("litellm.files.main.afile_content", new_callable=AsyncMock, side_effect=_file_content),
patch("litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02)),
redirect_stdout(buf),
):
await _parent_logging().async_success_handler(
result=_batch(),
batch_cost=1.5,
batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
batch_models=["gpt-4o"],
batch_successful_requests=1,
batch_failed_requests=1,
batch_prompt_cost=1.0,
batch_completion_cost=0.5,
)
finally:
litellm._async_success_callback = saved_success
litellm._async_failure_callback = saved_failure
def _openmeter_sinks() -> tuple[Any, RecordingHTTP]:
from litellm.integrations.openmeter import OpenMeterLogger
recorder = RecordingHTTP()
logger = OpenMeterLogger()
logger.async_http_handler = recorder
return logger, recorder
def _lago_sinks() -> tuple[Any, RecordingHTTP]:
from litellm.integrations.lago import LagoLogger
recorder = RecordingHTTP()
logger = LagoLogger()
logger.async_http_handler = recorder
return logger, recorder
@pytest.mark.asyncio
async def test_openmeter_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENMETER_API_KEY", "test-openmeter-key")
logger, recorder = _openmeter_sinks()
await _log_completed_batch(monkeypatch, [logger])
costs = [json.loads(p["body"])["data"]["cost"] for p in recorder.posts]
assert costs == [AGGREGATE_COST], (
f"OpenMeter metered per-line costs on top of the aggregate: {costs}; "
"line items must be billed only by the aggregate aretrieve_batch event"
)
@pytest.mark.asyncio
async def test_lago_bills_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LAGO_API_KEY", "test-lago-key")
monkeypatch.setenv("LAGO_API_BASE", "http://lago.invalid")
monkeypatch.setenv("LAGO_API_EVENT_CODE", "litellm-usage")
monkeypatch.setenv("LAGO_API_CHARGE_BY", "user_id")
logger, recorder = _lago_sinks()
await _log_completed_batch(monkeypatch, [logger])
costs = [json.loads(p["body"])["event"]["properties"]["response_cost"] for p in recorder.posts]
assert costs == [AGGREGATE_COST], (
f"Lago billed per-line costs on top of the aggregate: {costs}; "
"line items must be billed only by the aggregate aretrieve_batch event"
)
@pytest.mark.asyncio
async def test_datadog_cost_management_bills_only_the_aggregate_batch_event(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.integrations.datadog.datadog_cost_management import DatadogCostManagementLogger
monkeypatch.setenv("DD_API_KEY", "test-dd-key")
monkeypatch.setenv("DD_APP_KEY", "test-dd-app-key")
logger = DatadogCostManagementLogger(cost_tag_keys=[])
await _log_completed_batch(monkeypatch, [logger])
entries = list(logger.log_queue)
costs = [e.get("response_cost", 0) for e in entries]
assert costs == [AGGREGATE_COST], (
f"Datadog FOCUS BilledCost queued per-line entries on top of the aggregate: {costs}; "
"line items must be billed only by the aggregate aretrieve_batch event"
)
@pytest.mark.asyncio
async def test_newrelic_metrics_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.integrations.newrelic.newrelic_metrics import NewRelicMetricsLogger, build_metric_payload
logger = NewRelicMetricsLogger(newrelic_api_key="test-nr-key")
await _log_completed_batch(monkeypatch, [logger])
records = tuple(logger.log_queue)
costs = [r.response_cost for r in records]
assert costs == [AGGREGATE_COST], (
f"New Relic metered per-line costs on top of the aggregate: {costs}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)
now = time.time()
envelopes = build_metric_payload(records=records, window_start=now - 1, now=now)
cost_sum = 0.0
for envelope in envelopes:
for metric in envelope["metrics"]:
if "cost" in metric["name"]:
value = metric["value"]["sum"] if isinstance(metric["value"], dict) else metric["value"]
cost_sum += value
assert cost_sum == pytest.approx(AGGREGATE_COST)
@pytest.mark.asyncio
async def test_prometheus_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
prometheus_client = pytest.importorskip("prometheus_client")
from litellm.integrations.prometheus import PrometheusLogger
from prometheus_client import REGISTRY
for collector in list(REGISTRY._collector_to_names.keys()):
REGISTRY.unregister(collector)
logger = PrometheusLogger()
await _log_completed_batch(monkeypatch, [logger])
spend_samples = []
for metric in REGISTRY.collect():
if metric.name == "litellm_spend_metric":
spend_samples = [sample.value for sample in metric.samples if sample.name.endswith("_total")]
assert spend_samples, "litellm_spend_metric saw no samples at all"
assert sum(spend_samples) == pytest.approx(AGGREGATE_COST), (
f"Prometheus spend metric double-counted batch line items: {spend_samples}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)
def _otel_v1(monkeypatch: pytest.MonkeyPatch):
otel_sdk = pytest.importorskip("opentelemetry.sdk")
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from litellm.integrations.opentelemetry import OpenTelemetry as OTelV1, OpenTelemetryConfig
reader = InMemoryMetricReader()
span_exporter = InMemorySpanExporter()
tracer_provider = TracerProvider()
tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter))
logger = OTelV1(
config=OpenTelemetryConfig(exporter="console", enable_metrics=True),
callback_name="batch_sink_guard_v1",
tracer_provider=tracer_provider,
meter_provider=MeterProvider(metric_readers=[reader]),
)
return logger, reader, span_exporter
def _cost_points(reader: Any) -> list[float]:
data = reader.get_metrics_data()
points: list[float] = []
if data is None:
return points
for resource_metrics in data.resource_metrics:
for scope_metrics in resource_metrics.scope_metrics:
for metric in scope_metrics.metrics:
if not metric.name.endswith("cost"):
continue
for data_point in metric.data.data_points:
points.append(getattr(data_point, "sum", getattr(data_point, "value", None)))
return points
@pytest.mark.asyncio
async def test_otel_v1_meter_cost_skips_line_items_but_spans_stay(
monkeypatch: pytest.MonkeyPatch,
) -> None:
logger, reader, span_exporter = _otel_v1(monkeypatch)
await _log_completed_batch(monkeypatch, [logger])
costs = _cost_points(reader)
assert sum(costs) == pytest.approx(AGGREGATE_COST), (
f"OTel v1 gen_ai.usage.cost double-counted batch line items: {costs}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)
finished = span_exporter.get_finished_spans()
assert finished, "OTel v1 emitted no spans at all"
assert any("chatcmpl-line-1" in str(span.attributes) for span in finished), (
f"OTel v1 dropped the per-line span the feature exists to deliver: "
f"{[span.name for span in finished]}"
)
@pytest.mark.asyncio
async def test_otel_v2_meter_cost_skips_line_items(monkeypatch: pytest.MonkeyPatch) -> None:
pytest.importorskip("opentelemetry.sdk")
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
from opentelemetry.sdk.trace import TracerProvider
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
reader = InMemoryMetricReader()
logger = OpenTelemetryV2(
config=OpenTelemetryV2Config(exporter="console", enable_metrics=True),
callback_name="batch_sink_guard_v2",
tracer_provider=TracerProvider(),
meter_provider=MeterProvider(metric_readers=[reader]),
)
await _log_completed_batch(monkeypatch, [logger])
costs = _cost_points(reader)
assert sum(costs) == pytest.approx(AGGREGATE_COST), (
f"OTel v2 gen_ai.usage.cost double-counted batch line items: {costs}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)