mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
fix(interactions): bill background interactions once completed via cost polling
This commit is contained in:
parent
887fc0a73c
commit
59d4e52a3d
10 changed files with 519 additions and 59 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
131
litellm/interactions/background_cost_polling.py
Normal file
131
litellm/interactions/background_cost_polling.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
184
tests/test_litellm/interactions/test_background_cost_polling.py
Normal file
184
tests/test_litellm/interactions/test_background_cost_polling.py
Normal 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
|
||||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue