litellm/litellm/batches/batch_utils.py
devin-ai-integration[bot] 0ed1c08f02
feat(anthropic): workload identity federation and pluggable identity sources (#44448)
* 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>
2026-10-03 17:08:30 -07:00

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