mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
* feat(anthropic): workload identity federation and pluggable identity sources Backend half of #38818 (internal copy of the fork PR #38013), rebuilt as one commit on top of litellm_internal_staging without the dashboard changes. Deployments on anthropic/ without a static api_key can exchange an OIDC workload assertion for a short-lived sk-ant-oat01 token through a shared RFC 7523 JWT-bearer engine. The assertion comes from a mounted token file, an env token, a LiteLLM-signed issuer, or Keycloak, chosen per deployment, per named credential, or through ANTHROPIC_IDENTITY_SOURCE. The federation fields are server-owned: refused inline in request bodies and on POST /model/new, proxy-admin only on credentials, and the token exchange is pinned to api.anthropic.com unless LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS adds a host. GET /credentials/{name}/jwks exports the public key set of a LiteLLM-signed credential for the Claude Console. The OpenAI federation trio from #39613 rides along on the backend side with the same server-owned handling. Fixes #28607 Resolves LIT-6107 Co-authored-by: derhornspieler <15236687+derhornspieler@users.noreply.github.com> * fix(anthropic): let batch-result downloads mint from deployment params and accept host:port allowlist entries The files handler enabled workload identity on batch-result downloads but never received the deployment's litellm_params, so a deployment authenticating through a named credential could only mint from process-wide env vars. It now threads litellm_params through to the auth header the way the batch retrieve path already does. LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS entries written as host:port were read by urlsplit as a scheme, so the allowlist kept the raw entry while the exchange compared bare hostnames and refused the gateway. Entries are now parsed as network locations whether or not they carry a scheme. * fix(types): move the WIF kwargs key sets to a leaf module so the kwargs funnel imports without a cycle * test(anthropic): pin case-insensitive matching of WIF exchange-host allowlist entries * fix(anthropic): end workload identity federation errors without a period so the router suffix reads cleanly * fix(proxy): decrypt stored litellm_params before the WIF write gate * fix(proxy): hide WIF secret references from /health output * fix(proxy): keep the proxy error shape on credential endpoint refusals * fix(proxy): hide identity token file paths from /health output * fix(anthropic): rename the federation workspace param so Bedrock's anthropic_workspace_id keeps working The Bedrock Claude Platform route already reads anthropic_workspace_id from optional_params, so banning that spelling as a server-owned federation parameter broke a pre-existing client capability. The federation field is now anthropic_federation_workspace_id (env ANTHROPIC_FEDERATION_WORKSPACE_ID), which restores the base branch's behavior for Bedrock callers, drops the Bedrock-specific hint from the refusal message, and deletes the unconditional ban constant that no longer had a reader * fix(auth): share one exchanged token across workers reading the same assertion Anthropic accepts each identity assertion exactly once, so two uvicorn workers reading the same token file both minting from it means the second exchange is denied with jti_reused. Minted tokens now land in a per-user 0700 cache directory guarded by a file lock, so workers on the same host reuse one exchange until the token expires or the assertion rotates. A 401 is only retried when the re-read assertion actually differs, and the denial hint explains jti_reused. LITELLM_TOKEN_EXCHANGE_CACHE_DIR moves the cache and an empty value disables it * fix: keep anthropic federation from being shadowed or leaked An empty or whitespace-only ANTHROPIC_API_KEY counted as set, so a federated deployment sent an empty x-api-key on every call instead of minting a token. Blank values now read as unset, and a real static key on a federated deployment logs once that it outranks federation and nothing is being federated. The exchange-host allowlist matched hostnames only, so a second process on another port of an allowed host was trusted with the workload's identity token. An entry that names a port now trusts that port alone, while a bare host still trusts every port. The shared token store exists so the workers reading one projected token file do not each spend its single-use jti. A source that mints its own assertion per exchange shares nothing with another worker, so it no longer writes a live token to disk for a lookup that can never hit. * fix: unlink a staged token file a failed write leaves behind The 401 denial hint now also says federation ignores ANTHROPIC_WORKSPACE_ID, which the Bedrock Claude platform provider already reads. * refactor: move anthropic jwks derivation behind a provider-owned tagged union * fix: unlink the staged token file when its write fails at close A buffered write only reaches the disk when the handle closes, so a full disk surfaces at close and left the staging file behind holding a usable token. * fix(anthropic): close the staging descriptor before writing the shared token file * fix(wif): judge federation writes by what they set, not what is stored The admin gate read the stored deployment, so a team admin lost edit, delete and Test Connection on any deployment carrying federation params. It now returns early unless the submitted fields touch the federation surface, and a Test Connection probe that points the deployment at its own api_base is still refused, with the 403 no longer wrapped into a 500 The rest of the same review pass: POST /model/new refuses only a blocking value of `blocked`, so a client that always sends `blocked: false` is not turned away; a request body can no longer pick which federated identity to mint as by naming a stored credential; an advisory refresh the executor refuses disarms the entry instead of wedging the identity until the follower timeout; the static-key shadow warning resolves its env fallback inside the cache instead of once per request; credential writes drop nulls before storing them; the token exchange validates the endpoint URL before reading an assertion and keeps refusing redirects across a client heal; /health hides every server-owned federation field from non-admins; and the async create_file and create_batch paths say which setting is missing when the provider resolves no URL * fix(proxy): let a deployment write name a federated credential reject_federated_credential_reference runs from is_request_body_safe, which pre_db_read_auth_checks calls on every route, so it also fired on POST /model/new, /model/update, /model/{id}/update and /health/test_connection. A proxy admin could no longer attach a federated credential to a deployment over the API or the Admin UI, leaving a static config.yaml entry as the only way to configure the feature the rejection told the caller to go configure, and _reject_non_admin_wif_write never got to make the call it exists to make. is_request_body_safe now takes the route and skips only the credential-reference check on the routes that reach can_user_make_model_call. Federation fields typed inline into a body stay refused everywhere, and a call naming a federated credential still cannot pick the identity it mints as. * refactor(proxy): derive health display policy from the federation key sets The health check module hand-copied the five workload identity fields whose value is a credential, so a shared proxy surface named provider-specific parameters and a newly added secret-bearing field would have gone on being displayed until someone remembered both places WIF_SECRET_BEARING_KEYS now sits beside the key sets it splits out of, types/utils derives secret_bearing_wif_litellm_params from it, and the health layer splats that tuple the same way it already splats the admin-only one * fix(anthropic_wif): treat blank identity-source fields as unset * test(proxy): classify the federation params in the credential slot registry main's registry test (#43298) now fails the build for any credential-named deployment param without a classification. The five federation fields that carry a token, a token file path, or a signing or client secret reference are Unplanted, matching WIF_SECRET_BEARING_KEYS; the four remaining Keycloak settings name a URL, a client id, an auth method, or a scope and are NotSecret * fix(anthropic_wif): declare federation params as owned connection leaves and chart their metrics Register the 18 Anthropic and 3 OpenAI federation params as frozen ConnectionSettings leaves so the owned-kwarg registry, the kwargs funnel and the request-body ban list read one declaration. Pass the deployment api_base through to the count-tokens handler instead of a pre-suffixed URL, which doubled the /count_tokens path on main's prompt-cache predictor. Add the five litellm_anthropic_wif_* families to the all-metrics Grafana dashboard. * fix(credentials): gate PATCH on WIF fields resolved from model_id The credential PATCH handler checked server-owned workload identity federation fields only on the values the caller sent, while a body that named a deployment through model_id had its credential values resolved after that check. A non-admin could therefore copy a federated deployment's WIF fields onto an ordinary credential. Resolve the incoming values first and run the non-admin gate on them, matching the POST path * fix(anthropic): count tokens with ANTHROPIC_AUTH_TOKEN through the shared auth header Count-tokens walked its own credential ladder: a static key, else skip minting when ANTHROPIC_AUTH_TOKEN is set, else mint a federated token. With only the auth token set it forwarded nothing and the proxy silently fell back to its local tokenizer while chat on the same deployment authenticated with that token. The handler now takes the auth header that AnthropicModelInfo.aget_auth_header resolves, the same ladder chat, files, batches and skills use, and merges the oauth beta a minted or consumer token carries with the token-counting beta --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: derhornspieler <15236687+derhornspieler@users.noreply.github.com> Co-authored-by: mateo-berri <happymvw@gmail.com>
760 lines
31 KiB
Python
760 lines
31 KiB
Python
import json
|
|
from collections.abc import Iterable, Iterator, Mapping
|
|
from dataclasses import dataclass
|
|
from dataclasses import replace as dataclasses_replace
|
|
from enum import Enum
|
|
from typing import Any, Final, Literal
|
|
|
|
import litellm
|
|
from litellm._logging import verbose_logger
|
|
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
|
|
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
|
|
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
|
|
from litellm.llms.bedrock.batches.transformation import titan_embedding_usage_from_batch_output
|
|
from litellm.llms.vertex_ai.batches.transformation import (
|
|
is_native_vertex_batch_output_row,
|
|
native_vertex_batch_row_stats,
|
|
)
|
|
from litellm.types.llms.openai import Batch
|
|
from litellm.types.utils import ModelInfo, Usage
|
|
from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS
|
|
from litellm.utils import token_counter
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class BatchCostUsageResult:
|
|
"""Aggregate cost, usage, and per-line pass/fail counts for a completed batch."""
|
|
|
|
cost: float
|
|
usage: Usage
|
|
models: list[str]
|
|
successful_requests: int
|
|
failed_requests: int
|
|
prompt_cost: float = 0.0
|
|
completion_cost: float = 0.0
|
|
|
|
|
|
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
|
|
|
|
|
|
def _uses_native_vertex_output(
|
|
custom_llm_provider: str,
|
|
model_name: str | None,
|
|
first_row: Mapping[str, object] | None,
|
|
) -> bool:
|
|
if custom_llm_provider != "vertex_ai":
|
|
return False
|
|
if model_name and litellm.disable_vertex_batch_output_transformation:
|
|
return True
|
|
return first_row is not None and is_native_vertex_batch_output_row(first_row)
|
|
|
|
|
|
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
|
|
|
|
|
|
def batch_cost_is_final(batch: Batch) -> bool:
|
|
"""Whether this retrieve of the batch is the one to account its cost from.
|
|
|
|
A batch still in flight has nothing to price, and a "completed" batch can report
|
|
no output_file_id for a moment before the output populates; pricing either records
|
|
$0 under the batch's single spend row and pins it there. Final means a completed
|
|
batch whose output file has arrived or whose counts prove no line succeeded, or
|
|
any other terminal status (failed, cancelled, expired).
|
|
"""
|
|
if batch.status not in _TERMINAL_BATCH_STATUSES:
|
|
return False
|
|
if batch.status not in _COMPLETED_BATCH_STATUSES or batch.output_file_id is not None:
|
|
return True
|
|
request_counts: Final = batch.request_counts
|
|
return request_counts is not None and request_counts.total > 0 and request_counts.completed == 0
|
|
|
|
|
|
async def calculate_batch_cost_and_usage(
|
|
file_content_dictionary: list[dict],
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
|
|
model_name: str | None = None,
|
|
model_info: ModelInfo | None = None,
|
|
) -> BatchCostUsageResult:
|
|
"""
|
|
Calculate the cost and usage of a batch.
|
|
|
|
Args:
|
|
model_info: Optional deployment-level model info with custom batch
|
|
pricing. Threaded through to batch_cost_calculator so that
|
|
deployment-specific pricing (e.g. input_cost_per_token_batches)
|
|
is used instead of the global cost map.
|
|
"""
|
|
first_row: Final = file_content_dictionary[0] if file_content_dictionary else None
|
|
if _uses_native_vertex_output(custom_llm_provider, model_name, first_row):
|
|
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name, model_info=model_info)
|
|
|
|
return _aggregate_batch_cost_usage_models(
|
|
entries=file_content_dictionary,
|
|
custom_llm_provider=custom_llm_provider,
|
|
model_name=model_name,
|
|
model_info=model_info,
|
|
)
|
|
|
|
|
|
async def _handle_completed_batch(
|
|
batch: Batch,
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
|
|
model_name: str | None = None,
|
|
litellm_params: dict | None = None,
|
|
model_info: ModelInfo | None = None,
|
|
) -> BatchCostUsageResult:
|
|
"""Fetch a completed batch's output file and aggregate its cost, usage, and
|
|
models in a single pass over the JSONL lines, so the parsed file content is
|
|
never materialized in memory.
|
|
|
|
Args:
|
|
batch: The batch object
|
|
custom_llm_provider: The LLM provider
|
|
model_name: Optional model name
|
|
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
|
|
model_info: Optional deployment-level model info with custom pricing,
|
|
threaded through so a deployment's configured rates win over the
|
|
global cost map.
|
|
"""
|
|
# A completed batch whose request lines all failed has no output file - the
|
|
# results are written to a separate error_file_id and output_file_id is None.
|
|
# There is nothing to price or measure, so report an empty result set instead
|
|
# of calling _fetch_batch_output_file_content, which raises on a missing
|
|
# output file. Without this guard the logging worker crashes on every
|
|
# aretrieve_batch poll and the completed batch's zero-cost accounting is lost.
|
|
# The generic retrieval helper keeps raising for callers that explicitly ask
|
|
# for a missing output file.
|
|
if batch.output_file_id is None:
|
|
return BatchCostUsageResult(
|
|
cost=0.0,
|
|
usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
|
|
models=[],
|
|
successful_requests=0,
|
|
failed_requests=await count_error_file_failed_requests(
|
|
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
|
|
),
|
|
)
|
|
|
|
file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params)
|
|
error_file_failed_requests: Final = await count_error_file_failed_requests(
|
|
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
|
|
)
|
|
|
|
output_file_result: Final = (
|
|
calculate_vertex_ai_batch_cost_and_usage(
|
|
_iter_batch_output_entries(file_content), model_name, model_info=model_info
|
|
)
|
|
if _uses_native_vertex_output(
|
|
custom_llm_provider, model_name, next(_iter_batch_output_entries(file_content), None)
|
|
)
|
|
else _aggregate_batch_cost_usage_models(
|
|
entries=_iter_batch_output_entries(file_content),
|
|
custom_llm_provider=custom_llm_provider,
|
|
model_name=model_name,
|
|
model_info=model_info,
|
|
)
|
|
)
|
|
|
|
if not error_file_failed_requests:
|
|
return output_file_result
|
|
return dataclasses_replace(
|
|
output_file_result, failed_requests=output_file_result.failed_requests + error_file_failed_requests
|
|
)
|
|
|
|
|
|
class _LineOutcome(Enum):
|
|
"""A batch output line that yielded no billable stats."""
|
|
|
|
PROVIDER_FAILED = "provider_failed"
|
|
UNCOSTABLE = "uncostable"
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class _BatchOutputLineStats:
|
|
prompt_cost: float
|
|
completion_cost: float
|
|
prompt_tokens: int
|
|
completion_tokens: int
|
|
total_tokens: int
|
|
cache_read_tokens: int
|
|
cache_creation_tokens: int
|
|
reasoning_tokens: int
|
|
model: str | None
|
|
|
|
|
|
def _classify_output_line_stats(
|
|
entries: Iterable[dict],
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
|
|
model_name: str | None,
|
|
model_info: ModelInfo | None,
|
|
) -> Iterator[_BatchOutputLineStats | _LineOutcome]:
|
|
"""Classify every output line in a single pass, so counting failures never needs
|
|
a second read of a potentially huge output file. A line the provider reported as
|
|
failed yields ``PROVIDER_FAILED``; a successful line litellm could not price
|
|
yields ``UNCOSTABLE`` and still counts as a successful request billed at $0, so
|
|
the counts stay reconcilable with the provider's own ``request_counts``."""
|
|
for entry in entries:
|
|
if not _batch_response_was_successful(entry, custom_llm_provider):
|
|
yield _LineOutcome.PROVIDER_FAILED
|
|
continue
|
|
stats = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info)
|
|
yield stats if stats is not None else _LineOutcome.UNCOSTABLE
|
|
|
|
|
|
def _safe_output_line_stats(
|
|
entry: Mapping[str, object],
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
|
|
model_name: str | None,
|
|
model_info: ModelInfo | None,
|
|
) -> _BatchOutputLineStats | None:
|
|
"""Return the stats for one provider-successful batch output line, or None when
|
|
it cannot be costed, so a single bad line never aborts the whole batch's cost
|
|
accounting."""
|
|
custom_id: Final = entry.get("custom_id") if isinstance(entry, dict) else None
|
|
try:
|
|
return _compute_output_line_stats(entry, custom_llm_provider, model_name, model_info)
|
|
except Exception as e: # noqa: BLE001 # any single line's costing failure must not abort the whole batch
|
|
verbose_logger.warning(
|
|
"batch output line could not be costed, so it is billed at $0 and the rest of the batch "
|
|
"is still billed. custom_id=%s error=%s",
|
|
custom_id,
|
|
str(e),
|
|
)
|
|
return None
|
|
|
|
|
|
def _compute_output_line_stats(
|
|
entry: Mapping[str, object],
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
|
|
model_name: str | None,
|
|
model_info: ModelInfo | None,
|
|
) -> _BatchOutputLineStats:
|
|
response_body: Final = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
|
|
usage: Final = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
|
|
prompt_details: Final = parse_prompt_tokens_details(usage)
|
|
raw_model: Final = response_body.get("model")
|
|
response_model: Final = raw_model if isinstance(raw_model, str) and raw_model else None
|
|
completion_details: Final = usage.completion_tokens_details
|
|
line_prompt_cost, line_completion_cost = _output_line_cost(
|
|
response_body=response_body,
|
|
usage=usage,
|
|
custom_llm_provider=custom_llm_provider,
|
|
model_name=model_name,
|
|
response_model=response_model,
|
|
model_info=model_info,
|
|
)
|
|
return _BatchOutputLineStats(
|
|
prompt_cost=line_prompt_cost,
|
|
completion_cost=line_completion_cost,
|
|
prompt_tokens=usage.prompt_tokens,
|
|
completion_tokens=usage.completion_tokens,
|
|
total_tokens=usage.total_tokens,
|
|
cache_read_tokens=prompt_details["cache_hit_tokens"],
|
|
cache_creation_tokens=prompt_details["cache_creation_tokens"],
|
|
reasoning_tokens=(completion_details.reasoning_tokens if completion_details else None) or 0,
|
|
model=response_model,
|
|
)
|
|
|
|
|
|
def _ocr_usage_info_from_response_body(response_body: Mapping[str, object]) -> OCRUsageInfo | None:
|
|
"""OCR results report ``usage_info`` (pages) instead of ``usage`` (tokens); None for non-OCR lines."""
|
|
raw_usage_info: Final = response_body.get("usage_info")
|
|
if not isinstance(raw_usage_info, Mapping):
|
|
return None
|
|
return OCRUsageInfo.model_validate(raw_usage_info)
|
|
|
|
|
|
def _output_line_cost(
|
|
response_body: Mapping[str, object],
|
|
usage: Usage,
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
|
|
model_name: str | None,
|
|
response_model: str | None,
|
|
model_info: ModelInfo | None,
|
|
) -> tuple[float, float]:
|
|
"""(prompt_cost, completion_cost) for one output line, priced at batch rates."""
|
|
from litellm.cost_calculator import batch_cost_calculator, ocr_batch_cost
|
|
|
|
cost_model: Final = (
|
|
model_name if custom_llm_provider == "bedrock" and model_name else response_model or model_name or ""
|
|
)
|
|
ocr_usage: Final = _ocr_usage_info_from_response_body(response_body)
|
|
if ocr_usage is not None:
|
|
return ocr_batch_cost(
|
|
model=cost_model,
|
|
custom_llm_provider=custom_llm_provider,
|
|
usage_info=ocr_usage,
|
|
model_info=model_info,
|
|
)
|
|
return batch_cost_calculator(
|
|
usage=usage,
|
|
model=cost_model,
|
|
custom_llm_provider=custom_llm_provider,
|
|
model_info=model_info,
|
|
)
|
|
|
|
|
|
def _aggregate_batch_cost_usage_models(
|
|
entries: Iterable[dict],
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
|
|
model_name: str | None = None,
|
|
model_info: ModelInfo | None = None,
|
|
) -> BatchCostUsageResult:
|
|
"""Aggregate cost, usage, models, and pass/fail counts from batch output
|
|
entries in a single pass, holding one small stats record per line instead
|
|
of the parsed file."""
|
|
all_results: Final = tuple(_classify_output_line_stats(entries, custom_llm_provider, model_name, model_info))
|
|
line_stats: Final = tuple(result for result in all_results if isinstance(result, _BatchOutputLineStats))
|
|
failed_requests: Final = sum(1 for result in all_results if result is _LineOutcome.PROVIDER_FAILED)
|
|
successful_requests: Final = len(all_results) - failed_requests
|
|
|
|
cache_token_params: Final = {
|
|
key: tokens
|
|
for key, tokens in (
|
|
("cache_read_input_tokens", sum(stats.cache_read_tokens for stats in line_stats)),
|
|
("cache_creation_input_tokens", sum(stats.cache_creation_tokens for stats in line_stats)),
|
|
)
|
|
if tokens > 0
|
|
}
|
|
batch_usage: Final = Usage(
|
|
total_tokens=sum(stats.total_tokens for stats in line_stats),
|
|
prompt_tokens=sum(stats.prompt_tokens for stats in line_stats),
|
|
completion_tokens=sum(stats.completion_tokens for stats in line_stats),
|
|
reasoning_tokens=sum(stats.reasoning_tokens for stats in line_stats),
|
|
**cache_token_params,
|
|
)
|
|
batch_models: Final = [model_name] if model_name else [stats.model for stats in line_stats if stats.model]
|
|
total_prompt_cost: Final = sum((stats.prompt_cost for stats in line_stats), 0.0)
|
|
total_completion_cost: Final = sum((stats.completion_cost for stats in line_stats), 0.0)
|
|
total_cost: Final = total_prompt_cost + total_completion_cost
|
|
verbose_logger.debug(
|
|
"batch output aggregate: cost=%s usage=%s models=%s successful=%d failed=%d",
|
|
total_cost,
|
|
batch_usage,
|
|
batch_models,
|
|
successful_requests,
|
|
failed_requests,
|
|
)
|
|
return BatchCostUsageResult(
|
|
cost=total_cost,
|
|
usage=batch_usage,
|
|
models=batch_models,
|
|
successful_requests=successful_requests,
|
|
failed_requests=failed_requests,
|
|
prompt_cost=total_prompt_cost,
|
|
completion_cost=total_completion_cost,
|
|
)
|
|
|
|
|
|
def calculate_vertex_ai_batch_cost_and_usage(
|
|
vertex_ai_batch_responses: Iterable[dict],
|
|
model_name: str | None = None,
|
|
model_info: ModelInfo | None = None,
|
|
) -> BatchCostUsageResult:
|
|
"""
|
|
Cost and usage of a native Vertex predictions.jsonl, one
|
|
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
|
|
generateContent row or `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
|
|
embedding row per line. `model_name` (the deployment model) prices every row, else each row's own
|
|
`modelVersion` does; a row without a usable response counts as failed.
|
|
"""
|
|
from litellm.cost_calculator import batch_cost_calculator
|
|
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
|
|
|
row_stats: Final = tuple(
|
|
native_vertex_batch_row_stats(
|
|
row,
|
|
model_name,
|
|
model_info=model_info,
|
|
calculate_usage=VertexGeminiConfig._calculate_usage,
|
|
cost_calculator=batch_cost_calculator,
|
|
)
|
|
for row in vertex_ai_batch_responses
|
|
)
|
|
priced: Final = tuple(stats for stats in row_stats if stats is not None)
|
|
total_prompt_cost: Final = sum(stats.prompt_cost for stats in priced)
|
|
total_completion_cost: Final = sum(stats.completion_cost for stats in priced)
|
|
prompt_tokens: Final = sum(stats.usage.prompt_tokens for stats in priced)
|
|
completion_tokens: Final = sum(stats.usage.completion_tokens for stats in priced)
|
|
total_tokens: Final = sum(stats.total_tokens for stats in priced)
|
|
total_cost: Final = total_prompt_cost + total_completion_cost
|
|
verbose_logger.info(
|
|
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
|
|
total_cost,
|
|
prompt_tokens,
|
|
completion_tokens,
|
|
total_tokens,
|
|
len(priced),
|
|
len(row_stats) - len(priced),
|
|
)
|
|
|
|
return BatchCostUsageResult(
|
|
cost=total_cost,
|
|
usage=Usage(
|
|
total_tokens=total_tokens,
|
|
prompt_tokens=prompt_tokens,
|
|
completion_tokens=completion_tokens,
|
|
),
|
|
models=(
|
|
[model_name]
|
|
if model_name
|
|
else list(dict.fromkeys(stats.model for stats in priced if stats.model is not None))
|
|
),
|
|
successful_requests=len(priced),
|
|
failed_requests=len(row_stats) - len(priced),
|
|
prompt_cost=total_prompt_cost,
|
|
completion_cost=total_completion_cost,
|
|
)
|
|
|
|
|
|
def _provider_output_file_id(output_file_id: str) -> str:
|
|
"""
|
|
Resolve the file id the provider actually knows: unified ids yield their embedded
|
|
llm_output_file_id, model-encoded ids decode to the raw provider id, raw ids pass through.
|
|
"""
|
|
from litellm.proxy.openai_files_endpoints.common_utils import (
|
|
_is_base64_encoded_unified_file_id,
|
|
get_original_file_id,
|
|
)
|
|
|
|
unified_file_id: Final = _is_base64_encoded_unified_file_id(output_file_id)
|
|
if not unified_file_id:
|
|
return get_original_file_id(output_file_id)
|
|
try:
|
|
extracted: Final = unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
|
|
except (IndexError, AttributeError) as e:
|
|
verbose_logger.error(
|
|
"Failed to extract LLM output file ID from unified file ID: %s, error: %s",
|
|
output_file_id,
|
|
e,
|
|
)
|
|
return output_file_id
|
|
verbose_logger.debug("Extracted LLM output file ID from unified file ID: %s", extracted)
|
|
return extracted
|
|
|
|
|
|
async def _fetch_batch_managed_file_content(
|
|
file_id: str,
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
|
|
litellm_params: dict | None = None,
|
|
) -> bytes:
|
|
"""
|
|
Fetch a batch's output or error file and return its raw JSONL bytes.
|
|
|
|
Args:
|
|
file_id: The provider or unified (litellm-managed) file id to fetch
|
|
custom_llm_provider: The LLM provider
|
|
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
|
|
Required for Azure and other providers that need authentication
|
|
"""
|
|
from litellm.files.main import afile_content
|
|
|
|
# Build kwargs for afile_content with credentials from litellm_params
|
|
file_content_kwargs: Final = {
|
|
"file_id": _provider_output_file_id(file_id),
|
|
"custom_llm_provider": custom_llm_provider,
|
|
}
|
|
|
|
# Extract and add credentials for file access
|
|
credentials: Final = _extract_file_access_credentials(litellm_params)
|
|
file_content_kwargs.update(credentials)
|
|
|
|
_file_content: Final = await afile_content(**file_content_kwargs)
|
|
return _file_content.content
|
|
|
|
|
|
async def _fetch_batch_output_file_content(
|
|
batch: Batch,
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
|
|
litellm_params: dict | None = None,
|
|
) -> bytes:
|
|
"""
|
|
Fetch the batch output file and return its raw JSONL bytes
|
|
|
|
Args:
|
|
batch: The batch object
|
|
custom_llm_provider: The LLM provider
|
|
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
|
|
Required for Azure and other providers that need authentication
|
|
"""
|
|
if batch.output_file_id is None:
|
|
raise ValueError("Output file id is None cannot retrieve file content")
|
|
|
|
return await _fetch_batch_managed_file_content(
|
|
batch.output_file_id, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
|
|
)
|
|
|
|
|
|
async def count_error_file_failed_requests(
|
|
batch: Batch,
|
|
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
|
|
litellm_params: dict | None,
|
|
) -> int:
|
|
"""Count failed requests reported only in the batch's separate error file.
|
|
|
|
OpenAI-shaped batch providers write successful lines to ``output_file_id``
|
|
and per-request failures (e.g. a rejected param) to a distinct
|
|
``error_file_id`` - they never appear in the output file at all, so
|
|
counting failures from the output file alone silently undercounts them.
|
|
"""
|
|
if batch.error_file_id is None:
|
|
return 0
|
|
try:
|
|
error_file_content = await _fetch_batch_managed_file_content(
|
|
batch.error_file_id, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
|
|
)
|
|
except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch
|
|
verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e)
|
|
return 0
|
|
return sum(1 for _ in _iter_batch_input_lines(error_file_content))
|
|
|
|
|
|
def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
|
|
"""
|
|
Extract credentials from litellm_params for file access operations.
|
|
|
|
This method extracts relevant authentication and configuration parameters
|
|
needed for accessing files across different providers (Azure, Vertex AI, etc.).
|
|
|
|
Args:
|
|
litellm_params: Dictionary containing litellm parameters with credentials
|
|
|
|
Returns:
|
|
Dictionary containing only the credentials needed for file access
|
|
"""
|
|
credentials: Final = {}
|
|
|
|
if litellm_params:
|
|
# List of credential keys that should be passed to file operations
|
|
credential_keys: Final = (
|
|
"api_key",
|
|
"api_base",
|
|
"api_version",
|
|
"organization",
|
|
"azure_ad_token",
|
|
"azure_ad_token_provider",
|
|
"vertex_project",
|
|
"vertex_location",
|
|
"vertex_credentials",
|
|
"gcs_bucket_name",
|
|
"bucket_name",
|
|
"s3_endpoint_url",
|
|
"s3_region_name",
|
|
"timeout",
|
|
"max_retries",
|
|
"_litellm_internal_model_credentials",
|
|
*AWS_CREDENTIAL_KWARGS_KEYS,
|
|
# A federated deployment holds no api_key, so without these the fetch that reads a
|
|
# finished batch's output has nothing to authenticate with and its cost is never billed.
|
|
*sorted(ANTHROPIC_WIF_KWARGS_KEYS),
|
|
)
|
|
for key in credential_keys:
|
|
if key in litellm_params:
|
|
credentials[key] = litellm_params[key]
|
|
|
|
return credentials
|
|
|
|
|
|
def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]:
|
|
"""
|
|
Get the file content as a list of dictionaries from JSON Lines format,
|
|
skipping malformed lines
|
|
"""
|
|
return list(_iter_batch_output_entries(file_content))
|
|
|
|
|
|
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
|
|
"""
|
|
Yield non-empty JSONL lines (unparsed) one at a time, so a caller can parse
|
|
each row in its own try/except and a single malformed line cannot abort the
|
|
whole pass. Peak memory stays bounded for large batch files.
|
|
"""
|
|
start, length, newline = 0, len(file_content), ord("\n")
|
|
while start < length:
|
|
idx = file_content.find(newline, start)
|
|
if idx == -1:
|
|
chunk, start = file_content[start:], length
|
|
else:
|
|
chunk, start = file_content[start:idx], idx + 1
|
|
line = chunk.strip()
|
|
if line:
|
|
yield line
|
|
|
|
|
|
def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict]:
|
|
"""
|
|
Yield parsed batch output JSONL entries one at a time without materializing
|
|
the whole file as a list, so peak memory stays bounded. A malformed or
|
|
non-object line is skipped with a warning so one bad line never aborts the
|
|
whole batch's cost accounting.
|
|
"""
|
|
for line in _iter_batch_input_lines(file_content):
|
|
entry = _parse_batch_output_line(line)
|
|
if entry is not None:
|
|
yield entry
|
|
|
|
|
|
def _parse_batch_output_line(line: bytes) -> dict | None:
|
|
try:
|
|
parsed: Final[object] = json.loads(line)
|
|
except ValueError as e:
|
|
verbose_logger.warning("skipping malformed batch output line: %s", str(e))
|
|
return None
|
|
if isinstance(parsed, dict):
|
|
return parsed
|
|
verbose_logger.warning("skipping non-object batch output line of type %s", type(parsed).__name__)
|
|
return None
|
|
|
|
|
|
# A batch request's input tokens scale roughly with its serialized size, so this
|
|
# is a conservative per-row fallback when the token counter cannot measure a row.
|
|
_BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN: Final = 4
|
|
|
|
|
|
def _estimate_batch_entry_tokens(raw_line: bytes) -> int:
|
|
"""Conservative token estimate for a batch row the token counter cannot measure
|
|
(or that cannot be parsed). Keeps the batch token total non-zero so a crafted
|
|
row cannot evade the TPM limit, without hard-rejecting a legitimate batch."""
|
|
return max(1, len(raw_line) // _BATCH_TOKEN_ESTIMATE_BYTES_PER_TOKEN)
|
|
|
|
|
|
def _count_entry_tokens(
|
|
entry: dict,
|
|
model_name: str | None = None,
|
|
) -> int:
|
|
"""Token-count a single batch input entry's body (chat / text / embedding)."""
|
|
body: Final = entry.get("body", {}) or {}
|
|
model: Final = body.get("model", model_name or "")
|
|
|
|
messages: Final = body.get("messages")
|
|
if messages:
|
|
return token_counter(model=model, messages=messages)
|
|
|
|
prompt: Final = body.get("prompt")
|
|
if prompt:
|
|
return _count_prompt_or_input_tokens(model=model, value=prompt)
|
|
|
|
input_data: Final = body.get("input")
|
|
if input_data:
|
|
return _count_prompt_or_input_tokens(model=model, value=input_data)
|
|
|
|
return 0
|
|
|
|
|
|
def _count_prompt_or_input_tokens(model: str, value: object) -> int:
|
|
"""Token-count a ``prompt`` / ``input`` field that the OpenAI batch
|
|
schema allows in four shapes:
|
|
|
|
- ``str``: a single text prompt.
|
|
- ``list[str]``: multiple text prompts.
|
|
- ``list[int]``: a pre-tokenized prompt (each int counts as 1 token).
|
|
- ``list[list[int]]``: multiple pre-tokenized prompts.
|
|
|
|
Pre-fix only the string shapes were counted, so a caller could send
|
|
a large ``list[list[int]]`` payload and slip past TPM rate limits
|
|
with a recorded cost of zero tokens.
|
|
"""
|
|
if isinstance(value, str):
|
|
return token_counter(model=model, text=value)
|
|
if isinstance(value, list):
|
|
total = 0
|
|
for chunk in value:
|
|
if isinstance(chunk, str):
|
|
total += token_counter(model=model, text=chunk)
|
|
elif isinstance(chunk, int):
|
|
# Single pre-tokenized prompt at the top level: each
|
|
# int counts as one token.
|
|
total += 1
|
|
elif isinstance(chunk, list):
|
|
# Nested pre-tokenized prompt: every int contributes a
|
|
# token. Mixed string/int items still count.
|
|
total += sum(1 if isinstance(t, int) else 0 for t in chunk)
|
|
total += sum(token_counter(model=model, text=t) for t in chunk if isinstance(t, str))
|
|
return total
|
|
return 0
|
|
|
|
|
|
def _get_batch_job_usage_from_response_body(
|
|
response_body: Mapping[str, Any], custom_llm_provider: str = "openai"
|
|
) -> Usage:
|
|
"""
|
|
Get the tokens of a batch job from the response body
|
|
"""
|
|
if custom_llm_provider in ("anthropic", "bedrock"):
|
|
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
|
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
|
|
|
titan_usage: Final = (
|
|
titan_embedding_usage_from_batch_output(response_body) if custom_llm_provider == "bedrock" else None
|
|
)
|
|
if titan_usage is not None:
|
|
return titan_usage
|
|
usage_object: Final = response_body.get("usage", None) or {}
|
|
if custom_llm_provider == "bedrock" and AmazonConverseConfig.is_converse_usage_shape(usage_object):
|
|
return AmazonConverseConfig().usage_from_batch_output(usage_object)
|
|
anthropic_usage: Final = AnthropicConfig().calculate_usage(
|
|
usage_object=usage_object,
|
|
reasoning_content=None,
|
|
)
|
|
if usage_object and anthropic_usage.total_tokens == 0:
|
|
verbose_logger.warning(
|
|
"batch output line reported usage this parser does not understand, so it will be billed at $0. "
|
|
"provider=%s usage_keys=%s",
|
|
custom_llm_provider,
|
|
sorted(usage_object.keys()),
|
|
)
|
|
return anthropic_usage
|
|
from litellm.responses.utils import ResponseAPILoggingUtils
|
|
|
|
_usage_dict: Final = response_body.get("usage", None) or {}
|
|
if ResponseAPILoggingUtils._is_response_api_usage(_usage_dict):
|
|
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_usage_dict)
|
|
usage: Final[Usage] = Usage(**_usage_dict)
|
|
if custom_llm_provider == "xai":
|
|
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
|
|
|
XAIChatConfig.fold_reasoning_tokens_into_completion(usage)
|
|
return usage
|
|
|
|
|
|
def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> Mapping[str, Any]:
|
|
"""
|
|
Get the ``result`` object from a line of an Anthropic message batch results JSONL file.
|
|
|
|
Anthropic batch results lines look like:
|
|
``{"custom_id": ..., "result": {"type": "succeeded", "message": {..., "usage": {...}}}}``
|
|
"""
|
|
return batch_results_line.get("result", None) or {}
|
|
|
|
|
|
def _get_response_from_batch_job_output_file(
|
|
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
|
|
) -> Mapping[str, object]:
|
|
"""
|
|
Get the response from the batch job output file
|
|
"""
|
|
if custom_llm_provider == "anthropic":
|
|
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("message", None) or {}
|
|
if custom_llm_provider == "bedrock":
|
|
return batch_job_output_file.get("modelOutput", None) or {}
|
|
_response: Final[dict] = batch_job_output_file.get("response", None) or {}
|
|
_response_body: Final = _response.get("body", None) or {}
|
|
return _response_body
|
|
|
|
|
|
def _batch_response_was_successful(
|
|
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
|
|
) -> bool:
|
|
"""
|
|
Check if the batch job response was successful
|
|
|
|
OpenAI-shaped output rows report ``response.status_code == 200``; Anthropic
|
|
message batch results lines report ``result.type == "succeeded"``; Bedrock
|
|
batch output lines report ``modelOutput`` (and no ``error``).
|
|
"""
|
|
if custom_llm_provider == "anthropic":
|
|
return _get_anthropic_result_from_batch_results_line(batch_job_output_file).get("type") == "succeeded"
|
|
if custom_llm_provider == "bedrock":
|
|
return batch_job_output_file.get("modelOutput") is not None and batch_job_output_file.get("error") is None
|
|
_response: Final[dict] = batch_job_output_file.get("response", None) or {}
|
|
return _response.get("status_code", None) == 200
|