fix(interactions): bill background interactions once completed via cost polling

This commit is contained in:
mateo-berri 2026-07-14 20:05:33 -07:00
parent 887fc0a73c
commit 59d4e52a3d
10 changed files with 519 additions and 59 deletions

View file

@ -1457,6 +1457,17 @@ STALE_OBJECT_CLEANUP_BATCH_SIZE = max(1, int(os.getenv("STALE_OBJECT_CLEANUP_BAT
# installations with large numbers of stale managed objects).
_batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower()
PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true"
BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS = float(
os.getenv("BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS", 5)
)
BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS = float(
os.getenv("BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS", 60)
)
BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS = float(
os.getenv("BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS", 3600)
)
_background_interaction_cost_polling_env = os.getenv("BACKGROUND_INTERACTION_COST_POLLING_ENABLED", "true").lower()
BACKGROUND_INTERACTION_COST_POLLING_ENABLED = _background_interaction_cost_polling_env == "true"
PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605))
PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10

View file

@ -0,0 +1,131 @@
"""
Cost tracking for background interactions.
A create request with ``background=true`` returns ``in_progress`` with no
usage block, and GET polls are deliberately never billed (billing them would
double-charge every poll; the GET response also does not echo ``background``,
so a poll cannot be told apart from a re-fetch of an already-billed
interaction). The create call is therefore the only place that can own
billing: it schedules a poll task that fetches the interaction until it
reaches a terminal status and logs the final usage as a single success event
attributed to the original request.
"""
import asyncio
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Awaitable, Callable, Iterator, Optional
from litellm._logging import verbose_logger
from litellm.constants import (
BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS,
BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS,
BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS,
BACKGROUND_INTERACTION_COST_POLLING_ENABLED,
)
from litellm.types.interactions import InteractionsAPIResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
_TERMINAL_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete", "budget_exceeded"})
@dataclass(frozen=True, slots=True)
class BackgroundInteractionPollContext:
interaction_id: str
custom_llm_provider: str
logging_obj: "LiteLLMLoggingObj"
api_key: Optional[str] = None
api_base: Optional[str] = None
initial_interval_seconds: float = BACKGROUND_INTERACTION_COST_POLL_INITIAL_INTERVAL_SECONDS
max_interval_seconds: float = BACKGROUND_INTERACTION_COST_POLL_MAX_INTERVAL_SECONDS
timeout_seconds: float = BACKGROUND_INTERACTION_COST_POLL_TIMEOUT_SECONDS
FetchInteraction = Callable[[BackgroundInteractionPollContext], Awaitable[InteractionsAPIResponse]]
async def _fetch_interaction(context: BackgroundInteractionPollContext) -> InteractionsAPIResponse:
from litellm.interactions import aget
return await aget(
interaction_id=context.interaction_id,
custom_llm_provider=context.custom_llm_provider,
**{"api_key": context.api_key, "api_base": context.api_base, "no-log": True},
)
def _poll_intervals(initial: float, maximum: float, timeout: float) -> Iterator[float]:
elapsed = 0.0
interval = initial
while elapsed + interval <= timeout:
yield interval
elapsed += interval
interval = min(interval * 2, maximum)
async def poll_and_log_background_interaction_cost(
context: BackgroundInteractionPollContext,
fetch_interaction: FetchInteraction = _fetch_interaction,
) -> None:
for interval in _poll_intervals(
initial=context.initial_interval_seconds,
maximum=context.max_interval_seconds,
timeout=context.timeout_seconds,
):
await asyncio.sleep(interval)
try:
response = await fetch_interaction(context)
except Exception as e: # noqa: BLE001 # any fetch error must not kill the billing poll loop
verbose_logger.debug(
"Background interaction cost poll for %s failed, will retry: %s",
context.interaction_id,
e,
)
continue
if response.status not in _TERMINAL_STATUSES:
continue
if response.usage is not None:
await context.logging_obj.async_log_background_interaction_completion(result=response)
return
verbose_logger.warning(
"Gave up cost polling for background interaction %s after %ss; its usage will not be tracked",
context.interaction_id,
context.timeout_seconds,
)
_ACTIVE_POLL_TASKS: set["asyncio.Task[None]"] = set() # mutable-ok: asyncio requires strong refs to running tasks
def maybe_schedule_background_interaction_cost_polling(
response: Any,
create_kwargs: dict[str, Any],
custom_llm_provider: str,
) -> Optional["asyncio.Task[None]"]:
from litellm.litellm_core_utils.litellm_logging import Logging
if not BACKGROUND_INTERACTION_COST_POLLING_ENABLED:
return None
if not isinstance(response, InteractionsAPIResponse):
return None
if response.status != "in_progress" or not response.id:
return None
logging_obj = create_kwargs.get("litellm_logging_obj")
if not isinstance(logging_obj, Logging):
return None
try:
asyncio.get_running_loop()
except RuntimeError:
return None
context = BackgroundInteractionPollContext(
interaction_id=response.id,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
api_key=create_kwargs.get("api_key"),
api_base=create_kwargs.get("api_base"),
)
task = asyncio.create_task(poll_and_log_background_interaction_cost(context))
_ACTIVE_POLL_TASKS.add(task)
task.add_done_callback(_ACTIVE_POLL_TASKS.discard)
return task

View file

@ -39,6 +39,9 @@ from typing import Any, AsyncIterator, Coroutine, Dict, Iterator, List, Optional
import httpx
import litellm
from litellm.interactions.background_cost_polling import (
maybe_schedule_background_interaction_cost_polling,
)
from litellm.interactions.http_handler import interactions_http_handler
from litellm.interactions.utils import (
InteractionsAPIRequestUtils,
@ -170,6 +173,12 @@ async def acreate(
else:
response = init_response
maybe_schedule_background_interaction_cost_polling(
response=response,
create_kwargs=kwargs,
custom_llm_provider=custom_llm_provider,
)
return response # type: ignore
except Exception as e:
raise litellm.exception_type(

View file

@ -1895,7 +1895,11 @@ class Logging(LiteLLMLoggingBaseClass):
or isinstance(logging_result, FineTuningJob)
or isinstance(logging_result, LiteLLMBatch)
or isinstance(logging_result, ResponsesAPIResponse)
or (isinstance(logging_result, InteractionsAPIResponse) and self._is_interactions_create_call_type())
or (
isinstance(logging_result, InteractionsAPIResponse)
and logging_result.usage is not None
and self._is_interactions_create_call_type()
)
or isinstance(logging_result, OpenAIFileObject)
or isinstance(logging_result, LiteLLMRealtimeStreamLoggingObject)
or isinstance(logging_result, OpenAIModerationResponse)
@ -1921,6 +1925,13 @@ class Logging(LiteLLMLoggingBaseClass):
interaction. The proxy sets ``call_type`` from its route_type
(``create_interaction``/``acreate_interaction``); the SDK sets it from
the decorated function name (``create``/``acreate``).
Recognition additionally requires a usage block (checked at the call
site): a ``background=true`` create returns ``in_progress`` without
usage, and billing it would write a $0 spend log under the interaction
id that collides with the row the background poll task writes once the
interaction completes (see
``litellm.interactions.background_cost_polling``).
"""
return self.call_type in (
CallTypes.create_interaction.value,
@ -1929,6 +1940,20 @@ class Logging(LiteLLMLoggingBaseClass):
"acreate",
)
async def async_log_background_interaction_completion(
self,
result: InteractionsAPIResponse,
) -> None:
"""
Log the terminal result of a background interaction as a fresh success
event. The create request already ran success logging for its
``in_progress`` response (no usage, so no cost was tracked); clearing
the dedup flag lets the completed result flow through cost calculation
and spend tracking exactly once, spanning create to completion.
"""
self.model_call_details.pop("has_logged_async_success", None)
await self.async_success_handler(result=result)
def _flush_passthrough_collected_chunks_helper(
self,
raw_bytes: List[bytes],

View file

@ -279,6 +279,12 @@ class _ProxyDBLogger(CustomLogger):
await _release_budget_reservation(budget_reservation=budget_reservation)
else:
await _release_budget_reservation(budget_reservation=budget_reservation)
if _is_unbilled_in_progress_interaction(completion_response):
verbose_proxy_logger.debug(
"Cost tracking deferred for in-progress background interaction; "
"a poll task logs the final usage once it completes"
)
return
# Non-model call types (health checks, afile_delete) have no model or standard_logging_object.
# Use .get() for "stream" to avoid KeyError on health checks.
# WS session wrappers (_aresponses_websocket, _arealtime) also reach here with
@ -418,6 +424,12 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None:
bucket[key] = value
def _is_unbilled_in_progress_interaction(completion_response: Any) -> bool:
from litellm.types.interactions import InteractionsAPIResponse
return isinstance(completion_response, InteractionsAPIResponse) and completion_response.usage is None
def _should_track_cost_callback(
user_api_key: Optional[str],
user_id: Optional[str],

View file

@ -130,9 +130,7 @@ def classify_value(value: object, key: str = "scan") -> ValueClass:
return "plaintext"
if value.startswith(_V2_GCM_PREFIX):
return "migrated"
decrypted = decrypt_value_helper(
value=value, key=key, exception_type="debug", return_original_value=False
)
decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False)
if decrypted is None:
# Did not decrypt under nacl and has no v2 marker: legacy plaintext.
return "plaintext"
@ -151,9 +149,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object:
return value
if value.startswith(_V2_GCM_PREFIX):
return value # idempotent: already migrated
decrypted = decrypt_value_helper(
value=value, key=key, exception_type="debug", return_original_value=False
)
decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False)
if decrypted is None:
# Either legacy plaintext (no ciphertext to migrate) or corrupt. Either
# way, do not overwrite — preserve the value as stored.
@ -161,9 +157,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object:
return encrypt_value_helper(decrypted)
def reencrypt_selective_dict(
data: dict[str, object], sensitive_keys: list[str]
) -> dict[str, object]:
def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) -> dict[str, object]:
"""Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted.
Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is.
@ -212,9 +206,7 @@ async def _migrate_config_settings_row(
dict with selected sensitive fields (vantage_settings / cloudzero_settings).
"""
report = LocationReport(location=param_name)
record = await prisma_client.db.litellm_config.find_unique(
where={"param_name": param_name}
)
record = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name})
if record is None or record.param_value is None:
return report
@ -266,9 +258,7 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR
every present string field.
"""
report = LocationReport(location="sso_config")
record = await prisma_client.db.litellm_ssoconfig.find_unique(
where={"id": "sso_config"}
)
record = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"})
if record is None or record.sso_settings is None:
return report
@ -344,9 +334,7 @@ async def _migrate_callback_vars_table(
rows = await table.find_many()
for row in rows or []:
metadata = getattr(row, "metadata", None)
if not isinstance(metadata, dict) or (
"logging" not in metadata and "callback_settings" not in metadata
):
if not isinstance(metadata, dict) or ("logging" not in metadata and "callback_settings" not in metadata):
continue
# Classify every callback-var value directly (strip the litellm_enc::
@ -534,9 +522,7 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport:
"""Scan the ``environment_variables`` config row (``param_value`` dict)."""
report = LocationReport(location="config_environment_variables")
try:
record = await prisma_client.db.litellm_config.find_unique(
where={"param_name": "environment_variables"}
)
record = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"})
except Exception as e: # pragma: no cover - defensive
verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e))
return report
@ -557,11 +543,7 @@ async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]:
"""Read-only classification of every rotation-covered table. No writes."""
reports: list[LocationReport] = []
for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS:
reports.append(
await _scan_one_table(
prisma_client, location, db_attr, json_cols, scalar_cols
)
)
reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols))
reports.append(await _scan_config_env_vars(prisma_client))
return reports
@ -575,9 +557,7 @@ _VANTAGE_SENSITIVE = ["api_key", "integration_token"]
_CLOUDZERO_SENSITIVE = ["api_key"]
async def _migrate_covered_tables(
prisma_client: object, user_api_key_dict: object
) -> list[LocationReport]:
async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: object) -> list[LocationReport]:
"""Re-encrypt the tables already covered by ``_rotate_master_key`` (model
table, credentials, MCP credential/env tables, config environment_variables)
by running that orchestrator in *same-key* mode. With the AES gate on, the
@ -597,8 +577,7 @@ async def _migrate_covered_tables(
current_key = _get_salt_key()
if current_key is None:
raise RuntimeError(
"Cannot migrate covered tables: no salt key / master key is set. "
"Set LITELLM_SALT_KEY before migrating."
"Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating."
)
await _rotate_master_key(
prisma_client=cast("PrismaClient", prisma_client),
@ -648,19 +627,9 @@ async def migrate_encryption(
# Net-new walkers (items 3, 4, 11, 12, 13).
report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run))
report.add(
await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run
)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run
)
)
report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run))
report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run))
report.add(await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run))
report.add(await _migrate_sso_config(prisma_client, dry_run))
return report
@ -683,20 +652,10 @@ async def check_encryption(prisma_client: object) -> MigrationReport:
# Net-new walker locations, in dry-run (read-only) mode.
report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True))
report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True))
report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True))
report.add(
await _migrate_callback_vars_table(
prisma_client, "verification_token", dry_run=True
)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True
)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True
)
await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True)
)
report.add(await _migrate_sso_config(prisma_client, dry_run=True))
return report

View file

@ -0,0 +1,184 @@
import asyncio
import time
import pytest
from litellm.interactions.background_cost_polling import (
BackgroundInteractionPollContext,
maybe_schedule_background_interaction_cost_polling,
poll_and_log_background_interaction_cost,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
from litellm.types.interactions import InteractionsAPIResponse
USAGE_BLOCK = {
"total_tokens": 175,
"total_input_tokens": 100,
"input_tokens_by_modality": [{"modality": "text", "tokens": 100}],
"total_cached_tokens": 0,
"total_output_tokens": 50,
"output_tokens_by_modality": [{"modality": "text", "tokens": 50}],
"total_tool_use_tokens": 0,
"total_thought_tokens": 25,
}
def _logging_obj(call_type: str = "acreate_interaction") -> LitellmLogging:
logging_obj = LitellmLogging(
model="gemini-2.5-flash",
messages=[],
stream=False,
call_type=call_type,
start_time=time.time(),
litellm_call_id="bg-interactions-call-id",
function_id="bg-interactions-fn-id",
)
logging_obj.update_environment_variables(
litellm_params={},
optional_params={},
model="gemini-2.5-flash",
custom_llm_provider="gemini",
input="hi",
)
return logging_obj
def _context(logging_obj: LitellmLogging, timeout_seconds: float = 1.0) -> BackgroundInteractionPollContext:
return BackgroundInteractionPollContext(
interaction_id="interactions/bg-abc",
custom_llm_provider="gemini",
logging_obj=logging_obj,
initial_interval_seconds=0.001,
max_interval_seconds=0.002,
timeout_seconds=timeout_seconds,
)
def _response(status: str, with_usage: bool) -> InteractionsAPIResponse:
return InteractionsAPIResponse(
id="interactions/bg-abc",
model="gemini-2.5-flash",
status=status,
steps=[],
usage=dict(USAGE_BLOCK) if with_usage else None,
)
def _fetch_sequence(*responses):
remaining = list(responses)
calls = []
async def fetch(context):
calls.append(context.interaction_id)
item = remaining.pop(0) if len(remaining) > 1 else remaining[0]
if isinstance(item, Exception):
raise item
return item
return fetch, calls
@pytest.mark.asyncio
async def test_poller_bills_once_when_interaction_completes():
logging_obj = _logging_obj()
fetch, calls = _fetch_sequence(
_response("in_progress", with_usage=False),
_response("completed", with_usage=True),
)
await poll_and_log_background_interaction_cost(_context(logging_obj), fetch_interaction=fetch)
assert len(calls) == 2
assert logging_obj.model_call_details["response_cost"] > 0
assert logging_obj.model_call_details["standard_logging_object"]["total_tokens"] == 175
@pytest.mark.asyncio
async def test_poller_stops_without_billing_on_terminal_status_without_usage():
logging_obj = _logging_obj()
fetch, calls = _fetch_sequence(_response("failed", with_usage=False))
await poll_and_log_background_interaction_cost(_context(logging_obj), fetch_interaction=fetch)
assert len(calls) == 1
assert logging_obj.model_call_details.get("response_cost") is None
@pytest.mark.asyncio
async def test_poller_gives_up_after_timeout_without_billing():
logging_obj = _logging_obj()
fetch, calls = _fetch_sequence(_response("in_progress", with_usage=False))
await poll_and_log_background_interaction_cost(
_context(logging_obj, timeout_seconds=0.01),
fetch_interaction=fetch,
)
assert len(calls) >= 2
assert logging_obj.model_call_details.get("response_cost") is None
@pytest.mark.asyncio
async def test_poller_retries_after_fetch_error_and_still_bills():
logging_obj = _logging_obj()
fetch, calls = _fetch_sequence(
RuntimeError("transient network error"),
_response("completed", with_usage=True),
)
await poll_and_log_background_interaction_cost(_context(logging_obj), fetch_interaction=fetch)
assert len(calls) == 2
assert logging_obj.model_call_details["response_cost"] > 0
@pytest.mark.asyncio
async def test_schedule_creates_poll_task_for_in_progress_create():
logging_obj = _logging_obj()
task = maybe_schedule_background_interaction_cost_polling(
response=_response("in_progress", with_usage=False),
create_kwargs={"litellm_logging_obj": logging_obj},
custom_llm_provider="gemini",
)
assert isinstance(task, asyncio.Task)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
@pytest.mark.asyncio
@pytest.mark.parametrize(
"response,create_kwargs",
[
(_response("completed", with_usage=True), {"litellm_logging_obj": "placeholder"}),
(_response("in_progress", with_usage=False), {}),
("not a response", {"litellm_logging_obj": "placeholder"}),
],
)
async def test_schedule_skips_non_pollable_results(response, create_kwargs):
if create_kwargs.get("litellm_logging_obj") == "placeholder":
create_kwargs = {"litellm_logging_obj": _logging_obj()}
task = maybe_schedule_background_interaction_cost_polling(
response=response,
create_kwargs=create_kwargs,
custom_llm_provider="gemini",
)
assert task is None
@pytest.mark.asyncio
async def test_schedule_respects_kill_switch(monkeypatch):
import litellm.interactions.background_cost_polling as module
monkeypatch.setattr(module, "BACKGROUND_INTERACTION_COST_POLLING_ENABLED", False)
task = maybe_schedule_background_interaction_cost_polling(
response=_response("in_progress", with_usage=False),
create_kwargs={"litellm_logging_obj": _logging_obj()},
custom_llm_provider="gemini",
)
assert task is None

View file

@ -3812,10 +3812,67 @@ def test_interactions_response_is_recognized_for_logging(call_type):
from litellm.types.interactions import InteractionsAPIResponse
logging_obj = _interactions_logging_obj(stream=False, call_type=call_type)
response = InteractionsAPIResponse(id="interactions/abc", model="gemini-2.5-flash", status="completed")
response = InteractionsAPIResponse(
id="interactions/abc",
model="gemini-2.5-flash",
status="completed",
usage=dict(INTERACTIONS_USAGE_BLOCK),
)
assert logging_obj._is_recognized_call_type_for_logging(logging_result=response) is True
@pytest.mark.parametrize("call_type", ["acreate", "acreate_interaction"])
def test_in_progress_background_create_is_not_billed(call_type):
import datetime as dt
from litellm.types.interactions import InteractionsAPIResponse
logging_obj = _interactions_logging_obj(stream=False, call_type=call_type)
response = InteractionsAPIResponse(id="interactions/abc", model="gemini-2.5-flash", status="in_progress")
assert logging_obj._is_recognized_call_type_for_logging(logging_result=response) is False
logging_obj._success_handler_helper_fn(
result=response,
start_time=dt.datetime.now(),
end_time=dt.datetime.now(),
cache_hit=False,
)
assert logging_obj.model_call_details.get("response_cost") is None
assert logging_obj.model_call_details.get("standard_logging_object") is None
@pytest.mark.asyncio
async def test_background_interaction_completion_rebills_after_in_progress_success():
import datetime as dt
from litellm.types.interactions import InteractionsAPIResponse
logging_obj = _interactions_logging_obj(stream=False)
in_progress = InteractionsAPIResponse(id="interactions/abc", model="gemini-2.5-flash", status="in_progress")
await logging_obj.async_success_handler(
result=in_progress,
start_time=dt.datetime.now(),
end_time=dt.datetime.now(),
)
assert logging_obj.model_call_details.get("response_cost") is None
assert logging_obj.should_run_logging(event_type="async_success") is False
completed = InteractionsAPIResponse(
id="interactions/abc",
model="gemini-2.5-flash",
status="completed",
steps=[],
usage=dict(INTERACTIONS_USAGE_BLOCK),
)
await logging_obj.async_log_background_interaction_completion(result=completed)
assert logging_obj.model_call_details["response_cost"] > 0
assert logging_obj.model_call_details["standard_logging_object"]["total_tokens"] == 175
@pytest.mark.parametrize(
"call_type",
["aget", "get", "aget_interaction", "adelete_interaction", "acancel_interaction"],

View file

@ -604,6 +604,49 @@ async def test_track_cost_callback_skips_when_no_standard_logging_object():
mock_proxy_logging.failed_tracking_alert.assert_not_called()
@pytest.mark.asyncio
async def test_track_cost_callback_defers_in_progress_background_interaction():
"""
A background=true interaction create returns in_progress with no usage
block, so its success event has a model but no standard_logging_object.
The callback must skip quietly (billing happens later via the background
poll task) instead of raising 'Cost tracking failed' and alerting.
"""
from litellm.types.interactions import InteractionsAPIResponse
logger = _ProxyDBLogger()
kwargs = {
"call_type": "acreate_interaction",
"model": "gemini/gemini-3-flash-preview",
"litellm_call_id": "test-call-id",
"litellm_params": {},
"stream": False,
}
in_progress_response = InteractionsAPIResponse(
id="interactions/bg-abc",
model="gemini-3-flash-preview",
status="in_progress",
)
with patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
) as mock_proxy_logging:
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=in_progress_response,
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_proxy_logging.db_spend_update_writer.update_database.assert_not_called()
mock_proxy_logging.failed_tracking_alert.assert_not_called()
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj():
"""

View file

@ -3512,3 +3512,32 @@ def test_completion_cost_bills_interactions_api_response():
)
assert cost == pytest.approx(expected)
assert cost > 0
def test_completion_cost_bills_interactions_video_output_at_video_rate():
from litellm.types.interactions import InteractionsAPIResponse
model_info = litellm.get_model_info(model="gemini-omni-flash-preview", custom_llm_provider="gemini")
video_tokens = 5792 * 8
response = InteractionsAPIResponse(
id="interactions/video123",
model="gemini-omni-flash-preview",
status="completed",
steps=[],
usage={
"total_tokens": 10 + video_tokens,
"total_input_tokens": 10,
"input_tokens_by_modality": [{"modality": "text", "tokens": 10}],
"total_cached_tokens": 0,
"total_output_tokens": video_tokens,
"output_tokens_by_modality": [{"modality": "video", "tokens": video_tokens}],
"total_tool_use_tokens": 0,
"total_thought_tokens": 0,
},
)
cost = completion_cost(completion_response=response, custom_llm_provider="gemini")
expected = 10 * model_info["input_cost_per_token"] + video_tokens * model_info["output_cost_per_video_token"]
assert model_info["output_cost_per_video_token"] != model_info["output_cost_per_token"]
assert cost == pytest.approx(expected)