Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_watsonx_orchestrate_agents

This commit is contained in:
mateo-berri 2026-06-02 05:37:06 +00:00
commit a931b54e59
No known key found for this signature in database
118 changed files with 13482 additions and 632 deletions

View file

@ -63,3 +63,28 @@ jobs:
sha: commitHash,
});
core.info(`Created branch ${branchName} at ${commitHash}`);
- name: Create stable line branch
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
with:
script: |
const tag = process.env.TAG;
const commitHash = process.env.COMMIT_HASH;
const match = tag.match(/^v?(\d+)\.(\d+)\.0$/);
if (!match) {
core.info(`Tag ${tag} is not the X.Y.0 stable opener; skipping stable line branch`);
return;
}
const lineBranch = `stable/${match[1]}.${match[2]}.x`;
await github.rest.git.createRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `refs/heads/${lineBranch}`,
sha: commitHash,
});
core.info(`Created branch ${lineBranch} at ${commitHash}`);

View file

@ -15,6 +15,8 @@ When adding new features, add meaningful tests. Don't add tests that don't check
Same thing for bug fixes. The tests should make it so that this specific bug can never happen again without failing tests (i.e., regression)
`tests/test_litellm/` mirrors `litellm/` (see `tests/test_litellm/readme.md`). The default name is `test_<filename>.py` in the parallel path (`transformation.py` → `test_transformation.py`). Many provider dirs use a longer descriptive name instead (e.g. `test_anthropic_chat_transformation.py`) when `test_transformation.py` would be ambiguous across sibling folders or that name is already what the repo uses there; always match the existing test file in the directory you touch rather than introducing another style. Each `*_transformation.py` under `litellm/llms/{provider}/...` ideally has a matching test file in the parallel path. For bug fixes, do not create a new test file; add or extend a regression test in that existing mapped test file. Only create a new test file when adding a new feature (new provider, endpoint, or transformation module) that does not already have a mapped test file; then follow the naming pattern already used in that directory, or `test_<filename>.py` if you are the first test there. One focused regression test is better than many shallow ones.
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
Always use @.github/pull_request_template.md as a guide for your PR body

View file

@ -106,6 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
# Health & ops
"/health",
"/metrics",
"/watsonx"
)
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(

View file

@ -771,6 +771,7 @@ openai_compatible_endpoints: List = [
"https://api.moonshot.ai/v1",
"https://api.publicai.co/v1",
"https://api.synthetic.new/openai/v1",
"https://serverless.tensormesh.ai/v1",
"https://api.stima.tech/v1",
"https://nano-gpt.com/api/v1",
"https://api.poe.com/v1",
@ -820,6 +821,7 @@ openai_compatible_providers: List = [
"meta_llama",
"publicai", # PublicAI - JSON-configured provider
"synthetic", # Synthetic - JSON-configured provider
"tensormesh", # Tensormesh - JSON-configured provider
"apertis", # Apertis - JSON-configured provider
"nano-gpt", # Nano-GPT - JSON-configured provider
"poe", # Poe - JSON-configured provider
@ -855,6 +857,7 @@ openai_text_completion_compatible_providers: List = (
"moonshot",
"publicai",
"synthetic",
"tensormesh",
"apertis",
"nano-gpt",
"poe",
@ -868,6 +871,7 @@ openai_text_completion_compatible_providers: List = (
_openai_like_providers: List = [
"predibase",
"databricks",
"lemonade",
"watsonx",
] # private helper. similar to openai but require some custom auth / endpoint handling, so can't use the openai sdk
# well supported replicate llms

View file

@ -41,6 +41,7 @@ from litellm.integrations.datadog.datadog_handler import (
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.llms.custom_httpx.http_handler import (
MaskedHTTPStatusError,
_get_httpx_client,
get_async_httpx_client,
httpxSpecialProvider,
@ -68,6 +69,22 @@ DD_LOGGED_SUCCESS_SERVICE_TYPES = [
]
def _resolve_dd_batch_size() -> int:
raw = os.getenv("DD_BATCH_SIZE")
if raw is None:
return DD_MAX_BATCH_SIZE
try:
value = int(raw)
except ValueError:
verbose_logger.warning(
"Datadog: ignoring invalid DD_BATCH_SIZE=%r, using %s",
raw,
DD_MAX_BATCH_SIZE,
)
return DD_MAX_BATCH_SIZE
return max(1, min(value, DD_MAX_BATCH_SIZE))
class DataDogLogger(
CustomBatchLogger,
AdditionalLoggingUtils,
@ -128,7 +145,9 @@ class DataDogLogger(
asyncio.create_task(self.periodic_flush())
self.flush_lock = asyncio.Lock()
super().__init__(
**kwargs, flush_lock=self.flush_lock, batch_size=DD_MAX_BATCH_SIZE
**kwargs,
flush_lock=self.flush_lock,
batch_size=_resolve_dd_batch_size(),
)
except Exception as e:
verbose_logger.exception(
@ -339,28 +358,14 @@ class DataDogLogger(
"[DATADOG MOCK] Mock mode enabled - API calls will be intercepted"
)
response = await self.async_send_compressed_data(batch_to_send)
if response.status_code == 413:
verbose_logger.exception(DD_ERRORS.DATADOG_413_ERROR.value)
self.log_queue = batch_to_send + self.log_queue
return
response.raise_for_status()
if response.status_code != 202:
raise Exception(
f"Response from datadog API status_code: {response.status_code}, text: {response.text}"
)
undelivered = await self._send_with_413_split(batch_to_send)
if undelivered:
self.log_queue = undelivered + self.log_queue
if self.is_mock_mode:
verbose_logger.debug(
f"[DATADOG MOCK] Batch of {len(batch_to_send)} events successfully mocked"
)
else:
verbose_logger.debug(
"Datadog: Response from datadog API status_code: %s, text: %s",
response.status_code,
response.text,
)
except Exception as e:
self.log_queue = batch_to_send + self.log_queue
@ -368,6 +373,62 @@ class DataDogLogger(
f"Datadog Error sending batch API - {str(e)}\n{traceback.format_exc()}"
)
async def _send_with_413_split(self, batch: List) -> List:
"""
Send a batch, halving any sub-batch that 413s (payload too large) and retrying the
halves, since Datadog enforces a 5MB uncompressed limit per request.
A 413 surfaces as a raised MaskedHTTPStatusError (httpx raise_for_status), not a
returned response, so both paths are handled. A lone event that still 413s is
dropped to avoid wedging the queue on an undeliverable payload. Returns the events
that could not be delivered because of a non-413 (transient) error, so the caller
re-queues only those and never the events already accepted by Datadog.
"""
pending: List[List] = [batch]
while pending:
chunk = pending.pop()
if not chunk:
continue
try:
response = await self.async_send_compressed_data(chunk)
except Exception as e:
if isinstance(e, MaskedHTTPStatusError) and e.status_code == 413:
response = e.response
else:
verbose_logger.exception(
f"Datadog Error sending batch API - {str(e)}"
)
return self._undelivered(chunk, pending)
if response.status_code == 413:
if len(chunk) == 1:
verbose_logger.error(DD_ERRORS.DATADOG_413_ERROR.value)
continue
mid = len(chunk) // 2
pending.append(chunk[mid:])
pending.append(chunk[:mid])
continue
if response.status_code != 202:
verbose_logger.error(
"Datadog: unexpected response status_code=%s, text=%s",
response.status_code,
response.text,
)
return self._undelivered(chunk, pending)
verbose_logger.debug(
"Datadog: delivered %s events, status_code=%s, text=%s",
len(chunk),
response.status_code,
response.text,
)
return []
@staticmethod
def _undelivered(chunk: List, pending: List[List]) -> List:
return chunk + [event for remaining in reversed(pending) for event in remaining]
async def flush_queue(self):
if self.flush_lock is None:
return

View file

@ -95,7 +95,9 @@ class FocusTransformer:
pl.lit("Usage-Based").alias("ChargeFrequency"),
fmt(pl.col("ChargePeriodEnd")).alias("ChargePeriodEnd"),
fmt(pl.col("ChargePeriodStart")).alias("ChargePeriodStart"),
dec(pl.lit(1.0)).alias("ConsumedQuantity"),
dec(
pl.col("api_requests").cast(pl.Int64).cast(pl.Float64).fill_null(0.0)
).alias("ConsumedQuantity"),
pl.lit("Requests").alias("ConsumedUnit"),
dec(pl.col("spend").fill_null(0.0)).alias("ContractedCost"),
none_str.alias("ContractedUnitPrice"),
@ -107,7 +109,9 @@ class FocusTransformer:
none_str.alias("AvailabilityZone"),
pl.lit("USD").alias("PricingCurrency"),
none_str.alias("PricingCategory"),
dec(pl.lit(1.0)).alias("PricingQuantity"),
dec(
pl.col("api_requests").cast(pl.Int64).cast(pl.Float64).fill_null(0.0)
).alias("PricingQuantity"),
none_dec.alias("PricingCurrencyContractedUnitPrice"),
dec(pl.col("spend").fill_null(0.0)).alias("PricingCurrencyEffectiveCost"),
none_dec.alias("PricingCurrencyListUnitPrice"),

View file

@ -83,6 +83,19 @@ _VALID_CAPTURE_MODES = {
}
def _normalize_team_metadata_keys(value: Any) -> List[str]:
"""Coerce a team-metadata allowlist from a list or comma-separated string.
config.yaml passes a YAML list; an env var passes a comma-separated string.
Both collapse to a list of stripped, non-empty keys.
"""
if value is None:
return []
if isinstance(value, str):
return [item.strip() for item in value.split(",") if item.strip()]
return [str(item).strip() for item in value if str(item).strip()]
@dataclass
class OpenTelemetryConfig:
exporter: Union[str, SpanExporter] = "console"
@ -100,6 +113,10 @@ class OpenTelemetryConfig:
# One of NO_CONTENT, SPAN_ONLY, EVENT_ONLY, SPAN_AND_EVENT (or "true" as legacy alias).
capture_message_content: Optional[str] = None
semconv_stability_opt_in: Set[OTELSemconvCategory] = field(default_factory=set)
# Sub-keys of the team's free-form metadata stamped onto the inference span
# under ``litellm.team.metadata``. Empty by default so none of a team's
# metadata leaves the process until explicitly allowlisted.
baggage_team_metadata_keys: List[str] = field(default_factory=list)
def __post_init__(self) -> None:
# If endpoint is specified but exporter is still the default "console",
@ -130,6 +147,11 @@ class OpenTelemetryConfig:
self.semconv_stability_opt_in |= parse_semconv_opt_in(
os.getenv(OTEL_SEMCONV_STABILITY_OPT_IN_ENV)
)
self.baggage_team_metadata_keys = _normalize_team_metadata_keys(
self.baggage_team_metadata_keys
) or _normalize_team_metadata_keys(
os.getenv("LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS")
)
@classmethod
def from_env(cls):
@ -188,8 +210,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
meter_provider: Optional[Any] = None,
**kwargs,
):
team_metadata_keys_override = kwargs.pop("baggage_team_metadata_keys", None)
if config is None:
config = OpenTelemetryConfig.from_env()
if team_metadata_keys_override is not None:
config.baggage_team_metadata_keys = _normalize_team_metadata_keys(
team_metadata_keys_override
)
self.config = config
self.callback_name = callback_name
@ -1245,7 +1272,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
or {}
)
team_metadata = self._team_metadata_json(
raw_metadata.get("user_api_key_team_metadata")
raw_metadata.get("user_api_key_team_metadata"),
self.config.baggage_team_metadata_keys,
)
if team_metadata:
self.safe_set_attribute(
@ -1268,15 +1296,20 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
)
@staticmethod
def _team_metadata_json(value: Any) -> Optional[str]:
"""JSON-serialize a team's metadata dict for a single span attribute.
def _team_metadata_json(value: Any, allowed_keys: List[str]) -> Optional[str]:
"""JSON-serialize only the allowlisted sub-keys of a team's metadata.
Returns ``None`` for a missing, non-dict, or empty mapping so the
empty case is dropped rather than stamping a useless ``"{}"``.
Returns ``None`` when nothing is allowlisted or no allowlisted key is
present, so the empty case is dropped rather than stamping a useless
``"{}"`` (and so a team's metadata never leaves the process until an
operator opts each sub-key in via ``baggage_team_metadata_keys``).
"""
if not isinstance(value, dict) or not value:
if not isinstance(value, dict) or not value or not allowed_keys:
return None
return safe_dumps(value)
filtered = {key: value[key] for key in allowed_keys if key in value}
if not filtered:
return None
return safe_dumps(filtered)
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
duration_s = (end_time - start_time).total_seconds()

View file

@ -172,10 +172,13 @@ nothing here imports outside it:
`capture_span_content` gates whether prompt/response bodies may be written as
span attributes; it defaults **off** (`no_content`). The Baggage allowlists are
configurable, not hard-coded: set `LITELLM_OTEL_BAGGAGE_PROMOTED_KEYS` /
`LITELLM_OTEL_BAGGAGE_METADATA_KEYS` (comma-separated) as env vars, or
`baggage_promoted_keys` / `baggage_metadata_keys` (YAML lists) under
`callback_settings.otel` in `config.yaml` — the latter reach the config through
the logger's constructor kwargs.
`LITELLM_OTEL_BAGGAGE_METADATA_KEYS` /
`LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS` (comma-separated) as env vars, or
`baggage_promoted_keys` / `baggage_metadata_keys` /
`baggage_team_metadata_keys` (YAML lists) under `callback_settings.otel` in
`config.yaml` — the latter reach the config through the logger's constructor
kwargs. `baggage_team_metadata_keys` is empty by default, so none of a team's
free-form metadata is promoted until each sub-key is explicitly allowlisted.
- [`baggage.py`](./model/baggage.py) — the single definition of which request-identity
values are promoted into Baggage (so child spans inherit them) and under which
attribute keys.

View file

@ -254,6 +254,7 @@ class OpenTelemetryV2(CustomLogger):
data.request_model,
promoted_keys=tuple(self.config.baggage_promoted_keys),
metadata_keys=tuple(self.config.baggage_metadata_keys),
team_metadata_keys=tuple(self.config.baggage_team_metadata_keys),
)
if bag:
parent_ctx = set_request_baggage(bag, context=parent_ctx)
@ -380,6 +381,7 @@ class OpenTelemetryV2(CustomLogger):
model,
promoted_keys=tuple(self.config.baggage_promoted_keys),
metadata_keys=tuple(self.config.baggage_metadata_keys),
team_metadata_keys=tuple(self.config.baggage_team_metadata_keys),
)
if bag:
# Attach (no detach): the contextvar is scoped to this request's

View file

@ -6,26 +6,36 @@ LLM-call span so that child spans (guardrail, service) inherit them.
the allowlisted keys onto every span.
This module is the single place baggage is defined: ``_PROMOTABLE`` maps each
promotable attribute key to how its value is read, and the two ``*_KEYS``
defaults select what is promoted unless the config overrides them.
promotable attribute key to how its value is read, and the ``*_KEYS`` defaults
select what is promoted unless the config overrides them. ``TEAM_METADATA``'s
extractor filters the team's free-form metadata to the sub-keys an operator
allowlists via ``baggage_team_metadata_keys`` (default none), so the blob is
never promoted whole.
"""
from collections.abc import Callable
import json
from collections.abc import Callable, Mapping
from typing import Final
from litellm.integrations.otel.model.metadata import RequestIdentity
from litellm.integrations.otel.model.semconv import GenAI, LiteLLM
# Attribute key -> value extractor over (identity, request_model). The single
# definition of what may be promoted and under which key.
_PROMOTABLE: Final[dict[str, Callable[[RequestIdentity, str | None], str | None]]] = {
LiteLLM.TEAM_ID: lambda identity, model: identity.team_id,
LiteLLM.TEAM_ALIAS: lambda identity, model: identity.team_alias,
LiteLLM.TEAM_METADATA: lambda identity, model: identity.team_metadata,
LiteLLM.KEY_HASH: lambda identity, model: identity.key_hash,
LiteLLM.END_USER: lambda identity, model: identity.end_user,
GenAI.REQUEST_MODEL: lambda identity, model: model,
LiteLLM.PROVIDER_MODEL: lambda identity, model: identity.provider_model,
# Attribute key -> value extractor over (identity, request_model,
# team_metadata_keys). The single definition of what may be promoted and under
# which key. Only the ``TEAM_METADATA`` extractor consults team_metadata_keys
# (to filter the team's metadata to an allowlist); the rest ignore it.
_PROMOTABLE: Final[
dict[str, Callable[[RequestIdentity, str | None, tuple[str, ...]], str | None]]
] = {
LiteLLM.TEAM_ID: lambda identity, model, team_metadata_keys: identity.team_id,
LiteLLM.TEAM_ALIAS: lambda identity, model, team_metadata_keys: identity.team_alias,
LiteLLM.TEAM_METADATA: lambda identity, model, team_metadata_keys: _filtered_team_metadata_json(
identity.team_metadata, team_metadata_keys
),
LiteLLM.KEY_HASH: lambda identity, model, team_metadata_keys: identity.key_hash,
LiteLLM.END_USER: lambda identity, model, team_metadata_keys: identity.end_user,
GenAI.REQUEST_MODEL: lambda identity, model, team_metadata_keys: model,
LiteLLM.PROVIDER_MODEL: lambda identity, model, team_metadata_keys: identity.provider_model,
}
# Keys promoted by default (a subset of ``_PROMOTABLE``). ``END_USER`` is
@ -50,23 +60,31 @@ DEFAULT_BAGGAGE_METADATA_KEYS: Final[tuple[str, ...]] = (
"requester_ip_address",
)
# Sub-keys of the team's free-form metadata eligible for promotion under
# ``litellm.team.metadata``. Empty by default: a team's metadata can hold
# arbitrary operator data, so none of it is promoted until each key is
# explicitly allowlisted via ``config.baggage_team_metadata_keys``.
DEFAULT_BAGGAGE_TEAM_METADATA_KEYS: Final[tuple[str, ...]] = ()
def promoted_baggage(
identity: RequestIdentity,
request_model: str | None,
promoted_keys: tuple[str, ...],
metadata_keys: tuple[str, ...] = DEFAULT_BAGGAGE_METADATA_KEYS,
team_metadata_keys: tuple[str, ...] = DEFAULT_BAGGAGE_TEAM_METADATA_KEYS,
) -> dict[str, str]:
"""Identity values to write into Baggage, filtered to ``promoted_keys``.
``promoted_keys`` selects from ``_PROMOTABLE``; ``metadata_keys`` selects
sub-keys of ``identity.metadata`` to promote under ``litellm.metadata.*``.
Empty values are dropped.
sub-keys of ``identity.metadata`` to promote under ``litellm.metadata.*``;
``team_metadata_keys`` selects sub-keys of the team's metadata to promote
under ``litellm.team.metadata``. Empty values are dropped.
"""
out: dict[str, str] = {}
for key, extract in _PROMOTABLE.items():
if key in promoted_keys:
value = extract(identity, request_model)
value = extract(identity, request_model, team_metadata_keys)
if value:
out[key] = value
for meta_key in metadata_keys:
@ -74,3 +92,21 @@ def promoted_baggage(
if value:
out[f"{LiteLLM.METADATA_PREFIX}{meta_key}"] = value
return out
def _filtered_team_metadata_json(
metadata: Mapping[str, object] | None,
allowed_keys: tuple[str, ...],
) -> str | None:
"""JSON-serialize only the allowlisted sub-keys of a team's metadata.
Returns ``None`` when nothing is allowlisted or no allowlisted key is
present, so the empty case is dropped rather than promoting ``"{}"``. Keys
are sorted for a stable, diff-friendly value.
"""
if not isinstance(metadata, Mapping) or not allowed_keys:
return None
filtered = {key: metadata[key] for key in allowed_keys if key in metadata}
if not filtered:
return None
return json.dumps(filtered, default=str, sort_keys=True)

View file

@ -9,6 +9,7 @@ from typing_extensions import Annotated
from litellm.integrations.otel.model.baggage import (
BAGGAGE_PROMOTED_KEYS,
DEFAULT_BAGGAGE_METADATA_KEYS,
DEFAULT_BAGGAGE_TEAM_METADATA_KEYS,
)
#: Master feature-flag env var. The logger is inert until this is truthy.
@ -168,10 +169,25 @@ class OpenTelemetryV2Config(BaseSettings):
"``callback_settings.otel.baggage_metadata_keys`` in config.yaml."
),
)
baggage_team_metadata_keys: Annotated[List[str], NoDecode] = Field(
default_factory=lambda: list(DEFAULT_BAGGAGE_TEAM_METADATA_KEYS),
validation_alias=AliasChoices(
"baggage_team_metadata_keys", "LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS"
),
description=(
"Sub-keys of the team's free-form metadata promoted under "
"``litellm.team.metadata``. Empty by default so none of a team's "
"metadata leaves the process until explicitly allowlisted. Configure "
"via the ``LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS`` env var "
"(comma-separated) or "
"``callback_settings.otel.baggage_team_metadata_keys`` in config.yaml."
),
)
@field_validator(
"baggage_promoted_keys",
"baggage_metadata_keys",
"baggage_team_metadata_keys",
"mapper_names",
mode="before",
)

View file

@ -36,7 +36,6 @@ model. They coincide on the SDK path, which is correct.
from __future__ import annotations
import json
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Mapping, cast
@ -53,8 +52,10 @@ class RequestIdentity:
call_id: str | None = None
team_id: str | None = None
team_alias: str | None = None
# The team's free-form metadata dict, JSON-serialized (empty/missing -> None).
team_metadata: str | None = None
# The team's free-form metadata, carried raw (empty/missing -> None) and
# filtered to an operator allowlist only at Baggage-promotion time, so an
# unconfigured deployment never promotes any of it.
team_metadata: Mapping[str, Any] | None = None
key_hash: str | None = None
end_user: str | None = None
# The model litellm dispatched to the provider. Only known once the call
@ -86,7 +87,7 @@ class RequestIdentity:
or as_str(raw_meta.get("team_id")),
team_alias=as_str(raw_meta.get("user_api_key_team_alias"))
or as_str(raw_meta.get("team_alias")),
team_metadata=_team_metadata_json(
team_metadata=_team_metadata_dict(
raw_meta.get("user_api_key_team_metadata")
),
key_hash=as_str(raw_meta.get("user_api_key_hash")),
@ -121,7 +122,7 @@ class RequestIdentity:
return cls(
team_id=as_str(get("team_id")),
team_alias=as_str(get("team_alias")),
team_metadata=_team_metadata_json(get("team_metadata")),
team_metadata=_team_metadata_dict(get("team_metadata")),
key_hash=as_str(get("api_key")),
end_user=as_str(get("end_user_id")),
# ``provider_model`` is unknown at the auth boundary — routing hasn't
@ -300,16 +301,14 @@ def _model_info_id(model_info: object) -> str | None:
return None
def _team_metadata_json(value: object) -> str | None:
"""JSON-serialize a team's metadata dict for a single Baggage value.
def _team_metadata_dict(value: object) -> Mapping[str, Any] | None:
"""The team's free-form metadata as a raw mapping, or ``None`` when missing
or empty.
Returns ``None`` for a missing, non-dict, or empty mapping so the empty case
is dropped rather than promoting a useless ``"{}"``. Keys are sorted for a
stable, diff-friendly serialization.
Carried raw on the identity and filtered to an operator allowlist only at
Baggage-promotion time (see ``baggage.promoted_baggage``), so an empty case
is dropped rather than carrying a useless ``{}``.
"""
if not isinstance(value, Mapping) or not value:
return None
try:
return json.dumps(value, default=str, sort_keys=True)
except Exception:
return None
if isinstance(value, Mapping) and value:
return dict(value)
return None

View file

@ -511,6 +511,23 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_cached_tokens_metric"),
)
# Provider prompt-caching metrics
self.litellm_provider_cache_read_input_tokens_metric = self._counter_factory(
name="litellm_provider_cache_read_input_tokens_metric",
documentation="Total prompt/input tokens read from provider prompt cache (e.g. OpenAI/Anthropic/Gemini/Bedrock)",
labelnames=self.get_labels_for_metric(
"litellm_provider_cache_read_input_tokens_metric"
),
)
self.litellm_provider_cache_creation_input_tokens_metric = self._counter_factory(
name="litellm_provider_cache_creation_input_tokens_metric",
documentation="Total prompt/input tokens written to provider prompt cache (e.g. Anthropic/Bedrock)",
labelnames=self.get_labels_for_metric(
"litellm_provider_cache_creation_input_tokens_metric"
),
)
# User and Team count metrics
self.litellm_total_users_metric = self._gauge_factory(
"litellm_total_users",
@ -1458,11 +1475,11 @@ class PrometheusLogger(CustomLogger):
"""
cache_hit = standard_logging_payload.get("cache_hit")
# Only track if cache_hit has a definite value (True or False)
if cache_hit is None:
return
if cache_hit is True:
# Historically these metrics only tracked LiteLLM caching.
# Provider prompt-caching metrics are still emitted below.
pass
elif cache_hit is True:
# Increment cache hits counter
PrometheusLogger._inc_labeled_counter(
self,
@ -1493,6 +1510,51 @@ class PrometheusLogger(CustomLogger):
label_context=label_context,
)
# Provider prompt caching metrics are independent of LiteLLM cache_hit.
provider_cache_read_tokens = 0
provider_cache_creation_tokens = 0
usage_obj = (standard_logging_payload.get("metadata", {}) or {}).get(
"usage_object"
)
if isinstance(usage_obj, dict):
# Prefer explicit provider cache fields when available.
_read = usage_obj.get("cache_read_input_tokens")
_write = usage_obj.get("cache_creation_input_tokens")
if isinstance(_read, int):
provider_cache_read_tokens = _read
if isinstance(_write, int):
provider_cache_creation_tokens = _write
# Fallback to prompt_tokens_details.cached_tokens (common normalization point).
# Only fallback when the explicit field is genuinely absent (None).
if _read is None:
prompt_details = usage_obj.get("prompt_tokens_details")
if isinstance(prompt_details, dict):
cached_tokens = prompt_details.get("cached_tokens")
if isinstance(cached_tokens, int):
provider_cache_read_tokens = cached_tokens
if provider_cache_read_tokens > 0:
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_provider_cache_read_input_tokens_metric,
"litellm_provider_cache_read_input_tokens_metric",
enum_values,
label_context=label_context,
amount=float(provider_cache_read_tokens),
)
if provider_cache_creation_tokens > 0:
PrometheusLogger._inc_labeled_counter(
self,
self.litellm_provider_cache_creation_input_tokens_metric,
"litellm_provider_cache_creation_input_tokens_metric",
enum_values,
label_context=label_context,
amount=float(provider_cache_creation_tokens),
)
async def _increment_remaining_budget_metrics(
self,
user_api_team: Optional[str],

View file

@ -87,6 +87,7 @@ class ExceptionCheckers:
"is longer than the model's context length",
"input tokens exceed the configured limit",
"`inputs` tokens + `max_new_tokens` must be",
"exceeds the available context size", # llama.cpp/Lemonade
"exceeds the maximum number of tokens allowed", # Gemini
]
for substring in known_exception_substrings:
@ -891,12 +892,14 @@ def exception_type( # type: ignore # noqa: PLR0915
response=getattr(original_exception, "response", None),
litellm_debug_info=extra_information,
)
elif "model's maximum context limit" in error_str:
elif ExceptionCheckers.is_error_str_context_window_exceeded(error_str):
exception_mapping_worked = True
raise ContextWindowExceededError(
message=f"{custom_llm_provider.capitalize()}Exception: Context Window Error - {error_str}",
model=model,
llm_provider=custom_llm_provider,
response=getattr(original_exception, "response", None),
litellm_debug_info=extra_information,
)
elif "token_quota_reached" in error_str:
exception_mapping_worked = True

View file

@ -47,8 +47,9 @@ async def async_completion_with_fallbacks(**kwargs):
completion_kwargs = safe_deep_copy(base_kwargs)
# Handle dictionary fallback configurations
if isinstance(fallback, dict):
model = fallback.pop("model", original_model)
completion_kwargs.update(fallback)
fallback_config = safe_deep_copy(dict(fallback))
model = fallback_config.pop("model", original_model)
completion_kwargs.update(fallback_config)
else:
model = fallback

View file

@ -850,7 +850,47 @@ def extract_file_data(file_data: FileTypes) -> ExtractedFileData:
# ---------------------------------------------------------------------------
def unpack_defs(schema: dict, defs: dict) -> None:
def _estimate_json_bytes(obj: Any) -> int:
"""Estimate the JSON-serialised byte size of ``obj`` without materialising
JSON. Walks iteratively (no recursion stack risk).
String length is read via ``len()`` (O(1) on Python ``str``) so a target
containing a 100MB description costs ~one walk step, not a 100MB
serialisation. Escape sequences are not counted exactly, so this is an
approximation -- but always within a small constant factor of the real
serialised size, which is what a schema-bomb budget needs.
"""
total = 0
stack: list = [obj]
while stack:
x = stack.pop()
if isinstance(x, dict):
total += 2 # `{}`
for k, v in x.items():
total += len(str(k)) + 4 # `"k":,`
stack.append(v)
elif isinstance(x, list):
total += 2 # `[]`
total += max(0, len(x) - 1) # commas between items
stack.extend(x)
elif isinstance(x, str):
total += len(x) + 2
elif isinstance(x, bool): # bool subclasses int -- check first
total += 4 if x else 5
elif x is None:
total += 4
elif isinstance(x, (int, float)):
total += 24 # generous upper bound for stringified numbers
else:
total += 24
return total
def unpack_defs(
schema: dict,
defs: dict,
max_inlined_bytes: Optional[int] = None,
) -> None:
"""Expand *all* ``$ref`` entries pointing into ``$defs`` / ``definitions``.
This utility walks the entire schema tree (dicts and lists) so it naturally
@ -860,6 +900,15 @@ def unpack_defs(schema: dict, defs: dict) -> None:
It mutates *schema* in-place and does **not** return anything. The helper
keeps memory overhead low by resolving nodes as it encounters them rather
than materialising a fully dereferenced copy first.
``max_inlined_bytes`` caps the cumulative JSON-byte size of every target
that has been inlined and is checked *before* each ``copy.deepcopy``, so
an oversized expansion is rejected without first materialising it. A byte
bound is the universal measure of expansion -- it simultaneously caps
ref-count fan-out, node-count amplification, and scalar-byte amplification
(a target containing a large string, ``const``, or ``enum`` entry).
Defaults to ``None`` (unbounded) so existing callers are unaffected;
raises ``ValueError`` on overflow.
"""
import copy
@ -879,6 +928,7 @@ def unpack_defs(schema: dict, defs: dict) -> None:
queue: deque[
tuple[Any, Union[dict, list, None], Union[str, int, None], dict, set]
] = deque([(schema, None, None, root_defs, set())])
inlined_bytes = 0
while queue:
node, parent, key, active_defs, ref_chain = queue.popleft()
@ -899,6 +949,16 @@ def unpack_defs(schema: dict, defs: dict) -> None:
if target_schema is None:
continue
if max_inlined_bytes is not None:
inlined_bytes += _estimate_json_bytes(target_schema)
if inlined_bytes > max_inlined_bytes:
raise ValueError(
f"unpack_defs: inlined schema exceeded the "
f"{max_inlined_bytes:,}-byte budget. Refusing to "
f"deep-copy further to prevent schema-bomb "
f"resource exhaustion."
)
# Merge defs from the target to capture nested definitions
child_defs = {
**active_defs,
@ -946,6 +1006,61 @@ def unpack_defs(schema: dict, defs: dict) -> None:
queue.append((item, node, idx, active_defs, ref_chain))
def _has_legacy_defs(schema: object) -> bool:
if not isinstance(schema, dict):
return False
components = schema.get("components")
return "definitions" in schema or (
isinstance(components, dict) and isinstance(components.get("schemas"), dict)
)
# Schema-bomb budget for ``unpack_legacy_defs``: cap the cumulative JSON-byte
# size of every inlined target. A byte cap is the universal measure of
# expansion -- it simultaneously bounds ref-count fan-out, node-count
# amplification, and scalar-byte amplification (large ``description`` /
# ``const`` / ``enum`` values). Real-world MCP / OpenAPI-derived tool schemas
# inline well under 1MB; 10MB sits two orders of magnitude above that, well
# below memory-pressure territory, and rejects request-supplied bombs before
# the proxy materialises them.
_LEGACY_DEFS_MAX_INLINED_BYTES = 10_000_000
def unpack_legacy_defs(
schema: dict,
*,
copy: bool = False,
max_inlined_bytes: int = _LEGACY_DEFS_MAX_INLINED_BYTES,
) -> dict:
"""Inline ``$ref``s backed by draft-04 ``definitions`` / OpenAPI
``components.schemas``. ``$defs`` is left untouched.
Anthropic and Fireworks tool-schema resolvers only recognise ``$defs``;
legacy / OpenAPI def blocks are otherwise silently dropped and leave
dangling pointers. See https://github.com/BerriAI/litellm/issues/26692.
Mutates ``schema`` in place and returns it. Pass ``copy=True`` to deep-copy
first (only when there is actually work to do). ``max_inlined_bytes``
bounds the cumulative JSON-byte size of inlined targets so request-supplied
schemas cannot expand into a schema-bomb before reaching the upstream
provider -- raises ``ValueError`` on overflow.
"""
if not _has_legacy_defs(schema):
return schema
if copy:
import copy as _copy
schema = _copy.deepcopy(schema)
# On key collision, ``definitions`` wins over ``components.schemas`` --
# ``unpack_defs`` keys refs by last path segment so a single name can only
# resolve to one body, and ``definitions`` is the JSON-Schema-native
# namespace.
defs = schema.pop("components", {}).get("schemas") or {}
defs.update(schema.pop("definitions", None) or {})
unpack_defs(schema, defs, max_inlined_bytes=max_inlined_bytes)
return schema
def _get_image_mime_type_from_url(url: str) -> Optional[str]:
"""
Get mime type for common image URLs

View file

@ -29,6 +29,7 @@ from litellm.constants import (
RESPONSE_FORMAT_TOOL_NAME,
)
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_legacy_defs
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.anthropic import (
@ -680,6 +681,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if "properties" not in _input_schema:
_input_schema["properties"] = {}
# Inline legacy / OpenAPI $refs before the allow-list filter strips
# their backing def blocks (https://github.com/BerriAI/litellm/issues/26692).
_input_schema = unpack_legacy_defs(_input_schema, copy=True)
_allowed_properties = set(AnthropicInputSchema.__annotations__.keys())
input_schema_filtered = {
k: v for k, v in _input_schema.items() if k in _allowed_properties

View file

@ -1534,10 +1534,9 @@ class BaseAWSLLM:
)
sigv4 = SigV4Auth(credentials, service_name, aws_region_name)
if headers is not None:
headers = headers or {}
if not any(header_name.lower() == "content-type" for header_name in headers):
headers = {"Content-Type": "application/json", **headers}
else:
headers = {"Content-Type": "application/json"}
aws_signature_headers = self._filter_headers_for_aws_signature(headers)
request = AWSRequest(

View file

@ -233,6 +233,259 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
# example; add others here as they adopt the same schema.
CONVERSE_INVOKE_PROVIDERS = ("nova",)
# OpenAI batch URL that signals an embedding request. Per OpenAI Batch API
# spec, every JSONL record carries a `url` field; we use it as the
# authoritative signal to route the line to the embedding code path
# instead of inferring from the presence of `input` vs `messages`.
OPENAI_EMBEDDINGS_URL = "/v1/embeddings"
@staticmethod
def _is_embedding_record(openai_jsonl_record: Dict[str, Any]) -> bool:
"""
Decide whether an OpenAI batch JSONL line is an embedding request.
Precedence (strict - any explicit `url` short-circuits):
1. `url == "/v1/embeddings"` -> embedding. Authoritative per the
OpenAI Batch API spec.
2. Any other non-empty `url` (e.g. `/v1/chat/completions`) -> NOT
embedding. We trust the caller's explicit signal even if the
body would otherwise suggest embedding; misrouting a chat
record into the embedding transformer would corrupt the
modelInput, while a chat-shaped body sent to the chat path
either succeeds or fails cleanly inside that transformer.
3. `url` missing/empty -> fall back to body shape. Requires
`input` present AND `messages` absent so a malformed record
carrying both keys routes to the chat path (safer default:
Anthropic transforms ignore unknown top-level keys, whereas
the embedding transformer would silently drop the messages).
"""
url = openai_jsonl_record.get("url")
if url == BedrockFilesConfig.OPENAI_EMBEDDINGS_URL:
return True
if url:
return False
body = openai_jsonl_record.get("body", {})
if not isinstance(body, dict):
return False
return "input" in body and "messages" not in body
# Identifier for the Bedrock Titan v2 InvokeModel body schema as stored
# in `model_prices_and_context_window.json`. Centralized so future
# embedding-schema variants can add their own value
# (e.g. `cohere_v3`, `titan_g1`, `titan_multimodal`) without touching
# the detection logic.
_TITAN_V2_INVOCATION_SCHEMA = "titan_v2"
# Substring marker used as a fallback when the registry can't resolve
# the model id - notably cross-region inference profile prefixes
# (`us.amazon.titan-embed-text-v2:0`) and Bedrock ARN forms, which
# `get_model_info` doesn't normalize today.
_TITAN_V2_EMBED_MODEL_MARKER = "titan-embed-text-v2"
# Nested field name under `provider_specific_entry` that identifies the
# Bedrock InvokeModel body schema for batch inference.
# `provider_specific_entry` is the registry's escape hatch for fields
# `get_model_info` doesn't promote to top-level - exactly what we need
# here. Documented in the `sample_spec` entry of
# `model_prices_and_context_window.json` and surfaced by
# `get_model_info` (see `ModelInfo.provider_specific_entry`).
_BEDROCK_INVOCATION_SCHEMA_FIELD = "bedrock_invocation_schema"
@staticmethod
def _is_titan_v2_embed_model(model: str) -> bool:
"""
True iff `model` refers to Amazon Titan Text Embeddings V2.
Resolution order:
1. `model_prices_and_context_window.json` via `get_model_info`.
The Titan v2 registry entry carries an explicit
`provider_specific_entry.bedrock_invocation_schema` discriminator
(`"titan_v2"`). When the registry resolves the id we trust that
field as the source of truth - no hardcoded model-id comparison
needed.
2. Substring fallback (`titan-embed-text-v2` followed by `:`, `/`,
or end-of-string) for ids the registry can't normalize. This
catches cross-region inference profile prefixes
(`us.amazon.titan-embed-text-v2:0`) and Bedrock ARN forms; the
marker boundary check rejects lookalikes like
`titan-embed-text-v20` or `titan-embed-text-v2-experimental`.
Tolerant of common id shapes:
- "amazon.titan-embed-text-v2:0"
- "bedrock/amazon.titan-embed-text-v2:0"
- "us.amazon.titan-embed-text-v2:0" (cross-region inference profile)
- ARN forms ending in ".../amazon.titan-embed-text-v2:0"
"""
# Registry-driven path: when get_model_info resolves the id we trust
# the registry's discriminator. A resolved id with a different (or
# absent) schema value here is intentionally not given a substring
# second-chance - the registry is authoritative for ids it knows.
registry_schema = BedrockFilesConfig._lookup_provider_specific_field(
model, BedrockFilesConfig._BEDROCK_INVOCATION_SCHEMA_FIELD
)
if registry_schema is not None:
return registry_schema == BedrockFilesConfig._TITAN_V2_INVOCATION_SCHEMA
# Registry silence -> substring fallback for unmapped ids only.
normalized = model.lower()
if normalized.startswith("bedrock/"):
normalized = normalized[len("bedrock/") :]
marker = BedrockFilesConfig._TITAN_V2_EMBED_MODEL_MARKER
idx = normalized.find(marker)
if idx < 0:
return False
end = idx + len(marker)
return end == len(normalized) or normalized[end] in (":", "/")
@staticmethod
def _lookup_provider_specific_field(model_id: str, field: str) -> Optional[str]:
"""
Read a nested string field from the registry entry's
`provider_specific_entry` dict via `litellm.get_model_info`.
Returns the field's string value when:
- the registry resolves `model_id`,
- the entry exposes `provider_specific_entry` as a dict, and
- that dict has `field` mapped to a non-empty string.
Otherwise returns `None`.
Isolating this means feature detectors (Titan v2 today, future
Cohere Embed / Nova Multimodal branches) share one defensive
try/except shape instead of duplicating it. The `None` return
covers every realistic failure mode: `get_model_info` raises
(cross-region profile prefixes, Bedrock ARN forms, unreleased
models), returns a non-dict, has no `provider_specific_entry`, or
the requested field is missing / non-string / empty.
"""
try:
from litellm import get_model_info
info = get_model_info(model_id)
except Exception:
return None
if not isinstance(info, dict):
return None
provider_specific = info.get("provider_specific_entry")
if not isinstance(provider_specific, dict):
return None
value = provider_specific.get(field)
return value if isinstance(value, str) and value else None
@staticmethod
def _coerce_embedding_input_to_string(raw_input: Any, model: str = "") -> str:
"""
Normalize an OpenAI /v1/embeddings `input` field into the single
string that Bedrock Titan v2 InvokeModel expects in `inputText`.
Accepts: a string, or a single-element list containing one string.
Rejects (with actionable messages):
- None / missing -> ValueError
- Multi-element string lists -> ValueError, prompts caller to
emit one JSONL line per input
- Pre-tokenized inputs (List[int], List[List[int]]) -> NotImplementedError
- Any other type -> ValueError
Extracted so the validation can be exercised in isolation and so
future embedding-provider branches (Titan G1, Cohere) can reuse it
without duplicating the type-shaping logic.
"""
if raw_input is None:
raise ValueError(
"Embedding batch record is missing required `input` field: "
f"model={model}"
)
# Bedrock InvokeModel for Titan v2 takes exactly one string `inputText`
# per call. Pre-tokenized inputs and multi-element string lists are
# explicitly unsupported so callers emit one JSONL line per embedding
# instead of relying on us to silently fan out or concatenate.
if isinstance(raw_input, list):
if len(raw_input) == 1:
candidate = raw_input[0]
else:
raise ValueError(
"Bedrock batch embedding requires one input per JSONL "
"record. Got a list with "
f"{len(raw_input)} items for model={model}; emit one "
"JSONL line per input string instead."
)
else:
candidate = raw_input
# Catches pre-tokenized inputs (List[int] from OpenAI spec, or a
# single int slipping past the list-unwrap above).
# NOTE: bool is a subclass of int but treating True/False as a token
# is meaningless either way, so the broad check is fine.
if isinstance(candidate, (list, int)):
raise NotImplementedError(
"Bedrock Titan v2 batch embedding does not support "
"pre-tokenized integer inputs. Pass `input` as a string "
f"(model={model})."
)
if not isinstance(candidate, str):
raise ValueError(
"Bedrock batch embedding `input` must be a string (or a "
"single-element list of strings). Got type "
f"{type(candidate).__name__} for model={model}."
)
return candidate
def _map_openai_embedding_to_bedrock_params(
self,
openai_request_body: Dict[str, Any],
) -> Dict[str, Any]:
"""
Transform an OpenAI /v1/embeddings request body into the
Bedrock InvokeModel `modelInput` for embedding models that AWS
supports via batch inference (CreateModelInvocationJob).
Currently routes Amazon Titan Text Embeddings V2 only; other
embedding providers (Titan G1, Titan Multimodal, Cohere Embed,
Nova Multimodal Embeddings) raise NotImplementedError until they
get a dedicated branch. Splitting them keeps PR scope tight and
lets each model's request schema be exercised by its own tests.
AWS docs (Titan v2 InvokeModel body):
https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-titan-embed-text.html
"""
from litellm.llms.bedrock.embed.amazon_titan_v2_transformation import (
AmazonTitanV2Config,
)
_model = openai_request_body.get("model", "")
if not self._is_titan_v2_embed_model(_model):
# Refuse early instead of silently shaping the body for the wrong
# provider. The synchronous /v1/embeddings path supports more
# models, but each has a different InvokeModel schema; mapping
# them here without dedicated tests would risk corrupt batches.
raise NotImplementedError(
"Bedrock batch embedding currently supports only Amazon "
"Titan Text Embeddings V2 (model id contains "
f"'titan-embed-text-v2'). Got model={_model!r}. Track other "
"embedding models in https://github.com/BerriAI/litellm/issues."
)
input_text = self._coerce_embedding_input_to_string(
openai_request_body.get("input"), model=_model
)
# Map OpenAI-style params (dimensions, encoding_format) onto the
# Titan v2 schema (dimensions, embeddingTypes) via the embed config
# so this stays in sync with the synchronous /v1/embeddings path.
non_default_params = {
k: v for k, v in openai_request_body.items() if k not in ("model", "input")
}
titan_config = AmazonTitanV2Config()
inference_params = titan_config.map_openai_params(
non_default_params=non_default_params,
optional_params={},
)
return dict(
titan_config._transform_request(
input=input_text, inference_params=inference_params
)
)
def _map_openai_to_bedrock_params(
self,
openai_request_body: Dict[str, Any],
@ -349,10 +602,19 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
# Determine provider from model name
provider = self.get_bedrock_invoke_provider(model)
# Transform to Bedrock modelInput format
model_input = self._map_openai_to_bedrock_params(
openai_request_body=openai_body, provider=provider
)
# Route to the embedding transformer when the OpenAI batch line
# targets /v1/embeddings; otherwise fall back to the existing
# chat-completion path. We branch here (rather than inside
# `_map_openai_to_bedrock_params`) so the chat helper keeps its
# narrow contract and the embedding helper can evolve independently.
if self._is_embedding_record(_openai_jsonl_content):
model_input = self._map_openai_embedding_to_bedrock_params(
openai_request_body=openai_body
)
else:
model_input = self._map_openai_to_bedrock_params(
openai_request_body=openai_body, provider=provider
)
# Create Bedrock batch record
record_id = _openai_jsonl_content.get(

View file

@ -5,6 +5,7 @@ Common utilities, constants, and error handling for Black Forest Labs API.
"""
from typing import Dict
from urllib.parse import urlparse
from litellm.llms.base_llm.chat.transformation import BaseLLMException
@ -18,6 +19,42 @@ class BlackForestLabsError(BaseLLMException):
# API Constants
DEFAULT_API_BASE = "https://api.bfl.ai"
# BFL uses regional subdomains (e.g. gateway.bfl.ai) for polling URLs that
# differ from the submission host (api.bfl.ai). We validate against the
# registered domain rather than doing a strict same-origin check.
_BFL_REGISTERED_DOMAIN = "bfl.ai"
def assert_bfl_polling_url(polling_url: str) -> None:
"""Validate that a polling URL points to a BFL-controlled host.
BFL returns polling URLs on subdomains like ``gateway.bfl.ai`` that differ
from the submission host ``api.bfl.ai``. A strict same-origin check would
reject these legitimate URLs. Instead we verify the host is ``bfl.ai`` or
any subdomain of it, which keeps the SSRF guarantee (credentials only go
to BFL-controlled infrastructure) without false-positives on regional hosts.
Raises:
BlackForestLabsError: If the polling URL scheme or host is not trusted.
"""
parsed = urlparse(polling_url)
host = (parsed.hostname or "").lower()
if parsed.scheme != "https":
raise BlackForestLabsError(
status_code=502,
message="Rejected polling URL: scheme must be https",
)
if host != _BFL_REGISTERED_DOMAIN and not host.endswith(
"." + _BFL_REGISTERED_DOMAIN
):
raise BlackForestLabsError(
status_code=502,
message="Rejected polling URL: host is not within the bfl.ai domain",
)
# Polling configuration
DEFAULT_POLLING_INTERVAL = 1.5 # seconds
DEFAULT_MAX_POLLING_TIME = 300 # 5 minutes

View file

@ -15,7 +15,6 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -29,6 +28,7 @@ from ..common_utils import (
DEFAULT_MAX_POLLING_TIME,
DEFAULT_POLLING_INTERVAL,
BlackForestLabsError,
assert_bfl_polling_url,
)
from .transformation import BlackForestLabsImageEditConfig
@ -332,16 +332,11 @@ class BlackForestLabsImageEdit:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Reject polling URLs that don't belong to BFL-controlled infrastructure.
# BFL uses regional subdomains (e.g. gateway.bfl.ai) that differ from the
# submission host (api.bfl.ai), so we validate against the registered
# domain rather than doing a strict same-origin check. VERIA-51.
assert_bfl_polling_url(polling_url)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}
@ -428,16 +423,11 @@ class BlackForestLabsImageEdit:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Reject polling URLs that don't belong to BFL-controlled infrastructure.
# BFL uses regional subdomains (e.g. gateway.bfl.ai) that differ from the
# submission host (api.bfl.ai), so we validate against the registered
# domain rather than doing a strict same-origin check. VERIA-51.
assert_bfl_polling_url(polling_url)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}

View file

@ -15,7 +15,6 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -29,6 +28,7 @@ from ..common_utils import (
DEFAULT_MAX_POLLING_TIME,
DEFAULT_POLLING_INTERVAL,
BlackForestLabsError,
assert_bfl_polling_url,
)
from .transformation import BlackForestLabsImageGenerationConfig
@ -172,6 +172,10 @@ class BlackForestLabsImageGeneration:
raw_response=final_response,
model_response=model_response,
logging_obj=logging_obj,
request_data=data,
optional_params=optional_params,
litellm_params=litellm_params_dict,
encoding=None,
)
async def async_image_generation(
@ -274,6 +278,10 @@ class BlackForestLabsImageGeneration:
raw_response=final_response,
model_response=model_response,
logging_obj=logging_obj,
request_data=data,
optional_params=optional_params,
litellm_params=litellm_params_dict,
encoding=None,
)
def _poll_for_result_sync(
@ -318,16 +326,11 @@ class BlackForestLabsImageGeneration:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Reject polling URLs that don't belong to BFL-controlled infrastructure.
# BFL uses regional subdomains (e.g. gateway.bfl.ai) that differ from the
# submission host (api.bfl.ai), so we validate against the registered
# domain rather than doing a strict same-origin check. VERIA-51.
assert_bfl_polling_url(polling_url)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}
@ -414,16 +417,11 @@ class BlackForestLabsImageGeneration:
message="No polling_url in BFL response",
)
# Reject cross-origin polling URLs — the ``x-key`` auth header
# would otherwise leak to whatever URL the upstream returns.
# VERIA-51.
try:
assert_same_origin(polling_url, str(initial_response.request.url))
except SSRFError as ssrf_err:
raise BlackForestLabsError(
status_code=502,
message=f"Rejected polling URL: {ssrf_err}",
)
# Reject polling URLs that don't belong to BFL-controlled infrastructure.
# BFL uses regional subdomains (e.g. gateway.bfl.ai) that differ from the
# submission host (api.bfl.ai), so we validate against the registered
# domain rather than doing a strict same-origin check. VERIA-51.
assert_bfl_polling_url(polling_url)
# Get just the auth header for polling
polling_headers = {"x-key": headers.get("x-key", "")}

View file

@ -8,6 +8,7 @@ from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_legacy_defs
from litellm.litellm_core_utils.llm_response_utils.get_headers import (
get_response_headers,
)
@ -216,8 +217,13 @@ class FireworksAIConfig(OpenAIGPTConfig):
self, tools: List[OpenAIChatCompletionToolParam]
) -> List[OpenAIChatCompletionToolParam]:
for tool in tools:
if tool.get("type") == "function":
tool["function"].pop("strict", None)
if tool.get("type") != "function":
continue
function = tool["function"]
function.pop("strict", None)
params = function.get("parameters")
if isinstance(params, dict):
unpack_legacy_defs(params)
return tools
def _transform_messages_helper(

View file

@ -3,10 +3,12 @@ Translate from OpenAI's `/v1/chat/completions` to Lemonade's `/v1/chat/completio
"""
from typing import Any, List, Optional, Tuple, Union
from urllib.parse import quote
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
@ -18,6 +20,8 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig
class LemonadeChatConfig(OpenAILikeChatConfig):
_DEFAULT_API_KEY = "lemonade"
repeat_penalty: Optional[float] = None
functions: Optional[list] = None
logit_bias: Optional[dict] = None
@ -68,7 +72,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
This method queries the Lemonade /models endpoint to retrieve the list of available models.
Args:
api_key: Optional API key (Lemonade doesn't require authentication)
api_key: Optional API key for authenticated Lemonade servers
api_base: Optional API base URL (defaults to LEMONADE_API_BASE env var or http://localhost:8000)
Returns:
@ -87,6 +91,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
try:
response = litellm.module_level_client.get(
url=f"{api_base}/models",
headers=self._get_auth_headers(api_key),
)
except Exception as e:
raise ValueError(
@ -101,19 +106,131 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
model_list = response.json().get("data", [])
return ["lemonade/" + model["id"] for model in model_list]
@staticmethod
def _get_positive_int(value: Any) -> Optional[int]:
if isinstance(value, bool):
return None
if isinstance(value, int) and value > 0:
return value
if isinstance(value, str):
try:
parsed = int(value)
except ValueError:
return None
if parsed > 0:
return parsed
return None
@staticmethod
def _get_provider_specific_entry(model_info: dict) -> dict:
provider_specific_entry = model_info.get("provider_specific_entry")
if not isinstance(provider_specific_entry, dict):
provider_specific_entry = {}
else:
provider_specific_entry = provider_specific_entry.copy()
for key in ("recipe_options", "context_window", "max_context_window"):
if key in model_info:
provider_specific_entry[key] = model_info[key]
return provider_specific_entry
def _get_context_window(self, model_info: dict) -> Optional[int]:
provider_specific_entry = self._get_provider_specific_entry(model_info)
recipe_options = provider_specific_entry.get("recipe_options")
if not isinstance(recipe_options, dict):
recipe_options = {}
for value in (
recipe_options.get("ctx_size"),
model_info.get("max_input_tokens"),
provider_specific_entry.get("context_window"),
provider_specific_entry.get("max_context_window"),
):
parsed = self._get_positive_int(value)
if parsed is not None:
return parsed
return None
def _get_default_model_info(self, model: str) -> dict:
return {
"key": "lemonade/" + model,
"litellm_provider": "lemonade",
"mode": "chat",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"max_tokens": None,
"max_input_tokens": None,
"max_output_tokens": None,
}
def get_model_info(
self,
model: str,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> Any:
if model.startswith("lemonade/"):
model = model.split("/", 1)[1]
api_base, api_key = self._get_openai_compatible_provider_info(
api_base=api_base, api_key=api_key
)
encoded_model = quote(model, safe="")
try:
response = litellm.module_level_client.get(
url=f"{api_base}/models/{encoded_model}",
headers=self._get_auth_headers(api_key),
)
response.raise_for_status()
model_info = response.json()
except Exception:
verbose_logger.debug("LemonadeError: Could not get model info.")
return self._get_default_model_info(model)
max_input_tokens = self._get_context_window(model_info)
max_output_tokens = self._get_positive_int(model_info.get("max_output_tokens"))
max_tokens = self._get_positive_int(model_info.get("max_tokens"))
provider_specific_entry = self._get_provider_specific_entry(model_info)
model_info_response = self._get_default_model_info(model)
model_info_response.update(
{
"max_tokens": max_tokens or max_output_tokens,
"max_input_tokens": max_input_tokens,
"max_output_tokens": max_output_tokens,
}
)
if provider_specific_entry:
model_info_response["provider_specific_entry"] = provider_specific_entry
return model_info_response
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
# lemonade is openai compatible, we just need to set this to custom_openai and have the api_base be lemonade's endpoint
passed_api_base = api_base
api_base = (
api_base
or get_secret_str("LEMONADE_API_BASE")
or "http://localhost:8000/api/v1"
) # type: ignore
# Lemonade doesn't check the key
key = "lemonade"
key = self._DEFAULT_API_KEY
if passed_api_base is None or api_key:
key = (
api_key
or litellm.lemonade_key
or get_secret_str("LEMONADE_API_KEY")
or self._DEFAULT_API_KEY
)
return api_base, key
def _get_auth_headers(self, api_key: Optional[str]) -> dict:
if api_key is None or api_key == self._DEFAULT_API_KEY:
return {}
return {"Authorization": f"Bearer {api_key}"}
def transform_response(
self,
model: str,

View file

@ -1,4 +1,4 @@
from typing import List, Optional, Union
from typing import Any, List, Optional, Union
import httpx
@ -65,7 +65,8 @@ class OllamaModelInfo(BaseLLMModelInfo):
from litellm.secret_managers.main import get_secret_str
return (
os.environ.get("OLLAMA_API_KEY")
api_key
or os.environ.get("OLLAMA_API_KEY")
or litellm.api_key
or litellm.openai_key
or get_secret_str("OLLAMA_API_KEY")
@ -78,13 +79,31 @@ class OllamaModelInfo(BaseLLMModelInfo):
# env var OLLAMA_API_BASE or default
return api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434"
@classmethod
def get_server_api_base(cls, api_base: Optional[str] = None) -> str:
api_base = cls.get_api_base(api_base).rstrip("/")
for suffix in (
"/api/generate",
"/api/chat",
"/api/embed",
"/api/embeddings",
"/api/show",
"/api/tags",
):
if api_base.endswith(suffix):
return api_base[: -len(suffix)]
return api_base
def get_models(self, api_key=None, api_base: Optional[str] = None) -> List[str]:
"""
List all models available on the Ollama server via /api/tags endpoint.
"""
base = self.get_api_base(api_base)
api_key = self.get_api_key()
passed_api_base = api_base
base = self.get_server_api_base(api_base)
api_key = (
self.get_api_key(api_key) if passed_api_base is None or api_key else None
)
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
names: set[str] = set()
@ -126,6 +145,103 @@ class OllamaModelInfo(BaseLLMModelInfo):
result = sorted(names)
return result
@staticmethod
def _strip_ollama_model_prefix(model: str) -> str:
if model.startswith("ollama/") or model.startswith("ollama_chat/"):
return model.split("/", 1)[1]
return model
@staticmethod
def _is_static_ollama_model(model: str) -> bool:
from litellm import model_cost
stripped_model = OllamaModelInfo._strip_ollama_model_prefix(model)
potential_model_names = {
model,
stripped_model,
"ollama/" + stripped_model,
"ollama_chat/" + stripped_model,
}
model_cost_keys = {key.lower() for key in model_cost}
return any(name.lower() in model_cost_keys for name in potential_model_names)
@staticmethod
def _supports_function_calling(ollama_model_info: dict) -> bool:
_template: str = str(ollama_model_info.get("template", "") or "")
return "tools" in _template.lower()
@staticmethod
def _get_max_tokens(ollama_model_info: dict) -> Optional[int]:
_model_info: dict = ollama_model_info.get("model_info", {})
for key, value in _model_info.items():
if "context_length" in key:
return value
return None
def get_runtime_model_info(
self,
model: str,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> dict[str, Any]:
from litellm import module_level_client
model = self._strip_ollama_model_prefix(model)
passed_api_base = api_base
api_base = self.get_server_api_base(api_base)
api_key = (
self.get_api_key(api_key) if passed_api_base is None or api_key else None
)
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
try:
response = module_level_client.post(
url=f"{api_base}/api/show",
json={"name": model},
headers=headers,
)
response.raise_for_status()
except Exception:
verbose_logger.debug("OllamaError: Could not get model info.")
return {
"key": model,
"litellm_provider": "ollama",
"mode": "chat",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"max_tokens": None,
"max_input_tokens": None,
"max_output_tokens": None,
}
model_info = response.json()
max_tokens = self._get_max_tokens(model_info)
return {
"key": model,
"litellm_provider": "ollama",
"mode": "chat",
"supports_function_calling": self._supports_function_calling(model_info),
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"max_tokens": max_tokens,
"max_input_tokens": max_tokens,
"max_output_tokens": max_tokens,
}
def get_model_info(
self,
model: str,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> Optional[dict[str, Any]]:
if self._is_static_ollama_model(model):
return None
return self.get_runtime_model_info(
model=model, api_base=api_base, api_key=api_key
)
def validate_environment(
self,
headers: dict,

View file

@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional,
from httpx._models import Headers, Response
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
@ -17,19 +17,17 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
)
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUsageBlock
from litellm.types.utils import (
Delta,
GenericStreamingChunk,
ModelInfoBase,
ModelResponse,
ModelResponseStream,
ProviderField,
StreamingChoices,
)
from ..common_utils import OllamaError, _convert_image
from ..common_utils import OllamaError, OllamaModelInfo, _convert_image
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -224,59 +222,18 @@ class OllamaConfig(BaseConfig):
)
def get_model_info(
self, model: str, api_base: Optional[str] = None
) -> ModelInfoBase:
self,
model: str,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> Any:
"""
curl http://localhost:11434/api/show -d '{
"name": "mistral"
}'
"""
if model.startswith("ollama/") or model.startswith("ollama_chat/"):
model = model.split("/", 1)[1]
api_base = (
api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434"
)
api_key = self.get_api_key()
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
try:
response = litellm.module_level_client.post(
url=f"{api_base}/api/show",
json={"name": model},
headers=headers,
)
except Exception as e:
verbose_logger.debug(
"OllamaError: Could not get model info for %s from %s. Error: %s",
model,
api_base,
e,
)
return ModelInfoBase(
key=model,
litellm_provider="ollama",
mode="chat",
input_cost_per_token=0.0,
output_cost_per_token=0.0,
max_tokens=None,
max_input_tokens=None,
max_output_tokens=None,
)
model_info = response.json()
_max_tokens: Optional[int] = self._get_max_tokens(model_info)
return ModelInfoBase(
key=model,
litellm_provider="ollama",
mode="chat",
supports_function_calling=self._supports_function_calling(model_info),
input_cost_per_token=0.0,
output_cost_per_token=0.0,
max_tokens=_max_tokens,
max_input_tokens=_max_tokens,
max_output_tokens=_max_tokens,
return OllamaModelInfo().get_model_info(
model=model, api_base=api_base, api_key=api_key
)
def get_error_class(

View file

@ -376,6 +376,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
returned_tool_calls = guardrailed_inputs.get("tool_calls")
guardrailed_tool_calls: List[Dict[str, Any]] = (
cast(List[Dict[str, Any]], returned_tool_calls)
if isinstance(returned_tool_calls, list)
and len(returned_tool_calls) == len(tool_calls_to_check)
else tool_calls_to_check
)
# Step 3: Map guardrail responses back to original response structure
if guardrailed_texts and texts_to_check:
@ -386,10 +393,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
)
# Step 4: Apply guardrailed tool calls back to response
if tool_calls_to_check:
if guardrailed_tool_calls:
await self._apply_guardrail_responses_to_output_tool_calls(
response=response,
tool_calls=tool_calls_to_check,
tool_calls=guardrailed_tool_calls,
task_mappings=tool_call_task_mappings,
)
@ -748,10 +755,11 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
task_mappings: List[Tuple[int, int]],
) -> None:
"""
Apply guardrailed tool calls back to output response.
Apply guardrailed tool calls back to the output response.
The guardrail may have modified the tool_calls list in place,
so we apply the modified tool calls back to the original response.
The guardrail may return updated tool calls (either mutated in place or as
a new list), so we apply the provided tool calls back to the original
response.
Override this method to customize how tool call responses are applied.
"""

View file

@ -114,5 +114,14 @@
"param_mappings": {
"max_completion_tokens": "max_tokens"
}
},
"tensormesh": {
"base_url": "https://serverless.tensormesh.ai/v1",
"api_key_env": "TENSORMESH_INFERENCE_API_KEY",
"api_base_env": "TENSORMESH_SERVERLESS_BASE_URL",
"base_class": "openai_gpt",
"param_mappings": {
"max_completion_tokens": "max_tokens"
}
}
}

View file

@ -80,31 +80,47 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_params: dict,
) -> str:
"""
Get the Base endpoint for Vertex AI Search API
Get the Base endpoint for Vertex AI Search API.
Branches on whether a `vertex_engine_id` is configured:
- Engine ID present: route through the search app (engine) — required for website,
healthcare, and connector-based data stores. Note the serving config name differs
(`default_serving_config` vs `default_config` for direct data store search).
- Engine ID absent: query the data store directly via `vector_store_id`.
"""
if api_base:
return api_base.rstrip("/")
vertex_location = self.get_vertex_ai_location(litellm_params)
vertex_project = self.get_vertex_ai_project(litellm_params)
collection_id = (
litellm_params.get("vertex_collection_id") or "default_collection"
)
datastore_id = litellm_params.get("vector_store_id")
if not datastore_id:
raise ValueError("vector_store_id is required")
if api_base:
return api_base.rstrip("/")
encoded_collection_id = encode_url_path_segment(
collection_id, field_name="vertex_collection_id"
)
base = (
f"https://discoveryengine.googleapis.com/v1/"
f"projects/{vertex_project}/locations/{vertex_location}/"
f"collections/{encoded_collection_id}"
)
engine_id = litellm_params.get("vertex_engine_id")
if engine_id:
encoded_engine_id = encode_url_path_segment(
engine_id, field_name="vertex_engine_id"
)
return f"{base}/engines/{encoded_engine_id}/servingConfigs/default_serving_config"
datastore_id = litellm_params.get("vector_store_id")
if not datastore_id:
raise ValueError(
"vector_store_id is required when vertex_engine_id is not set"
)
encoded_datastore_id = encode_url_path_segment(
datastore_id, field_name="vector_store_id"
)
# Vertex AI Search API endpoint for search
return (
f"https://discoveryengine.googleapis.com/v1/"
f"projects/{vertex_project}/locations/{vertex_location}/"
f"collections/{encoded_collection_id}/dataStores/{encoded_datastore_id}/servingConfigs/default_config"
)
return f"{base}/dataStores/{encoded_datastore_id}/servingConfigs/default_config"
def transform_search_vector_store_request(
self,

View file

@ -41,6 +41,7 @@ class PartnerModelPrefixes(str, Enum):
MINIMAX_PREFIX = "minimaxai/"
MOONSHOT_PREFIX = "moonshotai/"
ZAI_PREFIX = "zai-org/"
GEMMA_MAAS_PREFIX = "google/gemma-"
class VertexAIPartnerModels(VertexBase):
@ -68,6 +69,7 @@ class VertexAIPartnerModels(VertexBase):
or model.startswith(PartnerModelPrefixes.MINIMAX_PREFIX)
or model.startswith(PartnerModelPrefixes.MOONSHOT_PREFIX)
or model.startswith(PartnerModelPrefixes.ZAI_PREFIX)
or model.startswith(PartnerModelPrefixes.GEMMA_MAAS_PREFIX)
):
return True
return False
@ -82,6 +84,7 @@ class VertexAIPartnerModels(VertexBase):
PartnerModelPrefixes.MINIMAX_PREFIX,
PartnerModelPrefixes.MOONSHOT_PREFIX,
PartnerModelPrefixes.ZAI_PREFIX,
PartnerModelPrefixes.GEMMA_MAAS_PREFIX,
]
if any(provider in model for provider in OPENAI_LIKE_VERTEX_PROVIDERS):
return True

View file

@ -0,0 +1,69 @@
from typing import TYPE_CHECKING, List, Optional, Tuple
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.watsonx.common_utils import IBMWatsonXMixin
if TYPE_CHECKING:
from httpx import URL
class WatsonxPassthroughConfig(IBMWatsonXMixin, BasePassthroughConfig):
"""
Watsonx-specific passthrough configuration.
"""
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
"""Check if request should be streamed"""
return request_data.get("stream", False)
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
endpoint: str,
request_query_params: Optional[dict],
litellm_params: dict,
) -> Tuple["URL", str]:
"""
Construct complete Watsonx URL with version parameter.
This ensures the version parameter is ALWAYS included in the URL,
solving the query parameter issue.
"""
base_target_url = str(self.get_api_base(api_base))
# Use the format_url helper to construct URL with query params
complete_url = self.format_url(
endpoint=endpoint,
base_target_url=base_target_url,
request_query_params=request_query_params,
)
return (complete_url, base_target_url)
@staticmethod
def get_api_base(
api_base: Optional[str] = None,
) -> Optional[str]:
return api_base or IBMWatsonXMixin()._get_base_url(api_base=api_base)
@staticmethod
def get_api_key(
api_key: Optional[str] = None,
) -> Optional[str]:
return (
api_key
or IBMWatsonXMixin.get_watsonx_credentials(
optional_params=dict(), api_base=None, api_key=api_key
)["api_key"]
)
@staticmethod
def get_base_model(model: str) -> Optional[str]:
return model
def get_models(
self, api_key: Optional[str] = None, api_base: Optional[str] = None
) -> List[str]:
return super().get_models(api_key, api_base)

View file

@ -577,7 +577,10 @@
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024
"output_vector_size": 1024,
"provider_specific_entry": {
"bedrock_invocation_schema": "titan_v2"
}
},
"amazon.titan-image-generator-v1": {
"input_cost_per_image": 0.0,
@ -8899,15 +8902,16 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -8920,15 +8924,16 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-east-1/claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9072,15 +9077,16 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9093,15 +9099,16 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-west-1/claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -24894,6 +24901,21 @@
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/ministral-8b-latest": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 1.5e-07,
"source": "https://mistral.ai/pricing",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/mistral-tiny": {
"input_cost_per_token": 2.5e-07,
"litellm_provider": "mistral",
@ -31843,19 +31865,21 @@
"supports_native_structured_output": true
},
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 7.2e-06,
"output_cost_per_token_above_200k_tokens": 2.7e-05,
"cache_creation_input_token_cost_above_200k_tokens": 9.0e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.44e-05,
"cache_read_input_token_cost_above_200k_tokens": 7.2e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -34944,6 +34968,22 @@
"us-central1"
]
},
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "vertex_ai-openai_models",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6e-07,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/maas/google/gemma-4-26b-a4b-it",
"supported_regions": [
"global"
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/openai/gpt-oss-120b-maas": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "vertex_ai-openai_models",
@ -41417,6 +41457,7 @@
},
"bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
"input_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock",
@ -41439,6 +41480,7 @@
},
"bedrock/us-gov-west-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
"input_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock",

View file

@ -67,6 +67,17 @@ def _prepare_mcp_server_data(
# ``alias=None`` is a valid request to clear the stored alias.
if data_dict.get("alias") is None and "alias" not in fields_set:
data_dict.pop("alias", None)
# Prisma ``allowed_tools`` is a required String[]; ``null`` is invalid.
# The UI sends null to clear a whitelist — treat that as ``[]``.
if "allowed_tools" in data_dict and data_dict["allowed_tools"] is None:
data_dict["allowed_tools"] = []
# Json map fields use ``@default("{}")``; explicit null means clear overrides.
for json_map_field in (
"tool_name_to_display_name",
"tool_name_to_description",
):
if json_map_field in data_dict and data_dict[json_map_field] is None:
data_dict[json_map_field] = {}
else:
data_dict = data.model_dump(exclude_none=True)
# Ensure alias is always present in the dict (even if None)
@ -93,13 +104,13 @@ def _prepare_mcp_server_data(
if data_dict.get("env") is not None:
data_dict["env"] = safe_dumps(data_dict["env"])
if data_dict.get("tool_name_to_display_name") is not None:
if "tool_name_to_display_name" in data_dict:
data_dict["tool_name_to_display_name"] = safe_dumps(
data_dict["tool_name_to_display_name"]
data_dict["tool_name_to_display_name"] or {}
)
if data_dict.get("tool_name_to_description") is not None:
if "tool_name_to_description" in data_dict:
data_dict["tool_name_to_description"] = safe_dumps(
data_dict["tool_name_to_description"]
data_dict["tool_name_to_description"] or {}
)
# mcp_access_groups is already List[str], no serialization needed

View file

@ -2429,7 +2429,13 @@ class MCPServerManager:
"""
Check if the tool is allowed or banned for the given server
"""
if server.allowed_tools:
from litellm.proxy._experimental.mcp_server.utils import (
server_applies_tool_allowlist,
)
if server_applies_tool_allowlist(server):
if not server.allowed_tools:
return False
return (
tool_name in server.allowed_tools
or f"{server.name}-{tool_name}" in server.allowed_tools

View file

@ -365,10 +365,9 @@ if MCP_AVAILABLE:
user_api_key_auth=user_api_key_auth,
)
# Filter tools based on allowed_tools configuration
# Only filter if allowed_tools is explicitly configured (not None and not empty)
if server.allowed_tools is not None and len(server.allowed_tools) > 0:
tools = filter_tools_by_allowed_tools(tools, server)
# Always apply allowed_tools/disallowed_tools so the blacklist is
# enforced even when no allowlist is set (matches the SSE/HTTP path).
tools = filter_tools_by_allowed_tools(tools, server)
# Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions
# This provides per-key/team/org control over which tools can be accessed

View file

@ -945,10 +945,16 @@ if MCP_AVAILABLE:
Returns:
Filtered list of tools
"""
from litellm.proxy._experimental.mcp_server.utils import (
server_applies_tool_allowlist,
)
tools_to_return = tools
# Filter by allowed_tools (whitelist)
if mcp_server.allowed_tools:
if server_applies_tool_allowlist(mcp_server):
if not mcp_server.allowed_tools:
return []
tools_to_return = [
tool
for tool in tools

View file

@ -2,6 +2,7 @@
MCP Server Utilities
"""
import json
import re
from typing import Any, Dict, Iterator, Mapping, Optional, Tuple, Union
@ -162,6 +163,36 @@ def lookup_mcp_server_auth_in_headers(
return None
MCP_TOOL_ALLOWLIST_ENFORCED_KEY = "tool_allowlist_enforced"
def _parse_mcp_info_dict(mcp_info: Any) -> Optional[Dict[str, Any]]:
if mcp_info is None:
return None
if isinstance(mcp_info, dict):
return mcp_info
if isinstance(mcp_info, str):
try:
parsed = json.loads(mcp_info)
except (ValueError, TypeError):
return None
return parsed if isinstance(parsed, dict) else None
return None
def is_server_tool_allowlist_enforced(mcp_server: Any) -> bool:
mcp_info = _parse_mcp_info_dict(getattr(mcp_server, "mcp_info", None))
if not mcp_info:
return False
return bool(mcp_info.get(MCP_TOOL_ALLOWLIST_ENFORCED_KEY))
def server_applies_tool_allowlist(mcp_server: Any) -> bool:
"""Whether server-level allowed_tools whitelist filtering is active."""
allowed_tools = getattr(mcp_server, "allowed_tools", None) or []
return is_server_tool_allowlist_enforced(mcp_server) or bool(allowed_tools)
def validate_and_normalize_mcp_server_payload(payload: Any) -> None:
"""
Validate and normalize MCP server payload fields (server_name and alias).

View file

@ -419,6 +419,7 @@ class LiteLLMRoutes(enum.Enum):
"/vllm",
"/mistral",
"/milvus",
"/watsonx",
]
#########################################################
@ -3901,7 +3902,9 @@ class LiteLLM_TeamMembership(LiteLLMPydanticObjectBase):
# Union so Pydantic picks Full when data has server-managed fields
# (/team/info) and Base when callers/tests construct with only
# user-settable fields.
litellm_budget_table: Optional[Union[LiteLLM_BudgetTableFull, LiteLLM_BudgetTable]]
litellm_budget_table: Optional[
Union[LiteLLM_BudgetTableFull, LiteLLM_BudgetTable]
] = None
def safe_get_team_member_rpm_limit(self) -> Optional[int]:
if self.litellm_budget_table is not None:

View file

@ -105,6 +105,8 @@ async def google_stream_generate_content(
if "model" not in data:
data["model"] = model_name
data["stream"] = True
# google-genai SDK (?alt=sse) must not receive OpenAI's data: [DONE] terminator.
data["_litellm_skip_openai_stream_done"] = True
processor = ProxyBaseLLMRequestProcessing(data=data)
try:

View file

@ -0,0 +1,37 @@
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .cato_networks import CatoNetworksGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
from litellm.proxy.guardrails.guardrail_hooks.cato_networks import (
CatoNetworksGuardrail,
)
_cato_callback = CatoNetworksGuardrail(
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
ssl_verify=getattr(litellm_params, "ssl_verify", None),
)
litellm.logging_callback_manager.add_litellm_callback(_cato_callback)
return _cato_callback
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.CATO_NETWORKS.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.CATO_NETWORKS.value: CatoNetworksGuardrail,
}

View file

@ -0,0 +1,635 @@
# +-------------------------------------------------------------+
#
# Use Cato Networks Guardrails for your LLM calls
# https://www.catonetworks.com/
#
# +-------------------------------------------------------------+
import asyncio
import contextlib
import json
import os
import ssl
from typing import TYPE_CHECKING, Any, AsyncGenerator, Optional, Type, Union
from fastapi import HTTPException
from pydantic import BaseModel
from websockets.asyncio.client import ClientConnection, connect
from websockets.exceptions import ConnectionClosed
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
get_ssl_configuration,
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import (
apply_redacted_messages_back,
build_inspection_messages,
)
from litellm.types.utils import (
CallTypesLiteral,
Choices,
EmbeddingResponse,
ImageResponse,
ModelResponse,
ModelResponseStream,
ResponsesAPIResponse,
)
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
class CatoNetworksGuardrailMissingSecrets(Exception):
pass
class CatoNetworksGuardrail(CustomGuardrail):
def __init__(
self, api_key: Optional[str] = None, api_base: Optional[str] = None, **kwargs
):
ssl_verify = kwargs.pop("ssl_verify", None)
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
params={"ssl_verify": ssl_verify} if ssl_verify is not None else None,
)
self.api_key = api_key or os.environ.get("CATO_API_KEY")
if not self.api_key:
msg = (
"Couldn't get Cato Networks api key, either set the `CATO_API_KEY` in the environment or "
"pass it as a parameter to the guardrail in the config file"
)
raise CatoNetworksGuardrailMissingSecrets(msg)
self.api_base = (
api_base
or os.environ.get("CATO_API_BASE")
or "https://api.aisec.catonetworks.com"
)
self.api_base = self.api_base.rstrip("/")
self.ws_api_base = self.api_base.replace("http://", "ws://").replace(
"https://", "wss://"
)
self._ws_connect_ssl_kwargs = self._build_ws_ssl_kwargs(
ssl_verify, self.ws_api_base
)
super().__init__(**kwargs)
@staticmethod
def _build_ws_ssl_kwargs(
ssl_verify: Optional[Union[bool, str]], ws_api_base: str
) -> dict:
"""Resolve the ``ssl`` argument for ``websockets.connect``. Mirrors the
``ssl_verify`` handling applied to the HTTP handler so a custom Cato instance
behind TLS honours the same verification settings for streaming."""
if ssl_verify is None or not ws_api_base.startswith("wss://"):
return {}
ssl_config = get_ssl_configuration(ssl_verify)
if ssl_config is False:
ssl_config = ssl.create_default_context()
ssl_config.check_hostname = False
ssl_config.verify_mode = ssl.CERT_NONE
return {"ssl": ssl_config}
@staticmethod
def _resolve_cato_user_email(user_api_key_dict: UserAPIKeyAuth) -> Optional[str]:
"""Only the key/JWT-bound user email is trusted. ``end_user_id`` is derived from
caller-supplied request fields (OpenAI ``user``, headers, metadata) and is spoofable,
so it must never be forwarded as the Cato user identity."""
return user_api_key_dict.user_email
@staticmethod
async def _cancel_background_task(task: asyncio.Task) -> None:
task.cancel()
with contextlib.suppress(asyncio.CancelledError, Exception):
await task
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: CallTypesLiteral,
) -> Union[Exception, str, dict, None]:
verbose_proxy_logger.debug("Inside Cato Pre-Call Hook")
return await self.call_cato_guardrail(
data,
hook="pre_call",
key_alias=user_api_key_dict.key_alias,
user_email=self._resolve_cato_user_email(user_api_key_dict),
)
async def async_moderation_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
call_type: CallTypesLiteral,
) -> Union[Exception, str, dict, None]:
verbose_proxy_logger.debug("Inside Cato Moderation Hook")
return await self.call_cato_guardrail(
data,
hook="moderation",
key_alias=user_api_key_dict.key_alias,
user_email=self._resolve_cato_user_email(user_api_key_dict),
)
@classmethod
def _inspection_messages(cls, data: dict) -> list:
"""Flatten multimodal list ``content`` into plain text so Cato inspects
every text fragment. Chat ``messages`` stay 1:1 with the request so
redacted results map back by index, and every other field the proxy
forwards to the model (Responses-API ``input``/``instructions``, legacy
completion ``prompt`` and tool/function/``response_format`` schema strings)
is appended as synthetic messages so blocked text cannot bypass inspection
by hiding in one of them."""
flattened = []
for message in data.get("messages") or []:
if isinstance(message, dict) and isinstance(message.get("content"), list):
parts = build_inspection_messages({"messages": [message]})
flattened.append(
{**message, "content": parts[0]["content"] if parts else ""}
)
else:
flattened.append(message)
for _field, messages in cls._extra_inspection_sources(data):
flattened.extend(messages)
return flattened
@staticmethod
def _prompt_inspection_messages(prompt: Any) -> list:
"""Synthetic user messages for a legacy completion ``prompt`` (a string
or a list of string prompts)."""
if isinstance(prompt, str):
return [{"role": "user", "content": prompt}] if prompt else []
if isinstance(prompt, list):
return [
{"role": "user", "content": part}
for part in prompt
if isinstance(part, str) and part
]
return []
@staticmethod
def _iter_schema_string_refs(data: dict):
"""Yield ``(container, key)`` for every non-empty schema string the proxy
forwards to the model inside tool/function and structured-output schemas:
each ``tools[].function`` and legacy ``functions[]`` entry plus the
``response_format`` JSON schema, walked recursively for the free-text and
value strings a caller could hide blocked text in (``description``,
``title``, ``const``, ``default`` and every ``enum``/``examples`` item).
Blocked text in any of them must be inspected and redacted like any other
prompt."""
scalar_keys = ("description", "title", "const", "default")
list_keys = ("enum", "examples")
stack: list = []
for tool in data.get("tools") or []:
if isinstance(tool, dict) and isinstance(tool.get("function"), dict):
stack.append(tool["function"])
for function in data.get("functions") or []:
if isinstance(function, dict):
stack.append(function)
response_format = data.get("response_format")
if isinstance(response_format, dict):
stack.append(response_format)
stack.reverse()
while stack:
node = stack.pop()
if isinstance(node, dict):
for key in scalar_keys:
value = node.get(key)
if isinstance(value, str) and value:
yield node, key
for key in list_keys:
items = node.get(key)
if isinstance(items, list):
for idx, item in enumerate(items):
if isinstance(item, str) and item:
yield items, idx
stack.extend(reversed(list(node.values())))
elif isinstance(node, list):
stack.extend(reversed(node))
@classmethod
def _extra_inspection_sources(cls, data: dict) -> list:
"""Text the proxy forwards to the model outside chat ``messages``:
Responses-API ``input`` and ``instructions``, legacy completion
``prompt`` and tool/function/``response_format`` schema strings. Returned
as ``(field, messages)`` in a fixed order so the anonymize path can slice
redactions back to the field they came from."""
sources: list = []
input_messages = build_inspection_messages({"input": data.get("input")})
if input_messages:
sources.append(("input", input_messages))
instructions = data.get("instructions")
if isinstance(instructions, str) and instructions:
sources.append(
("instructions", [{"role": "system", "content": instructions}])
)
prompt_messages = cls._prompt_inspection_messages(data.get("prompt"))
if prompt_messages:
sources.append(("prompt", prompt_messages))
schema_strings = [
{"role": "system", "content": container[key]}
for container, key in cls._iter_schema_string_refs(data)
]
if schema_strings:
sources.append(("schema_strings", schema_strings))
return sources
async def call_cato_guardrail(
self,
data: dict,
hook: str,
key_alias: Optional[str],
user_email: Optional[str] = None,
) -> dict:
call_id = data.get("litellm_call_id")
headers = self._build_cato_headers(
hook=hook,
key_alias=key_alias,
user_email=user_email,
litellm_call_id=call_id,
)
response = await self.async_handler.post(
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": self._inspection_messages(data)},
)
response.raise_for_status()
res = response.json()
required_action = res.get("required_action")
action_type = required_action and required_action.get("action_type", None)
if action_type is None:
verbose_proxy_logger.debug("Cato: No required action specified")
return data
if action_type == "monitor_action":
verbose_proxy_logger.info("Cato: monitor action")
elif action_type == "block_action":
self._handle_block_action(res.get("analysis_result", {}), required_action)
elif action_type == "anonymize_action":
return self._anonymize_request(res, data)
else:
verbose_proxy_logger.error(f"Cato: {action_type} action")
return data
def _handle_block_action(self, analysis_result: Any, required_action: Any) -> None:
detection_message = required_action.get("detection_message", None)
verbose_proxy_logger.info(
"Cato: Violation detected enabled policies: {policies}".format(
policies=list(analysis_result.get("policy_drill_down", {}).keys()),
),
)
raise HTTPException(status_code=400, detail=detection_message)
def _anonymize_request(self, res: Any, data: dict) -> dict:
verbose_proxy_logger.info("Cato: anonymize action")
redacted_chat = res.get("redacted_chat")
if not redacted_chat:
return data
redacted_messages = redacted_chat.get("all_redacted_messages") or []
original_messages = data.get("messages")
offset = 0
if original_messages:
data["messages"] = [
(
{**original, "content": redacted_messages[idx]["content"]}
if idx < len(redacted_messages)
and redacted_messages[idx].get("content") is not None
else original
)
for idx, original in enumerate(original_messages)
]
offset = len(original_messages)
for field, messages in self._extra_inspection_sources(data):
redacted_slice = redacted_messages[offset : offset + len(messages)]
offset += len(messages)
if redacted_slice:
self._apply_extra_redaction(data, field, redacted_slice)
return data
@classmethod
def _apply_extra_redaction(cls, data: dict, field: str, redacted: list) -> None:
if field == "input":
input_only = {"input": data["input"]}
apply_redacted_messages_back(input_only, redacted)
data["input"] = input_only["input"]
elif field == "instructions":
if redacted[0].get("content") is not None:
data["instructions"] = redacted[0]["content"]
elif field == "prompt":
cls._apply_prompt_redaction(data, redacted)
elif field == "schema_strings":
cls._apply_schema_string_redaction(data, redacted)
@classmethod
def _apply_schema_string_redaction(cls, data: dict, redacted: list) -> None:
redactions = iter(redacted)
for container, key in cls._iter_schema_string_refs(data):
replacement = next(redactions, None)
if replacement is not None and replacement.get("content") is not None:
container[key] = replacement["content"]
@staticmethod
def _apply_prompt_redaction(data: dict, redacted: list) -> None:
contents = [m.get("content") for m in redacted if isinstance(m, dict)]
prompt = data.get("prompt")
if isinstance(prompt, str):
if contents and contents[0] is not None:
data["prompt"] = contents[0]
return
if isinstance(prompt, list):
new_prompt = list(prompt)
redactions = iter(contents)
for idx, part in enumerate(new_prompt):
if isinstance(part, str) and part:
replacement = next(redactions, None)
if replacement is not None:
new_prompt[idx] = replacement
data["prompt"] = new_prompt
async def call_cato_guardrail_on_output(
self,
request_data: dict,
output: str,
hook: str,
key_alias: Optional[str],
user_email: Optional[str] = None,
) -> Optional[dict]:
call_id = request_data.get("litellm_call_id")
inspection_messages = self._inspection_messages(request_data)
assistant_index = len(inspection_messages)
response = await self.async_handler.post(
f"{self.api_base}/fw/v1/analyze",
headers=self._build_cato_headers(
hook=hook,
key_alias=key_alias,
user_email=user_email,
litellm_call_id=call_id,
),
json={
"messages": inspection_messages
+ [{"role": "assistant", "content": output}]
},
)
response.raise_for_status()
res = response.json()
required_action = res.get("required_action")
action_type = required_action and required_action.get("action_type", None)
if action_type and action_type == "block_action":
self._handle_block_action_on_output(
res.get("analysis_result", {}), required_action
)
redacted_chat = res.get("redacted_chat", None)
if action_type and action_type == "anonymize_action" and redacted_chat:
all_redacted = redacted_chat.get("all_redacted_messages") or []
if assistant_index < len(all_redacted):
redacted_output = all_redacted[assistant_index].get("content")
if redacted_output is not None:
return {"redacted_output": redacted_output}
return None
def _handle_block_action_on_output(
self, analysis_result: Any, required_action: Any
) -> None:
detection_message = required_action.get("detection_message", None)
verbose_proxy_logger.info(
"Cato: detected: {detected}, enabled policies: {policies}".format(
detected=True,
policies=list(analysis_result.get("policy_drill_down", {}).keys()),
),
)
raise HTTPException(status_code=400, detail=detection_message)
def _build_cato_headers(
self,
*,
hook: str,
key_alias: Optional[str],
user_email: Optional[str],
litellm_call_id: Optional[str],
):
"""
A helper function to build the http headers that are required by Cato guardrails.
"""
return (
{
"Authorization": f"Bearer {self.api_key}",
# Used by Cato Networks to apply only the guardrails that should be applied in a specific request phase.
"x-cato-litellm-hook": hook,
# Used by Cato Networks to track LiteLLM version and provide backward compatibility.
"x-cato-litellm-version": litellm_version,
}
# Used by Cato Networks to track together single call input and output
| ({"x-cato-call-id": litellm_call_id} if litellm_call_id else {})
# Used by Cato Networks to track guardrails violations by user.
| ({"x-cato-user-email": user_email} if user_email else {})
| (
{
# Used by Cato Networks apply only the guardrails that are associated with the key alias.
"x-cato-gateway-key-alias": key_alias,
}
if key_alias
else {}
)
)
@staticmethod
def _output_fragments(message: Any) -> list:
"""Assistant text the proxy returns to the caller: ``content`` plus every
``tool_calls[].function.arguments`` string, each tagged with where a
redaction must be written back. ``content`` is only included when present
so a tool-call-only choice keeps its ``None`` content (the text-vs-tool-call
signal downstream consumers rely on) while its arguments are still inspected."""
fragments: list = []
if message.content is not None:
fragments.append((("content", None), message.content))
for idx, tool_call in enumerate(message.tool_calls or []):
function = getattr(tool_call, "function", None)
arguments = getattr(function, "arguments", None)
if isinstance(arguments, str) and arguments:
fragments.append((("tool_call", idx), arguments))
return fragments
@staticmethod
def _apply_output_fragment(message: Any, target: tuple, redacted: str) -> None:
kind, idx = target
if kind == "content":
message.content = redacted
else:
message.tool_calls[idx].function.arguments = redacted
@staticmethod
def _responses_output_field(item: Any, key: str) -> Any:
return item.get(key) if isinstance(item, dict) else getattr(item, key, None)
@classmethod
def _responses_output_fragments(cls, response: ResponsesAPIResponse) -> list:
"""Assistant text the Responses API returns to the caller: every
``output_text`` content block plus every function-call ``arguments``
string, each paired with the ``(container, key)`` a Cato redaction is
written back to. Output items and their content may be pydantic objects
or plain dicts, so both access patterns are handled."""
fragments: list = []
for item in response.output or []:
item_type = cls._responses_output_field(item, "type")
if item_type == "function_call":
arguments = cls._responses_output_field(item, "arguments")
if isinstance(arguments, str) and arguments:
fragments.append((item, "arguments", arguments))
elif item_type == "message":
for content in cls._responses_output_field(item, "content") or []:
if cls._responses_output_field(content, "type") != "output_text":
continue
text = cls._responses_output_field(content, "text")
if isinstance(text, str) and text:
fragments.append((content, "text", text))
return fragments
@staticmethod
def _apply_responses_output_fragment(
container: Any, key: str, redacted: str
) -> None:
if isinstance(container, dict):
container[key] = redacted
else:
setattr(container, key, redacted)
async def _inspect_output_text(
self,
data: dict,
text: str,
user_api_key_dict: UserAPIKeyAuth,
user_email: Optional[str],
) -> Optional[str]:
"""Run the Cato output guardrail on a single assistant text fragment.
Raises on a block action and returns the redacted replacement, or
``None`` when the fragment must be left unchanged."""
cato_output_guardrail_result = await self.call_cato_guardrail_on_output(
data,
text,
hook="output",
key_alias=user_api_key_dict.key_alias,
user_email=user_email,
)
if cato_output_guardrail_result:
return cato_output_guardrail_result.get("redacted_output")
return None
async def async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse],
) -> Any:
user_email = self._resolve_cato_user_email(user_api_key_dict)
if isinstance(response, ModelResponse) and response.choices:
for choice in response.choices:
if not isinstance(choice, Choices):
continue
for target, text in self._output_fragments(choice.message):
redacted_output = await self._inspect_output_text(
data, text, user_api_key_dict, user_email
)
if redacted_output is not None:
self._apply_output_fragment(
choice.message, target, redacted_output
)
elif isinstance(response, ResponsesAPIResponse):
for container, key, text in self._responses_output_fragments(response):
redacted_output = await self._inspect_output_text(
data, text, user_api_key_dict, user_email
)
if redacted_output is not None:
self._apply_responses_output_fragment(
container, key, redacted_output
)
return response
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response,
request_data: dict,
) -> AsyncGenerator[ModelResponseStream, None]:
from litellm.proxy.proxy_server import StreamingCallbackError
user_email = self._resolve_cato_user_email(user_api_key_dict)
call_id = request_data.get("litellm_call_id")
async with connect(
f"{self.ws_api_base}/fw/v1/analyze/stream",
additional_headers=self._build_cato_headers(
hook="output",
key_alias=user_api_key_dict.key_alias,
user_email=user_email,
litellm_call_id=call_id,
),
**self._ws_connect_ssl_kwargs,
) as websocket:
sender = asyncio.create_task(
self.forward_the_stream_to_cato(websocket, response)
)
try:
while True:
raw_message = await self._await_cato_message(websocket, sender)
result = json.loads(raw_message)
if verified_chunk := result.get("verified_chunk"):
yield ModelResponseStream.model_validate(verified_chunk)
continue
if result.get("done"):
return
if blocking_message := result.get("blocking_message"):
raise StreamingCallbackError(blocking_message)
verbose_proxy_logger.error(
f"Unknown message received from Cato: {result}"
)
return
finally:
await self._cancel_background_task(sender)
async def _await_cato_message(
self, websocket: ClientConnection, sender: asyncio.Task
) -> Any:
"""Wait for the next Cato message, surfacing a dead forwarding task instead of blocking."""
from litellm.proxy.proxy_server import StreamingCallbackError
recv_task = asyncio.ensure_future(websocket.recv())
pending = {recv_task, sender} if not sender.done() else {recv_task}
await asyncio.wait(pending, return_when=asyncio.FIRST_COMPLETED)
if sender.done() and (sender_exc := sender.exception()) is not None:
await self._cancel_background_task(recv_task)
raise StreamingCallbackError(
"Cato guardrail upstream stream failed"
) from sender_exc
try:
return await recv_task
except ConnectionClosed as exc:
raise StreamingCallbackError(
"Cato guardrail connection closed unexpectedly"
) from exc
async def forward_the_stream_to_cato(
self,
websocket: ClientConnection,
response_iter: AsyncGenerator[Any, None],
) -> None:
async for chunk in response_iter:
if isinstance(chunk, BaseModel):
chunk = chunk.model_dump_json()
elif not isinstance(chunk, (str, bytes)):
chunk = json.dumps(chunk)
await websocket.send(chunk)
await websocket.send(json.dumps({"done": True}))
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.cato_networks import (
CatoNetworksGuardrailConfigModel,
)
return CatoNetworksGuardrailConfigModel

View file

@ -328,7 +328,24 @@ class ContentFilterGuardrail(CustomGuardrail):
return result
@staticmethod
def _resolve_category_file_path(file_path: str) -> str:
def _assert_within_categories_dir(path: str, categories_dir: str) -> None:
"""Raise ValueError if path escapes the categories directory."""
resolved = os.path.realpath(path)
allowed = os.path.realpath(categories_dir)
try:
common = os.path.commonpath([resolved, allowed])
except ValueError:
# commonpath() raises ValueError on Windows when paths span different drives
raise ValueError(
f"Category file path '{path}' is outside the allowed categories directory"
)
if common != allowed:
raise ValueError(
f"Category file path '{path}' is outside the allowed "
f"categories directory '{categories_dir}'"
)
def _resolve_category_file_path(self, file_path: str) -> str:
"""
Resolve a category file path that may be relative.
@ -339,12 +356,17 @@ class ContentFilterGuardrail(CustomGuardrail):
file isn't found.
Resolution order:
1. Return as-is if absolute or already exists.
2. Try joining the full path relative to this module's directory.
1. Return as-is if absolute or already exists (jailed to module dir).
2. Try joining the full path relative to this module's directory (jailed).
3. Progressively strip leading path components and try each suffix
relative to this module's directory (handles paths like
"litellm/proxy/.../policy_templates/file.yaml" by finding the
"policy_templates/file.yaml" suffix that exists).
relative to this module's directory (jailed).
The directory jail can be disabled for deployments that legitimately
store category files outside the package (e.g. mounted volumes) by
setting the environment variable
``LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS=true``. Use only in
trusted environments where the proxy configuration cannot be influenced
by untrusted input.
Args:
file_path: The file path to resolve (absolute or relative).
@ -352,15 +374,33 @@ class ContentFilterGuardrail(CustomGuardrail):
Returns:
The resolved absolute-ish path, or the original path if
resolution fails (caller should check existence).
"""
if os.path.isabs(file_path) or os.path.exists(file_path):
return file_path
Raises:
ValueError: If the resolved path escapes the module directory
and ``LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS`` is not set.
"""
module_dir = os.path.dirname(__file__)
allow_external = (
os.environ.get("LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS", "").lower()
== "true"
)
if os.path.isabs(file_path) or os.path.exists(file_path):
if not allow_external:
self._assert_within_categories_dir(file_path, module_dir)
else:
verbose_proxy_logger.warning(
"LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS is set — "
"skipping directory jail for category_file '%s'",
file_path,
)
return file_path
# Try the full relative path joined to the module directory
candidate = os.path.join(module_dir, file_path)
if os.path.exists(candidate):
if not allow_external:
self._assert_within_categories_dir(candidate, module_dir)
return candidate
# Progressively strip leading components to find a matching suffix
@ -369,8 +409,17 @@ class ContentFilterGuardrail(CustomGuardrail):
suffix = os.path.join(*parts[i:])
candidate = os.path.join(module_dir, suffix)
if os.path.exists(candidate):
if not allow_external:
self._assert_within_categories_dir(candidate, module_dir)
return candidate
# File not found via any resolution strategy — jail the module-relative
# path anyway to reject traversal attempts (e.g. "../../../../etc/passwd")
# regardless of CWD or whether the target file exists.
if not allow_external:
self._assert_within_categories_dir(
os.path.join(module_dir, file_path), module_dir
)
return file_path
def _load_categories(self, categories: List[ContentFilterCategoryConfig]) -> None:
@ -395,6 +444,13 @@ class ContentFilterGuardrail(CustomGuardrail):
)
continue
# Prevent path traversal via category_name (e.g. "../../etc/passwd")
if not re.match(r"^[a-zA-Z0-9_\-]+$", category_name):
verbose_proxy_logger.warning(
f"Category name '{category_name}' contains invalid characters, skipping"
)
continue
enabled = cat_config.get("enabled", True)
action = cat_config.get("action")
severity_threshold = (
@ -411,7 +467,13 @@ class ContentFilterGuardrail(CustomGuardrail):
# Load category file (custom or default)
if custom_file:
category_file_path = self._resolve_category_file_path(custom_file)
try:
category_file_path = self._resolve_category_file_path(custom_file)
except ValueError as e:
verbose_proxy_logger.warning(
f"Category {category_name}: invalid category_file path, skipping. {e}"
)
continue
else:
# Try .yaml first, then .json (e.g. harm_toxic_abuse.json)
yaml_path = os.path.join(categories_dir, f"{category_name}.yaml")

View file

@ -140,7 +140,12 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
self.fallback_on_error = fallback_on_error
self.timeout = timeout
# Coerce defensively. The dashboard UI persists this field as a JSON
# string, and Pydantic extras (the path that splats model_dump into
# this handler) preserve whatever type the user supplied. A string
# value would otherwise reach httpx, which raises TypeError on its
# internal '<=' comparison and surfaces as a misleading api_error.
self.timeout = float(timeout) if timeout is not None else 10.0
# Tri-state: None = not set (default-on for Anthropic), True = explicit on, False = explicit off
self.experimental_use_latest_role_message_only: Optional[bool] = kwargs.get(

View file

@ -0,0 +1,34 @@
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .vigil_guard import VigilGuardGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
_vigil_guard_callback = VigilGuardGuardrail(
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
unreachable_fallback=litellm_params.unreachable_fallback,
timeout=litellm_params.timeout,
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_vigil_guard_callback)
return _vigil_guard_callback
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.VIGIL_GUARD.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.VIGIL_GUARD.value: VigilGuardGuardrail,
}

View file

@ -0,0 +1,485 @@
from json import JSONDecodeError
from typing import (
TYPE_CHECKING,
Any,
Awaitable,
Dict,
List,
Literal,
Optional,
Protocol,
Tuple,
Type,
cast,
)
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import GuardrailRaisedException
from litellm.exceptions import Timeout as LiteLLMTimeout
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
)
from litellm.types.proxy.guardrails.guardrail_hooks.base import (
GuardrailConfigModel,
)
_ANALYZE_ENDPOINT = "/v1/guard/analyze"
_DEFAULT_VIGIL_TIMEOUT = httpx.Timeout(10.0, connect=5.0)
_BLOCK_REASON_MAX_CHARS = 500
_METADATA_STRING_MAX_CHARS = 500
_METADATA_ARRAY_MAX_ITEMS = 10
_VALID_DECISIONS = ("ALLOWED", "SANITIZED", "BLOCKED")
_TRANSIENT_STATUS_CODES = frozenset({429, 502, 503, 504})
_METADATA_ALLOWLIST = (
"model",
"model_group",
"provider",
"region",
"deployment",
"user",
"user_id",
"session_id",
"conversation_id",
"request_id",
"tenant_id",
"org_id",
)
_FallbackMode = Literal["fail_closed", "fail_open"]
class _AsyncPostHandler(Protocol):
def post(
self,
*,
url: str,
headers: Dict[str, str],
json: Dict[str, Any],
timeout: httpx.Timeout,
) -> Awaitable[httpx.Response]: ...
class VigilGuardMissingConfig(ValueError):
pass
class VigilGuardGuardrail(CustomGuardrail):
def __init__(
self,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
unreachable_fallback: Optional[str] = None,
timeout: Optional[float] = None,
async_handler: Optional[_AsyncPostHandler] = None,
**kwargs: Any,
) -> None:
resolved_base = api_base or get_secret_str("VIGIL_GUARD_URL")
if not resolved_base:
raise VigilGuardMissingConfig(
"Vigil Guard api_base is required. Set api_base in the guardrail "
"config or the VIGIL_GUARD_URL environment variable."
)
self.api_base = resolved_base.rstrip("/")
resolved_key = api_key or get_secret_str("VIGIL_GUARD_API_KEY")
if not resolved_key:
raise VigilGuardMissingConfig(
"Vigil Guard api_key is required. Set api_key in the guardrail "
"config or the VIGIL_GUARD_API_KEY environment variable."
)
self.api_key = resolved_key
fallback = (unreachable_fallback or "fail_closed").lower()
self.unreachable_fallback: _FallbackMode = (
"fail_open" if fallback == "fail_open" else "fail_closed"
)
self.timeout: httpx.Timeout = (
_DEFAULT_VIGIL_TIMEOUT
if timeout is None
else httpx.Timeout(timeout, connect=min(timeout, 5.0))
)
self.async_handler: _AsyncPostHandler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
super().__init__(**kwargs)
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
VigilGuardGuardrailConfigModel,
)
return VigilGuardGuardrailConfigModel
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
texts = inputs.get("texts") or []
has_text = any(isinstance(text, str) and text.strip() for text in texts)
tool_call_args = (
self._tool_call_arguments(inputs.get("tool_calls"))
if input_type == "response"
else []
)
if not has_text and not tool_call_args:
return inputs
source = "user_input" if input_type == "request" else "model_output"
metadata = self._collect_metadata(request_data, logging_obj)
result_texts: List[str] = []
for index, text in enumerate(texts):
if not isinstance(text, str) or not text.strip():
result_texts.append(text)
continue
try:
analysis = await self._analyze(
text=text, source=source, metadata=metadata
)
except (
httpx.HTTPError,
LiteLLMTimeout,
JSONDecodeError,
OSError,
) as exc:
return self._handle_backend_failure(
exc,
inputs,
source,
result_texts + list(texts[index:]),
inputs.get("tool_calls"),
)
decision = analysis.get("decision") if isinstance(analysis, dict) else None
if decision not in _VALID_DECISIONS:
verbose_proxy_logger.error(
"Vigil Guard unrecognized decision for guardrail_name=%s "
"source=%s: %r",
self.guardrail_name,
source,
decision,
)
if self.unreachable_fallback == "fail_open":
return self._build_output(
inputs,
result_texts + list(texts[index:]),
inputs.get("tool_calls"),
)
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message="Vigil Guard returned an unrecognized decision.",
should_wrap_with_default_message=False,
)
if decision == "BLOCKED":
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=self._build_block_reason(analysis),
should_wrap_with_default_message=False,
)
if decision == "SANITIZED":
result_texts.append(self._resolve_sanitized_text(text, analysis))
else:
result_texts.append(text)
result_tool_calls = inputs.get("tool_calls")
for tc_index, arguments in tool_call_args:
try:
analysis = await self._analyze(
text=arguments, source=source, metadata=metadata
)
except (
httpx.HTTPError,
LiteLLMTimeout,
JSONDecodeError,
OSError,
) as exc:
return self._handle_backend_failure(
exc, inputs, source, result_texts, result_tool_calls
)
decision = analysis.get("decision") if isinstance(analysis, dict) else None
if decision not in _VALID_DECISIONS:
verbose_proxy_logger.error(
"Vigil Guard unrecognized decision for guardrail_name=%s "
"source=%s: %r",
self.guardrail_name,
source,
decision,
)
if self.unreachable_fallback == "fail_open":
return self._build_output(inputs, result_texts, result_tool_calls)
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message="Vigil Guard returned an unrecognized decision.",
should_wrap_with_default_message=False,
)
if decision == "BLOCKED":
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message=self._build_block_reason(analysis),
should_wrap_with_default_message=False,
)
if decision == "SANITIZED":
result_tool_calls = self._set_tool_call_arguments(
result_tool_calls,
tc_index,
self._resolve_sanitized_text(arguments, analysis),
)
return self._build_output(inputs, result_texts, result_tool_calls)
def _handle_backend_failure(
self,
exc: Exception,
inputs: GenericGuardrailAPIInputs,
source: str,
final_texts: List[Any],
final_tool_calls: Any,
) -> GenericGuardrailAPIInputs:
if self.unreachable_fallback == "fail_open":
verbose_proxy_logger.error(
"Vigil Guard backend failure with fail_open; allowing request "
"unscanned. guardrail_name=%s source=%s error=%s",
self.guardrail_name,
source,
str(exc),
)
return self._build_output(inputs, final_texts, final_tool_calls)
verbose_proxy_logger.error(
"Vigil Guard backend failure with fail_closed; blocking request. "
"guardrail_name=%s source=%s error=%s",
self.guardrail_name,
source,
str(exc),
)
raise GuardrailRaisedException(
guardrail_name=self.guardrail_name,
message="Vigil Guard backend unreachable; request blocked by fail_closed policy.",
should_wrap_with_default_message=False,
) from exc
@staticmethod
def _build_output(
inputs: GenericGuardrailAPIInputs,
final_texts: List[Any],
final_tool_calls: Any,
) -> GenericGuardrailAPIInputs:
# When nothing was changed, return the input shape verbatim so the guardrail
# logs "allow" rather than "mask". When a text or a tool-call argument was
# changed (sanitized), return only the remap-relevant keys and drop
# structured_messages so a stale, unsanitized payload cannot reach the model.
texts_changed = final_texts != (inputs.get("texts") or [])
tool_calls_changed = final_tool_calls != inputs.get("tool_calls")
if not texts_changed and not tool_calls_changed:
return cast(GenericGuardrailAPIInputs, dict(inputs))
guardrailed: GenericGuardrailAPIInputs = {"texts": final_texts}
if "images" in inputs:
guardrailed["images"] = inputs["images"]
if "tools" in inputs:
guardrailed["tools"] = inputs["tools"]
if tool_calls_changed:
guardrailed["tool_calls"] = final_tool_calls
return guardrailed
@staticmethod
def _tool_call_arguments(tool_calls: Any) -> List[Tuple[int, str]]:
pairs: List[Tuple[int, str]] = []
if isinstance(tool_calls, list):
for index, tool_call in enumerate(tool_calls):
function = (
tool_call.get("function") if isinstance(tool_call, dict) else None
)
arguments = (
function.get("arguments") if isinstance(function, dict) else None
)
if isinstance(arguments, str) and arguments.strip():
pairs.append((index, arguments))
return pairs
@staticmethod
def _set_tool_call_arguments(
tool_calls: Any, index: int, arguments: str
) -> List[Any]:
updated = list(tool_calls)
tool_call = dict(updated[index])
function = dict(tool_call.get("function") or {})
function["arguments"] = arguments
tool_call["function"] = function
updated[index] = tool_call
return updated
async def _analyze(
self, text: str, source: str, metadata: Dict[str, Any]
) -> Dict[str, Any]:
payload = {
"text": text,
"source": source,
"mode": "full",
"metadata": metadata,
}
endpoint = f"{self.api_base}{_ANALYZE_ENDPOINT}"
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
response = await self._post_with_retry(endpoint, headers, payload)
return response.json()
async def _post_with_retry(
self, endpoint: str, headers: Dict[str, str], payload: Dict[str, Any]
) -> httpx.Response:
for attempt in range(2):
try:
response = await self.async_handler.post(
url=endpoint,
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
return response
except Exception as exc:
if attempt == 0 and self._is_transient(exc):
verbose_proxy_logger.debug(
"Vigil Guard transient failure; retrying once: %s",
type(exc).__name__,
)
continue
raise
raise AssertionError("unreachable") # pragma: no cover
@staticmethod
def _is_transient(exc: Exception) -> bool:
if isinstance(exc, httpx.HTTPStatusError):
return exc.response.status_code in _TRANSIENT_STATUS_CODES
return isinstance(
exc,
(
httpx.ConnectError,
httpx.ConnectTimeout,
httpx.ReadTimeout,
httpx.RemoteProtocolError,
LiteLLMTimeout,
),
)
@staticmethod
def _build_block_reason(analysis: Dict[str, Any]) -> str:
for key in ("blockMessage", "decisionReason"):
value = analysis.get(key)
if isinstance(value, str) and value.strip():
return value.strip()[:_BLOCK_REASON_MAX_CHARS]
categories = analysis.get("categories")
if isinstance(categories, list):
names = [c for c in categories if isinstance(c, str) and c.strip()]
if names:
return ", ".join(names)[:_BLOCK_REASON_MAX_CHARS]
return "Blocked by policy"
@staticmethod
def _resolve_sanitized_text(original: str, analysis: Dict[str, Any]) -> str:
for key in ("sanitizedText", "outputText"):
value = analysis.get(key)
if isinstance(value, str):
return value
return original
def _collect_metadata(
self, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]
) -> Dict[str, Any]:
sources: List[dict] = []
if isinstance(request_data, dict):
sources.append(request_data)
for nested_key in ("metadata", "litellm_metadata"):
nested = request_data.get(nested_key)
if isinstance(nested, dict):
sources.append(nested)
collected: Dict[str, Any] = {}
for field in _METADATA_ALLOWLIST:
for source in sources:
if field in source and source[field] is not None:
clamped = self._clamp_metadata_value(source[field])
if clamped is not None:
collected[field] = clamped
break
call_id = self._extract_call_id(request_data, logging_obj)
if call_id:
collected["litellm_call_id"] = call_id
return collected
@staticmethod
def _clamp_metadata_value(value: Any) -> Any:
if isinstance(value, bool):
return None
if isinstance(value, str):
return value[:_METADATA_STRING_MAX_CHARS]
if isinstance(value, (int, float)):
return value
if isinstance(value, list):
clamped: List[Any] = []
for item in value[:_METADATA_ARRAY_MAX_ITEMS]:
if isinstance(item, bool):
continue
if isinstance(item, str):
clamped.append(item[:_METADATA_STRING_MAX_CHARS])
elif isinstance(item, (int, float)):
clamped.append(item)
return clamped or None
return None
@staticmethod
def _extract_call_id(
request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]
) -> Optional[str]:
if logging_obj is not None:
call_id = getattr(logging_obj, "litellm_call_id", None)
if isinstance(call_id, str) and call_id:
return call_id
if isinstance(request_data, dict):
call_id = request_data.get("litellm_call_id")
if isinstance(call_id, str) and call_id:
return call_id
metadata = request_data.get("metadata")
if isinstance(metadata, dict):
nested = metadata.get("litellm_call_id")
if isinstance(nested, str) and nested:
return nested
return None

View file

@ -217,7 +217,15 @@ def initialize_panw_prisma_airs(litellm_params, guardrail):
mask_response_content=getattr(litellm_params, "mask_response_content", False),
app_name=getattr(litellm_params, "app_name", None),
fallback_on_error=getattr(litellm_params, "fallback_on_error", "block"),
timeout=float(getattr(litellm_params, "timeout", 10.0)),
# `timeout` is now declared on BaseLitellmParams (Optional[float] = None),
# so the attribute always exists. The Pydantic validator on LitellmParams
# coerces strings to float, but None still means "use handler default" —
# guard against float(None) here.
timeout=(
float(getattr(litellm_params, "timeout", None))
if getattr(litellm_params, "timeout", None) is not None
else 10.0
),
violation_message_template=litellm_params.violation_message_template,
)
litellm.logging_callback_manager.add_litellm_callback(_panw_callback)

View file

@ -17,7 +17,7 @@ Quick summary:
- async_log_success_event() fires on GET /v1/batches/{id} (batch completion)
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union
from fastapi import HTTPException
from pydantic import BaseModel
@ -25,12 +25,13 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import (
_extract_file_access_credentials,
_get_batch_job_input_file_usage,
_get_file_content_as_dictionary,
_get_models_from_batch_input_file_content,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -97,6 +98,276 @@ class _PROXY_BatchRateLimiter(CustomLogger):
"""
self.internal_usage_cache = internal_usage_cache
self.parallel_request_limiter = parallel_request_limiter
self._warned_unsupported_model_skip = False
def _get_file_bound_batch_model(self, data: Dict) -> Optional[str]:
"""Resolve the model bound to the batch input file ID.
``create_batch`` routes a file-bound id (model-embedded ``file-...`` or
unified managed file) on that bound model and ignores the top-level
``model``, so this is the authoritative routing model whenever the file
binds one. The provider is then read from that deployment's trusted
credentials for the provider-level skip decision.
"""
input_file_id = data.get("input_file_id")
if not isinstance(input_file_id, str) or not input_file_id:
return None
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
decode_model_from_file_id,
get_models_from_unified_file_id,
)
model_from_file_id = decode_model_from_file_id(input_file_id)
if model_from_file_id:
return model_from_file_id
unified_file_id = _is_base64_encoded_unified_file_id(input_file_id)
if unified_file_id:
target_model_names = get_models_from_unified_file_id(unified_file_id)
if target_model_names:
return target_model_names[0]
return None
def _get_batch_routing_model(self, data: Dict) -> Optional[str]:
"""Resolve the deployment/model used for this batch from request data.
Mirrors ``create_batch`` routing precedence: a model bound to the input
file id wins over the top-level ``model``, because the batch endpoint
ignores the top-level model for file-bound ids. Resolving the provider
skip from the top-level model first would let a caller point ``model``
at a skip-listed provider while the file routes a rate-limited one.
"""
file_bound_model = self._get_file_bound_batch_model(data)
if file_bound_model:
return file_bound_model
model = data.get("model")
if isinstance(model, str) and model:
return model
return None
def _resolve_batch_provider(self, batch_model: Optional[str]) -> Optional[str]:
"""Resolve the provider from the deployment that serves ``batch_model``.
The provider is read from trusted router credentials rather than the
user-supplied ``custom_llm_provider`` request field, so a caller cannot
spoof a skip-listed provider to bypass batch rate limiting.
"""
if not batch_model:
return None
from litellm.proxy.openai_files_endpoints.common_utils import (
get_credentials_for_model,
)
from litellm.proxy.proxy_server import llm_router
if llm_router is None:
return None
try:
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=batch_model,
operation_context="batch input file read (rate limiting)",
)
except HTTPException:
return None
provider = credentials.get("custom_llm_provider")
return provider if isinstance(provider, str) and provider else None
def _create_batch_rate_limit_descriptors(
self,
user_api_key_dict: UserAPIKeyAuth,
data: Dict,
) -> List["RateLimitDescriptor"]:
return self.parallel_request_limiter._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data=data,
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
def _should_skip_batch_input_file_processing(
self,
data: Dict,
user_api_key_dict: UserAPIKeyAuth,
) -> Tuple[bool, Optional[List["RateLimitDescriptor"]]]:
"""
Skip downloading batch input files when the operator disabled batch
input-file rate limiting, when the batch runs entirely on a skip-listed
provider, or when there is nothing to enforce (no applicable rate
limits).
A skip is only honored for keys with unrestricted model access. When
the key has a model allowlist, the JSONL must still be downloaded so
``_enforce_batch_file_model_access`` can validate every ``body.model``
entry, otherwise a restricted key could smuggle unauthorized models
into the file via an admin-configured skip.
The skip is never keyed on a specific model name. The models a batch
actually runs are its JSONL ``body.model`` entries, and any model
identifier the caller can influence (the top-level ``model`` or the
unsigned model embedded in a ``file-...`` id) can be pointed at a
skip-listed deployment while the file routes a different, rate-limited
model. The provider skip is safe because the provider is read from the
routing deployment's trusted credentials and the batch is constrained
to run on that provider.
Returns ``(should_skip, descriptors)`` where ``descriptors`` is the
rate-limit descriptor list computed for the no-limits check, so the
caller can reuse it for counter enforcement without recomputing.
"""
from litellm.proxy.proxy_server import general_settings
self._warn_if_unsupported_model_skip_configured(general_settings)
if self._key_requires_batch_model_access_check(user_api_key_dict):
return False, None
if general_settings.get("disable_batch_input_file_rate_limiting") is True:
return True, None
skip_providers = (
general_settings.get("skip_batch_input_file_rate_limiting_for_providers")
or []
)
if skip_providers:
batch_provider = self._resolve_batch_provider(
self._get_batch_routing_model(data)
)
if batch_provider and batch_provider in skip_providers:
verbose_proxy_logger.debug(
f"Skipping batch input file processing for provider={batch_provider}"
)
return True, None
descriptors = self._create_batch_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data=data,
)
if not self._has_applicable_batch_rate_limits(descriptors):
verbose_proxy_logger.debug(
"Skipping batch input file processing: no rate limits configured"
)
return True, None
return False, descriptors
def _warn_if_unsupported_model_skip_configured(
self, general_settings: Dict
) -> None:
"""Warn once that ``skip_batch_input_file_rate_limiting_for_models`` is a no-op.
A per-model skip is intentionally not honored because the model a batch
runs on is caller-influenced and can be pointed at a skip-listed
deployment while the JSONL routes a different, rate-limited model.
"""
if self._warned_unsupported_model_skip:
return
if general_settings.get("skip_batch_input_file_rate_limiting_for_models"):
self._warned_unsupported_model_skip = True
verbose_proxy_logger.warning(
"general_settings.skip_batch_input_file_rate_limiting_for_models is not "
"supported and has no effect. Use "
"skip_batch_input_file_rate_limiting_for_providers or "
"disable_batch_input_file_rate_limiting instead."
)
@staticmethod
def _key_requires_batch_model_access_check(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""True when the key may only call a subset of models (JSONL must be checked)."""
models = user_api_key_dict.models or []
if "*" in models:
return False
if SpecialModelNames.all_proxy_models.value in models:
return False
if user_api_key_dict.access_group_ids:
return True
if not models:
return False
return True
@staticmethod
def _has_applicable_batch_rate_limits(
descriptors: List["RateLimitDescriptor"],
) -> bool:
for descriptor in descriptors:
rate_limit = descriptor.get("rate_limit") or {}
if (
rate_limit.get("requests_per_unit") is not None
or rate_limit.get("tokens_per_unit") is not None
or rate_limit.get("max_parallel_requests") is not None
):
return True
return False
def _resolve_batch_input_file_fetch_params(
self,
file_id: str,
custom_llm_provider: str,
data: Dict,
) -> Tuple[str, Dict[str, Any]]:
"""
Map proxy-facing file IDs to provider file IDs and credentials.
Model-embedded IDs (``file-<base64>``) are not unified managed-file IDs;
without decoding them, ``afile_content`` is called with the encoded ID
and the upstream provider returns 404.
"""
from litellm.proxy.openai_files_endpoints.common_utils import (
decode_model_from_file_id,
get_credentials_for_model,
get_original_file_id,
)
from litellm.proxy.proxy_server import llm_router
fetch_kwargs: Dict[str, Any] = {
"custom_llm_provider": custom_llm_provider,
}
model_from_file_id = decode_model_from_file_id(file_id)
if model_from_file_id:
if llm_router is not None:
try:
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=model_from_file_id,
operation_context="batch input file read (rate limiting)",
)
fetch_kwargs.update(_extract_file_access_credentials(credentials))
fetch_kwargs["model"] = model_from_file_id
provider = credentials.get("custom_llm_provider")
if provider:
fetch_kwargs["custom_llm_provider"] = provider
except HTTPException:
pass
return get_original_file_id(file_id), fetch_kwargs
request_model = data.get("model")
if isinstance(request_model, str) and request_model and llm_router is not None:
try:
credentials = get_credentials_for_model(
llm_router=llm_router,
model_id=request_model,
operation_context="batch input file read (rate limiting)",
)
fetch_kwargs.update(_extract_file_access_credentials(credentials))
fetch_kwargs["model"] = request_model
provider = credentials.get("custom_llm_provider")
if provider:
fetch_kwargs["custom_llm_provider"] = provider
except HTTPException:
pass
return file_id, fetch_kwargs
def _raise_rate_limit_error(
self,
@ -163,6 +434,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
user_api_key_dict: UserAPIKeyAuth,
data: Dict,
batch_usage: BatchFileUsage,
descriptors: Optional[List["RateLimitDescriptor"]] = None,
) -> None:
"""
Atomically check + increment rate-limit counters by the batch amounts.
@ -171,14 +443,15 @@ class _PROXY_BatchRateLimiter(CustomLogger):
case no counter is modified. Backed by `atomic_check_and_increment_by_n`
which uses a Redis Lua script when available (multi-process atomic) and
falls back to a per-process asyncio.Lock + in-memory operation.
``descriptors`` may be passed in by the pre-call hook to reuse the list
already computed when deciding whether to skip file processing.
"""
descriptors = self.parallel_request_limiter._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data=data,
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
if descriptors is None:
descriptors = self._create_batch_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data=data,
)
increment: Dict[Literal["requests", "tokens"], int] = {
"requests": batch_usage.request_count,
@ -211,6 +484,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
data: Optional[Dict] = None,
) -> BatchFileUsage:
"""
Count number of requests and tokens in a batch input file.
@ -238,14 +512,27 @@ class _PROXY_BatchRateLimiter(CustomLogger):
user_api_key_dict=user_api_key_dict,
)
else:
provider_file_id, fetch_kwargs = (
self._resolve_batch_input_file_fetch_params(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
data=data or {},
)
)
# For non-managed files, use the standard litellm.afile_content
file_content = await litellm.afile_content(
file_id=file_id,
custom_llm_provider=custom_llm_provider,
file_id=provider_file_id,
user_api_key_dict=user_api_key_dict,
**fetch_kwargs,
)
file_content_as_dict = _get_file_content_as_dictionary(file_content.content)
file_content_bytes = getattr(file_content, "content", None)
if not isinstance(file_content_bytes, bytes):
raise ValueError(
f"Expected bytes content from file retrieval for {file_id}, "
f"got {type(file_content_bytes)}"
)
file_content_as_dict = _get_file_content_as_dictionary(file_content_bytes)
# Validate every model named in the batch JSONL against the
# caller's per-key model allowlist. Without this, a caller
@ -441,6 +728,14 @@ class _PROXY_BatchRateLimiter(CustomLogger):
)
return data
should_skip, batch_rate_limit_descriptors = (
self._should_skip_batch_input_file_processing(
data=data, user_api_key_dict=user_api_key_dict
)
)
if should_skip:
return data
# Get custom_llm_provider for token counting
custom_llm_provider = data.get("custom_llm_provider", "openai")
@ -452,6 +747,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
file_id=input_file_id,
custom_llm_provider=custom_llm_provider,
user_api_key_dict=user_api_key_dict,
data=data,
)
verbose_proxy_logger.debug(
@ -469,6 +765,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
user_api_key_dict=user_api_key_dict,
data=data,
batch_usage=batch_usage,
descriptors=batch_rate_limit_descriptors,
)
verbose_proxy_logger.debug(

View file

@ -2433,3 +2433,89 @@ def create_generic_websocket_passthrough_endpoint(
_forward_headers=forward_headers,
cost_per_request=cost_per_request,
)
@router.api_route(
"/watsonx/{endpoint:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
tags=["Watsonx Pass-through", "pass-through"],
)
async def watsonx_proxy_route(
endpoint: str,
request: Request,
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Watsonx pass-through endpoint.
Allows using Watsonx APIs with automatic IAM token management and version parameter injection.
Example:
POST /watsonx/ml/v1/text/tokenization
POST /watsonx/ml/v1/text/generation
"""
# Direct passthrough with WatsonxPassthroughConfig
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
provider_config = ProviderConfigManager.get_provider_passthrough_config(
provider=LlmProviders.WATSONX,
model="",
)
if provider_config is None:
raise HTTPException(
status_code=404, detail="Watsonx passthrough config not found"
)
# Get complete URL with version parameter
complete_url, _ = provider_config.get_complete_url(
api_base=None,
api_key=None,
model="",
endpoint=endpoint,
request_query_params=None,
litellm_params={},
)
# Get auth headers with IAM token
auth_headers = provider_config.validate_environment(
headers={},
model="",
messages=[],
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
# Check for streaming
is_streaming_request = False
if request.method == "POST":
if "multipart/form-data" not in request.headers.get("content-type", ""):
_request_body = await request.json()
else:
_request_body = await get_form_data(request)
if _request_body.get("stream"):
is_streaming_request = True
request_query_params = dict(request.query_params)
if request_query_params.get("version") is None:
request_query_params["version"] = litellm.WATSONX_DEFAULT_API_VERSION
# Create pass-through endpoint
endpoint_func = create_pass_through_route(
endpoint=endpoint,
target=str(complete_url),
custom_headers=auth_headers,
is_streaming_request=is_streaming_request,
custom_llm_provider="watsonx",
query_params=request_query_params,
)
return await endpoint_func(
request,
fastapi_response,
user_api_key_dict,
)

View file

@ -7072,11 +7072,12 @@ async def async_data_generator( # noqa: PLR0915
# still flush their post-stream logging.
ProxyLogging._fire_deferred_stream_logging(request_data)
# Streaming is done, yield the [DONE] chunk
if error_message is not None:
yield error_message
done_message = "[DONE]"
yield f"data: {done_message}\n\n"
# OpenAI-compatible streams terminate with data: [DONE]; Google GenAI (?alt=sse) does not.
if not request_data.get("_litellm_skip_openai_stream_done"):
done_message = "[DONE]"
yield f"data: {done_message}\n\n"
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.proxy_server.async_data_generator(): Exception occured - {}".format(

View file

@ -36,9 +36,18 @@ router = APIRouter()
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def spend_key_fn():
async def spend_key_fn(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
View all keys created, ordered by spend
View keys created, ordered by spend.
- Admin callers (PROXY_ADMIN / PROXY_ADMIN_VIEW_ONLY) see every key in
the database.
- All other callers (INTERNAL_USER / INTERNAL_USER_VIEW_ONLY, etc.) are
scoped to keys they own (``user_id == caller``). A caller with no
``user_id`` has no scope and receives an empty list rather than the
full table.
Example Request:
```
@ -55,8 +64,17 @@ async def spend_key_fn():
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
)
key_info = await prisma_client.get_data(table_name="key", query_type="find_all")
return key_info
if _is_admin_view_safe(user_api_key_dict=user_api_key_dict):
return await prisma_client.get_data(table_name="key", query_type="find_all")
caller_user_id = user_api_key_dict.user_id
if not caller_user_id:
return []
return await prisma_client.get_data(
table_name="key",
query_type="find_all",
user_id=caller_user_id,
)
except Exception as e:
raise HTTPException(
@ -85,9 +103,19 @@ async def spend_user_fn(
default=None,
description="Get User Table row for user_id",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
View all users created, ordered by spend
View users created, ordered by spend.
- Admin callers (PROXY_ADMIN / PROXY_ADMIN_VIEW_ONLY) see every user, or
a specific user when ``user_id`` is supplied.
- All other callers may only read their own row. If they supply a
``user_id`` query parameter that does not match their authenticated
``user_id`` the request is rejected with HTTP 403; supplying their
own id (or none at all) returns just their row. A caller with no
``user_id`` on their key has no scope and receives an empty list
rather than the full table.
Example Request:
```
@ -109,6 +137,17 @@ async def spend_user_fn(
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
)
if not _is_admin_view_safe(user_api_key_dict=user_api_key_dict):
caller_user_id = user_api_key_dict.user_id
if not caller_user_id:
return []
if user_id is not None and user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": "Not authorized to view spend for another user."},
)
user_id = caller_user_id
if user_id is not None:
user_info = await prisma_client.get_data(
table_name="user", query_type="find_unique", user_id=user_id
@ -123,6 +162,8 @@ async def spend_user_fn(
_strip_password_from_users(result)
return result
except HTTPException:
raise
except Exception as e:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,

View file

@ -41,6 +41,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import (
from litellm.types.proxy.guardrails.guardrail_hooks.qohash import (
QostodianNexusConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
VigilGuardGuardrailConfigModel,
)
"""
Pydantic object defining how to set guardrails on litellm proxy
@ -67,6 +70,7 @@ class SupportedGuardrailIntegrations(Enum):
HIDE_SECRETS = "hide-secrets"
HIDDENLAYER = "hiddenlayer"
AIM = "aim"
CATO_NETWORKS = "cato_networks"
PANGEA = "pangea"
CROWDSTRIKE_AIDR = "crowdstrike_aidr"
LASSO = "lasso"
@ -102,6 +106,7 @@ class SupportedGuardrailIntegrations(Enum):
LLM_AS_A_JUDGE = "llm_as_a_judge"
QOSTODIAN_NEXUS = "qostodian_nexus"
RUBRIK = "rubrik"
VIGIL_GUARD = "vigil_guard"
class Role(Enum):
@ -757,6 +762,15 @@ class BaseLitellmParams(
description="Python-like code containing the apply_guardrail function for custom guardrail logic",
)
timeout: Optional[float] = Field(
default=None,
description=(
"Per-request timeout for the guardrail provider API call (seconds). "
"Accepts int, float, or numeric string; coerced to float on load. "
"Each guardrail handler chooses its own default when unset."
),
)
model_config = ConfigDict(extra="allow", protected_namespaces=())
@ -790,6 +804,7 @@ class LitellmParams(
BlockCodeExecutionGuardrailConfigModel,
HiddenlayerGuardrailConfigModel,
QostodianNexusConfigModel,
VigilGuardGuardrailConfigModel,
):
guardrail: str = Field(description="The type of guardrail integration to use")
mode: Union[str, List[str], Mode] = Field(
@ -813,6 +828,18 @@ class LitellmParams(
return [x.lower() if isinstance(x, str) else x for x in v]
return v
@field_validator("timeout", mode="before", check_fields=False)
@classmethod
def coerce_timeout(cls, v):
"""Accept string-valued timeouts (dashboard UI sends JSON strings)
and coerce to float before any handler reads the value."""
if v is None or v == "":
return None
try:
return float(v)
except (TypeError, ValueError) as e:
raise ValueError(f"timeout must be numeric, got {v!r}") from e
def __init__(self, **kwargs):
default_on = kwargs.pop("default_on", None)
if default_on is not None:

View file

@ -238,6 +238,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_cache_hits_metric",
"litellm_cache_misses_metric",
"litellm_cached_tokens_metric",
# Provider prompt-caching metrics (e.g. OpenAI/Anthropic/Bedrock/Gemini)
"litellm_provider_cache_read_input_tokens_metric",
"litellm_provider_cache_creation_input_tokens_metric",
"litellm_deployment_tpm_limit",
"litellm_deployment_rpm_limit",
"litellm_remaining_api_key_requests_for_model",
@ -655,6 +658,10 @@ class PrometheusMetricLabels:
litellm_cache_misses_metric = _cache_metric_labels
litellm_cached_tokens_metric = _cache_metric_labels
# Provider prompt-caching metrics - track tokens read/written to provider caches
litellm_provider_cache_read_input_tokens_metric = _cache_metric_labels
litellm_provider_cache_creation_input_tokens_metric = _cache_metric_labels
# Metrics whose emission paths supply org context (used by get_labels)
_org_label_metrics: ClassVar[frozenset] = frozenset(
{
@ -672,7 +679,6 @@ class PrometheusMetricLabels:
"litellm_output_tokens_metric",
}
)
# Managed batch metrics
_batch_user_labels = [
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,

View file

@ -0,0 +1,20 @@
from typing import Optional
from pydantic import Field
from .base import GuardrailConfigModel
class CatoNetworksGuardrailConfigModel(GuardrailConfigModel):
api_key: Optional[str] = Field(
default=None,
description="The API key for the Cato Networks guardrail. If not provided, the `CATO_API_KEY` environment variable is checked.",
)
api_base: Optional[str] = Field(
default=None,
description="The API base for the Cato Networks guardrail. Default is https://api.aisec.catonetworks.com. Also checks if the `CATO_API_BASE` environment variable is set.",
)
@staticmethod
def ui_friendly_name() -> str:
return "Cato Networks Guardrail"

View file

@ -0,0 +1,26 @@
from typing import Optional
from pydantic import Field
from .base import GuardrailConfigModel
class VigilGuardGuardrailConfigModel(GuardrailConfigModel):
api_base: Optional[str] = Field(
default=None,
description=(
"Vigil Guard API base URL. "
"Falls back to the VIGIL_GUARD_URL environment variable."
),
)
api_key: Optional[str] = Field(
default=None,
description=(
"Vigil Guard API key. "
"Falls back to the VIGIL_GUARD_API_KEY environment variable."
),
)
@staticmethod
def ui_friendly_name() -> str:
return "Vigil Guard"

View file

@ -3180,6 +3180,7 @@ all_litellm_params = (
"allowed_openai_params",
"litellm_session_id",
"use_litellm_proxy",
"use_chat_completions_api",
"prompt_label",
"shared_session",
"search_tool_name",
@ -3364,6 +3365,7 @@ class LlmProviders(str, Enum):
POE = "poe"
CHUTES = "chutes"
XIAOMI_MIMO = "xiaomi_mimo"
TENSORMESH = "tensormesh"
LITELLM_AGENT = "litellm_agent"
CURSOR = "cursor"
BEDROCK_MANTLE = "bedrock_mantle"

View file

@ -5443,7 +5443,7 @@ def _invalidate_model_cost_lowercase_map() -> None:
_model_cost_mutation_generation += 1
# Clear LRU caches that depend on model_cost data
get_model_info.cache_clear()
_cached_get_model_info.cache_clear()
_cached_get_model_info_helper.cache_clear()
@ -5680,7 +5680,9 @@ def _cached_get_model_info_helper(
Speed Optimization to hit high RPS
"""
return _get_model_info_helper(
model=model, custom_llm_provider=custom_llm_provider, api_base=api_base
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
)
@ -5720,6 +5722,7 @@ def _get_model_info_helper( # noqa: PLR0915
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> ModelInfoBase:
"""
Helper for 'get_model_info'. Separated out to avoid infinite loop caused by returning 'supported_openai_param's
@ -5754,6 +5757,31 @@ def _get_model_info_helper( # noqa: PLR0915
split_model = potential_model_names["split_model"]
custom_llm_provider = potential_model_names["custom_llm_provider"]
#########################
provider_config: Optional[BaseLLMModelInfo] = None
if custom_llm_provider and custom_llm_provider in LlmProvidersSet:
provider_config = ProviderConfigManager.get_provider_model_info(
model=model, provider=LlmProviders(custom_llm_provider)
)
if provider_config is not None:
provider_get_model_info = getattr(provider_config, "get_model_info", None)
if callable(provider_get_model_info):
try:
provider_model_info = provider_get_model_info(
model=model,
api_base=api_base,
api_key=api_key,
)
if provider_model_info is not None:
return provider_model_info
except Exception as e:
verbose_logger.warning(
"Could not get dynamic model info for model=%s, provider=%s; "
"falling back to the static cost map: %s",
model,
custom_llm_provider,
e,
)
if custom_llm_provider == "huggingface":
max_tokens = _get_max_position_embeddings(model_name=model)
return ModelInfoBase(
@ -5774,10 +5802,6 @@ def _get_model_info_helper( # noqa: PLR0915
supports_computer_use=None,
supports_pdf_input=None,
)
elif (
custom_llm_provider == "ollama" or custom_llm_provider == "ollama_chat"
) and not _is_potential_model_name_in_model_cost(potential_model_names):
return litellm.OllamaConfig().get_model_info(model, api_base=api_base)
else:
"""
Check if: (in order of specificity)
@ -6064,11 +6088,53 @@ def _get_model_info_helper( # noqa: PLR0915
)
def _build_model_info(
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> ModelInfo:
supported_openai_params = litellm.get_supported_openai_params(
model=model, custom_llm_provider=custom_llm_provider
)
_model_info = _get_model_info_helper(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
)
provider_info = get_provider_info(
model=model, custom_llm_provider=custom_llm_provider
)
if provider_info:
for key, value in provider_info.items():
if value is not None:
_model_info[key] = value # type: ignore
# if verbose_logger.isEnabledFor(logging.DEBUG):
# verbose_logger.debug(f"model_info: {_model_info}")
return ModelInfo(**_model_info, supported_openai_params=supported_openai_params)
@lru_cache(maxsize=DEFAULT_MAX_LRU_CACHE_SIZE)
def _cached_get_model_info(
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
) -> ModelInfo:
return _build_model_info(
model=model, custom_llm_provider=custom_llm_provider, api_base=api_base
)
def get_model_info(
model: str,
custom_llm_provider: Optional[str] = None,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> ModelInfo:
"""
Get a dict for the maximum tokens (context window), input_cost_per_token, output_cost_per_token for a given model.
@ -6140,32 +6206,15 @@ def get_model_info(
"supported_openai_params": ["temperature", "max_tokens", "top_p", "frequency_penalty", "presence_penalty"]
}
"""
supported_openai_params = litellm.get_supported_openai_params(
model=model, custom_llm_provider=custom_llm_provider
)
# api_key is a per-caller credential, not part of the model identity, so it is
# kept out of the cache key; explicit keys are resolved without the cache.
if api_key is not None:
return _build_model_info(model, custom_llm_provider, api_base, api_key)
return _cached_get_model_info(model, custom_llm_provider, api_base)
_model_info = _get_model_info_helper(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
)
provider_info = get_provider_info(
model=model, custom_llm_provider=custom_llm_provider
)
if provider_info:
for key, value in provider_info.items():
if value is not None:
_model_info[key] = value # type: ignore
# if verbose_logger.isEnabledFor(logging.DEBUG):
# verbose_logger.debug(f"model_info: {_model_info}")
returned_model_info = ModelInfo(
**_model_info, supported_openai_params=supported_openai_params
)
return returned_model_info
get_model_info.cache_clear = _cached_get_model_info.cache_clear # type: ignore[attr-defined]
get_model_info.cache_info = _cached_get_model_info.cache_info # type: ignore[attr-defined]
def json_schema_type(python_type_name: str):
@ -8936,6 +8985,12 @@ class ProviderConfigManager:
)
return AzurePassthroughConfig()
elif LlmProviders.WATSONX == provider:
from litellm.llms.watsonx.passthrough.transformation import (
WatsonxPassthroughConfig,
)
return WatsonxPassthroughConfig()
return None
@staticmethod

View file

@ -577,7 +577,10 @@
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024
"output_vector_size": 1024,
"provider_specific_entry": {
"bedrock_invocation_schema": "titan_v2"
}
},
"amazon.titan-image-generator-v1": {
"input_cost_per_image": 0.0,
@ -8899,15 +8902,16 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -8920,15 +8924,16 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-east-1/claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9072,15 +9077,16 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -9093,15 +9099,16 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-west-1/claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"litellm_provider": "bedrock",
"max_input_tokens": 200000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -24702,6 +24709,21 @@
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/ministral-8b-latest": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "mistral",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 1.5e-07,
"source": "https://mistral.ai/pricing",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"mistral/mistral-tiny": {
"input_cost_per_token": 2.5e-07,
"litellm_provider": "mistral",
@ -31718,19 +31740,21 @@
"supports_native_structured_output": true
},
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
"input_cost_per_token": 3.6e-06,
"input_cost_per_token_above_200k_tokens": 7.2e-06,
"output_cost_per_token_above_200k_tokens": 2.7e-05,
"cache_creation_input_token_cost_above_200k_tokens": 9.0e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.44e-05,
"cache_read_input_token_cost_above_200k_tokens": 7.2e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"output_cost_per_token": 1.8e-05,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -34828,6 +34852,22 @@
"us-central1"
]
},
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "vertex_ai-openai_models",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 6e-07,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/maas/google/gemma-4-26b-a4b-it",
"supported_regions": [
"global"
],
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/openai/gpt-oss-120b-maas": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "vertex_ai-openai_models",
@ -41301,6 +41341,7 @@
},
"bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
"input_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock",
@ -41323,6 +41364,7 @@
},
"bedrock/us-gov-west-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
"input_cost_per_token": 1.2e-06,
"litellm_provider": "bedrock",

View file

@ -2079,6 +2079,24 @@
"a2a": false
}
},
"tensormesh": {
"display_name": "Tensormesh (`tensormesh`)",
"url": "https://docs.litellm.ai/docs/providers/tensormesh",
"endpoints": {
"chat_completions": true,
"messages": true,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false,
"text_completion": true
}
},
"text-completion-codestral": {
"display_name": "Text Completion Codestral (`text-completion-codestral`)",
"url": "https://docs.litellm.ai/docs/providers/codestral",

View file

@ -97,7 +97,7 @@ async def test_gemini_3_responses_api_with_thought_signatures():
pytest.skip("GEMINI_API_KEY not set")
litellm.set_verbose = False
request_model = "gemini/gemini-3-pro-preview"
request_model = "gemini/gemini-3.1-pro-preview"
tools = [
{
@ -197,7 +197,7 @@ async def test_gemini_3_responses_api_streaming_with_thought_signatures():
pytest.skip("GEMINI_API_KEY not set")
litellm.set_verbose = False
request_model = "gemini/gemini-3-pro-preview"
request_model = "gemini/gemini-3.1-pro-preview"
tools = [
{

View file

@ -1862,9 +1862,11 @@ async def test_get_tools_for_single_server():
)
from mcp.types import Tool as MCPTool
# Create a mock server
# Create a mock server (pin allowlist fields; MagicMock auto-attrs are truthy)
mock_server = MagicMock()
mock_server.mcp_info = {"server_name": "zapier"}
mock_server.allowed_tools = None
mock_server.disallowed_tools = None
# Create mock tools
mock_tools = [
@ -1899,6 +1901,44 @@ async def test_get_tools_for_single_server():
assert result[0].mcp_info == {"server_name": "zapier"}
@pytest.mark.asyncio
async def test_get_tools_for_single_server_applies_disallowed_tools_without_allowlist():
"""REST listing must honor disallowed_tools even when no allowlist is set."""
from litellm.proxy._experimental.mcp_server.rest_endpoints import (
_get_tools_for_single_server,
)
from mcp.types import Tool as MCPTool
mock_server = MagicMock()
mock_server.mcp_info = {"server_name": "zapier"}
mock_server.name = "zapier"
mock_server.server_id = "zapier"
mock_server.allowed_tools = None
mock_server.disallowed_tools = ["send_email"]
mock_tools = [
MCPTool(
name="send_email",
description="Send an email",
inputSchema={"type": "object"},
),
MCPTool(
name="read_email",
description="Read an email",
inputSchema={"type": "object"},
),
]
with patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager"
) as mock_manager:
mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools)
result = await _get_tools_for_single_server(mock_server, "Bearer test_token")
assert [tool.name for tool in result] == ["read_email"]
@pytest.mark.asyncio
async def test_list_tool_rest_api_with_server_specific_auth():
"""Test list_tool_rest_api with server-specific auth headers."""

View file

@ -1,10 +1,49 @@
from unittest.mock import AsyncMock, Mock, patch
import httpx
import pytest
from httpx import Request, Response
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.types.integrations.datadog import DatadogPayload
from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError
from litellm.types.integrations.datadog import DD_MAX_BATCH_SIZE, DatadogPayload
def _payloads(n):
return [
DatadogPayload(
ddsource="litellm",
ddtags="env:test",
hostname="host",
message=f'{{"event": {i}}}',
service="svc",
status="info",
)
for i in range(n)
]
def _raised_413():
request = Request("POST", "https://example.com")
response = Response(413, request=request, text="Payload Too Large")
return MaskedHTTPStatusError(
httpx.HTTPStatusError("413", request=request, response=response)
)
def _make_send(max_ok, delivered, *, raise_413=True):
"""Datadog double: 413 batches larger than max_ok, 202 (recording delivery) otherwise."""
async def _send(data):
request = Request("POST", "https://example.com")
if len(data) > max_ok:
if raise_413:
raise _raised_413()
return Response(413, request=request, text="Payload Too Large")
delivered.extend(event["message"] for event in data)
return Response(202, request=request, text="Accepted")
return _send
@pytest.fixture
@ -75,40 +114,152 @@ async def test_failure_hook_threshold_flush_uses_flush_queue(datadog_env):
@pytest.mark.asyncio
async def test_async_send_batch_requeues_events_on_413(datadog_env):
async def test_413_splits_oversized_batch_and_delivers_every_event(datadog_env):
"""A raised 413 (the real httpx path) halves the batch until each piece is accepted."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = [
DatadogPayload(
ddsource="litellm",
ddtags="env:test",
hostname="host",
message=f'{{"event": {i}}}',
service="svc",
status="info",
logger.log_queue = _payloads(4)
delivered: list = []
logger.async_send_compressed_data = AsyncMock(side_effect=_make_send(1, delivered))
await logger.async_send_batch()
assert sorted(delivered) == [f'{{"event": {i}}}' for i in range(4)]
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_413_does_not_requeue_oversized_batch(datadog_env):
"""Regression for the infinite 413 loop: an undeliverable batch must not be re-queued."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = _payloads(4)
logger.async_send_compressed_data = AsyncMock(side_effect=_make_send(0, []))
await logger.async_send_batch()
await logger.async_send_batch()
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_413_drops_single_oversized_event(datadog_env):
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = _payloads(1)
send = AsyncMock(side_effect=_make_send(0, []))
logger.async_send_compressed_data = send
await logger.async_send_batch()
assert send.await_count == 1
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_413_returned_response_also_splits(datadog_env):
"""Defensive path: a 413 returned (not raised) is handled the same way."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = _payloads(4)
delivered: list = []
logger.async_send_compressed_data = AsyncMock(
side_effect=_make_send(1, delivered, raise_413=False)
)
await logger.async_send_batch()
assert sorted(delivered) == [f'{{"event": {i}}}' for i in range(4)]
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_partial_delivery_then_transient_error_requeues_only_undelivered(
datadog_env,
):
"""A transient error after a partial split delivery must not duplicate delivered events."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = _payloads(4)
delivered: list = []
async def _send(data):
messages = [event["message"] for event in data]
if len(data) > 2:
raise _raised_413()
if messages == ['{"event": 2}', '{"event": 3}']:
raise RuntimeError("transient network error")
delivered.extend(messages)
return Response(
202, request=Request("POST", "https://example.com"), text="Accepted"
)
for i in range(2)
logger.async_send_compressed_data = AsyncMock(side_effect=_send)
await logger.async_send_batch()
assert delivered == ['{"event": 0}', '{"event": 1}']
assert [event["message"] for event in logger.log_queue] == [
'{"event": 2}',
'{"event": 3}',
]
@pytest.mark.asyncio
async def test_unexpected_non_202_status_requeues(datadog_env):
"""A non-413, non-202 response is treated as undelivered and re-queued."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = _payloads(2)
logger.async_send_compressed_data = AsyncMock(
return_value=Response(
413,
request=Request("POST", "https://example.com"),
text="Payload Too Large",
200, request=Request("POST", "https://example.com"), text="OK"
)
)
await logger.async_send_batch()
assert logger.async_send_compressed_data.await_count == 1
assert len(logger.log_queue) == 2
assert [event["message"] for event in logger.log_queue] == [
'{"event": 0}',
'{"event": 1}',
]
@pytest.mark.parametrize(
"value, expected",
[
("50", 50),
("1", 1),
("0", 1),
("-5", 1),
(str(DD_MAX_BATCH_SIZE + 100), DD_MAX_BATCH_SIZE),
("not_an_int", DD_MAX_BATCH_SIZE),
],
)
def test_dd_batch_size_env_resolution(monkeypatch, value, expected):
monkeypatch.setenv("DD_API_KEY", "test_api_key")
monkeypatch.setenv("DD_SITE", "test.datadoghq.com")
monkeypatch.setenv("DD_BATCH_SIZE", value)
with patch("asyncio.create_task"):
logger = DataDogLogger()
assert logger.batch_size == expected
def test_dd_batch_size_defaults_to_max(monkeypatch):
monkeypatch.setenv("DD_API_KEY", "test_api_key")
monkeypatch.setenv("DD_SITE", "test.datadoghq.com")
monkeypatch.delenv("DD_BATCH_SIZE", raising=False)
with patch("asyncio.create_task"):
logger = DataDogLogger()
assert logger.batch_size == DD_MAX_BATCH_SIZE
@pytest.mark.asyncio
async def test_async_send_batch_handles_empty_queue(datadog_env):
with patch("asyncio.create_task"):

View file

@ -0,0 +1,69 @@
"""Tests for FocusTransformer — ConsumedQuantity / PricingQuantity correctness."""
from __future__ import annotations
from decimal import Decimal
import polars as pl
from litellm.integrations.focus.transformer import FocusTransformer
def _base_row(**overrides) -> dict:
row = {
"date": "2026-05-25",
"user_id": "u1",
"api_key": "sk-test",
"api_key_alias": "my-key",
"model": "gpt-4o",
"model_group": "openai",
"custom_llm_provider": "openai",
"spend": 0.05,
"api_requests": 3,
"team_id": "team1",
"team_alias": "Engineering",
"user_email": "user@example.com",
}
row.update(overrides)
return row
def _transform(rows: list[dict]) -> pl.DataFrame:
frame = pl.DataFrame(rows, infer_schema_length=None)
return FocusTransformer().transform(frame)
def test_consumed_quantity_reflects_api_requests():
result = _transform([_base_row(api_requests=7)])
assert result["ConsumedQuantity"][0] == Decimal("7.000000")
def test_pricing_quantity_reflects_api_requests():
result = _transform([_base_row(api_requests=7)])
assert result["PricingQuantity"][0] == Decimal("7.000000")
def test_null_api_requests_falls_back_to_zero_not_one():
"""Rows with NULL api_requests (old schema rows) must produce 0, not 1."""
result = _transform([_base_row(api_requests=None)])
assert result["ConsumedQuantity"][0] == Decimal("0.000000")
assert result["PricingQuantity"][0] == Decimal("0.000000")
def test_zero_api_requests_stays_zero():
result = _transform([_base_row(api_requests=0)])
assert result["ConsumedQuantity"][0] == Decimal("0.000000")
assert result["PricingQuantity"][0] == Decimal("0.000000")
def test_bigint_api_requests_cast_correctly():
"""api_requests comes from Postgres as BigInt — large values must not overflow."""
result = _transform([_base_row(api_requests=1_000_000)])
assert result["ConsumedQuantity"][0] == Decimal("1000000.000000")
assert result["PricingQuantity"][0] == Decimal("1000000.000000")
def test_consumed_and_pricing_quantity_match():
"""ConsumedQuantity and PricingQuantity must always be equal."""
result = _transform([_base_row(api_requests=42)])
assert result["ConsumedQuantity"][0] == result["PricingQuantity"][0]

View file

@ -77,28 +77,67 @@ def test_identity_promoted_onto_every_span():
assert span.attributes.get(GenAI.REQUEST_MODEL) == "gpt-4o"
def test_team_metadata_and_provider_model_promoted():
"""The team's metadata dict (JSON) and the provider/underlying model name are
promoted onto every span, alongside the user-facing ``gen_ai.request.model``."""
def test_team_metadata_promoted_only_for_allowlisted_subkeys():
"""Allowlisted team-metadata sub-keys are promoted (JSON) onto every span;
non-allowlisted sub-keys are excluded, alongside the provider/underlying
model name and the user-facing ``gen_ai.request.model``."""
import json
engine, exporter = _engine_and_exporter()
data = LLMCallSpanData.from_standard_logging_payload(_payload())
bag = promoted_baggage(data.identity, data.request_model, BAGGAGE_PROMOTED_KEYS)
bag = promoted_baggage(
data.identity,
data.request_model,
BAGGAGE_PROMOTED_KEYS,
team_metadata_keys=("tier",),
)
ctx = ctx_mod.set_request_baggage(bag)
engine.emit(SpanRole.SERVICE, ServiceSpanData("redis", call_type="set"), ctx)
(span,) = exporter.get_finished_spans()
# team metadata: the whole dict, JSON-serialized into one value
assert json.loads(span.attributes[LiteLLM.TEAM_METADATA]) == {
"tier": "gold",
"cost_center": "42",
}
# only the allowlisted sub-key is promoted; ``cost_center`` is excluded
assert json.loads(span.attributes[LiteLLM.TEAM_METADATA]) == {"tier": "gold"}
# provider model is distinct from the user-facing request model
assert span.attributes.get(LiteLLM.PROVIDER_MODEL) == "azure/my-deployment"
assert span.attributes.get(GenAI.REQUEST_MODEL) == "gpt-4o"
def test_team_metadata_not_promoted_by_default():
"""The default allowlist is empty, so a team's metadata is never promoted
even though its dict is present on the request."""
data = LLMCallSpanData.from_standard_logging_payload(_payload())
# raw dict is carried on the identity for promotion-time filtering
assert data.identity.team_metadata == {"tier": "gold", "cost_center": "42"}
bag = promoted_baggage(data.identity, data.request_model, BAGGAGE_PROMOTED_KEYS)
assert LiteLLM.TEAM_METADATA not in bag
def test_team_metadata_dropped_when_no_allowlisted_key_present():
"""An allowlist that matches no present sub-key drops team_metadata rather
than promoting a useless ``{}``."""
data = LLMCallSpanData.from_standard_logging_payload(_payload())
bag = promoted_baggage(
data.identity,
data.request_model,
BAGGAGE_PROMOTED_KEYS,
team_metadata_keys=("absent_key",),
)
assert LiteLLM.TEAM_METADATA not in bag
def test_team_metadata_not_promoted_when_key_excluded_from_promoted_keys():
"""Even with sub-keys allowlisted, team_metadata stays off the wire when
``litellm.team.metadata`` itself isn't in ``promoted_keys``."""
data = LLMCallSpanData.from_standard_logging_payload(_payload())
bag = promoted_baggage(
data.identity,
data.request_model,
(LiteLLM.TEAM_ID,),
team_metadata_keys=("tier",),
)
assert LiteLLM.TEAM_METADATA not in bag
def test_empty_team_metadata_is_dropped():
"""An absent/empty team_metadata dict must not promote a useless ``"{}"``."""
payload = _payload()

View file

@ -96,8 +96,12 @@ def _kwargs(payload=None):
}
def _logger(legacy_compat=True):
cfg = OpenTelemetryV2Config(exporter="in_memory", legacy_compat=legacy_compat)
def _logger(legacy_compat=True, team_metadata_keys=None):
cfg = OpenTelemetryV2Config(
exporter="in_memory",
legacy_compat=legacy_compat,
baggage_team_metadata_keys=team_metadata_keys or [],
)
exporter = InMemorySpanExporter()
tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter)
return OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider), exporter
@ -541,8 +545,9 @@ class _Auth:
def test_provider_model_and_team_metadata_on_real_boundary_flow():
"""End-to-end on the proxy boundary path (the gap a pure-emitter test misses):
- ``litellm.team.metadata`` is known at auth, so it rides identity Baggage
seeded there onto EVERY span (server + LLM call).
- ``litellm.team.metadata`` (filtered to the allowlisted sub-keys) is known
at auth, so it rides identity Baggage seeded there onto EVERY span
(server + LLM call).
- ``litellm.provider.model`` is only known once routing picks a deployment
(in the payload at close), AFTER the auth seed and AFTER the boundary span
starts — so it can't ride Baggage. It's stamped directly on the LLM-call
@ -550,7 +555,7 @@ def test_provider_model_and_team_metadata_on_real_boundary_flow():
"""
import json
logger, exporter = _logger()
logger, exporter = _logger(team_metadata_keys=["tier", "cost_center"])
server = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)

View file

@ -23,6 +23,7 @@ from litellm.integrations.opentelemetry import (
OpenTelemetry,
OpenTelemetryConfig,
OTELSemconvCategory,
_normalize_team_metadata_keys,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -5187,8 +5188,13 @@ class TestOpenTelemetryInferenceIdentityAttributes(unittest.TestCase):
},
}
def _otel_with_team_metadata_keys(self, keys):
return OpenTelemetry(
config=OpenTelemetryConfig(baggage_team_metadata_keys=keys)
)
def test_all_identity_attributes_stamped(self):
otel = OpenTelemetry()
otel = self._otel_with_team_metadata_keys(["tier", "cost_center"])
span, exp = self._span()
otel.set_attributes(span, self._kwargs(), {"model": "azure/gpt-4o"})
attrs = self._attr(span, exp)
@ -5201,6 +5207,33 @@ class TestOpenTelemetryInferenceIdentityAttributes(unittest.TestCase):
assert attrs["litellm.model_group"] == "gpt-4o"
assert attrs["litellm.provider.model"] == "azure/my-deployment"
def test_team_metadata_defaults_to_none_stamped(self):
"""With no allowlist configured (the default), a team's metadata must
never be stamped, even when present on the request."""
otel = OpenTelemetry()
span, exp = self._span()
otel.set_attributes(span, self._kwargs(), {"model": "azure/gpt-4o"})
assert "litellm.team.metadata" not in self._attr(span, exp)
def test_only_allowlisted_team_metadata_keys_stamped(self):
"""Sub-keys outside the allowlist are excluded from the stamped value."""
otel = self._otel_with_team_metadata_keys(["tier"])
span, exp = self._span()
otel.set_attributes(span, self._kwargs(), {"model": "azure/gpt-4o"})
assert json.loads(self._attr(span, exp)["litellm.team.metadata"]) == {
"tier": "gold"
}
def test_team_metadata_allowlist_from_config_yaml_kwarg(self):
"""callback_settings.otel.baggage_team_metadata_keys arrives as a kwarg
and must drive the allowlist."""
otel = OpenTelemetry(baggage_team_metadata_keys=["cost_center"])
span, exp = self._span()
otel.set_attributes(span, self._kwargs(), {"model": "azure/gpt-4o"})
assert json.loads(self._attr(span, exp)["litellm.team.metadata"]) == {
"cost_center": "42"
}
def test_provider_model_falls_back_to_payload_model(self):
"""Without hidden_params.litellm_model_name the dispatched model is
the payload model (the SDK path, where no router renaming happened)."""
@ -5229,8 +5262,52 @@ class TestOpenTelemetryInferenceIdentityAttributes(unittest.TestCase):
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
assert "http.route" not in self._attr(span, exp)
def test_team_metadata_json_helper_non_dict(self):
assert OpenTelemetry._team_metadata_json(None) is None
assert OpenTelemetry._team_metadata_json("not-a-dict") is None
assert OpenTelemetry._team_metadata_json({}) is None
assert json.loads(OpenTelemetry._team_metadata_json({"a": 1})) == {"a": 1}
def test_team_metadata_json_helper(self):
keys = ["a", "b"]
assert OpenTelemetry._team_metadata_json(None, keys) is None
assert OpenTelemetry._team_metadata_json("not-a-dict", keys) is None
assert OpenTelemetry._team_metadata_json({}, keys) is None
# empty allowlist -> nothing stamped, even with data present
assert OpenTelemetry._team_metadata_json({"a": 1}, []) is None
# no allowlisted key present -> dropped, not a useless "{}"
assert OpenTelemetry._team_metadata_json({"c": 1}, keys) is None
# only allowlisted sub-keys survive
assert json.loads(
OpenTelemetry._team_metadata_json({"a": 1, "c": 2}, keys)
) == {"a": 1}
class TestOpenTelemetryTeamMetadataKeysConfig(unittest.TestCase):
def test_normalize_from_csv_string(self):
# comma-separated env var: strip whitespace and drop empties
assert _normalize_team_metadata_keys("tier, cost_center , ,") == [
"tier",
"cost_center",
]
def test_normalize_from_list(self):
assert _normalize_team_metadata_keys(["tier", " cost_center ", ""]) == [
"tier",
"cost_center",
]
def test_normalize_none(self):
assert _normalize_team_metadata_keys(None) == []
def test_config_reads_csv_env_var(self):
with patch.dict(
"os.environ",
{"LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS": "tier, cost_center"},
):
assert OpenTelemetryConfig().baggage_team_metadata_keys == [
"tier",
"cost_center",
]
def test_explicit_keys_win_over_env_var(self):
with patch.dict(
"os.environ",
{"LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS": "from_env"},
):
cfg = OpenTelemetryConfig(baggage_team_metadata_keys=["from_arg"])
assert cfg.baggage_team_metadata_keys == ["from_arg"]

View file

@ -35,6 +35,8 @@ class TestPrometheusCacheMetrics:
assert "litellm_cache_hits_metric" in defined_metrics
assert "litellm_cache_misses_metric" in defined_metrics
assert "litellm_cached_tokens_metric" in defined_metrics
assert "litellm_provider_cache_read_input_tokens_metric" in defined_metrics
assert "litellm_provider_cache_creation_input_tokens_metric" in defined_metrics
def test_cache_metric_labels_defined(self):
"""Test that cache metric labels are properly defined"""
@ -44,6 +46,13 @@ class TestPrometheusCacheMetrics:
assert hasattr(PrometheusMetricLabels, "litellm_cache_hits_metric")
assert hasattr(PrometheusMetricLabels, "litellm_cache_misses_metric")
assert hasattr(PrometheusMetricLabels, "litellm_cached_tokens_metric")
assert hasattr(
PrometheusMetricLabels, "litellm_provider_cache_read_input_tokens_metric"
)
assert hasattr(
PrometheusMetricLabels,
"litellm_provider_cache_creation_input_tokens_metric",
)
# Verify labels include expected keys
expected_labels = [
@ -59,6 +68,14 @@ class TestPrometheusCacheMetrics:
assert label in PrometheusMetricLabels.litellm_cache_hits_metric
assert label in PrometheusMetricLabels.litellm_cache_misses_metric
assert label in PrometheusMetricLabels.litellm_cached_tokens_metric
assert (
label
in PrometheusMetricLabels.litellm_provider_cache_read_input_tokens_metric
)
assert (
label
in PrometheusMetricLabels.litellm_provider_cache_creation_input_tokens_metric
)
def test_increment_cache_metrics_on_cache_hit(self, sample_enum_values):
"""Test that cache hit increments the correct metrics"""
@ -76,12 +93,20 @@ class TestPrometheusCacheMetrics:
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"cache_read_input_tokens": 25,
"cache_creation_input_tokens": 10,
}
},
}
# Create mock metrics
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
@ -114,6 +139,14 @@ class TestPrometheusCacheMetrics:
# Verify cache misses metric was NOT called
mock_logger.litellm_cache_misses_metric.labels.assert_not_called()
# Verify provider prompt caching metrics were incremented
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels().inc.assert_called_once_with(
25
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels().inc.assert_called_once_with(
10
)
def test_increment_cache_metrics_on_cache_miss(self, sample_enum_values):
"""Test that cache miss increments the correct metrics"""
# Create mock for PrometheusLogger instance
@ -129,12 +162,20 @@ class TestPrometheusCacheMetrics:
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
# Explicit provider field absent -> fallback should use prompt_tokens_details.cached_tokens
"prompt_tokens_details": {"cached_tokens": 20},
}
},
}
# Create mock metrics
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
@ -162,6 +203,61 @@ class TestPrometheusCacheMetrics:
mock_logger.litellm_cache_hits_metric.labels.assert_not_called()
mock_logger.litellm_cached_tokens_metric.labels.assert_not_called()
# Provider prompt caching metrics should still be emitted
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels().inc.assert_called_once_with(
20
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels.assert_not_called()
def test_provider_cache_read_does_not_fallback_on_explicit_zero(
self, sample_enum_values
):
"""Explicit cache_read_input_tokens=0 must not trigger fallback to cached_tokens."""
mock_logger = MagicMock()
from litellm.integrations.prometheus import PrometheusLogger
standard_logging_payload = {
"cache_hit": False,
"total_tokens": 100,
"prompt_tokens": 50,
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"cache_read_input_tokens": 0,
"prompt_tokens_details": {"cached_tokens": 20},
}
},
}
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
"hashed_api_key",
"api_key_alias",
"team",
"team_alias",
"end_user",
"user",
]
)
PrometheusLogger._increment_cache_metrics(
mock_logger,
standard_logging_payload=standard_logging_payload,
enum_values=sample_enum_values,
)
# Should not emit read metric, because explicit provider value is zero.
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels.assert_not_called()
def test_increment_cache_metrics_when_cache_hit_is_none(self, sample_enum_values):
"""Test that no metrics are incremented when cache_hit is None"""
# Create mock for PrometheusLogger instance
@ -177,12 +273,19 @@ class TestPrometheusCacheMetrics:
"completion_tokens": 50,
"model_group": "openai",
"request_tags": [],
"metadata": {
"usage_object": {
"cache_read_input_tokens": 25,
}
},
}
# Create mock metrics
mock_logger.litellm_cache_hits_metric = MagicMock()
mock_logger.litellm_cache_misses_metric = MagicMock()
mock_logger.litellm_cached_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_read_input_tokens_metric = MagicMock()
mock_logger.litellm_provider_cache_creation_input_tokens_metric = MagicMock()
mock_logger.get_labels_for_metric = MagicMock(
return_value=[
"model",
@ -207,6 +310,12 @@ class TestPrometheusCacheMetrics:
mock_logger.litellm_cache_misses_metric.labels.assert_not_called()
mock_logger.litellm_cached_tokens_metric.labels.assert_not_called()
# Provider prompt caching metrics should still be emitted
mock_logger.litellm_provider_cache_read_input_tokens_metric.labels().inc.assert_called_once_with(
25
)
mock_logger.litellm_provider_cache_creation_input_tokens_metric.labels.assert_not_called()
if __name__ == "__main__":
pytest.main([__file__, "-v"])

View file

@ -546,3 +546,178 @@ class TestExtractFileDataBareStr:
extracted = extract_file_data(("foo.txt", b"raw bytes content"))
assert extracted.get("filename") == "foo.txt"
assert extracted.get("content") == b"raw bytes content"
class TestUnpackLegacyDefs:
"""Cover the public ``unpack_legacy_defs`` helper directly so the no-op
branches (non-dict input, schema with no legacy/OpenAPI defs) are exercised
without needing a provider-specific entry point.
"""
@pytest.mark.parametrize(
"value",
[None, [], "string-not-a-dict", 42, 1.5, True, set(), tuple()],
)
def test_non_dict_returns_unchanged_no_op(self, value):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
# Should never raise; returns the input unchanged.
assert unpack_legacy_defs(value) is value
assert unpack_legacy_defs(value, copy=True) is value
def test_dict_without_legacy_defs_is_no_op(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
schema = {
"type": "object",
"properties": {"a": {"$ref": "#/$defs/A"}},
"$defs": {"A": {"type": "string"}},
}
snapshot = json.loads(json.dumps(schema))
# No `definitions` and no `components.schemas` -> early return, no work.
out = unpack_legacy_defs(schema)
assert out is schema
assert schema == snapshot, "schema mutated despite no legacy defs"
def test_components_with_no_schemas_block_is_no_op(self):
"""``components`` without a ``schemas`` sub-key must not be popped."""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
schema = {
"type": "object",
"properties": {"a": {"type": "string"}},
"components": {"securitySchemes": {"foo": "bar"}},
}
snapshot = json.loads(json.dumps(schema))
unpack_legacy_defs(schema)
assert schema == snapshot, "components without schemas was incorrectly popped"
def test_legitimate_schema_within_budget_succeeds(self):
"""A flat schema with many distinct ``$ref``s into small targets must
inline cleanly under the default budget -- the budget rejects bombs,
not legitimately-shaped schemas.
"""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
n = 200
schema = {
"type": "object",
"properties": {f"f{i}": {"$ref": f"#/definitions/T{i}"} for i in range(n)},
"definitions": {f"T{i}": {"type": "string"} for i in range(n)},
}
out = unpack_legacy_defs(schema)
assert "definitions" not in out
for i in range(n):
assert out["properties"][f"f{i}"] == {"type": "string"}
# Schema-bomb amplification vectors. ``max_inlined_bytes`` is the universal
# measure of expansion: every other dimension (ref count, node count,
# scalar size) reduces to bytes-on-the-wire, so a single byte budget
# closes all three vectors at once.
def test_rejects_fan_out_bomb(self):
"""Each level multiplies refs (cycle detection only stops re-entry
along the *same* path). Must trip the byte budget."""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
depth, fanout = 12, 2 # 2**12 = 4096 leaves
definitions = {
f"L{i}": {
"type": "object",
"properties": {
f"x{j}": {"$ref": f"#/definitions/L{i + 1}"} for j in range(fanout)
},
}
for i in range(depth)
}
definitions[f"L{depth}"] = {"type": "string"}
schema = {
"type": "object",
"properties": {"root": {"$ref": "#/definitions/L0"}},
"definitions": definitions,
}
with pytest.raises(ValueError, match="byte budget"):
unpack_legacy_defs(schema, max_inlined_bytes=100_000)
def test_rejects_target_amplification_bomb(self):
"""Few refs each deep-copying one large target -- bounded total
expanded bytes catches it even though ref count is small."""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
big = {
"type": "object",
"properties": {f"p{i}": {"type": "string"} for i in range(100)},
}
schema = {
"type": "object",
"properties": {f"r{i}": {"$ref": "#/definitions/Big"} for i in range(50)},
"definitions": {"Big": big},
}
with pytest.raises(ValueError, match="byte budget"):
unpack_legacy_defs(schema, max_inlined_bytes=10_000)
def test_rejects_scalar_byte_amplification_bomb(self):
"""Many ``$ref``s to a target containing one large scalar (e.g. a
long ``description``, ``const`` value, or ``enum`` entry). A
node-counter would treat this as 1 node per resolution and miss it;
a byte budget catches the actual wire-size amplification.
"""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
big_description = "x" * 100_000 # 100KB string
schema = {
"type": "object",
"properties": {f"r{i}": {"$ref": "#/definitions/Big"} for i in range(50)},
"definitions": {
"Big": {"type": "string", "description": big_description},
},
}
# 50 refs * ~100KB string == ~5MB cumulative; 1MB budget trips.
with pytest.raises(ValueError, match="byte budget"):
unpack_legacy_defs(schema, max_inlined_bytes=1_000_000)
def test_budget_does_not_trip_for_legitimate_large_schema(self):
"""An OpenAPI-derived tool with ~50 small targets must inline cleanly
under the default ``max_inlined_bytes`` budget."""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
unpack_legacy_defs,
)
schema = {
"type": "object",
"properties": {
f"r{i}": {"$ref": f"#/components/schemas/T{i}"} for i in range(50)
},
"components": {
"schemas": {
f"T{i}": {
"type": "object",
"properties": {f"p{j}": {"type": "string"} for j in range(5)},
}
for i in range(50)
}
},
}
out = unpack_legacy_defs(schema)
assert "components" not in out
assert out["properties"]["r0"]["properties"]["p0"] == {"type": "string"}

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.exception_mapping_utils import (
exception_type,
extract_and_raise_litellm_exception,
)
from litellm.llms.openai.common_utils import OpenAIError
# Test cases for is_error_str_context_window_exceeded
# Tuple format: (error_message, expected_result)
@ -41,6 +42,10 @@ context_window_test_cases = [
"`inputs` tokens + `max_new_tokens` must be <= 4096",
True,
),
(
"request (67311 tokens) exceeds the available context size (65536 tokens), try increasing it",
True,
),
# Gemini 2.5/3 format
(
"The input token count exceeds the maximum number of tokens allowed 1048576.",
@ -182,7 +187,6 @@ class TestExceptionCheckers:
]
for error_str in positive_cases:
print("testing positive case=", error_str)
result = ExceptionCheckers.is_azure_content_policy_violation_error(
error_str
)
@ -255,6 +259,33 @@ def test_gemini_context_window_error_mapping(
)
def test_lemonade_context_window_error_mapping():
"""Lemonade's llama.cpp backend should map context overflows to LiteLLM's standard error."""
model = "lemonade/Qwen3.6-35B-A3B-GGUF"
error_message = (
'{"error":{"code":"context_length_exceeded","message":"request '
"(80010 tokens) exceeds the available context size (65536 tokens), "
'try increasing it","status_code":400,"type":"invalid_request_error"}}'
)
original_exception = OpenAIError(
status_code=400,
message=error_message,
headers={},
)
with pytest.raises(litellm.ContextWindowExceededError) as excinfo:
exception_type(
model=model,
original_exception=original_exception,
custom_llm_provider="lemonade",
)
assert excinfo.value.status_code == 400
assert excinfo.value.llm_provider == "lemonade"
assert excinfo.value.model == model
# Test cases for Vertex AI RateLimitError mapping
# As per https://github.com/BerriAI/litellm/issues/16189
vertex_rate_limit_test_cases = [

View file

@ -0,0 +1,43 @@
import pytest
import litellm
from litellm.litellm_core_utils.fallback_utils import async_completion_with_fallbacks
@pytest.mark.asyncio
async def test_fallback_dict_not_mutated(monkeypatch):
fallback_dict = {"model": "fallback-model", "temperature": 0.2}
original_fallback_dict = dict(fallback_dict)
attempted_models: list[str] = []
async def _fake_acompletion(*, model: str, **kwargs):
attempted_models.append(model)
if model == "primary-model":
raise Exception("primary failed")
return {"model": model, "temperature": kwargs.get("temperature")}
monkeypatch.setattr(litellm, "acompletion", _fake_acompletion)
# Call 1: primary fails, fallback dict succeeds
response_1 = await async_completion_with_fallbacks(
model="primary-model",
kwargs={"fallbacks": [fallback_dict]},
)
assert response_1["model"] == "fallback-model"
assert fallback_dict == original_fallback_dict
# Call 2: re-use the same dict object; it should still work and remain unchanged
response_2 = await async_completion_with_fallbacks(
model="primary-model",
kwargs={"fallbacks": [fallback_dict]},
)
assert response_2["model"] == "fallback-model"
assert fallback_dict == original_fallback_dict
assert attempted_models == [
"primary-model",
"fallback-model",
"primary-model",
"fallback-model",
]

View file

@ -4889,3 +4889,204 @@ def test_sanitize_tool_names_in_request_no_tools_is_noop():
forward, reverse = AnthropicConfig._sanitize_tool_names_in_request({"tools": []})
assert forward == {}
assert reverse == {}
# -----------------------------------------------------------------------------
# Regression tests for legacy / OpenAPI $ref defs in tool input_schema.
#
# Anthropic only resolves `$defs` (JSON Schema 2020-12). Tools coming from MCP
# servers (legacy `definitions`) or OpenAPI-derived gateways like AWS
# AgentCore (`components.schemas`) used to silently lose their def blocks
# while keeping dangling `$ref`s, causing upstream 400s. See
# https://github.com/BerriAI/litellm/issues/26692.
# -----------------------------------------------------------------------------
def _assert_no_unresolved_refs(input_schema: dict) -> None:
import json
blob = json.dumps(input_schema)
assert "$ref" not in blob, f"unresolved $ref in transformed input_schema: {blob}"
def test_map_tool_helper_inlines_components_schemas_refs():
"""OpenAPI `components.schemas` $refs (AgentCore-style) must be inlined."""
config = AnthropicConfig()
tool = {
"type": "function",
"function": {
"name": "slides_presentations_create",
"description": "Create a Google Slides presentation",
"parameters": {
"type": "object",
"properties": {
"body": {"$ref": "#/components/schemas/Presentation"},
},
"required": ["body"],
"components": {
"schemas": {
"Presentation": {
"type": "object",
"properties": {
"title": {"type": "string"},
"presentationId": {"type": "string"},
},
}
}
},
},
},
}
transformed, _ = config._map_tool_helper(tool)
assert transformed is not None
schema = transformed["input_schema"]
_assert_no_unresolved_refs(schema)
assert schema["properties"]["body"] == {
"type": "object",
"properties": {
"title": {"type": "string"},
"presentationId": {"type": "string"},
},
}
# The OpenAPI components block is not part of Anthropic's allow-list and
# must not be forwarded.
assert "components" not in schema
def test_map_tool_helper_inlines_legacy_definitions_refs():
"""Legacy draft-04 `definitions` $refs (DevRev MCP-style) must be inlined."""
config = AnthropicConfig()
tool = {
"type": "function",
"function": {
"name": "create_thing",
"description": "Create a thing",
"parameters": {
"type": "object",
"properties": {
"thing": {"$ref": "#/definitions/Thing"},
},
"definitions": {
"Thing": {
"type": "object",
"properties": {"id": {"type": "string"}},
}
},
},
},
}
transformed, _ = config._map_tool_helper(tool)
assert transformed is not None
schema = transformed["input_schema"]
_assert_no_unresolved_refs(schema)
assert schema["properties"]["thing"] == {
"type": "object",
"properties": {"id": {"type": "string"}},
}
assert "definitions" not in schema
def test_map_tool_helper_preserves_native_dollar_defs():
"""`$defs` is JSON Schema 2020-12 native; Anthropic resolves it itself.
Re-implementation must not pop or unpack `$defs`.
"""
config = AnthropicConfig()
tool = {
"type": "function",
"function": {
"name": "native_defs_tool",
"description": "",
"parameters": {
"type": "object",
"properties": {"a": {"$ref": "#/$defs/A"}},
"$defs": {"A": {"type": "string"}},
},
},
}
transformed, _ = config._map_tool_helper(tool)
assert transformed is not None
schema = transformed["input_schema"]
assert schema["$defs"] == {"A": {"type": "string"}}
assert schema["properties"]["a"] == {"$ref": "#/$defs/A"}
def test_map_tool_helper_does_not_mutate_caller_dict():
"""Caller-supplied tool dict must not be mutated by the inlining step."""
import copy
config = AnthropicConfig()
tool = {
"type": "function",
"function": {
"name": "create_thing",
"description": "Create a thing",
"parameters": {
"type": "object",
"properties": {"thing": {"$ref": "#/definitions/Thing"}},
"definitions": {
"Thing": {
"type": "object",
"properties": {"id": {"type": "string"}},
}
},
},
},
}
snapshot = copy.deepcopy(tool)
config._map_tool_helper(tool)
assert tool == snapshot, "caller's tool dict was mutated in place"
def test_map_tool_helper_collision_prefers_definitions_over_components_schemas():
"""If both `definitions.X` and `components.schemas.X` exist with the same
name, prefer the `definitions` body. ``unpack_defs`` keys refs by last path
segment so only one body can win; pick the JSON-Schema-native one.
This locks in the residual limitation as a deliberate contract: a ref
written as ``#/components/schemas/X`` will *also* resolve to the
``definitions`` body when both namespaces define ``X``. Cross-namespace
disambiguation would require teaching ``unpack_defs`` to key by full ref
path, which is out of scope here.
"""
config = AnthropicConfig()
tool = {
"type": "function",
"function": {
"name": "collision_tool",
"description": "",
"parameters": {
"type": "object",
"properties": {
"from_definitions": {"$ref": "#/definitions/Thing"},
"from_components": {"$ref": "#/components/schemas/Thing"},
},
"definitions": {
"Thing": {"type": "string", "description": "from-definitions"},
},
"components": {
"schemas": {
"Thing": {"type": "integer", "description": "from-components"},
}
},
},
},
}
transformed, _ = config._map_tool_helper(tool)
assert transformed is not None
expected = {"type": "string", "description": "from-definitions"}
# Direct ref resolves to the `definitions` body (the documented winner).
assert transformed["input_schema"]["properties"]["from_definitions"] == expected
# Cross-namespace ref *also* resolves to the `definitions` body because
# ``unpack_defs`` keys by last path segment -- documented limitation.
assert transformed["input_schema"]["properties"]["from_components"] == expected

View file

@ -0,0 +1,3 @@
{"recordId": "embed-1", "modelInput": {"inputText": "Hello world"}}
{"recordId": "embed-2", "modelInput": {"inputText": "Another document to embed", "dimensions": 512}}
{"recordId": "embed-3", "modelInput": {"inputText": "Single element list", "embeddingTypes": ["binary"]}}

View file

@ -0,0 +1,3 @@
{"custom_id": "embed-1", "method": "POST", "url": "/v1/embeddings", "body": {"model": "bedrock/amazon.titan-embed-text-v2:0", "input": "Hello world"}}
{"custom_id": "embed-2", "method": "POST", "url": "/v1/embeddings", "body": {"model": "bedrock/amazon.titan-embed-text-v2:0", "input": "Another document to embed", "dimensions": 512}}
{"custom_id": "embed-3", "method": "POST", "url": "/v1/embeddings", "body": {"model": "bedrock/amazon.titan-embed-text-v2:0", "input": ["Single element list"], "encoding_format": "base64"}}

View file

@ -426,7 +426,7 @@ class TestBedrockFilesTransformation:
"s3_bucket_name": "litellm-batch-352026",
"s3_region_name": "us-gov-west-1",
}
# aws_region_name set to something different — s3_region_name must still win
# aws_region_name set to something different - s3_region_name must still win
optional_params = {"aws_region_name": "us-east-1"}
captured_optional_params: dict = {}
@ -482,3 +482,630 @@ class TestBedrockFilesTransformation:
assert "messages" in model_input
assert "max_tokens" in model_input
assert model_input["max_tokens"] == 10
class TestBedrockFilesEmbeddingTransformation:
"""
Tests for routing OpenAI /v1/embeddings batch JSONL records through the
Titan v2 transformer so AWS Bedrock's CreateModelInvocationJob receives
a valid modelInput body.
Scope is intentionally Titan v2 only - other embedding models will get
their own follow-up PRs/tests so each schema is exercised in isolation.
"""
def test_titan_v2_embedding_jsonl_matches_fixture(self):
"""Round-trip the input fixture against the expected Bedrock output."""
import json
import os
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
here = os.path.dirname(__file__)
with open(os.path.join(here, "input_batch_embeddings.jsonl")) as f:
openai_jsonl = [json.loads(line) for line in f if line.strip()]
with open(os.path.join(here, "expected_bedrock_batch_embeddings.jsonl")) as f:
expected = [json.loads(line) for line in f if line.strip()]
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
openai_jsonl
)
assert result == expected
def test_titan_v2_simple_string_input(self):
"""Single string `input` maps to `{"inputText": <str>}` with no extras."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": "Hello",
},
}
]
)
assert result == [{"recordId": "e1", "modelInput": {"inputText": "Hello"}}]
def test_titan_v2_dimensions_and_encoding_format(self):
"""OpenAI `dimensions` / `encoding_format` map to Titan v2 schema."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": "Hi",
"dimensions": 256,
"encoding_format": "float",
},
}
]
)
model_input = result[0]["modelInput"]
assert model_input["inputText"] == "Hi"
assert model_input["dimensions"] == 256
assert model_input["embeddingTypes"] == ["float"]
def test_embedding_routing_falls_back_to_body_shape(self):
"""Records without `url` still route via `input` presence."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": "Hello",
},
}
]
)
assert result[0]["modelInput"] == {"inputText": "Hello"}
def test_embedding_single_element_list_input_is_accepted(self):
"""A single-element list maps to the same shape as a bare string."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": ["only one"],
},
}
]
)
assert result[0]["modelInput"]["inputText"] == "only one"
def test_embedding_multi_input_list_raises(self):
"""Multi-element `input` lists are rejected with a clear message."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
with pytest.raises(ValueError, match="one input per JSONL record"):
config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": ["a", "b"],
},
}
]
)
def test_embedding_missing_input_raises(self):
"""A record routed to /v1/embeddings without `input` is an error."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
with pytest.raises(ValueError, match="missing required `input`"):
config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {"model": "bedrock/amazon.titan-embed-text-v2:0"},
}
]
)
def test_mixed_chat_and_embedding_in_same_batch(self):
"""Chat and embedding records in the same JSONL each take their path."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "chat-1",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"messages": [{"role": "user", "content": "Hi"}],
"max_tokens": 5,
},
},
{
"custom_id": "embed-1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": "Hi",
},
},
]
)
assert result[0]["recordId"] == "chat-1"
assert "messages" in result[0]["modelInput"]
assert result[0]["modelInput"]["anthropic_version"] == "bedrock-2023-05-31"
assert result[1]["recordId"] == "embed-1"
assert result[1]["modelInput"] == {"inputText": "Hi"}
def test_unsupported_embedding_model_raises_not_implemented(self):
"""Cohere/Nova/Titan-G1 embed get a clear NotImplementedError, not a corrupt body."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
for unsupported_model in (
"bedrock/cohere.embed-english-v3",
"bedrock/amazon.titan-embed-text-v1",
"bedrock/amazon.titan-embed-image-v1",
"bedrock/amazon.nova-2-multimodal-embeddings-v1:0",
):
with pytest.raises(NotImplementedError, match="titan-embed-text-v2"):
config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {"model": unsupported_model, "input": "Hi"},
}
]
)
def test_titan_v2_model_name_variants_route_correctly(self):
"""All common Titan v2 model id shapes route through the embedding path."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
for model_id in (
"amazon.titan-embed-text-v2:0",
"bedrock/amazon.titan-embed-text-v2:0",
"us.amazon.titan-embed-text-v2:0",
"bedrock/us.amazon.titan-embed-text-v2:0",
):
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {"model": model_id, "input": "Hi"},
}
]
)
assert result[0]["modelInput"] == {
"inputText": "Hi"
}, f"model id {model_id} did not route to Titan v2 embedding path"
def test_pretokenized_input_list_of_ints_raises(self):
"""`input: List[int]` (pre-tokenized) is rejected, not silently mis-shaped."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
with pytest.raises(
(NotImplementedError, ValueError), match=r"pre-tokenized|one input per"
):
config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": [1, 2, 3],
},
}
]
)
def test_pretokenized_single_wrapped_list_raises(self):
"""`input: List[List[int]]` with one element is rejected as pre-tokenized."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
with pytest.raises(NotImplementedError, match="pre-tokenized"):
config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {
"model": "bedrock/amazon.titan-embed-text-v2:0",
"input": [[1, 2, 3]],
},
}
]
)
def test_record_with_both_input_and_messages_routes_to_chat(self):
"""If a record has both fields, chat wins (safer default - see helper docstring)."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "ambiguous-1",
"body": {
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"messages": [{"role": "user", "content": "Hi"}],
"input": "this should be ignored by chat path",
"max_tokens": 5,
},
}
]
)
assert "messages" in result[0]["modelInput"]
assert "inputText" not in result[0]["modelInput"]
def test_url_embeddings_with_missing_input_raises_not_chat_error(self):
"""url says embed, body lacks input → embedding-path error, not chat-path crash."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
config = BedrockFilesConfig()
with pytest.raises(ValueError, match="missing required `input`"):
config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "e1",
"method": "POST",
"url": "/v1/embeddings",
"body": {"model": "bedrock/amazon.titan-embed-text-v2:0"},
}
]
)
def test_titan_v2_marker_boundary_rejects_lookalikes(self):
"""The marker must end at `:`, `/`, or end-of-string to avoid false positives."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
# Look-alikes that must NOT route through the Titan v2 path
for model in (
"bedrock/amazon.titan-embed-text-v20:0",
"bedrock/amazon.titan-embed-text-v2-experimental:0",
"bedrock/amazon.titan-embed-text-v2foo",
):
assert not BedrockFilesConfig._is_titan_v2_embed_model(
model
), f"{model} unexpectedly matched the Titan v2 marker"
# Real Titan v2 ids that MUST match
for model in (
"amazon.titan-embed-text-v2:0",
"bedrock/amazon.titan-embed-text-v2:0",
"us.amazon.titan-embed-text-v2:0",
"arn:aws:bedrock:us-east-1:123:foundation-model/amazon.titan-embed-text-v2:0",
):
assert BedrockFilesConfig._is_titan_v2_embed_model(
model
), f"{model} unexpectedly missed the Titan v2 marker"
def test_titan_v2_accepted_when_registry_schema_field_matches(self, mocker):
"""Registry-driven happy path: nested
`provider_specific_entry.bedrock_invocation_schema == "titan_v2"`
is the authoritative signal."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
mocker.patch(
"litellm.get_model_info",
return_value={
"provider_specific_entry": {"bedrock_invocation_schema": "titan_v2"}
},
)
assert BedrockFilesConfig._is_titan_v2_embed_model(
"amazon.titan-embed-text-v2:0"
)
def test_titan_v2_rejected_when_registry_schema_field_differs(self, mocker):
"""Registry resolves with a different schema value (e.g. a hypothetical
Cohere Embed entry) -> reject. Registry is authoritative; no substring
second-chance for ids the registry knows."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
mocker.patch(
"litellm.get_model_info",
return_value={
"provider_specific_entry": {"bedrock_invocation_schema": "cohere_v3"}
},
)
# Even though the model id looks like Titan v2, the registry says
# otherwise and we trust it.
assert not BedrockFilesConfig._is_titan_v2_embed_model(
"amazon.titan-embed-text-v2:0"
)
def test_titan_v2_falls_back_to_marker_when_registry_lacks_schema_field(
self, mocker
):
"""Registry resolves but the entry has no
`provider_specific_entry.bedrock_invocation_schema` field yet (e.g.
a stale local registry) -> fall through to substring."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
# No provider_specific_entry at all
mocker.patch(
"litellm.get_model_info",
return_value={"mode": "embedding"},
)
assert BedrockFilesConfig._is_titan_v2_embed_model(
"amazon.titan-embed-text-v2:0"
)
# provider_specific_entry present but missing the schema key
mocker.patch(
"litellm.get_model_info",
return_value={
"mode": "embedding",
"provider_specific_entry": {"unrelated": "value"},
},
)
assert BedrockFilesConfig._is_titan_v2_embed_model(
"amazon.titan-embed-text-v2:0"
)
def test_titan_v2_accepted_when_registry_silent(self, mocker):
"""Marker-only match is fine for ids the registry can't resolve
(cross-region profile prefixes, ARN forms)."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
mocker.patch("litellm.get_model_info", side_effect=Exception("not mapped"))
assert BedrockFilesConfig._is_titan_v2_embed_model(
"us.amazon.titan-embed-text-v2:0"
)
assert BedrockFilesConfig._is_titan_v2_embed_model(
"arn:aws:bedrock:us-east-1:123:foundation-model/amazon.titan-embed-text-v2:0"
)
def test_lookup_provider_specific_field_helper(self, mocker):
"""Direct coverage of the nested registry field helper."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
# Happy path: returns the nested field's string value
mocker.patch(
"litellm.get_model_info",
return_value={
"provider_specific_entry": {"bedrock_invocation_schema": "titan_v2"}
},
)
assert (
BedrockFilesConfig._lookup_provider_specific_field(
"anything", "bedrock_invocation_schema"
)
== "titan_v2"
)
# Registry raises -> None
mocker.patch("litellm.get_model_info", side_effect=Exception("not mapped"))
assert (
BedrockFilesConfig._lookup_provider_specific_field("anything", "any")
is None
)
# Registry returns non-dict -> None
mocker.patch("litellm.get_model_info", return_value="not a dict")
assert (
BedrockFilesConfig._lookup_provider_specific_field("anything", "any")
is None
)
# Registry returns dict without provider_specific_entry -> None
mocker.patch("litellm.get_model_info", return_value={"mode": "embedding"})
assert (
BedrockFilesConfig._lookup_provider_specific_field(
"anything", "bedrock_invocation_schema"
)
is None
)
# provider_specific_entry exists but isn't a dict -> None
mocker.patch(
"litellm.get_model_info",
return_value={"provider_specific_entry": "not a dict"},
)
assert (
BedrockFilesConfig._lookup_provider_specific_field(
"anything", "bedrock_invocation_schema"
)
is None
)
# provider_specific_entry dict missing the requested field -> None
mocker.patch(
"litellm.get_model_info",
return_value={"provider_specific_entry": {"unrelated": "x"}},
)
assert (
BedrockFilesConfig._lookup_provider_specific_field(
"anything", "bedrock_invocation_schema"
)
is None
)
# Non-string nested value -> None
mocker.patch(
"litellm.get_model_info",
return_value={"provider_specific_entry": {"bedrock_invocation_schema": 42}},
)
assert (
BedrockFilesConfig._lookup_provider_specific_field(
"anything", "bedrock_invocation_schema"
)
is None
)
# Empty-string nested value -> None
mocker.patch(
"litellm.get_model_info",
return_value={"provider_specific_entry": {"bedrock_invocation_schema": ""}},
)
assert (
BedrockFilesConfig._lookup_provider_specific_field(
"anything", "bedrock_invocation_schema"
)
is None
)
def test_is_embedding_record_helper(self):
"""Helper detects embeddings via `url` first, then by body shape."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
assert BedrockFilesConfig._is_embedding_record(
{"url": "/v1/embeddings", "body": {"input": "x"}}
)
# body-only fallback
assert BedrockFilesConfig._is_embedding_record({"body": {"input": "x"}})
# chat shape
assert not BedrockFilesConfig._is_embedding_record(
{"url": "/v1/chat/completions", "body": {"messages": []}}
)
# ambiguous body without `input` is treated as not-embedding
assert not BedrockFilesConfig._is_embedding_record({"body": {}})
def test_explicit_chat_url_with_input_body_short_circuits_to_chat(self):
"""Explicit url=/v1/chat/completions wins even if body looks like embedding.
Without this short-circuit, a chat record whose body happens to carry
`input` (and no `messages`) would be mis-routed to the embedding
transformer, corrupting the modelInput.
"""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
# Direct helper assertion
assert not BedrockFilesConfig._is_embedding_record(
{
"url": "/v1/chat/completions",
"body": {
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"input": "this would mis-route under the old precedence",
},
}
)
# End-to-end: a record like this routes through the chat path. We
# just need to make sure we DON'T silently produce an inputText
# body and call it a chat completion.
config = BedrockFilesConfig()
result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content(
[
{
"custom_id": "explicit-chat-with-input",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"messages": [{"role": "user", "content": "Hi"}],
"input": "should not become inputText",
"max_tokens": 5,
},
}
]
)
model_input = result[0]["modelInput"]
assert (
"inputText" not in model_input
), "explicit chat URL must not produce an embedding-shaped modelInput"
def test_coerce_embedding_input_helper_isolated(self):
"""Direct coverage of the extracted input-normalization helper."""
import pytest
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
# Happy paths
assert BedrockFilesConfig._coerce_embedding_input_to_string("hello") == "hello"
assert (
BedrockFilesConfig._coerce_embedding_input_to_string(["hello"]) == "hello"
)
# Error paths
with pytest.raises(ValueError, match="missing required `input`"):
BedrockFilesConfig._coerce_embedding_input_to_string(None, model="m")
with pytest.raises(ValueError, match="one input per JSONL record"):
BedrockFilesConfig._coerce_embedding_input_to_string(["a", "b"])
# A multi-element list of ints is rejected as "one input per JSONL
# record" too - we can't tell if it's pre-tokenized or "3 strings"
# without more context, so the most-actionable error wins.
with pytest.raises(ValueError, match="one input per JSONL record"):
BedrockFilesConfig._coerce_embedding_input_to_string([1, 2, 3])
# Single-element list wrapping a token list -> pre-tokenized error.
with pytest.raises(NotImplementedError, match="pre-tokenized"):
BedrockFilesConfig._coerce_embedding_input_to_string([[1, 2, 3]])
# Single-element list wrapping a bare int -> pre-tokenized error.
with pytest.raises(NotImplementedError, match="pre-tokenized"):
BedrockFilesConfig._coerce_embedding_input_to_string([42])
with pytest.raises(ValueError, match="must be a string"):
BedrockFilesConfig._coerce_embedding_input_to_string({"unsupported": True})
def test_other_non_embedding_urls_route_to_chat(self):
"""Any non-/v1/embeddings url short-circuits to chat path."""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
# /v1/completions (legacy completions endpoint)
assert not BedrockFilesConfig._is_embedding_record(
{"url": "/v1/completions", "body": {"input": "x"}}
)
# Arbitrary unknown url - caller's explicit signal still wins
assert not BedrockFilesConfig._is_embedding_record(
{"url": "/v1/responses", "body": {"input": "x"}}
)

View file

@ -1,8 +1,9 @@
import json
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import httpx
import pytest
from botocore.credentials import Credentials
def _anthropic_response(url: str) -> httpx.Response:
@ -310,3 +311,54 @@ async def test_anthropic_messages_routes_bedrock_claude_platform_to_messages_api
assert requests[0]["body"]["messages"] == [{"role": "user", "content": "hello"}]
assert requests[0]["body"]["max_tokens"] == 10
assert requests[0]["body"]["model"] == "claude-sonnet-4-6"
def test_sigv4_no_duplicate_content_type_when_caller_sets_lowercase():
"""
Regression: get_anthropic_headers() supplies "content-type" (lowercase).
_sign_request() used to prepend "Content-Type" (uppercase), leaving both
keys in the dict. botocore joins them into "application/json, application/json"
in the canonical string, while the wire request sends only one value → 401.
Fix: prepend with lowercase "content-type" so **headers overwrites it when
the caller already set it.
"""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
llm = BaseAWSLLM()
mock_credentials = Credentials("key", "secret", "token")
mock_sigv4 = MagicMock()
captured: list[dict] = []
def fake_aws_request(method, url, data, headers):
captured.append(dict(headers))
req = MagicMock()
req.headers = {"Authorization": "AWS4-HMAC-SHA256 Credential=test"}
req.body = data.encode() if isinstance(data, str) else data
return req
with (
patch("botocore.auth.SigV4Auth", return_value=mock_sigv4),
patch("botocore.awsrequest.AWSRequest", side_effect=fake_aws_request),
patch.object(llm, "get_credentials", return_value=mock_credentials),
patch.object(llm, "_get_aws_region_name", return_value="us-east-1"),
):
llm._sign_request(
service_name="aws-external-anthropic",
headers={"content-type": "application/json"},
optional_params={"aws_region_name": "us-east-1"},
request_data={
"model": "claude-sonnet-4-6",
"messages": [],
"max_tokens": 10,
},
api_base="https://aws-external-anthropic.us-east-1.api.aws/v1/messages",
)
signed = captured[0]
ct_keys = [k for k in signed if k.lower() == "content-type"]
assert ct_keys == ["content-type"], (
f"Expected exactly one 'content-type' key, got {ct_keys}. "
"Duplicate keys produce 'application/json, application/json' in the "
"SigV4 canonical string and cause a 401."
)

View file

@ -0,0 +1,67 @@
"""
Tests for Black Forest Labs common_utils — specifically assert_bfl_polling_url.
BFL uses regional subdomains (e.g. gateway.bfl.ai) for polling URLs that
differ from the submission host (api.bfl.ai). These tests verify that the
domain-aware check accepts legitimate BFL subdomains while still rejecting
off-domain and non-HTTPS URLs.
"""
import pytest
from litellm.llms.black_forest_labs.common_utils import (
BlackForestLabsError,
assert_bfl_polling_url,
)
class TestAssertBflPollingUrl:
# --- should pass ---
def test_exact_registered_domain(self):
assert_bfl_polling_url("https://bfl.ai/v1/get_result?id=abc")
def test_api_subdomain(self):
assert_bfl_polling_url("https://api.bfl.ai/v1/get_result?id=abc")
def test_gateway_subdomain(self):
# BFL uses gateway.bfl.ai for polling — this was the original bug trigger
assert_bfl_polling_url("https://gateway.bfl.ai/v1/get_result?id=abc")
def test_regional_subdomain(self):
assert_bfl_polling_url("https://eu.api.bfl.ai/v1/get_result?id=abc")
def test_deep_subdomain(self):
assert_bfl_polling_url("https://region.gateway.bfl.ai/poll?id=xyz")
# --- should raise BlackForestLabsError ---
def test_rejects_http_scheme(self):
# HTTP must be rejected — x-key would be forwarded in plaintext
with pytest.raises(BlackForestLabsError, match="scheme must be https"):
assert_bfl_polling_url("http://api.bfl.ai/v1/get_result?id=abc")
def test_rejects_off_domain(self):
with pytest.raises(BlackForestLabsError, match="host is not within"):
assert_bfl_polling_url("https://evil.com/steal-key")
def test_rejects_lookalike_domain(self):
with pytest.raises(BlackForestLabsError, match="host is not within"):
assert_bfl_polling_url("https://notbfl.ai/v1/get_result?id=abc")
def test_rejects_bfl_ai_as_suffix_only(self):
# "fakebfl.ai" must not match — the check is on registered domain boundary
with pytest.raises(BlackForestLabsError, match="host is not within"):
assert_bfl_polling_url("https://fakebfl.ai/v1/get_result?id=abc")
def test_rejects_bfl_in_path(self):
with pytest.raises(BlackForestLabsError, match="host is not within"):
assert_bfl_polling_url("https://evil.com/bfl.ai/steal")
def test_rejects_ftp_scheme(self):
with pytest.raises(BlackForestLabsError, match="scheme must be https"):
assert_bfl_polling_url("ftp://api.bfl.ai/v1/get_result?id=abc")
def test_rejects_javascript_scheme(self):
with pytest.raises(BlackForestLabsError, match="scheme must be https"):
assert_bfl_polling_url("javascript://api.bfl.ai/alert(1)")

View file

@ -329,3 +329,170 @@ def test_transform_messages_helper_strips_thinking_blocks():
)
assert "thinking_blocks" not in out[1]
assert out[1]["content"] == "I can help."
# -----------------------------------------------------------------------------
# Regression tests for legacy / OpenAPI $ref defs in tool parameters.
#
# Fireworks (like Anthropic) only resolves `$defs` (JSON Schema 2020-12). Tools
# coming from MCP servers (legacy `definitions`) or OpenAPI-derived gateways
# such as AWS AgentCore (`components.schemas`) used to leave dangling `$ref`
# pointers, causing upstream "Error resolving schema reference" failures. See
# https://github.com/BerriAI/litellm/issues/26692.
# -----------------------------------------------------------------------------
def _assert_no_unresolved_refs(parameters: dict) -> None:
blob = json.dumps(parameters)
assert "$ref" not in blob, f"unresolved $ref in transformed parameters: {blob}"
def test_transform_tools_inlines_components_schemas_refs():
"""OpenAPI `components.schemas` $refs (AgentCore-style) must be inlined."""
config = FireworksAIConfig()
tools = [
{
"type": "function",
"function": {
"name": "slides_presentations_create",
"description": "Create a Google Slides presentation",
"parameters": {
"type": "object",
"properties": {
"body": {"$ref": "#/components/schemas/Presentation"},
},
"required": ["body"],
"components": {
"schemas": {
"Presentation": {
"type": "object",
"properties": {
"title": {"type": "string"},
"presentationId": {"type": "string"},
},
}
}
},
},
},
}
]
out = config._transform_tools(tools)
params = out[0]["function"]["parameters"]
_assert_no_unresolved_refs(params)
assert params["properties"]["body"] == {
"type": "object",
"properties": {
"title": {"type": "string"},
"presentationId": {"type": "string"},
},
}
assert "components" not in params
def test_transform_tools_inlines_legacy_definitions_refs():
"""Legacy draft-04 `definitions` $refs must be inlined."""
config = FireworksAIConfig()
tools = [
{
"type": "function",
"function": {
"name": "create_thing",
"description": "Create a thing",
"parameters": {
"type": "object",
"properties": {"thing": {"$ref": "#/definitions/Thing"}},
"definitions": {
"Thing": {
"type": "object",
"properties": {"id": {"type": "string"}},
}
},
},
},
}
]
out = config._transform_tools(tools)
params = out[0]["function"]["parameters"]
_assert_no_unresolved_refs(params)
assert params["properties"]["thing"] == {
"type": "object",
"properties": {"id": {"type": "string"}},
}
assert "definitions" not in params
def test_transform_tools_preserves_native_dollar_defs():
"""`$defs` is JSON Schema 2020-12 native; Fireworks resolves it itself."""
config = FireworksAIConfig()
tools = [
{
"type": "function",
"function": {
"name": "native_defs_tool",
"description": "",
"parameters": {
"type": "object",
"properties": {"a": {"$ref": "#/$defs/A"}},
"$defs": {"A": {"type": "string"}},
},
},
}
]
out = config._transform_tools(tools)
params = out[0]["function"]["parameters"]
assert params["$defs"] == {"A": {"type": "string"}}
assert params["properties"]["a"] == {"$ref": "#/$defs/A"}
def test_transform_tools_skips_non_function_tools():
"""Non-``function`` tools (e.g. provider-native tool types) must pass
through ``_transform_tools`` untouched -- no ``strict`` pop, no $ref
inlining, no error.
"""
config = FireworksAIConfig()
non_function_tool = {
"type": "code_interpreter",
"code_interpreter": {"some": "config"},
}
function_tool = {
"type": "function",
"function": {
"name": "create_thing",
"description": "Create a thing",
"parameters": {
"type": "object",
"properties": {"thing": {"$ref": "#/definitions/Thing"}},
"definitions": {
"Thing": {
"type": "object",
"properties": {"id": {"type": "string"}},
}
},
},
"strict": True,
},
}
out = config._transform_tools([non_function_tool, function_tool])
# Non-function tool is preserved verbatim.
assert out[0] == {
"type": "code_interpreter",
"code_interpreter": {"some": "config"},
}
# Function tool still goes through both transformations: `strict` popped
# and the legacy $ref inlined.
assert "strict" not in out[1]["function"]
inlined = out[1]["function"]["parameters"]
assert "definitions" not in inlined
assert inlined["properties"]["thing"] == {
"type": "object",
"properties": {"id": {"type": "string"}},
}

View file

@ -1,17 +1,14 @@
import json
import os
import sys
import pytest
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
import litellm
from litellm.llms.lemonade.chat.transformation import LemonadeChatConfig
from litellm.types.utils import ModelResponse
import httpx
def test_lemonade_config_initialization():
@ -28,8 +25,11 @@ def test_lemonade_config_initialization():
assert config.repeat_penalty == 1.1
def test_get_openai_compatible_provider_info():
def test_get_openai_compatible_provider_info(monkeypatch):
"""Test the provider info method returns correct API base and key"""
monkeypatch.delenv("LEMONADE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
@ -40,8 +40,11 @@ def test_get_openai_compatible_provider_info():
assert key == "lemonade"
def test_get_openai_compatible_provider_info_with_custom_base():
def test_get_openai_compatible_provider_info_with_custom_base(monkeypatch):
"""Test the provider info method with custom API base"""
monkeypatch.delenv("LEMONADE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
custom_api_base = "https://custom.lemonade.ai/v1"
@ -53,6 +56,335 @@ def test_get_openai_compatible_provider_info_with_custom_base():
assert key == "lemonade"
def test_get_openai_compatible_provider_info_with_api_key_env(monkeypatch):
"""Test the provider info method reads Lemonade's API key from the environment."""
monkeypatch.setenv("LEMONADE_API_KEY", "test-key")
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base=None, api_key=None
)
assert api_base == "http://localhost:8000/api/v1"
assert key == "test-key"
def test_get_openai_compatible_provider_info_skips_env_key_for_custom_base(
monkeypatch,
):
"""Test that caller-supplied bases do not receive server-side Lemonade keys."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="https://attacker.example/v1", api_key=None
)
assert api_base == "https://attacker.example/v1"
assert key == "lemonade"
assert config._get_auth_headers(key) == {}
def test_get_openai_compatible_provider_info_uses_explicit_key_for_custom_base(
monkeypatch,
):
"""Test that explicitly supplied Lemonade keys are sent to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="https://lemonade.example/v1", api_key="explicit-lemonade-key"
)
assert api_base == "https://lemonade.example/v1"
assert key == "explicit-lemonade-key"
assert config._get_auth_headers(key) == {
"Authorization": "Bearer explicit-lemonade-key"
}
def test_get_openai_compatible_provider_info_empty_key_does_not_leak_to_custom_base(
monkeypatch,
):
"""An empty explicit key must not fall back to server-side Lemonade creds for a custom base."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="https://attacker.example/v1", api_key=""
)
assert api_base == "https://attacker.example/v1"
assert key == "lemonade"
assert config._get_auth_headers(key) == {}
def test_get_openai_compatible_provider_info_ignores_global_api_key(monkeypatch):
"""Test that Lemonade discovery does not send unrelated global API keys."""
monkeypatch.delenv("LEMONADE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", "global-openai-key")
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="http://lemonade.test/v1", api_key=None
)
assert api_base == "http://lemonade.test/v1"
assert key == "lemonade"
assert config._get_auth_headers(key) == {}
def test_get_models_does_not_leak_lemonade_key_to_custom_base(monkeypatch):
"""Test Lemonade discovery does not send server-side keys to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {"data": []}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
models = config.get_models(api_base="https://attacker.example/v1")
assert models == []
assert mock_get.call_args.kwargs["headers"] == {}
def test_get_model_info_uses_loaded_context_size():
"""Test that Lemonade model info prefers the effective loaded ctx_size."""
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"recipe_options": {"ctx_size": 65536},
"max_context_window": 262144,
}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
model_info = config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["key"] == "lemonade/Qwen3.6-35B-A3B-GGUF"
assert model_info["litellm_provider"] == "lemonade"
assert model_info["max_input_tokens"] == 65536
assert model_info["provider_specific_entry"] == {
"recipe_options": {"ctx_size": 65536},
"max_context_window": 262144,
}
assert "supports_function_calling" not in model_info
assert "supports_response_schema" not in model_info
assert "supports_tool_choice" not in model_info
assert mock_get.call_args.kwargs["headers"] == {}
def test_get_model_info_falls_back_when_server_unavailable():
"""Test that Lemonade metadata lookup failures return safe defaults."""
config = LemonadeChatConfig()
with patch.object(
litellm.module_level_client, "get", side_effect=Exception("boom")
):
model_info = config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["key"] == "lemonade/Qwen3.6-35B-A3B-GGUF"
assert model_info["litellm_provider"] == "lemonade"
assert model_info["mode"] == "chat"
assert model_info["input_cost_per_token"] == 0.0
assert model_info["output_cost_per_token"] == 0.0
assert model_info["max_tokens"] is None
assert model_info["max_input_tokens"] is None
assert model_info["max_output_tokens"] is None
assert "supports_function_calling" not in model_info
assert "supports_response_schema" not in model_info
assert "supports_tool_choice" not in model_info
def test_get_model_info_reads_context_from_provider_specific_entry():
"""Test that Lemonade model info uses provider-specific runtime metadata."""
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"provider_specific_entry": {
"recipe_options": {"ctx_size": "32768"},
"max_context_window": 262144,
},
}
with patch.object(litellm.module_level_client, "get", return_value=response):
model_info = config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["max_input_tokens"] == 32768
assert model_info["provider_specific_entry"] == {
"recipe_options": {"ctx_size": "32768"},
"max_context_window": 262144,
}
def test_get_model_info_sends_lemonade_api_key_for_configured_base(monkeypatch):
"""Test that Lemonade model info uses auth for configured servers."""
monkeypatch.setenv("LEMONADE_API_KEY", "test-key")
monkeypatch.setenv("LEMONADE_API_BASE", "http://lemonade.test/v1")
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"recipe_options": {"ctx_size": 65536},
}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
)
assert mock_get.call_args.kwargs["headers"] == {"Authorization": "Bearer test-key"}
def test_get_model_info_sends_explicit_lemonade_api_key_for_custom_base(monkeypatch):
"""Test that Lemonade model info sends explicitly supplied auth to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-key")
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"recipe_options": {"ctx_size": 65536},
}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
api_key="explicit-test-key",
)
assert mock_get.call_args.kwargs["headers"] == {
"Authorization": "Bearer explicit-test-key"
}
def test_litellm_get_model_info_does_not_leak_lemonade_key_to_custom_base(
monkeypatch,
):
"""Test top-level model info does not send server-side keys to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"max_input_tokens": 65536,
"max_context_window": 262144,
}
litellm.get_model_info.cache_clear()
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
try:
model_info = litellm.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="https://attacker.example/v1",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 65536
assert mock_get.call_args.kwargs["headers"] == {}
def test_litellm_get_model_info_forwards_explicit_lemonade_key_to_custom_base(
monkeypatch,
):
"""Top-level model info must forward an explicit api_key to the supplied base."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"max_input_tokens": 65536,
}
litellm.get_model_info.cache_clear()
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
try:
model_info = litellm.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="https://lemonade.example/v1",
api_key="explicit-lemonade-key",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 65536
assert mock_get.call_args.kwargs["headers"] == {
"Authorization": "Bearer explicit-lemonade-key"
}
def test_litellm_get_model_info_uses_lemonade_api_base():
"""Test that LiteLLM model info is wired to Lemonade's model metadata API."""
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"max_input_tokens": 65536,
"max_context_window": 262144,
}
litellm.get_model_info.cache_clear()
with patch.object(litellm.module_level_client, "get", return_value=response):
try:
model_info = litellm.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 65536
assert response.raise_for_status.called
assert response.json.called
def test_transform_response():
"""Test the response transformation adds lemonade prefix to model name"""
config = LemonadeChatConfig()

View file

@ -1,6 +1,5 @@
import os
import sys
from unittest.mock import patch
import pytest
@ -23,6 +22,7 @@ if "httpx" not in sys.modules:
sys.modules["httpx"] = httpx_mod
import httpx
import litellm
from litellm.llms.ollama.common_utils import OllamaModelInfo
@ -105,6 +105,68 @@ class TestOllamaModelInfo:
"Authorization": "Bearer test_api_key"
}
def test_get_models_does_not_leak_server_key_to_provided_api_base(
self, monkeypatch
):
"""Model discovery should not send server-side keys to caller-supplied bases."""
call_headers = []
def mock_get(url, headers):
call_headers.append(headers)
return DummyResponse({"models": []}, status_code=200)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
monkeypatch.setattr(litellm, "openai_key", "global-openai-key")
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models(api_base="https://attacker.example")
assert models == []
assert call_headers[0] == {}
def test_get_models_uses_explicit_api_key_for_provided_api_base(self, monkeypatch):
"""Model discovery should send an explicitly supplied key to the provided base."""
call_headers = []
def mock_get(url, headers):
call_headers.append(headers)
return DummyResponse({"models": []}, status_code=200)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models(
api_base="https://ollama.example",
api_key="explicit-api-key",
)
assert models == []
assert call_headers[0] == {"Authorization": "Bearer explicit-api-key"}
def test_get_models_empty_key_does_not_leak_to_provided_api_base(
self, monkeypatch
):
"""An empty explicit key must not fall back to server-side creds for a custom base."""
call_headers = []
def mock_get(url, headers):
call_headers.append(headers)
return DummyResponse({"models": []}, status_code=200)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
monkeypatch.setattr(litellm, "openai_key", "global-openai-key")
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models(api_base="https://attacker.example", api_key="")
assert models == []
assert call_headers[0] == {}
def test_get_models_from_list_response(self, monkeypatch):
"""
When the /api/tags endpoint returns a list of dicts,
@ -190,7 +252,7 @@ class TestOllamaGetModelInfo:
config = OllamaConfig()
result = config.get_model_info(
"llama3", api_base="http://my-remote-server:11434"
"my-custom-model", api_base="http://my-remote-server:11434"
)
assert captured_urls[0] == "http://my-remote-server:11434/api/show"
@ -200,6 +262,181 @@ class TestOllamaGetModelInfo:
"""When no api_base is passed, should fall back to OLLAMA_API_BASE env var."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
captured_urls = []
captured_headers = []
def mock_post(url, json, headers=None):
captured_urls.append(url)
captured_headers.append(headers)
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_BASE", "http://env-server:11434")
monkeypatch.setenv("OLLAMA_API_KEY", "env-api-key")
config = OllamaConfig()
config.get_model_info("my-custom-model")
assert captured_urls[0] == "http://env-server:11434/api/show"
assert captured_headers[0] == {"Authorization": "Bearer env-api-key"}
def test_get_model_info_uses_explicit_api_key_for_provided_api_base(
self, monkeypatch
):
"""When api_key is explicit, model info should send it to the provided api_base."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
config = OllamaConfig()
config.get_model_info(
"my-custom-model",
api_base="http://my-remote-server:11434",
api_key="explicit-api-key",
)
assert captured_headers[0] == {"Authorization": "Bearer explicit-api-key"}
def test_get_model_info_empty_key_does_not_leak_to_provided_api_base(
self, monkeypatch
):
"""An empty explicit key must not fall back to server-side creds for a custom base."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
monkeypatch.setattr(litellm, "openai_key", "global-openai-key")
config = OllamaConfig()
config.get_model_info(
"my-custom-model",
api_base="https://attacker.example",
api_key="",
)
assert captured_headers[0] == {}
def test_litellm_get_model_info_does_not_leak_server_key_to_provided_api_base(
self, monkeypatch
):
"""Global model info should not send server-side keys to caller-supplied bases."""
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
return DummyResponse(
{
"template": "{{ .System }} tools {{ .Prompt }}",
"model_info": {"llama.context_length": 32768},
},
status_code=200,
)
litellm.get_model_info.cache_clear()
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
monkeypatch.setattr(litellm, "openai_key", "global-openai-key")
try:
model_info = litellm.get_model_info(
"ollama/unknown-model",
api_base="https://attacker.example",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 32768
assert captured_headers[0] == {}
def test_litellm_get_model_info_forwards_explicit_api_key_to_provided_base(
self, monkeypatch
):
"""An explicit api_key passed to litellm.get_model_info must reach the provided base."""
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
return DummyResponse(
{
"template": "{{ .System }} tools {{ .Prompt }}",
"model_info": {"llama.context_length": 32768},
},
status_code=200,
)
litellm.get_model_info.cache_clear()
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
try:
model_info = litellm.get_model_info(
"ollama/unknown-model",
api_base="https://ollama.example",
api_key="explicit-api-key",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 32768
assert captured_headers[0] == {"Authorization": "Bearer explicit-api-key"}
def test_litellm_get_model_info_does_not_cache_on_api_key(self, monkeypatch):
"""Regression: api_key must not be part of the get_model_info cache key.
Distinct api_keys for the same (model, api_base) must not each create their
own cache entry (which would churn the shared LRU cache), and every explicit
key must still reach the backend rather than be served from a result cached
with a different key.
"""
from litellm.utils import _cached_get_model_info
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
return DummyResponse(
{
"template": "{{ .System }} tools {{ .Prompt }}",
"model_info": {"llama.context_length": 32768},
},
status_code=200,
)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
litellm.get_model_info.cache_clear()
try:
for api_key in ("key-one", "key-two", "key-three"):
litellm.get_model_info(
"ollama/unknown-model",
api_base="https://ollama.example",
api_key=api_key,
)
assert _cached_get_model_info.cache_info().currsize <= 1
assert captured_headers == [
{"Authorization": "Bearer key-one"},
{"Authorization": "Bearer key-two"},
{"Authorization": "Bearer key-three"},
]
finally:
litellm.get_model_info.cache_clear()
def test_get_model_info_normalizes_generate_api_base(self, monkeypatch):
"""When completion passes the final generate URL, model info should use the server base."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
captured_urls = []
def mock_post(url, json, headers=None):
@ -207,12 +444,13 @@ class TestOllamaGetModelInfo:
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_BASE", "http://env-server:11434")
config = OllamaConfig()
config.get_model_info("llama3")
config.get_model_info(
"my-custom-model", api_base="http://localhost:11434/api/generate"
)
assert captured_urls[0] == "http://env-server:11434/api/show"
assert captured_urls[0] == "http://localhost:11434/api/show"
def test_get_model_info_graceful_fallback_on_connection_error(self, monkeypatch):
"""When the Ollama server is unreachable, should return defaults instead of raising."""
@ -225,14 +463,42 @@ class TestOllamaGetModelInfo:
monkeypatch.delenv("OLLAMA_API_BASE", raising=False)
config = OllamaConfig()
result = config.get_model_info("llama3", api_base="http://unreachable:11434")
result = config.get_model_info(
"my-custom-model", api_base="http://unreachable:11434"
)
assert result["key"] == "llama3"
assert result["key"] == "my-custom-model"
assert result["litellm_provider"] == "ollama"
assert result["input_cost_per_token"] == 0.0
assert result["output_cost_per_token"] == 0.0
assert result["max_tokens"] is None
def test_get_model_info_graceful_fallback_on_http_error_status(self, monkeypatch):
"""A non-2xx /api/show response must fall back to defaults, not parse the error body."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
def mock_post(url, json, headers=None):
return DummyResponse(
{
"template": "{{ .System }} tools {{ .Prompt }}",
"model_info": {"llama.context_length": 8192},
},
status_code=404,
)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
config = OllamaConfig()
result = config.get_model_info(
"my-custom-model", api_base="http://localhost:11434"
)
assert result["key"] == "my-custom-model"
assert result["litellm_provider"] == "ollama"
assert result["max_tokens"] is None
assert result["max_input_tokens"] is None
assert "supports_function_calling" not in result
def test_get_model_info_strips_ollama_prefix(self, monkeypatch):
"""Should strip 'ollama/' or 'ollama_chat/' prefix from model name."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
@ -246,11 +512,72 @@ class TestOllamaGetModelInfo:
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
config = OllamaConfig()
config.get_model_info("ollama/llama3", api_base="http://localhost:11434")
assert captured_json[0]["name"] == "llama3"
config.get_model_info(
"ollama/my-custom-model", api_base="http://localhost:11434"
)
assert captured_json[0]["name"] == "my-custom-model"
config.get_model_info("ollama_chat/llama3", api_base="http://localhost:11434")
assert captured_json[1]["name"] == "llama3"
config.get_model_info(
"ollama_chat/my-custom-model", api_base="http://localhost:11434"
)
assert captured_json[1]["name"] == "my-custom-model"
def test_get_model_info_skips_network_for_static_model(self, monkeypatch):
"""Statically-priced models must not trigger an /api/show network call."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
def mock_post(url, json, headers=None):
raise AssertionError("Static Ollama model should not query /api/show")
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
config = OllamaConfig()
assert config.get_model_info("ollama/llama2") is None
def test_litellm_get_model_info_uses_provider_hook_for_unknown_model(
self, monkeypatch
):
"""Unmapped Ollama models should use the provider-level dynamic hook."""
captured_json = []
def mock_post(url, json, headers=None):
captured_json.append(json)
return DummyResponse(
{
"template": "{{ .System }} tools {{ .Prompt }}",
"model_info": {"llama.context_length": 32768},
},
status_code=200,
)
litellm.get_model_info.cache_clear()
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
try:
model_info = litellm.get_model_info(
"ollama/unknown-model", api_base="http://localhost:11434"
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 32768
assert model_info["supports_function_calling"] is True
assert captured_json[0]["name"] == "unknown-model"
def test_litellm_get_model_info_keeps_static_map_for_known_model(self, monkeypatch):
"""Mapped Ollama models should keep using the static model map."""
def mock_post(url, json, headers=None):
raise AssertionError("Static Ollama model should not query /api/show")
litellm.get_model_info.cache_clear()
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
try:
model_info = litellm.get_model_info("ollama/llama2")
finally:
litellm.get_model_info.cache_clear()
assert model_info["key"] == "ollama/llama2"
assert model_info["litellm_provider"] == "ollama"
class TestOllamaAuthHeaders:

View file

@ -8,8 +8,7 @@ with guardrail transformations, including tool calls.
import json
import os
import sys
from typing import Any, List, Literal, Optional, Tuple
from unittest.mock import AsyncMock, MagicMock
from typing import Any, Literal, Optional
import pytest
@ -84,6 +83,70 @@ class MockGuardrail(CustomGuardrail):
return result
class MockCopiedToolCallGuardrail(CustomGuardrail):
"""Mock guardrail that returns copied tool calls instead of mutating inputs."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tool_calls = inputs.get("tool_calls", [])
copied_tool_calls = []
for tool_call in tool_calls:
copied = dict(tool_call)
function = dict(copied["function"])
function["arguments"] = json.dumps({"email": "[EMAIL]"})
copied["function"] = function
copied_tool_calls.append(copied)
return GenericGuardrailAPIInputs(
texts=inputs.get("texts", []),
tool_calls=copied_tool_calls,
)
class MockNonListToolCallGuardrail(CustomGuardrail):
"""Mock guardrail that returns tool_calls as a non-list envelope on the response
path, as some released guardrails do when they assign a detection API JSON dict."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
result = GenericGuardrailAPIInputs(texts=inputs.get("texts", []))
result["tool_calls"] = {"verdict": "allow", "detections": []} # type: ignore
return result
class MockMisalignedToolCallGuardrail(CustomGuardrail):
"""Mock guardrail that returns a tool_calls list whose length differs from the
input, so it cannot be applied positionally onto the response."""
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
tool_calls = inputs.get("tool_calls", [])
shortened = []
if tool_calls:
first = dict(tool_calls[0])
first["function"] = {"name": "x", "arguments": json.dumps({"x": 1})}
shortened.append(first)
return GenericGuardrailAPIInputs(
texts=inputs.get("texts", []),
tool_calls=shortened,
)
class TestOpenAIChatCompletionsHandlerToolsInput:
"""Test input processing with tools (function definitions)"""
@ -740,6 +803,131 @@ class TestOpenAIChatCompletionsHandlerToolCallsOutput:
assert response.model == "gpt-4o-mini"
assert response.choices[0].finish_reason == "tool_calls"
@pytest.mark.asyncio
async def test_output_response_uses_returned_guardrailed_tool_calls(self):
"""Test returned tool_calls are remapped even when guardrail does not mutate inputs."""
handler = OpenAIChatCompletionsHandler()
guardrail = MockCopiedToolCallGuardrail(guardrail_name="test")
response = ModelResponse(
id="chatcmpl-tool-copy",
created=1234567890,
model="gpt-4",
object="chat.completion",
choices=[
Choices(
finish_reason="tool_calls",
index=0,
message=Message(
content=None,
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
id="call_email",
type="function",
function=Function(
name="send_email",
arguments=json.dumps({"email": "john@example.com"}),
),
)
],
),
)
],
)
await handler.process_output_response(response, guardrail)
response_tool_call = response.choices[0].message.tool_calls[0]
assert response_tool_call.function.name == "send_email"
assert json.loads(response_tool_call.function.arguments) == {"email": "[EMAIL]"}
@pytest.mark.asyncio
async def test_output_response_ignores_non_list_returned_tool_calls(self):
"""A guardrail returning tool_calls as a non-list (e.g. a detection-API envelope
dict) must not crash the remap; the original arguments are preserved."""
handler = OpenAIChatCompletionsHandler()
guardrail = MockNonListToolCallGuardrail(guardrail_name="test")
original = json.dumps({"email": "john@example.com"})
response = ModelResponse(
id="chatcmpl-nonlist",
created=1234567890,
model="gpt-4",
object="chat.completion",
choices=[
Choices(
finish_reason="tool_calls",
index=0,
message=Message(
content=None,
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
id="call_email",
type="function",
function=Function(
name="send_email", arguments=original
),
)
],
),
)
],
)
await handler.process_output_response(response, guardrail)
response_tool_call = response.choices[0].message.tool_calls[0]
assert response_tool_call.function.arguments == original
@pytest.mark.asyncio
async def test_output_response_ignores_misaligned_returned_tool_calls(self):
"""A guardrail returning a tool_calls list of a different length than the input
cannot be applied positionally; the handler falls back and preserves the
original arguments instead of writing onto the wrong tool call."""
handler = OpenAIChatCompletionsHandler()
guardrail = MockMisalignedToolCallGuardrail(guardrail_name="test")
first_args = json.dumps({"email": "a@example.com"})
second_args = json.dumps({"email": "b@example.com"})
response = ModelResponse(
id="chatcmpl-misaligned",
created=1234567890,
model="gpt-4",
object="chat.completion",
choices=[
Choices(
finish_reason="tool_calls",
index=0,
message=Message(
content=None,
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
id="call_1",
type="function",
function=Function(
name="send_email", arguments=first_args
),
),
ChatCompletionMessageToolCall(
id="call_2",
type="function",
function=Function(
name="send_email", arguments=second_args
),
),
],
),
)
],
)
await handler.process_output_response(response, guardrail)
tool_calls = response.choices[0].message.tool_calls
assert tool_calls[0].function.arguments == first_args
assert tool_calls[1].function.arguments == second_args
class MockPassThroughGuardrail(CustomGuardrail):
"""Mock guardrail that passes through without blocking - for testing streaming fallback behavior"""
@ -765,7 +953,7 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
This test verifies the fix for the bug where accessing chunk.choices[0]
would raise IndexError when a streaming chunk has an empty choices list.
"""
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
from litellm.types.utils import ModelResponseStream
handler = OpenAIChatCompletionsHandler()
guardrail = MockPassThroughGuardrail(guardrail_name="test")

View file

@ -0,0 +1,74 @@
"""
Regression test for issue #28146.
`use_chat_completions_api` is a LiteLLM-internal control flag (it forces the
/responses -> /chat/completions bridge). When set as a model-level param in the
proxy config, it must never be forwarded to the upstream provider's request
body. OpenAI/Anthropic reject unknown body params with HTTP 400.
"""
import os
import sys
from unittest.mock import MagicMock
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
from litellm.types.utils import all_litellm_params
from litellm.utils import get_non_default_completion_params
def test_use_chat_completions_api_is_a_known_litellm_param():
assert "use_chat_completions_api" in all_litellm_params
def test_use_chat_completions_api_not_forwarded_as_provider_param():
forwarded = get_non_default_completion_params(
{"use_chat_completions_api": True, "temperature": 0.5}
)
assert "use_chat_completions_api" not in forwarded
def test_completion_does_not_leak_flag_into_provider_request_body():
mock_response = MagicMock()
mock_response.model_dump.return_value = {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 1234567890,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 1,
"completion_tokens": 1,
"total_tokens": 2,
},
}
mock_raw_response = MagicMock()
mock_raw_response.headers = {}
mock_raw_response.parse.return_value = mock_response
mock_client = MagicMock()
mock_client.chat.completions.with_raw_response.create.return_value = (
mock_raw_response
)
litellm.completion(
model="openai/gpt-4o-mini",
messages=[{"role": "user", "content": "hi"}],
use_chat_completions_api=True,
api_key="sk-test",
client=mock_client,
)
create_kwargs = (
mock_client.chat.completions.with_raw_response.create.call_args.kwargs
)
assert "use_chat_completions_api" not in create_kwargs
assert "use_chat_completions_api" not in (create_kwargs.get("extra_body") or {})

View file

@ -0,0 +1,84 @@
"""
Tests for Tensormesh provider configuration and integration.
"""
import litellm
class TestTensormeshProviderConfig:
"""Test Tensormesh provider configuration"""
def test_tensormesh_in_provider_list(self):
"""Test that tensormesh is in the provider list"""
from litellm import LlmProviders
assert hasattr(LlmProviders, "TENSORMESH")
assert LlmProviders.TENSORMESH.value == "tensormesh"
assert "tensormesh" in litellm.provider_list
def test_tensormesh_json_config_exists(self):
"""Test that tensormesh is configured in providers.json"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
assert JSONProviderRegistry.exists("tensormesh")
tensormesh = JSONProviderRegistry.get("tensormesh")
assert tensormesh is not None
assert tensormesh.base_url == "https://serverless.tensormesh.ai/v1"
assert tensormesh.api_key_env == "TENSORMESH_INFERENCE_API_KEY"
assert tensormesh.api_base_env == "TENSORMESH_SERVERLESS_BASE_URL"
assert tensormesh.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_tensormesh_provider_resolution(self):
"""Test that provider resolution finds tensormesh and the default base URL"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="tensormesh/openai/gpt-oss-120b",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "openai/gpt-oss-120b"
assert provider == "tensormesh"
assert api_base == "https://serverless.tensormesh.ai/v1"
def test_tensormesh_api_base_override(self):
"""Test that an explicit api_base / api_key overrides the serverless default"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="tensormesh/openai/gpt-oss-120b",
custom_llm_provider=None,
api_base="https://custom.example.com/v1",
api_key="sk-test",
)
assert provider == "tensormesh"
assert api_base == "https://custom.example.com/v1"
assert api_key == "sk-test"
def test_tensormesh_text_completion_enabled(self):
"""Tensormesh is wired for the /completions (text completion) route,
matching the text_completion flag in provider_endpoints_support.json."""
assert "tensormesh" in litellm.openai_text_completion_compatible_providers
def test_tensormesh_router_config(self):
"""Test that tensormesh can be used in Router configuration"""
from litellm import Router
router = Router(
model_list=[
{
"model_name": "tensormesh-chat",
"litellm_params": {
"model": "tensormesh/openai/gpt-oss-120b",
"api_key": "test-key",
},
}
]
)
assert len(router.model_list) == 1
assert router.model_list[0]["model_name"] == "tensormesh-chat"

View file

@ -1326,6 +1326,63 @@ def test_vertex_ai_zai_is_partner_model():
assert VertexAIPartnerModels.is_vertex_partner_model("zai-org/glm-4.7-maas")
def test_vertex_ai_gemma_maas_is_partner_model():
"""
Ensure Gemma MaaS models are detected as Vertex AI partner models so they
route through the OpenAI-compatible /endpoints/openapi path (not the
legacy non-gemini path or the vertex_ai/gemma/ predict-endpoint handler).
"""
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIPartnerModels,
)
assert VertexAIPartnerModels.is_vertex_partner_model(
"google/gemma-4-26b-a4b-it-maas"
)
def test_vertex_ai_gemma_maas_uses_openai_handler():
"""
Ensure Gemma MaaS partner models re-use the OpenAI-format handler.
"""
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIPartnerModels,
)
assert VertexAIPartnerModels.should_use_openai_handler(
"google/gemma-4-26b-a4b-it-maas"
)
def test_vertex_ai_gemma_maas_routes_to_partner_models():
"""
Regression guard for owtaylor's worry that Gemma MaaS could be misrouted as
a gemma model. get_vertex_ai_model_route must return PARTNER_MODELS, never
GEMMA, MODEL_GARDEN, or NON_GEMINI.
"""
from litellm.llms.vertex_ai.common_utils import (
VertexAIModelRoute,
get_vertex_ai_model_route,
)
route = get_vertex_ai_model_route("google/gemma-4-26b-a4b-it-maas")
assert route == VertexAIModelRoute.PARTNER_MODELS
def test_vertex_ai_google_gemini_not_detected_as_gemma_maas():
"""
Negative: adding the "google/gemma-" prefix must not widen detection to
other google/* models like google/gemini-* (which should keep flowing
through the gemini route, not partner_models).
"""
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIPartnerModels,
)
assert not VertexAIPartnerModels.is_vertex_partner_model("google/gemini-1.5-pro")
assert not VertexAIPartnerModels.should_use_openai_handler("google/gemini-1.5-pro")
def test_build_vertex_schema_empty_properties():
"""
Test _build_vertex_schema handles empty properties objects correctly.

View file

@ -38,3 +38,91 @@ def test_should_reject_dot_segment_vertex_search_vector_store_id():
"vector_store_id": "..",
},
)
def test_should_use_engines_url_when_engine_id_provided():
config = VertexSearchAPIVectorStoreConfig()
url = config.get_complete_url(
api_base=None,
litellm_params={
"vertex_project": "test-project",
"vertex_location": "global",
"vertex_engine_id": "test-engine_1234",
},
)
assert url == (
"https://discoveryengine.googleapis.com/v1/"
"projects/test-project/locations/global/"
"collections/default_collection/engines/test-engine_1234/servingConfigs/default_serving_config"
)
def test_engine_id_takes_precedence_over_vector_store_id():
config = VertexSearchAPIVectorStoreConfig()
url = config.get_complete_url(
api_base=None,
litellm_params={
"vertex_project": "test-project",
"vertex_location": "global",
"vertex_engine_id": "test-engine_1234",
"vector_store_id": "ignored-when-engine-set",
},
)
assert "/engines/test-engine_1234/" in url
assert "/dataStores/" not in url
assert url.endswith("/servingConfigs/default_serving_config")
def test_should_encode_vertex_engine_id_in_complete_url():
config = VertexSearchAPIVectorStoreConfig()
url = config.get_complete_url(
api_base=None,
litellm_params={
"vertex_project": "test-project",
"vertex_location": "global",
"vertex_engine_id": "../../engines/other?x=1#frag",
},
)
assert url == (
"https://discoveryengine.googleapis.com/v1/"
"projects/test-project/locations/global/"
"collections/default_collection/engines/..%2F..%2Fengines%2Fother%3Fx%3D1%23frag/servingConfigs/default_serving_config"
)
def test_should_reject_dot_segment_vertex_engine_id():
config = VertexSearchAPIVectorStoreConfig()
with pytest.raises(
ValueError, match="vertex_engine_id cannot be a dot path segment"
):
config.get_complete_url(
api_base=None,
litellm_params={
"vertex_project": "test-project",
"vertex_location": "global",
"vertex_engine_id": "..",
},
)
def test_should_raise_when_neither_engine_id_nor_vector_store_id_provided():
config = VertexSearchAPIVectorStoreConfig()
with pytest.raises(
ValueError,
match="vector_store_id is required when vertex_engine_id is not set",
):
config.get_complete_url(
api_base=None,
litellm_params={
"vertex_project": "test-project",
"vertex_location": "global",
},
)

View file

@ -0,0 +1,441 @@
"""
Tests for Vertex AI Gemma MaaS models that route through the partner-models
OpenAI-compatible path (https://aiplatform.googleapis.com/.../endpoints/openapi).
These tests verify that:
1. The correct global URL is constructed (https://aiplatform.googleapis.com)
2. get_vertex_region resolves to "global" when model_cost says so
3. acompletion() goes through the OpenAI-compatible handler and hits
/endpoints/openapi/chat/completions
4. Function-calling payloads (tools + tool_choice) pass through unchanged
5. Vision/image_url payloads pass through unchanged
"""
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(
0, os.path.abspath("../../../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.llms.vertex_ai import VertexPartnerProvider
# ---------------------------------------------------------------------------
# Model-cost entry used by all tests that need the model to be known
# ---------------------------------------------------------------------------
_GEMMA_MODEL_COST_ENTRY = {
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
"litellm_provider": "vertex_ai-openai_models",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"input_cost_per_token": 1.5e-07,
"output_cost_per_token": 6e-07,
"supported_regions": ["global"],
"supports_function_calling": True,
"supports_tool_choice": True,
"supports_vision": True,
}
}
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture(autouse=True)
def _reset_litellm_http_client_cache():
"""Ensure each test gets a fresh async HTTP client mock."""
from litellm import in_memory_llm_clients_cache
in_memory_llm_clients_cache.flush_cache()
@pytest.fixture(autouse=True)
def clean_vertex_env():
"""Clear Google/Vertex AI environment variables before each test to prevent test isolation issues."""
saved_env = {}
env_vars_to_clear = [
"GOOGLE_APPLICATION_CREDENTIALS",
"GOOGLE_CLOUD_PROJECT",
"VERTEXAI_PROJECT",
"VERTEX_PROJECT",
"VERTEX_LOCATION",
"VERTEX_AI_PROJECT",
]
for var in env_vars_to_clear:
if var in os.environ:
saved_env[var] = os.environ[var]
del os.environ[var]
yield
for var, value in saved_env.items():
os.environ[var] = value
# ---------------------------------------------------------------------------
# Unit tests: region and URL construction
# ---------------------------------------------------------------------------
class TestVertexBaseGetVertexRegionGemma:
"""Test the get_vertex_region method for Gemma MaaS via model_cost lookup."""
def test_global_model_no_user_region_returns_global(self):
vertex_base = VertexBase()
with patch.dict(
litellm.model_cost,
{
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
"supported_regions": ["global"]
}
},
clear=False,
):
result = vertex_base.get_vertex_region(
vertex_region=None,
model="google/gemma-4-26b-a4b-it-maas",
)
assert result == "global"
def test_global_model_with_unsupported_user_region_overrides(self):
vertex_base = VertexBase()
with patch.dict(
litellm.model_cost,
{
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
"supported_regions": ["global"]
}
},
clear=False,
):
result = vertex_base.get_vertex_region(
vertex_region="us-central1",
model="google/gemma-4-26b-a4b-it-maas",
)
assert result == "global"
class TestCreateVertexURLGemma:
"""Test that create_vertex_url produces the expected OpenAI-compatible URL.
Gemma MaaS models reach this code path via should_use_openai_handler(), which
selects VertexPartnerProvider.llama for all OpenAI-compatible partners including
Gemma. test_gemma_routes_through_openai_handler() guards that mapping so the
URL-format tests below are meaningful regression guards for the Gemma path.
"""
def test_gemma_routes_through_openai_handler(self):
"""Gemma MaaS must be routed through the OpenAI-compatible handler.
This is what causes VertexPartnerProvider.llama to be selected downstream,
which in turn generates the /endpoints/openapi URL shape. If this mapping
ever changes, the URL-shape tests below become misleading.
"""
assert VertexAIPartnerModels.should_use_openai_handler(
"google/gemma-4-26b-a4b-it-maas"
), "Gemma MaaS must use the OpenAI-compatible handler (VertexPartnerProvider.llama path)"
def test_global_location_url_format(self):
# VertexPartnerProvider.llama is correct: Gemma MaaS reaches create_vertex_url
# via should_use_openai_handler() → partner = VertexPartnerProvider.llama.
# See test_gemma_routes_through_openai_handler for the routing guard.
url = VertexBase.create_vertex_url(
vertex_location="global",
vertex_project="test-project",
partner=VertexPartnerProvider.llama,
stream=False,
model="google/gemma-4-26b-a4b-it-maas",
)
assert url.startswith("https://aiplatform.googleapis.com")
assert "global-aiplatform.googleapis.com" not in url
assert "/locations/global/" in url
assert url.endswith("/endpoints/openapi/chat/completions")
def test_regional_location_url_format(self):
url = VertexBase.create_vertex_url(
vertex_location="us-central1",
vertex_project="test-project",
partner=VertexPartnerProvider.llama,
stream=False,
model="google/gemma-4-26b-a4b-it-maas",
)
assert url.startswith("https://us-central1-aiplatform.googleapis.com")
assert "/locations/us-central1/" in url
assert url.endswith("/endpoints/openapi/chat/completions")
# ---------------------------------------------------------------------------
# Capability-flag tests: verify get_model_info surfaces the advertised flags
# ---------------------------------------------------------------------------
def test_gemma_maas_supports_function_calling():
"""supports_function_calling=true in model_cost must be surfaced by the utility."""
with patch.dict(litellm.model_cost, _GEMMA_MODEL_COST_ENTRY, clear=False):
assert (
litellm.utils.supports_function_calling(
model="vertex_ai/google/gemma-4-26b-a4b-it-maas"
)
is True
)
def test_gemma_maas_supports_vision():
"""supports_vision=true in model_cost must be surfaced by the utility."""
with patch.dict(litellm.model_cost, _GEMMA_MODEL_COST_ENTRY, clear=False):
assert (
litellm.utils.supports_vision(
model="vertex_ai/google/gemma-4-26b-a4b-it-maas"
)
is True
)
# ---------------------------------------------------------------------------
# Integration tests: verify payloads reach the global OpenAI endpoint
#
# Patch target note (P1): AsyncHTTPHandler is patched at its *definition* site
# (litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler). This works
# correctly because the client is created by get_async_httpx_client(), which is
# also defined in http_handler.py and calls AsyncHTTPHandler(...) using the
# module-local name — so the patch intercepts instantiation there.
# llm_http_handler.py only imports the class for type annotations; it never
# instantiates it directly. Confirmed: without the mock the test raises
# AuthenticationError, proving the assertion would never silently pass against
# an un-mocked real call.
# ---------------------------------------------------------------------------
_MOCK_RESPONSE_JSON = {
"id": "chatcmpl-gemma-test",
"object": "chat.completion",
"created": 1234567890,
"model": "google/gemma-4-26b-a4b-it-maas",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Hello! How can I help you today?",
},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
}
@pytest.mark.asyncio
async def test_vertex_ai_gemma_global_endpoint_url():
"""
End-to-end: acompletion on vertex_ai/google/gemma-4-26b-a4b-it-maas should
POST to the global endpoints/openapi/chat/completions URL.
"""
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = _MOCK_RESPONSE_JSON
mock_vertexai = MagicMock()
mock_vertexai.preview = MagicMock()
with (
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler"
) as mock_http_handler,
patch(
"litellm.llms.vertex_ai.vertex_ai_partner_models.main.VertexAIPartnerModels._ensure_access_token",
return_value=("fake-token", "test-project"),
),
patch.dict(
"sys.modules",
{"vertexai": mock_vertexai, "vertexai.preview": mock_vertexai.preview},
),
patch.dict(
litellm.model_cost,
{
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
"supported_regions": ["global"]
}
},
clear=False,
),
):
mock_http_handler.return_value.post = AsyncMock(return_value=mock_response)
response = await litellm.acompletion(
model="vertex_ai/google/gemma-4-26b-a4b-it-maas",
messages=[{"role": "user", "content": "Hello"}],
vertex_ai_project="test-project",
)
mock_http_handler.return_value.post.assert_called_once()
call_args = mock_http_handler.return_value.post.call_args
called_url = call_args.kwargs["url"]
assert called_url.startswith("https://aiplatform.googleapis.com")
assert "global-aiplatform.googleapis.com" not in called_url
assert "/locations/global/" in called_url
assert "/endpoints/openapi/chat/completions" in called_url
assert response.model == "google/gemma-4-26b-a4b-it-maas"
@pytest.mark.asyncio
async def test_vertex_ai_gemma_function_calling_passthrough():
"""
Tools and tool_choice defined in the acompletion call must appear in the
JSON body POSTed to the global endpoints/openapi/chat/completions URL.
This confirms that supports_function_calling=true is backed by real
pass-through behaviour and that callers gating on get_model_info won't
silently send unsupported requests.
"""
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Return the current weather for a city.",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
},
}
]
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = _MOCK_RESPONSE_JSON
mock_vertexai = MagicMock()
mock_vertexai.preview = MagicMock()
with (
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler"
) as mock_http_handler,
patch(
"litellm.llms.vertex_ai.vertex_ai_partner_models.main.VertexAIPartnerModels._ensure_access_token",
return_value=("fake-token", "test-project"),
),
patch.dict(
"sys.modules",
{"vertexai": mock_vertexai, "vertexai.preview": mock_vertexai.preview},
),
patch.dict(litellm.model_cost, _GEMMA_MODEL_COST_ENTRY, clear=False),
):
mock_http_handler.return_value.post = AsyncMock(return_value=mock_response)
await litellm.acompletion(
model="vertex_ai/google/gemma-4-26b-a4b-it-maas",
messages=[{"role": "user", "content": "What's the weather in Paris?"}],
tools=tools,
tool_choice="auto",
vertex_ai_project="test-project",
)
mock_http_handler.return_value.post.assert_called_once()
call_args = mock_http_handler.return_value.post.call_args
# Must route to the global OpenAI-compatible endpoint
called_url = call_args.kwargs["url"]
assert called_url.startswith("https://aiplatform.googleapis.com"), called_url
assert "/endpoints/openapi/chat/completions" in called_url, called_url
# Tools and tool_choice must be forwarded in the request body
body = json.loads(call_args.kwargs["data"])
assert "tools" in body, f"'tools' key missing from request body: {body}"
assert body["tools"][0]["function"]["name"] == "get_weather"
assert "tool_choice" in body, f"'tool_choice' missing from request body: {body}"
assert body["tool_choice"] == "auto"
@pytest.mark.asyncio
async def test_vertex_ai_gemma_vision_passthrough():
"""
An image_url content part must survive transformation and appear in the
JSON body POSTed to the global endpoints/openapi/chat/completions URL.
This confirms that supports_vision=true is backed by real pass-through
behaviour and that callers gating on get_model_info won't silently send
unsupported multimodal requests.
"""
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this image."},
{
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
},
},
],
}
]
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = _MOCK_RESPONSE_JSON
mock_vertexai = MagicMock()
mock_vertexai.preview = MagicMock()
with (
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler"
) as mock_http_handler,
patch(
"litellm.llms.vertex_ai.vertex_ai_partner_models.main.VertexAIPartnerModels._ensure_access_token",
return_value=("fake-token", "test-project"),
),
patch.dict(
"sys.modules",
{"vertexai": mock_vertexai, "vertexai.preview": mock_vertexai.preview},
),
patch.dict(litellm.model_cost, _GEMMA_MODEL_COST_ENTRY, clear=False),
):
mock_http_handler.return_value.post = AsyncMock(return_value=mock_response)
await litellm.acompletion(
model="vertex_ai/google/gemma-4-26b-a4b-it-maas",
messages=messages,
vertex_ai_project="test-project",
)
mock_http_handler.return_value.post.assert_called_once()
call_args = mock_http_handler.return_value.post.call_args
# Must still route to the global OpenAI-compatible endpoint
called_url = call_args.kwargs["url"]
assert called_url.startswith("https://aiplatform.googleapis.com"), called_url
assert "/endpoints/openapi/chat/completions" in called_url, called_url
# The image_url content part must be present in the forwarded body
body = json.loads(call_args.kwargs["data"])
user_msg = next(m for m in body["messages"] if m["role"] == "user")
content = user_msg["content"]
assert isinstance(content, list), f"Expected list content, got: {content}"
image_parts = [p for p in content if p.get("type") == "image_url"]
assert image_parts, f"No image_url part in forwarded message content: {content}"

View file

@ -0,0 +1,282 @@
"""
Unit tests for WatsonxPassthroughConfig transformation.
Tests the Watsonx-specific passthrough configuration including URL construction,
streaming detection, and authentication handling.
"""
import os
import sys
from unittest.mock import MagicMock, patch
import httpx
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
import litellm
from litellm.llms.watsonx.passthrough.transformation import WatsonxPassthroughConfig
class TestWatsonxPassthroughConfig:
"""Tests for WatsonxPassthroughConfig class."""
def test_is_streaming_request_true(self):
"""Test that streaming is detected when stream=True in request data."""
config = WatsonxPassthroughConfig()
request_data = {"stream": True, "input": "test"}
result = config.is_streaming_request(
endpoint="ml/v1/text/generation", request_data=request_data
)
assert result is True
def test_is_streaming_request_false(self):
"""Test that streaming is not detected when stream=False in request data."""
config = WatsonxPassthroughConfig()
request_data = {"stream": False, "input": "test"}
result = config.is_streaming_request(
endpoint="ml/v1/text/generation", request_data=request_data
)
assert result is False
def test_is_streaming_request_missing_stream_key(self):
"""Test that streaming defaults to False when stream key is missing."""
config = WatsonxPassthroughConfig()
request_data = {"input": "test"}
result = config.is_streaming_request(
endpoint="ml/v1/text/generation", request_data=request_data
)
assert result is False
def test_get_complete_url_with_api_base(self):
"""Test URL construction with explicit api_base."""
config = WatsonxPassthroughConfig()
api_base = "https://us-south.ml.cloud.ibm.com"
endpoint = "ml/v1/text/generation"
request_query_params = {"version": "2024-03-19"}
complete_url, base_target_url = config.get_complete_url(
api_base=api_base,
api_key=None,
model="ibm/granite-13b-chat-v2",
endpoint=endpoint,
request_query_params=request_query_params,
litellm_params={},
)
assert isinstance(complete_url, httpx.URL)
assert str(complete_url).startswith(api_base)
assert endpoint in str(complete_url)
assert "version=2024-03-19" in str(complete_url)
assert base_target_url == api_base
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
def test_get_complete_url_with_env_api_base(self, mock_get_secret):
"""Test URL construction with api_base from environment."""
config = WatsonxPassthroughConfig()
env_api_base = "https://eu-de.ml.cloud.ibm.com"
mock_get_secret.return_value = env_api_base
endpoint = "ml/v1/text/tokenization"
request_query_params = {"version": "2024-03-19"}
complete_url, base_target_url = config.get_complete_url(
api_base=None,
api_key=None,
model="ibm/granite-13b-chat-v2",
endpoint=endpoint,
request_query_params=request_query_params,
litellm_params={},
)
assert isinstance(complete_url, httpx.URL)
assert str(complete_url).startswith(env_api_base)
assert endpoint in str(complete_url)
assert base_target_url == env_api_base
def test_get_complete_url_with_query_params(self):
"""Test that query parameters are correctly added to URL."""
config = WatsonxPassthroughConfig()
api_base = "https://us-south.ml.cloud.ibm.com"
endpoint = "ml/v1/text/generation"
request_query_params = {
"version": "2024-03-19",
}
complete_url, _ = config.get_complete_url(
api_base=api_base,
api_key=None,
model="ibm/granite-13b-chat-v2",
endpoint=endpoint,
request_query_params=request_query_params,
litellm_params={},
)
url_str = str(complete_url)
assert "version=2024-03-19" in url_str
def test_get_complete_url_without_query_params(self):
"""Test URL construction without query parameters."""
config = WatsonxPassthroughConfig()
api_base = "https://us-south.ml.cloud.ibm.com"
endpoint = "ml/v1/models"
complete_url, base_target_url = config.get_complete_url(
api_base=api_base,
api_key=None,
model="",
endpoint=endpoint,
request_query_params=None,
litellm_params={},
)
assert isinstance(complete_url, httpx.URL)
assert str(complete_url) == f"{api_base}/{endpoint}"
assert base_target_url == api_base
assert "version=2024-03-19" not in str(complete_url)
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
def test_get_api_base_with_explicit_value(self, mock_get_secret):
"""Test get_api_base returns explicit value when provided."""
explicit_base = "https://custom.watsonx.com"
result = WatsonxPassthroughConfig.get_api_base(api_base=explicit_base)
assert result == explicit_base
mock_get_secret.assert_not_called()
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
def test_get_api_base_from_environment(self, mock_get_secret):
"""Test get_api_base retrieves from environment when not provided."""
env_base = "https://env.watsonx.com"
mock_get_secret.return_value = env_base
result = WatsonxPassthroughConfig.get_api_base(api_base=None)
assert result == env_base
mock_get_secret.assert_called_once_with("WATSONX_API_BASE")
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
def test_get_api_key_with_explicit_value(self, mock_get_secret):
"""Test get_api_key returns explicit value when provided."""
explicit_key = "test-api-key-123"
result = WatsonxPassthroughConfig.get_api_key(api_key=explicit_key)
assert result == explicit_key
mock_get_secret.assert_not_called()
@patch("litellm.llms.watsonx.common_utils.get_secret_str")
def test_get_api_key_from_environment(self, mock_get_secret):
"""Test get_api_key retrieves from environment when not provided."""
env_key = "env-api-key-456"
mock_get_secret.return_value = env_key
result = WatsonxPassthroughConfig.get_api_key(api_key=None)
assert result == env_key
mock_get_secret.assert_any_call("WATSONX_APIKEY")
def test_get_base_model_returns_model(self):
"""Test get_base_model returns the model as-is."""
model = "ibm/granite-13b-chat-v2"
result = WatsonxPassthroughConfig.get_base_model(model)
assert result == model
def test_get_base_model_with_deployment(self):
"""Test get_base_model with deployment model."""
model = "deployment/test-deployment-id"
result = WatsonxPassthroughConfig.get_base_model(model)
assert result == model
def test_get_complete_url_with_different_endpoints(self):
"""Test URL construction with various endpoint paths."""
config = WatsonxPassthroughConfig()
api_base = "https://us-south.ml.cloud.ibm.com"
endpoints = [
"ml/v1/text/generation",
"ml/v1/text/tokenization",
"ml/v1/deployments/test-id/text/generation",
"ml/v1/models",
"ml/v1/foundation_model_specs",
]
for endpoint in endpoints:
complete_url, base_target_url = config.get_complete_url(
api_base=api_base,
api_key=None,
model="",
endpoint=endpoint,
request_query_params={"version": "2024-03-19"},
litellm_params={},
)
assert isinstance(complete_url, httpx.URL)
assert endpoint in str(complete_url)
assert base_target_url == api_base
def test_get_complete_url_preserves_query_param_order(self):
"""Test that query parameters maintain their values correctly."""
config = WatsonxPassthroughConfig()
api_base = "https://us-south.ml.cloud.ibm.com"
endpoint = "ml/v1/text/generation"
request_query_params = {
"version": "2024-03-19",
"project_id": "abc-123",
"space_id": "xyz-789",
}
complete_url, _ = config.get_complete_url(
api_base=api_base,
api_key=None,
model="",
endpoint=endpoint,
request_query_params=request_query_params,
litellm_params={},
)
url_str = str(complete_url)
# Verify all params are present
assert "version=2024-03-19" in url_str
assert "project_id=abc-123" in url_str
assert "space_id=xyz-789" in url_str
def test_is_streaming_request_with_various_stream_values(self):
"""Test streaming detection with different stream value types."""
config = WatsonxPassthroughConfig()
# Test with boolean True
assert config.is_streaming_request("endpoint", {"stream": True}) is True
# Test with boolean False
assert config.is_streaming_request("endpoint", {"stream": False}) is False
# Test with string "true" (truthy string)
result = config.is_streaming_request("endpoint", {"stream": "true"})
assert result == "true" # Returns the value as-is from .get()
# Test with integer 1 (truthy)
result = config.is_streaming_request("endpoint", {"stream": 1})
assert result == 1
# Test with integer 0 (falsy)
result = config.is_streaming_request("endpoint", {"stream": 0})
assert result == 0
# Test with None
result = config.is_streaming_request("endpoint", {"stream": None})
assert result is None
# Test with empty dict (defaults to False)
assert config.is_streaming_request("endpoint", {}) is False

View file

@ -68,6 +68,34 @@ async def test_partial_update_omits_unset_defaultful_fields():
)
@pytest.mark.asyncio
async def test_partial_update_null_tool_name_maps_clear_to_empty_json():
"""Explicit null on Json map fields must clear overrides (UI legacy)."""
data = UpdateMCPServerRequest(
server_id="my-test-server",
tool_name_to_display_name=None,
tool_name_to_description=None,
)
data_dict = await _run_update(data)
assert data_dict["tool_name_to_display_name"] == "{}"
assert data_dict["tool_name_to_description"] == "{}"
@pytest.mark.asyncio
async def test_partial_update_null_allowed_tools_clears_whitelist():
"""Explicit null must clear the whitelist (UI legacy); Prisma requires []."""
data = UpdateMCPServerRequest(
server_id="my-test-server",
allowed_tools=None,
)
data_dict = await _run_update(data)
assert data_dict["allowed_tools"] == []
@pytest.mark.asyncio
async def test_partial_update_preserves_http_transport():
"""The reported prod incident: a PUT without transport must not flip http->sse."""

View file

@ -4184,6 +4184,85 @@ def test_filter_tools_by_allowed_tools_no_filter():
assert len(filtered_tools) == 2
def test_filter_tools_enforced_empty_allowlist_blocks_all():
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server.server import (
filter_tools_by_allowed_tools,
)
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
tools = [
Tool(
name="read_wiki_structure",
title=None,
description="",
inputSchema={"type": "object"},
outputSchema=None,
annotations=None,
),
]
server = MCPServer(
server_id="deepwiki",
name="deepwiki",
transport=MCPTransport.http,
allowed_tools=[],
mcp_info={"tool_allowlist_enforced": True},
)
assert filter_tools_by_allowed_tools(tools, server) == []
def test_filter_tools_legacy_empty_allowlist_allows_all():
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server.server import (
filter_tools_by_allowed_tools,
)
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
tools = [
Tool(
name="read_wiki_structure",
title=None,
description="",
inputSchema={"type": "object"},
outputSchema=None,
annotations=None,
),
]
server = MCPServer(
server_id="legacy",
name="legacy",
transport=MCPTransport.http,
allowed_tools=[],
mcp_info=None,
)
assert len(filter_tools_by_allowed_tools(tools, server)) == 1
def test_check_allowed_or_banned_tools_enforced_empty_denies_calls():
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager = MCPServerManager.__new__(MCPServerManager)
server = MCPServer(
server_id="deepwiki",
name="deepwiki",
transport=MCPTransport.http,
allowed_tools=[],
mcp_info={"tool_allowlist_enforced": True},
)
assert manager.check_allowed_or_banned_tools("read_wiki_structure", server) is False
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token():
"""
@ -4540,9 +4619,9 @@ class TestEnsureUpstreamInitializeInstructionsCached:
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
server
)
assert create.await_count == 1, (
"Second probe within cooldown must not reconnect to upstream"
)
assert (
create.await_count == 1
), "Second probe within cooldown must not reconnect to upstream"
assert (
"empty-server"
not in global_mcp_server_manager._upstream_initialize_instructions_by_server_id
@ -4567,7 +4646,9 @@ class TestEnsureUpstreamInitializeInstructionsCached:
server = _make_instruction_server(server_id="boom-server", instructions=None)
fake_client = MagicMock()
fake_client.run_with_session = AsyncMock(side_effect=RuntimeError("upstream down"))
fake_client.run_with_session = AsyncMock(
side_effect=RuntimeError("upstream down")
)
fake_client._last_initialize_instructions = None
create = AsyncMock(return_value=fake_client)
@ -4579,9 +4660,9 @@ class TestEnsureUpstreamInitializeInstructionsCached:
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
server
)
assert create.await_count == 1, (
"Second probe within cooldown must not reconnect after failure"
)
assert (
create.await_count == 1
), "Second probe within cooldown must not reconnect after failure"
assert (
"boom-server"
not in global_mcp_server_manager._upstream_initialize_instructions_by_server_id

File diff suppressed because it is too large Load diff

View file

@ -5375,5 +5375,93 @@ class TestPanwAirsDualScanIndependence:
assert mcp_call.get("content") is None
class TestPanwAirsTimeoutCoercion:
"""Regression tests for string-valued timeout handling.
Before the fix, a string `timeout` (which is what the dashboard UI persists
and what raw YAML preserves if quoted) survived into httpx, which raised
`TypeError: '<=' not supported between instances of 'str' and 'int'`. The
broad except in apply_guardrail swallowed it and the proxy returned a
misleading 500 'Security scan failed - request blocked for safety'.
"""
def test_handler_coerces_string_timeout_to_float(self):
handler = make_handler(timeout="30")
assert handler.timeout == 30.0
assert isinstance(handler.timeout, float)
def test_handler_accepts_int_timeout(self):
handler = make_handler(timeout=15)
assert handler.timeout == 15.0
def test_handler_accepts_float_timeout(self):
handler = make_handler(timeout=7.5)
assert handler.timeout == 7.5
def test_handler_none_timeout_falls_back_to_default(self):
handler = make_handler(timeout=None)
assert handler.timeout == 10.0
def test_handler_omitted_timeout_uses_default(self):
handler = make_handler()
assert handler.timeout == 10.0
def test_litellm_params_coerces_string_timeout(self):
"""Boundary validation: the Pydantic model itself should normalize
string timeouts before any handler reads the value via model_dump()."""
params = LitellmParams(
guardrail="panw_prisma_airs",
mode="pre_call",
api_key="test_key",
profile_name="test_profile",
timeout="30",
)
assert params.timeout == 30.0
assert isinstance(params.timeout, float)
def test_litellm_params_rejects_garbage_timeout(self):
with pytest.raises(ValueError):
LitellmParams(
guardrail="panw_prisma_airs",
mode="pre_call",
api_key="test_key",
profile_name="test_profile",
timeout="not-a-number",
)
def test_litellm_params_empty_string_timeout_becomes_none(self):
"""Empty-string timeout (which the dashboard form can send) should
be coerced to None, not crash, and not produce float('')."""
params = LitellmParams(
guardrail="panw_prisma_airs",
mode="pre_call",
api_key="test_key",
profile_name="test_profile",
timeout="",
)
assert params.timeout is None
def test_legacy_initializer_handles_unset_timeout(self):
"""Regression guard: with timeout now a declared Optional[float] = None
on BaseLitellmParams, the legacy panw initializer at
guardrail_initializers.py:220 must not crash on float(None) when the
caller omits timeout entirely."""
from litellm.proxy.guardrails.guardrail_initializers import (
initialize_panw_prisma_airs,
)
params = LitellmParams(
guardrail="panw_prisma_airs",
mode="pre_call",
api_key="test_key",
profile_name="test_profile",
# timeout intentionally omitted - field defaults to None
)
guardrail_config = {"guardrail_name": "test_legacy"}
handler = initialize_panw_prisma_airs(params, guardrail_config)
# Default fallback applied, not crashed on float(None)
assert handler.timeout == 10.0
if __name__ == "__main__":
pytest.main([__file__, "-v"])

View file

@ -0,0 +1,900 @@
import json
import logging
import ssl
from types import SimpleNamespace
from typing import Any, List
import httpx
import pytest
from litellm.exceptions import GuardrailRaisedException
from litellm.exceptions import Timeout as LiteLLMTimeout
from litellm.proxy.guardrails.guardrail_hooks.vigil_guard import (
VigilGuardGuardrail,
guardrail_class_registry,
guardrail_initializer_registry,
initialize_guardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.vigil_guard.vigil_guard import (
_DEFAULT_VIGIL_TIMEOUT,
VigilGuardMissingConfig,
)
from litellm.types.guardrails import LitellmParams, SupportedGuardrailIntegrations
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
VigilGuardGuardrailConfigModel,
)
_ENDPOINT = "https://vigil.test/v1/guard/analyze"
def _resp(body: dict, status_code: int = 200) -> httpx.Response:
return httpx.Response(
status_code=status_code,
json=body,
request=httpx.Request("POST", _ENDPOINT),
)
class FakeHandler:
def __init__(self, items: List[Any]):
self._items = list(items)
self.calls: List[SimpleNamespace] = []
async def post(self, *, url, headers, json, timeout=None): # noqa: A002
self.calls.append(
SimpleNamespace(url=url, headers=headers, json=json, timeout=timeout)
)
if not self._items:
raise AssertionError("FakeHandler ran out of programmed responses")
item = self._items.pop(0)
if isinstance(item, BaseException):
raise item
return item
def _make_guardrail(
handler: FakeHandler,
*,
unreachable_fallback="fail_closed",
api_base="https://vigil.test",
api_key="vg_secret_key_123",
guardrail_name="vigil-guard",
timeout=None,
) -> VigilGuardGuardrail:
return VigilGuardGuardrail(
api_base=api_base,
api_key=api_key,
unreachable_fallback=unreachable_fallback,
timeout=timeout,
async_handler=handler,
guardrail_name=guardrail_name,
event_hook="pre_call",
default_on=True,
)
def _transient_exceptions() -> List[BaseException]:
req = httpx.Request("POST", _ENDPOINT)
return [
httpx.ConnectError("boom", request=req),
httpx.ConnectTimeout("boom", request=req),
httpx.ReadTimeout("boom", request=req),
httpx.RemoteProtocolError("boom", request=req),
LiteLLMTimeout(message="t", model="m", llm_provider="vigil_guard"),
]
def test_requires_api_base(monkeypatch):
monkeypatch.delenv("VIGIL_GUARD_URL", raising=False)
monkeypatch.delenv("VIGIL_GUARD_API_KEY", raising=False)
with pytest.raises(VigilGuardMissingConfig):
VigilGuardGuardrail(api_key="k", async_handler=FakeHandler([]))
def test_requires_api_key(monkeypatch):
monkeypatch.delenv("VIGIL_GUARD_API_KEY", raising=False)
with pytest.raises(VigilGuardMissingConfig):
VigilGuardGuardrail(
api_base="https://vigil.test", async_handler=FakeHandler([])
)
def test_trailing_slash_stripped():
g = _make_guardrail(FakeHandler([]), api_base="https://vigil.test/")
assert g.api_base == "https://vigil.test"
def test_env_fallback(monkeypatch):
monkeypatch.setenv("VIGIL_GUARD_URL", "https://env.vigil.test")
monkeypatch.setenv("VIGIL_GUARD_API_KEY", "env_key")
g = VigilGuardGuardrail(
async_handler=FakeHandler([]),
guardrail_name="vg",
event_hook="pre_call",
default_on=True,
)
assert g.api_base == "https://env.vigil.test"
assert g.api_key == "env_key"
def test_default_unreachable_fallback_is_fail_closed():
g = _make_guardrail(FakeHandler([]), unreachable_fallback=None)
assert g.unreachable_fallback == "fail_closed"
def test_explicit_fail_open_is_stored():
g = _make_guardrail(FakeHandler([]), unreachable_fallback="fail_open")
assert g.unreachable_fallback == "fail_open"
def test_unknown_fallback_defaults_to_fail_closed():
g = _make_guardrail(FakeHandler([]), unreachable_fallback="weird")
assert g.unreachable_fallback == "fail_closed"
async def test_allowed_preserves_full_input_shape_and_logs_allow():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
structured = [{"role": "user", "content": "hello"}]
inputs = {"texts": ["hello"], "structured_messages": structured, "model": "gpt-4o"}
request_data = {"metadata": {}}
out = await g.apply_guardrail(
inputs=inputs, request_data=request_data, input_type="request", logging_obj=None
)
assert out["texts"] == ["hello"]
assert out["structured_messages"] is structured
assert out["model"] == "gpt-4o"
assert out is not inputs
assert inputs["structured_messages"] is structured
assert len(handler.calls) == 1
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert entries[0]["guardrail_response"] == "allow"
async def test_sanitized_replaces_text():
handler = FakeHandler(
[_resp({"decision": "SANITIZED", "sanitizedText": "[REDACTED]"})]
)
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["my ssn is 123"]}, request_data={}, input_type="request"
)
assert out["texts"] == ["[REDACTED]"]
@pytest.mark.parametrize(
"body,expected",
[
(
{
"decision": "SANITIZED",
"sanitizedText": "S",
"outputText": "O",
},
"S",
),
({"decision": "SANITIZED", "outputText": "O"}, "O"),
({"decision": "SANITIZED", "sanitizedText": 123, "outputText": "O"}, "O"),
({"decision": "SANITIZED", "sanitizedText": ""}, ""),
({"decision": "SANITIZED"}, "orig"),
],
)
async def test_sanitized_precedence(body, expected):
handler = FakeHandler([_resp(body)])
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["orig"]}, request_data={}, input_type="request"
)
assert out["texts"] == [expected]
async def test_blocked_raises_guardrail_exception_with_400():
handler = FakeHandler([_resp({"decision": "BLOCKED", "blockMessage": "nope"})])
g = _make_guardrail(handler)
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": ["bad"]}, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert exc_info.value.guardrail_name == "vigil-guard"
assert exc_info.value.message == "nope"
@pytest.mark.parametrize(
"body,expected",
[
(
{
"decision": "BLOCKED",
"blockMessage": "bm",
"decisionReason": "dr",
"categories": ["c1"],
},
"bm",
),
({"decision": "BLOCKED", "blockMessage": " ", "decisionReason": "dr"}, "dr"),
(
{"decision": "BLOCKED", "decisionReason": "dr", "categories": ["c1", "c2"]},
"dr",
),
({"decision": "BLOCKED", "categories": ["c1", "c2"]}, "c1, c2"),
({"decision": "BLOCKED"}, "Blocked by policy"),
],
)
async def test_block_reason_precedence(body, expected):
handler = FakeHandler([_resp(body)])
g = _make_guardrail(handler)
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert exc_info.value.message == expected
async def test_block_reason_is_clamped_to_500_chars():
handler = FakeHandler([_resp({"decision": "BLOCKED", "blockMessage": "x" * 600})])
g = _make_guardrail(handler)
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert "x" * 500 in exc_info.value.message
assert "x" * 501 not in exc_info.value.message
async def test_empty_and_whitespace_texts_skip_analyze():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["", " ", "real"]}, request_data={}, input_type="request"
)
assert out["texts"] == ["", " ", "real"]
assert len(handler.calls) == 1
assert handler.calls[0].json["text"] == "real"
async def test_no_scannable_text_returns_inputs_unchanged():
handler = FakeHandler([])
g = _make_guardrail(handler)
inputs = {"texts": ["", " "], "structured_messages": [{"role": "user"}]}
out = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request")
assert out is inputs
assert len(handler.calls) == 0
async def test_multi_text_preserves_length_and_order():
handler = FakeHandler(
[
_resp({"decision": "ALLOWED"}),
_resp({"decision": "SANITIZED", "sanitizedText": "B-clean"}),
_resp({"decision": "ALLOWED"}),
]
)
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["A", "B", "C"]}, request_data={}, input_type="request"
)
assert out["texts"] == ["A", "B-clean", "C"]
assert len(handler.calls) == 3
async def test_one_blocked_text_blocks_the_whole_call():
handler = FakeHandler(
[
_resp({"decision": "ALLOWED"}),
_resp({"decision": "BLOCKED", "blockMessage": "bad second"}),
]
)
g = _make_guardrail(handler)
with pytest.raises(GuardrailRaisedException):
await g.apply_guardrail(
inputs={"texts": ["ok", "bad"]}, request_data={}, input_type="request"
)
async def test_request_source_is_user_input():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert handler.calls[0].json["source"] == "user_input"
async def test_response_source_is_model_output():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="response"
)
assert handler.calls[0].json["source"] == "model_output"
async def test_sanitized_returns_canonical_shape_and_logs_mask():
handler = FakeHandler(
[_resp({"decision": "SANITIZED", "sanitizedText": "[REDACTED]"})]
)
g = _make_guardrail(handler)
tools = [{"type": "function", "function": {"name": "f"}}]
inputs = {
"texts": ["my ssn is 123"],
"images": ["img1"],
"tools": tools,
"tool_calls": [{"id": "1"}],
"structured_messages": [{"role": "user", "content": "my ssn is 123"}],
"model": "gpt-4o",
}
request_data = {"metadata": {}}
out = await g.apply_guardrail(
inputs=inputs, request_data=request_data, input_type="request"
)
assert out["texts"] == ["[REDACTED]"]
assert out["images"] == ["img1"]
assert out["tools"] == tools
assert set(out.keys()) == {"texts", "images", "tools"}
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert entries[0]["guardrail_response"] == "mask"
async def test_empty_images_and_tools_are_preserved_when_present():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["x"], "images": [], "tools": []},
request_data={},
input_type="request",
)
assert set(out.keys()) == {"texts", "images", "tools"}
assert out["images"] == []
assert out["tools"] == []
async def test_logging_obj_none_supported():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request", logging_obj=None
)
assert out["texts"] == ["x"]
async def test_standard_guardrail_logging_remains_active():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
request_data = {"metadata": {}}
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data=request_data, input_type="request"
)
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert len(entries) == 1
assert entries[0]["guardrail_name"] == "vigil-guard"
assert entries[0]["guardrail_status"] == "success"
async def test_request_url_headers_and_body():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler, api_base="https://vigil.test", api_key="vg_secret")
await g.apply_guardrail(
inputs={"texts": ["hello"]}, request_data={}, input_type="request"
)
call = handler.calls[0]
assert call.url == "https://vigil.test/v1/guard/analyze"
assert call.headers["Authorization"] == "Bearer vg_secret"
assert call.headers["Content-Type"] == "application/json"
assert call.json["text"] == "hello"
assert call.json["mode"] == "full"
assert set(call.json.keys()) == {"text", "source", "mode", "metadata"}
assert "metadata" in call.json
async def test_default_timeout_forwarded_when_unset():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
assert g.timeout == _DEFAULT_VIGIL_TIMEOUT
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert handler.calls[0].timeout == _DEFAULT_VIGIL_TIMEOUT
async def test_configured_timeout_forwarded_to_handler():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler, timeout=30)
expected = httpx.Timeout(30, connect=5.0)
assert g.timeout == expected
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert handler.calls[0].timeout == expected
def test_short_timeout_caps_connect():
g = _make_guardrail(FakeHandler([]), timeout=2)
assert g.timeout == httpx.Timeout(2, connect=2.0)
def test_initialize_guardrail_forwards_timeout():
lp = LitellmParams(
guardrail="vigil_guard",
mode="pre_call",
api_base="https://vigil.test",
api_key="k",
timeout="30",
)
cb = initialize_guardrail(lp, {"guardrail_name": "vg"})
assert cb.timeout == httpx.Timeout(30, connect=5.0)
async def test_api_key_only_in_header_never_in_payload():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler, api_key="super_secret_key")
await g.apply_guardrail(
inputs={"texts": ["hello"]},
request_data={"metadata": {"user_id": "u"}},
input_type="request",
)
call = handler.calls[0]
assert "super_secret_key" not in json.dumps(call.json)
assert call.headers["Authorization"] == "Bearer super_secret_key"
@pytest.mark.parametrize("code", [429, 502, 503, 504])
async def test_retry_once_on_transient_status(code):
handler = FakeHandler([_resp({}, status_code=code), _resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert out["texts"] == ["x"]
assert len(handler.calls) == 2
@pytest.mark.parametrize("exc", _transient_exceptions())
async def test_retry_once_on_transient_exception(exc):
handler = FakeHandler([exc, _resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
out = await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert out["texts"] == ["x"]
assert len(handler.calls) == 2
@pytest.mark.parametrize(
"exc, expected",
[
(RuntimeError("boom"), RuntimeError),
(
httpx.WriteError("boom", request=httpx.Request("POST", _ENDPOINT)),
GuardrailRaisedException,
),
],
)
async def test_no_retry_on_non_transient_exception(exc, expected):
handler = FakeHandler([exc])
g = _make_guardrail(handler)
with pytest.raises(expected):
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert len(handler.calls) == 1
@pytest.mark.parametrize("code", [400, 401, 403, 404, 422])
async def test_no_retry_on_non_429_4xx(code):
handler = FakeHandler([_resp({}, status_code=code)])
g = _make_guardrail(handler)
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert len(handler.calls) == 1
async def test_fail_closed_raises_after_exhausted_retry(caplog):
handler = FakeHandler([_resp({}, status_code=503), _resp({}, status_code=503)])
g = _make_guardrail(handler)
with (
caplog.at_level(logging.ERROR),
pytest.raises(GuardrailRaisedException) as exc_info,
):
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert len(handler.calls) == 2
assert any("fail_closed" in record.message for record in caplog.records)
assert any("vigil-guard" in record.message for record in caplog.records)
@pytest.mark.parametrize("exc", _transient_exceptions())
async def test_fail_closed_raises_controlled_block_on_transport_error(exc, caplog):
handler = FakeHandler([exc, exc])
g = _make_guardrail(handler)
with (
caplog.at_level(logging.ERROR),
pytest.raises(GuardrailRaisedException) as exc_info,
):
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert exc_info.value.guardrail_name == "vigil-guard"
assert exc_info.value.__cause__ is exc
assert any("fail_closed" in record.message for record in caplog.records)
async def test_fail_open_returns_inputs_unchanged_on_backend_error(caplog):
handler = FakeHandler([_resp({}, status_code=503), _resp({}, status_code=503)])
g = _make_guardrail(handler, unreachable_fallback="fail_open")
structured = [{"role": "user", "content": "x"}]
inputs = {"texts": ["x"], "structured_messages": structured}
request_data = {"metadata": {}}
with caplog.at_level(logging.ERROR):
out = await g.apply_guardrail(
inputs=inputs, request_data=request_data, input_type="request"
)
assert out is not inputs
assert out["texts"] == ["x"]
assert out["structured_messages"] == structured
assert len(handler.calls) == 2
assert any("fail_open" in record.message for record in caplog.records)
assert any("vigil-guard" in record.message for record in caplog.records)
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert entries[0]["guardrail_response"] == "allow"
@pytest.mark.parametrize("exc", [ssl.SSLError("tls failed"), OSError("network down")])
async def test_fail_open_returns_inputs_unchanged_on_transport_error(exc):
handler = FakeHandler([exc])
g = _make_guardrail(handler, unreachable_fallback="fail_open")
inputs = {"texts": ["x"]}
out = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request")
assert out is not inputs
assert out["texts"] == ["x"]
assert len(handler.calls) == 1
@pytest.mark.parametrize(
"exc",
[
TypeError("bug"),
KeyError("bug"),
AttributeError("bug"),
],
)
async def test_fail_open_does_not_swallow_programming_errors(exc):
handler = FakeHandler([exc])
g = _make_guardrail(handler, unreachable_fallback="fail_open")
with pytest.raises(type(exc)):
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert len(handler.calls) == 1
async def test_invalid_decision_fail_closed_raises(caplog):
handler = FakeHandler([_resp({"decision": "MAYBE"})])
g = _make_guardrail(handler)
with (
caplog.at_level(logging.ERROR),
pytest.raises(GuardrailRaisedException) as exc_info,
):
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert exc_info.value.status_code == 400
assert "MAYBE" not in exc_info.value.message
assert any("MAYBE" in record.message for record in caplog.records)
async def test_invalid_decision_fail_open_returns_inputs():
handler = FakeHandler([_resp({"decision": "MAYBE"})])
g = _make_guardrail(handler, unreachable_fallback="fail_open")
out = await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={}, input_type="request"
)
assert out["texts"] == ["x"]
async def test_fail_open_multi_text_preserves_earlier_sanitization():
handler = FakeHandler(
[
_resp({"decision": "SANITIZED", "sanitizedText": "[REDACTED]"}),
_resp({}, status_code=503),
_resp({}, status_code=503),
]
)
g = _make_guardrail(handler, unreachable_fallback="fail_open")
request_data = {"metadata": {}}
out = await g.apply_guardrail(
inputs={"texts": ["my ssn is 123", "second"]},
request_data=request_data,
input_type="request",
)
assert out["texts"] == ["[REDACTED]", "second"]
assert len(handler.calls) == 3
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert entries[0]["guardrail_response"] == "mask"
def _tool_call(arguments, name="f", tc_id="1"):
return {
"id": tc_id,
"type": "function",
"function": {"name": name, "arguments": arguments},
}
async def test_response_tool_call_arguments_allowed_unchanged():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
tcs = [_tool_call('{"q": "weather"}')]
out = await g.apply_guardrail(
inputs={"texts": [], "tool_calls": tcs}, request_data={}, input_type="response"
)
assert handler.calls[0].json["text"] == '{"q": "weather"}'
assert handler.calls[0].json["source"] == "model_output"
assert out["tool_calls"] == tcs
async def test_response_tool_call_arguments_sanitized_in_place():
handler = FakeHandler(
[_resp({"decision": "SANITIZED", "sanitizedText": '{"email": "[EMAIL]"}'})]
)
g = _make_guardrail(handler)
tcs = [_tool_call('{"email": "john@example.com"}', name="send_mail")]
inputs = {"texts": [], "tool_calls": tcs}
out = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="response")
assert out["tool_calls"][0]["function"]["arguments"] == '{"email": "[EMAIL]"}'
assert out["tool_calls"][0]["function"]["name"] == "send_mail"
# original inputs are not mutated in place
assert inputs["tool_calls"][0]["function"]["arguments"] == (
'{"email": "john@example.com"}'
)
async def test_response_tool_call_arguments_blocked_raises():
handler = FakeHandler(
[_resp({"decision": "BLOCKED", "blockMessage": "tool blocked"})]
)
g = _make_guardrail(handler)
tcs = [_tool_call('{"x": "bad"}')]
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": [], "tool_calls": tcs},
request_data={},
input_type="response",
)
assert exc_info.value.status_code == 400
assert exc_info.value.message == "tool blocked"
async def test_request_tool_calls_are_not_scanned():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
tcs = [_tool_call('{"x": "y"}')]
await g.apply_guardrail(
inputs={"texts": ["hello"], "tool_calls": tcs},
request_data={},
input_type="request",
)
assert len(handler.calls) == 1
assert handler.calls[0].json["text"] == "hello"
async def test_tool_call_scan_backend_failure_fail_closed_raises():
handler = FakeHandler([_resp({}, status_code=503), _resp({}, status_code=503)])
g = _make_guardrail(handler)
tcs = [_tool_call('{"x": "y"}')]
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": [], "tool_calls": tcs},
request_data={},
input_type="response",
)
assert exc_info.value.status_code == 400
assert len(handler.calls) == 2
async def test_tool_call_scan_backend_failure_fail_open_passes_through():
handler = FakeHandler([_resp({}, status_code=503), _resp({}, status_code=503)])
g = _make_guardrail(handler, unreachable_fallback="fail_open")
tcs = [_tool_call('{"x": "y"}')]
out = await g.apply_guardrail(
inputs={"texts": [], "tool_calls": tcs}, request_data={}, input_type="response"
)
assert out["tool_calls"] == tcs
async def test_response_tool_call_unrecognized_decision_fail_closed_raises():
handler = FakeHandler([_resp({"decision": "MAYBE"})])
g = _make_guardrail(handler)
tcs = [_tool_call('{"x": "y"}')]
with pytest.raises(GuardrailRaisedException) as exc_info:
await g.apply_guardrail(
inputs={"texts": [], "tool_calls": tcs},
request_data={},
input_type="response",
)
assert exc_info.value.status_code == 400
async def test_response_tool_call_unrecognized_decision_fail_open_passes_through():
handler = FakeHandler([_resp({"decision": "MAYBE"})])
g = _make_guardrail(handler, unreachable_fallback="fail_open")
tcs = [_tool_call('{"x": "y"}')]
out = await g.apply_guardrail(
inputs={"texts": [], "tool_calls": tcs}, request_data={}, input_type="response"
)
assert out["tool_calls"] == tcs
async def test_metadata_allowlist_and_clamping():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
request_data = {
"model": "gpt-4o",
"metadata": {
"user_id": "u1",
"tenant_id": "t1",
"secret_unlisted": "should_not_forward",
"session_id": "s" * 600,
"org_id": ["a"] * 20,
"request_id": True,
"conversation_id": 7,
},
}
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data=request_data, input_type="request"
)
md = handler.calls[0].json["metadata"]
assert md["model"] == "gpt-4o"
assert md["user_id"] == "u1"
assert md["tenant_id"] == "t1"
assert "secret_unlisted" not in md
assert len(md["session_id"]) == 500
assert len(md["org_id"]) == 10
assert "request_id" not in md
assert md["conversation_id"] == 7
async def test_metadata_source_precedence_and_litellm_metadata_fallback():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
request_data = {
"user_id": "top",
"metadata": {"user_id": "nested"},
"litellm_metadata": {"tenant_id": "lm-tenant"},
}
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data=request_data, input_type="request"
)
md = handler.calls[0].json["metadata"]
assert md["user_id"] == "top"
assert md["tenant_id"] == "lm-tenant"
async def test_metadata_uses_later_source_when_earlier_value_is_unclampable():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
request_data = {
"user_id": {"drop": "dicts are not forwarded"},
"metadata": {"user_id": "nested"},
}
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data=request_data, input_type="request"
)
assert handler.calls[0].json["metadata"]["user_id"] == "nested"
async def test_metadata_array_items_are_clamped_and_filtered():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
request_data = {
"metadata": {
"org_id": ["z" * 600, 123, True, {"drop": 1}, None],
},
}
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data=request_data, input_type="request"
)
assert handler.calls[0].json["metadata"]["org_id"] == ["z" * 500, 123]
async def test_metadata_array_with_no_supported_items_is_dropped():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"metadata": {"org_id": [{"drop": 1}, None]}},
input_type="request",
)
assert "org_id" not in handler.calls[0].json["metadata"]
async def test_call_id_forwarded_from_logging_obj():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
logging_obj = SimpleNamespace(litellm_call_id="call-123")
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={},
input_type="request",
logging_obj=logging_obj,
)
assert handler.calls[0].json["metadata"]["litellm_call_id"] == "call-123"
async def test_call_id_forwarded_from_request_data():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"litellm_call_id": "rd-1"},
input_type="request",
logging_obj=None,
)
assert handler.calls[0].json["metadata"]["litellm_call_id"] == "rd-1"
async def test_call_id_forwarded_from_request_metadata():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"metadata": {"litellm_call_id": "md-1"}},
input_type="request",
logging_obj=None,
)
assert handler.calls[0].json["metadata"]["litellm_call_id"] == "md-1"
async def test_call_id_logging_obj_takes_precedence():
handler = FakeHandler([_resp({"decision": "ALLOWED"})])
g = _make_guardrail(handler)
logging_obj = SimpleNamespace(litellm_call_id="log-1")
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"litellm_call_id": "rd-1"},
input_type="request",
logging_obj=logging_obj,
)
assert handler.calls[0].json["metadata"]["litellm_call_id"] == "log-1"
def test_enum_value():
assert SupportedGuardrailIntegrations.VIGIL_GUARD.value == "vigil_guard"
def test_config_model_ui_name_and_instantiation():
assert VigilGuardGuardrailConfigModel.ui_friendly_name() == "Vigil Guard"
model = VigilGuardGuardrailConfigModel(api_base="https://x", api_key="k")
assert model.api_base == "https://x"
def test_get_config_model_returns_config_model():
g = _make_guardrail(FakeHandler([]))
assert g.get_config_model() is VigilGuardGuardrailConfigModel
def test_registries_expose_initializer_and_class():
assert "vigil_guard" in guardrail_initializer_registry
assert guardrail_class_registry["vigil_guard"] is VigilGuardGuardrail
def test_litellm_params_includes_config_model():
assert VigilGuardGuardrailConfigModel in LitellmParams.__mro__
def test_config_driven_initialization_creates_callback():
lp = LitellmParams(
guardrail="vigil_guard",
mode="pre_call",
api_base="https://vigil.test",
api_key="k",
)
cb = initialize_guardrail(lp, {"guardrail_name": "vg"})
assert isinstance(cb, VigilGuardGuardrail)
assert cb.unreachable_fallback == "fail_closed"

View file

@ -0,0 +1,213 @@
import os
from unittest.mock import patch
import pytest
class TestContentFilterPathTraversal:
"""Tests that _resolve_category_file_path rejects path traversal."""
def _get_guardrail(self):
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
return ContentFilterGuardrail.__new__(ContentFilterGuardrail)
def test_traversal_via_relative_dotdot_raises(self):
guardrail = self._get_guardrail()
with pytest.raises(ValueError, match="outside the allowed categories"):
guardrail._resolve_category_file_path("../../../../etc/passwd")
def test_traversal_via_absolute_path_raises(self):
guardrail = self._get_guardrail()
with pytest.raises(ValueError, match="outside the allowed categories"):
guardrail._resolve_category_file_path("/etc/passwd")
def test_valid_category_file_inside_categories_dir_allowed(self):
guardrail = self._get_guardrail()
categories_dir = os.path.join(
os.path.dirname(
__import__(
"litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter",
fromlist=["content_filter"],
).__file__
),
"categories",
)
valid_file = os.path.join(categories_dir, "harmful_self_harm.yaml")
if not os.path.exists(valid_file):
pytest.skip("harmful_self_harm.yaml not present in this environment")
result = guardrail._resolve_category_file_path(valid_file)
assert result == valid_file
def test_invalid_category_name_skipped(self):
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
guardrail = ContentFilterGuardrail.__new__(ContentFilterGuardrail)
guardrail.loaded_categories = {}
guardrail.severity_threshold = "medium"
guardrail.category_keywords = {}
guardrail.always_block_category_keywords = {}
guardrail.conditional_categories = {}
# category name with path traversal chars must be skipped, not crash
guardrail._load_categories([{"category": "../../etc/passwd", "enabled": True}])
assert "../../etc/passwd" not in guardrail.loaded_categories
def test_category_name_with_slash_skipped(self):
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
guardrail = ContentFilterGuardrail.__new__(ContentFilterGuardrail)
guardrail.loaded_categories = {}
guardrail.severity_threshold = "medium"
guardrail.category_keywords = {}
guardrail.always_block_category_keywords = {}
guardrail.conditional_categories = {}
guardrail._load_categories(
[{"category": "foo/../../etc/passwd", "enabled": True}]
)
assert "foo/../../etc/passwd" not in guardrail.loaded_categories
def test_assert_within_categories_dir_blocks_parent_traversal(self):
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
categories_dir = os.path.join(
os.path.dirname(
__import__(
"litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter",
fromlist=["content_filter"],
).__file__
),
"categories",
)
with pytest.raises(ValueError, match="outside the allowed categories"):
ContentFilterGuardrail._assert_within_categories_dir(
"/etc/passwd", categories_dir
)
def test_assert_within_categories_dir_allows_valid_file(self, tmp_path):
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
categories_dir = str(tmp_path)
valid_file = str(tmp_path / "test.yaml")
# Should not raise
ContentFilterGuardrail._assert_within_categories_dir(valid_file, categories_dir)
def test_assert_within_categories_dir_commonpath_raises_valueerror(self, tmp_path):
"""Cover the except-ValueError branch (Windows cross-drive paths)."""
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
categories_dir = str(tmp_path)
valid_file = str(tmp_path / "test.yaml")
with patch(
"os.path.commonpath", side_effect=ValueError("Paths on different drives")
):
with pytest.raises(
ValueError, match="outside the allowed categories directory"
):
ContentFilterGuardrail._assert_within_categories_dir(
valid_file, categories_dir
)
def test_resolve_category_file_path_direct_join_hit(self):
"""Cover the first-join-attempt success branch (lines 383-384)."""
guardrail = self._get_guardrail()
# "categories/<file>" joined directly to module_dir resolves to an existing file.
categories_dir = os.path.join(
os.path.dirname(
__import__(
"litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter",
fromlist=["content_filter"],
).__file__
),
"categories",
)
yaml_files = [f for f in os.listdir(categories_dir) if f.endswith(".yaml")]
if not yaml_files:
pytest.skip("No category YAML files present in this environment")
relative_path = os.path.join("categories", yaml_files[0])
result = guardrail._resolve_category_file_path(relative_path)
assert os.path.isabs(result) or os.path.exists(result)
def test_resolve_category_file_path_component_strip_hit(self):
"""Cover the component-stripping loop success branch (lines 392-393)."""
guardrail = self._get_guardrail()
categories_dir = os.path.join(
os.path.dirname(
__import__(
"litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter",
fromlist=["content_filter"],
).__file__
),
"categories",
)
yaml_files = [f for f in os.listdir(categories_dir) if f.endswith(".yaml")]
if not yaml_files:
pytest.skip("No category YAML files present in this environment")
# Prefix with a fake leading component so the first-join attempt misses,
# but stripping that component reveals categories/<file> which exists.
prefixed_path = "some_prefix/categories/" + yaml_files[0]
result = guardrail._resolve_category_file_path(prefixed_path)
assert os.path.isabs(result) or os.path.exists(result)
def test_load_categories_traversal_category_file_skipped(self):
"""Cover the except-ValueError branch in _load_categories (lines 451-454)."""
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
guardrail = ContentFilterGuardrail.__new__(ContentFilterGuardrail)
guardrail.loaded_categories = {}
guardrail.severity_threshold = "medium"
guardrail.category_keywords = {}
guardrail.always_block_category_keywords = {}
guardrail.conditional_categories = {}
# A traversal path in category_file must be skipped (not crash) via ValueError.
guardrail._load_categories(
[
{
"category": "valid_name",
"enabled": True,
"category_file": "../../../../etc/passwd",
}
]
)
assert "valid_name" not in guardrail.loaded_categories
def test_allow_external_paths_env_var_bypasses_jail(self, tmp_path):
"""LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS=true skips the directory jail."""
import os as _os
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
guardrail = ContentFilterGuardrail.__new__(ContentFilterGuardrail)
# Create a real file outside the module directory (simulates mounted volume).
external_file = tmp_path / "external_categories.yaml"
external_file.write_text("category_name: test\n")
with patch.dict(
_os.environ, {"LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS": "true"}
):
# Should return the path without raising ValueError.
result = guardrail._resolve_category_file_path(str(external_file))
assert result == str(external_file)
def test_traversal_blocked_when_allow_external_not_set(self):
"""Without the env var the jail still blocks traversal paths."""
import os as _os
guardrail = self._get_guardrail()
with patch.dict(_os.environ, {}, clear=False):
_os.environ.pop("LITELLM_CONTENT_FILTER_ALLOW_EXTERNAL_PATHS", None)
with pytest.raises(ValueError, match="outside the allowed categories"):
guardrail._resolve_category_file_path("/etc/passwd")

View file

@ -259,6 +259,188 @@ async def test_pre_call_allows_authorized_model_in_batch_file():
)
@pytest.mark.asyncio
async def test_pre_call_skips_file_fetch_when_disabled_in_general_settings():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["*"])
with patch(
"litellm.proxy.proxy_server.general_settings",
{"disable_batch_input_file_rate_limiting": True},
):
result = await rate_limiter.async_pre_call_hook(
user_api_key_dict=user,
cache=MagicMock(),
data={"input_file_id": "file-abc123"},
call_type="acreate_batch",
)
assert result == {"input_file_id": "file-abc123"}
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called()
@pytest.mark.asyncio
async def test_pre_call_skips_file_fetch_for_configured_provider():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["*"])
data = {"input_file_id": "file-abc123", "model": "my-vllm-model"}
with (
patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]},
),
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
return_value={"custom_llm_provider": "hosted_vllm"},
),
patch("litellm.afile_content", new=AsyncMock()) as mock_afile_content,
):
result = await rate_limiter.async_pre_call_hook(
user_api_key_dict=user,
cache=MagicMock(),
data=data,
call_type="acreate_batch",
)
assert result == data
# A real skip must short-circuit before any file download or rate-limit
# work — assert the skip happened rather than the hook's error-recovery
# path (which also returns data unchanged).
mock_afile_content.assert_not_awaited()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.assert_not_called()
@pytest.mark.asyncio
async def test_pre_call_does_not_skip_for_spoofed_provider():
"""The provider skip is resolved from trusted deployment credentials, so a
user-supplied ``custom_llm_provider`` that is not backed by the routing
deployment must not trigger a skip: the input file must still be fetched
and the rate-limit counters incremented."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
# An applicable rate limit keeps the no-limits shortcut from firing, so the
# only thing that could prevent the fetch below is the provider skip. If the
# spoofed ``custom_llm_provider`` were honored, afile_content would never be
# awaited.
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 100}}
]
rate_limiter.parallel_request_limiter.atomic_check_and_increment_by_n = AsyncMock(
return_value={"overall_code": "OK", "statuses": []}
)
user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice", models=["*"])
mock_router = MagicMock()
mock_router.model_list = []
mock_router.resolve_model_name_from_model_id.return_value = "my-openai-model"
mock_content = MagicMock()
mock_content.content = (
b'{"body": {"model": "my-openai-model", '
b'"messages": [{"role": "user", "content": "hi"}]}}\n'
)
with (
patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]},
),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
return_value={"custom_llm_provider": "openai"},
),
patch(
"litellm.afile_content", new=AsyncMock(return_value=mock_content)
) as mock_afile_content,
):
await rate_limiter.async_pre_call_hook(
user_api_key_dict=user,
cache=MagicMock(),
data={
"input_file_id": "file-abc123",
"model": "my-openai-model",
"custom_llm_provider": "hosted_vllm",
},
call_type="acreate_batch",
)
# The spoofed provider did not short-circuit the skip decision: the file was
# fetched and the counters were incremented.
mock_afile_content.assert_awaited_once()
rate_limiter.parallel_request_limiter.atomic_check_and_increment_by_n.assert_awaited_once()
@pytest.mark.asyncio
async def test_count_input_file_usage_decodes_model_embedded_file_id():
import base64
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
original_file_id = "file-provider-xyz"
encoded_payload = (
base64.urlsafe_b64encode(
f"litellm:{original_file_id};model,my-vllm-batch".encode()
)
.decode()
.rstrip("=")
)
encoded_file_id = f"file-{encoded_payload}"
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
mock_content = MagicMock()
mock_content.content = b'{"custom_id": "1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "my-vllm-batch", "messages": [{"role": "user", "content": "hi"}]}}\n'
with (
patch(
"litellm.afile_content",
new=AsyncMock(return_value=mock_content),
) as mock_afile_content,
patch(
"litellm.proxy.proxy_server.llm_router",
MagicMock(),
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
return_value={
"api_key": "test-key",
"api_base": "http://vllm:8000/v1",
"custom_llm_provider": "hosted_vllm",
},
),
):
await rate_limiter.count_input_file_usage(
file_id=encoded_file_id,
custom_llm_provider="openai",
user_api_key_dict=UserAPIKeyAuth(api_key="sk-ok", user_id="alice"),
data={},
)
mock_afile_content.assert_awaited_once()
assert mock_afile_content.await_args.kwargs["file_id"] == original_file_id
assert mock_afile_content.await_args.kwargs["custom_llm_provider"] == "hosted_vllm"
@pytest.mark.asyncio
async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias():
"""After replace_model_in_jsonl, body.model is the provider id (e.g. gpt-5.5).
@ -323,3 +505,524 @@ async def test_pre_call_skips_check_when_no_models_present():
user_api_key_dict=user,
file_content_as_dict=[{"body": {}}],
)
# ---------------------------------------------------------------------------
# Skip-path helpers
# ---------------------------------------------------------------------------
def _make_rate_limiter():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
return _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
def test_get_batch_routing_model_uses_request_model_for_plain_file():
rate_limiter = _make_rate_limiter()
assert (
rate_limiter._get_batch_routing_model({"model": "gpt-4o-mini"}) == "gpt-4o-mini"
)
def test_get_batch_routing_model_prefers_file_bound_over_request_model():
"""``create_batch`` routes a model-embedded file id on its bound model and
ignores the top-level ``model``. The skip decision must use the same
precedence, otherwise a caller could point ``model`` at a skip-listed
provider while the file routes a rate-limited one."""
import base64
rate_limiter = _make_rate_limiter()
encoded = (
base64.urlsafe_b64encode(b"litellm:file-xyz;model,vllm-batch")
.decode()
.rstrip("=")
)
assert (
rate_limiter._get_batch_routing_model(
{"input_file_id": f"file-{encoded}", "model": "gpt-4o-mini"}
)
== "vllm-batch"
)
def test_get_batch_routing_model_returns_none_without_model_or_file():
rate_limiter = _make_rate_limiter()
assert rate_limiter._get_batch_routing_model({}) is None
assert rate_limiter._get_batch_routing_model({"input_file_id": ""}) is None
def test_get_batch_routing_model_decodes_model_embedded_file_id():
import base64
rate_limiter = _make_rate_limiter()
encoded = (
base64.urlsafe_b64encode(b"litellm:file-xyz;model,vllm-batch")
.decode()
.rstrip("=")
)
assert (
rate_limiter._get_batch_routing_model({"input_file_id": f"file-{encoded}"})
== "vllm-batch"
)
def test_get_batch_routing_model_uses_unified_file_id_target():
rate_limiter = _make_rate_limiter()
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils.decode_model_from_file_id",
return_value=None,
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
return_value="unified-id",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_models_from_unified_file_id",
return_value=["model-a", "model-b"],
),
):
assert (
rate_limiter._get_batch_routing_model({"input_file_id": "file-managed"})
== "model-a"
)
def test_key_requires_batch_model_access_check_branches():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
check = _PROXY_BatchRateLimiter._key_requires_batch_model_access_check
assert check(UserAPIKeyAuth(api_key="sk", models=["*"])) is False
assert check(UserAPIKeyAuth(api_key="sk", models=["all-proxy-models"])) is False
assert (
check(UserAPIKeyAuth(api_key="sk", models=[], access_group_ids=["grp"])) is True
)
assert check(UserAPIKeyAuth(api_key="sk", models=[])) is False
assert check(UserAPIKeyAuth(api_key="sk", models=["gpt-4o-mini"])) is True
# Wildcard / all-proxy-models grant access to every model, so
# can_key_call_model passes any model regardless of access groups (which
# only ever widen access). Such keys must not be forced to download and
# validate the JSONL even when access_group_ids are also present.
assert (
check(UserAPIKeyAuth(api_key="sk", models=["*"], access_group_ids=["grp"]))
is False
)
assert (
check(
UserAPIKeyAuth(
api_key="sk", models=["all-proxy-models"], access_group_ids=["grp"]
)
)
is False
)
# A concrete model allowlist is still a subset even with access groups.
assert (
check(
UserAPIKeyAuth(
api_key="sk", models=["gpt-4o-mini"], access_group_ids=["grp"]
)
)
is True
)
def test_has_applicable_batch_rate_limits():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
has_limits = _PROXY_BatchRateLimiter._has_applicable_batch_rate_limits
assert has_limits([{"rate_limit": {"tokens_per_unit": 100}}]) is True
assert has_limits([{"rate_limit": {"requests_per_unit": 5}}]) is True
assert has_limits([{"rate_limit": {"max_parallel_requests": 2}}]) is True
assert has_limits([{"rate_limit": {}}, {}]) is False
def test_should_skip_returns_false_when_key_needs_model_access_check():
rate_limiter = _make_rate_limiter()
user = UserAPIKeyAuth(api_key="sk", models=["gpt-4o-mini"])
should_skip, descriptors = rate_limiter._should_skip_batch_input_file_processing(
data={"input_file_id": "file-abc"}, user_api_key_dict=user
)
assert should_skip is False
assert descriptors is None
def test_should_skip_ignores_client_supplied_metadata_flag():
"""A caller must not be able to bypass batch rate limits by setting
``litellm_metadata.skip_batch_input_file_rate_limiting`` in the request
body. The skip decision is server-controlled only, so with applicable rate
limits the JSONL is still processed despite the client flag."""
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
with patch("litellm.proxy.proxy_server.general_settings", {}):
should_skip, descriptors = (
rate_limiter._should_skip_batch_input_file_processing(
data={
"input_file_id": "file-abc",
"litellm_metadata": {"skip_batch_input_file_rate_limiting": True},
},
user_api_key_dict=user,
)
)
assert should_skip is False
def test_should_not_skip_for_forged_model_embedded_file_id():
"""A ``file-<base64>`` id embeds an unsigned model name the caller fully
controls, so a caller can re-encode any accessible provider file id with a
skip-listed model while the JSONL still routes rate-limited ``body.model``
entries. The per-model skip must therefore never fire: with applicable rate
limits, a forged skip-listed file-bound model still falls through to file
processing and counter enforcement."""
import base64
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
encoded = (
base64.urlsafe_b64encode(b"litellm:file-xyz;model,gpt-4o-mini")
.decode()
.rstrip("=")
)
with patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_models": ["gpt-4o-mini"]},
):
should_skip, descriptors = (
rate_limiter._should_skip_batch_input_file_processing(
data={"input_file_id": f"file-{encoded}"},
user_api_key_dict=user,
)
)
assert should_skip is False
assert descriptors is not None
def test_should_not_skip_for_skip_listed_top_level_model():
"""A caller must not bypass batch rate limits by naming a skip-listed model
in the top-level ``model`` while routing a different model through the JSONL
``body.model`` entries. No per-model skip exists, so a skip-listed model over
a plain file still gets processed."""
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
with patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_models": ["gpt-4o-mini"]},
):
should_skip, descriptors = (
rate_limiter._should_skip_batch_input_file_processing(
data={"model": "gpt-4o-mini", "input_file_id": "file-abc"},
user_api_key_dict=user,
)
)
assert should_skip is False
def test_should_not_skip_when_file_bound_provider_is_rate_limited():
"""A caller must not bypass batch rate limits by pointing the top-level
``model`` at a skip-listed provider while the model-embedded ``input_file_id``
routes to a rate-limited provider. ``create_batch`` runs the batch on the
file-bound model, so the skip decision must resolve the provider from that
model and still process the file when its provider is not skip-listed."""
import base64
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
encoded = (
base64.urlsafe_b64encode(b"litellm:file-orig;model,vllm-batch")
.decode()
.rstrip("=")
)
def _creds(model_id, **kwargs):
provider = "hosted_vllm" if model_id == "vllm-batch" else "openai"
return {"custom_llm_provider": provider}
with (
patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_providers": ["openai"]},
),
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
side_effect=_creds,
),
):
should_skip, descriptors = (
rate_limiter._should_skip_batch_input_file_processing(
data={"input_file_id": f"file-{encoded}", "model": "gpt-skip"},
user_api_key_dict=user,
)
)
assert should_skip is False
assert descriptors is not None
def test_should_skip_when_file_bound_provider_is_skip_listed():
"""The provider skip must still fire when the model the batch actually runs
on (the file-bound model) resolves to a skip-listed provider, even if the
top-level ``model`` resolves to a different, non-skipped provider."""
import base64
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
encoded = (
base64.urlsafe_b64encode(b"litellm:file-orig;model,vllm-batch")
.decode()
.rstrip("=")
)
def _creds(model_id, **kwargs):
provider = "hosted_vllm" if model_id == "vllm-batch" else "openai"
return {"custom_llm_provider": provider}
with (
patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_providers": ["hosted_vllm"]},
),
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
side_effect=_creds,
),
):
should_skip, descriptors = (
rate_limiter._should_skip_batch_input_file_processing(
data={"input_file_id": f"file-{encoded}", "model": "gpt-skip"},
user_api_key_dict=user,
)
)
assert should_skip is True
def test_warns_once_for_unsupported_model_skip_setting():
"""Operators who set the no-op per-model skip key get a single warning so a
misconfigured deployment does not silently leave batch limits unenforced."""
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
with (
patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_models": ["gpt-4o-mini"]},
),
patch(
"litellm.proxy.hooks.batch_rate_limiter.verbose_proxy_logger"
) as mock_logger,
):
for _ in range(3):
rate_limiter._should_skip_batch_input_file_processing(
data={"model": "gpt-4o-mini", "input_file_id": "file-abc"},
user_api_key_dict=user,
)
assert mock_logger.warning.call_count == 1
assert (
"skip_batch_input_file_rate_limiting_for_models"
in mock_logger.warning.call_args[0][0]
)
def test_no_warning_when_model_skip_setting_absent():
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"requests_per_unit": 5}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
with (
patch(
"litellm.proxy.proxy_server.general_settings",
{"skip_batch_input_file_rate_limiting_for_providers": ["openai"]},
),
patch(
"litellm.proxy.hooks.batch_rate_limiter.verbose_proxy_logger"
) as mock_logger,
):
rate_limiter._should_skip_batch_input_file_processing(
data={"model": "gpt-4o-mini", "input_file_id": "file-abc"},
user_api_key_dict=user,
)
mock_logger.warning.assert_not_called()
def test_should_skip_when_no_rate_limits_configured():
rate_limiter = _make_rate_limiter()
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {}}
]
user = UserAPIKeyAuth(api_key="sk", models=["*"])
with patch("litellm.proxy.proxy_server.general_settings", {}):
should_skip, descriptors = (
rate_limiter._should_skip_batch_input_file_processing(
data={"model": "gpt-4o-mini", "input_file_id": "file-abc"},
user_api_key_dict=user,
)
)
assert should_skip is True
assert descriptors is None
def test_should_not_skip_and_reuses_descriptors_when_limits_present():
rate_limiter = _make_rate_limiter()
descriptors = [{"rate_limit": {"tokens_per_unit": 100}}]
rate_limiter.parallel_request_limiter._create_rate_limit_descriptors.return_value = (
descriptors
)
user = UserAPIKeyAuth(api_key="sk", models=["*"])
with patch("litellm.proxy.proxy_server.general_settings", {}):
should_skip, returned = rate_limiter._should_skip_batch_input_file_processing(
data={"model": "gpt-4o-mini", "input_file_id": "file-abc"},
user_api_key_dict=user,
)
assert should_skip is False
assert returned is descriptors
def test_resolve_fetch_params_uses_request_model_credentials():
rate_limiter = _make_rate_limiter()
with (
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
return_value={
"api_key": "k",
"api_base": "http://vllm:8000/v1",
"custom_llm_provider": "hosted_vllm",
},
),
):
provider_file_id, fetch_kwargs = (
rate_limiter._resolve_batch_input_file_fetch_params(
file_id="file-plain-openai",
custom_llm_provider="openai",
data={"model": "my-vllm-batch"},
)
)
assert provider_file_id == "file-plain-openai"
assert fetch_kwargs["model"] == "my-vllm-batch"
assert fetch_kwargs["custom_llm_provider"] == "hosted_vllm"
assert fetch_kwargs["api_base"] == "http://vllm:8000/v1"
def test_resolve_fetch_params_fails_open_on_credential_lookup_error():
rate_limiter = _make_rate_limiter()
with (
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
side_effect=HTTPException(status_code=404, detail="no creds"),
),
):
provider_file_id, fetch_kwargs = (
rate_limiter._resolve_batch_input_file_fetch_params(
file_id="file-plain-openai",
custom_llm_provider="openai",
data={"model": "my-vllm-batch"},
)
)
assert provider_file_id == "file-plain-openai"
assert fetch_kwargs == {"custom_llm_provider": "openai"}
def test_resolve_fetch_params_model_embedded_fails_open_on_credential_error():
import base64
rate_limiter = _make_rate_limiter()
encoded = (
base64.urlsafe_b64encode(b"litellm:file-orig;model,vllm-batch")
.decode()
.rstrip("=")
)
encoded_file_id = f"file-{encoded}"
get_credentials = MagicMock(
side_effect=HTTPException(status_code=404, detail="no creds")
)
with (
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model",
get_credentials,
),
):
provider_file_id, fetch_kwargs = (
rate_limiter._resolve_batch_input_file_fetch_params(
file_id=encoded_file_id,
custom_llm_provider="openai",
data={},
)
)
get_credentials.assert_called_once()
assert provider_file_id == "file-orig"
assert fetch_kwargs == {"custom_llm_provider": "openai"}
@pytest.mark.asyncio
async def test_check_and_increment_computes_descriptors_when_not_passed():
from litellm.proxy.hooks.batch_rate_limiter import (
BatchFileUsage,
_PROXY_BatchRateLimiter,
)
parallel_request_limiter = MagicMock()
parallel_request_limiter._create_rate_limit_descriptors.return_value = [
{"rate_limit": {"tokens_per_unit": 100}}
]
parallel_request_limiter.atomic_check_and_increment_by_n = AsyncMock(
return_value={"overall_code": "OK", "statuses": []}
)
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=parallel_request_limiter,
)
await rate_limiter._check_and_increment_batch_counters(
user_api_key_dict=UserAPIKeyAuth(api_key="sk", models=["*"]),
data={"model": "gpt-4o-mini"},
batch_usage=BatchFileUsage(total_tokens=10, request_count=1),
descriptors=None,
)
parallel_request_limiter._create_rate_limit_descriptors.assert_called_once()
@pytest.mark.asyncio
async def test_count_input_file_usage_raises_on_non_bytes_content():
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
bad_content = MagicMock()
bad_content.content = "not-bytes"
with patch("litellm.afile_content", new=AsyncMock(return_value=bad_content)):
with pytest.raises(ValueError, match="Expected bytes content"):
await rate_limiter.count_input_file_usage(
file_id="file-plain",
custom_llm_provider="openai",
user_api_key_dict=UserAPIKeyAuth(api_key="sk", models=["*"]),
data={},
)

View file

@ -0,0 +1,444 @@
"""
Unit tests for watsonx_proxy_route endpoint.
Tests the Watsonx pass-through endpoint that handles automatic IAM token management
and version parameter injection.
"""
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import pytest
from fastapi import HTTPException, Request, Response
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
watsonx_proxy_route,
)
class TestWatsonxProxyRoute:
"""Tests for the Watsonx pass-through route."""
@pytest.mark.asyncio
async def test_watsonx_proxy_route_success_non_streaming(self):
"""Test successful non-streaming request through Watsonx proxy route."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "application/json"}
mock_request.json = AsyncMock(return_value={"stream": False, "input": "test"})
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock provider config
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token"
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(
return_value={"model_id": "ibm/granite-13b-chat-v2", "results": []}
)
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
result = await watsonx_proxy_route(
endpoint="ml/v1/text/generation",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify provider config was called correctly
mock_provider_config.get_complete_url.assert_called_once()
mock_provider_config.validate_environment.assert_called_once()
# Verify create_pass_through_route was called with correct parameters
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert call_args["endpoint"] == "ml/v1/text/generation"
assert (
call_args["target"]
== "https://us-south.ml.cloud.ibm.com/ml/v1/text/generation"
)
assert (
call_args["custom_headers"]["Authorization"] == "Bearer test-iam-token"
)
assert call_args["is_streaming_request"] is False
assert call_args["custom_llm_provider"] == "watsonx"
assert (
call_args["query_params"]["version"]
== litellm.WATSONX_DEFAULT_API_VERSION
)
# Verify endpoint function was called
mock_endpoint_func.assert_called_once_with(
mock_request, mock_response, mock_user_api_key_dict
)
assert result == {"model_id": "ibm/granite-13b-chat-v2", "results": []}
@pytest.mark.asyncio
async def test_watsonx_proxy_route_success_streaming(self):
"""Test successful streaming request through Watsonx proxy route."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "application/json"}
mock_request.json = AsyncMock(return_value={"stream": True, "input": "test"})
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock provider config
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation_stream",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token"
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(return_value="streaming_response")
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
result = await watsonx_proxy_route(
endpoint="ml/v1/text/generation_stream",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify create_pass_through_route was called with streaming enabled
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert call_args["is_streaming_request"] is True
assert result == "streaming_response"
@pytest.mark.asyncio
async def test_watsonx_proxy_route_get_request(self):
"""Test GET request through Watsonx proxy route."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.query_params = {"project_id": "test-project"}
mock_request.headers = {}
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock provider config
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
"https://us-south.ml.cloud.ibm.com/ml/v1/models",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token"
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(return_value={"resources": []})
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
result = await watsonx_proxy_route(
endpoint="ml/v1/models",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify is_streaming_request is False for GET requests
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert call_args["is_streaming_request"] is False
assert result == {"resources": []}
@pytest.mark.asyncio
async def test_watsonx_proxy_route_multipart_form_data(self):
"""Test multipart/form-data request through Watsonx proxy route."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "multipart/form-data; boundary=----"}
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock form data
mock_form_data = {"file": "test_file", "stream": False}
# Mock provider config
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
"https://us-south.ml.cloud.ibm.com/ml/v1/text/tokenization",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token"
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(return_value={"token_count": 10})
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_form_data",
return_value=mock_form_data,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
result = await watsonx_proxy_route(
endpoint="ml/v1/text/tokenization",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify is_streaming_request is False for non-streaming form data
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert call_args["is_streaming_request"] is False
assert result == {"token_count": 10}
@pytest.mark.asyncio
async def test_watsonx_proxy_route_no_provider_config(self):
"""Test that HTTPException is raised when provider config is not found."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "application/json"}
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=None,
),
):
with pytest.raises(HTTPException) as exc_info:
await watsonx_proxy_route(
endpoint="ml/v1/text/generation",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
assert exc_info.value.status_code == 404
assert exc_info.value.detail == "Watsonx passthrough config not found"
@pytest.mark.asyncio
async def test_watsonx_proxy_route_version_parameter_injection(self):
"""Test that version parameter is correctly injected into query params."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "application/json"}
mock_request.json = AsyncMock(return_value={"input": "test"})
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock provider config
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token"
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(return_value={})
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
await watsonx_proxy_route(
endpoint="ml/v1/text/generation",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify version parameter is injected
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert "query_params" in call_args
assert "version" in call_args["query_params"]
assert (
call_args["query_params"]["version"]
== litellm.WATSONX_DEFAULT_API_VERSION
)
@pytest.mark.asyncio
async def test_watsonx_proxy_route_custom_headers_from_validate_environment(self):
"""Test that custom headers from validate_environment are passed through."""
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "application/json"}
mock_request.json = AsyncMock(return_value={"input": "test"})
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock provider config with custom headers
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
"https://us-south.ml.cloud.ibm.com/ml/v1/text/generation",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token",
"X-Custom-Header": "custom-value",
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(return_value={})
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
await watsonx_proxy_route(
endpoint="ml/v1/text/generation",
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify custom headers are passed through
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert "custom_headers" in call_args
assert (
call_args["custom_headers"]["Authorization"] == "Bearer test-iam-token"
)
assert call_args["custom_headers"]["X-Custom-Header"] == "custom-value"
@pytest.mark.asyncio
async def test_watsonx_proxy_route_different_endpoints(self):
"""Test various Watsonx endpoint paths."""
endpoints = [
"ml/v1/text/generation",
"ml/v1/text/tokenization",
"ml/v1/deployments/test-deployment/text/generation",
"ml/v1/models",
]
for endpoint_path in endpoints:
# Setup mocks
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = {}
mock_request.headers = {"content-type": "application/json"}
mock_request.json = AsyncMock(return_value={"input": "test"})
mock_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
# Mock provider config
mock_provider_config = MagicMock()
mock_provider_config.get_complete_url.return_value = (
f"https://us-south.ml.cloud.ibm.com/{endpoint_path}",
{},
)
mock_provider_config.validate_environment.return_value = {
"Authorization": "Bearer test-iam-token"
}
# Mock endpoint function
mock_endpoint_func = AsyncMock(return_value={})
with (
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_passthrough_config",
return_value=mock_provider_config,
),
patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
return_value=mock_endpoint_func,
) as mock_create_route,
):
await watsonx_proxy_route(
endpoint=endpoint_path,
request=mock_request,
fastapi_response=mock_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify endpoint is passed correctly
mock_create_route.assert_called_once()
call_args = mock_create_route.call_args[1]
assert call_args["endpoint"] == endpoint_path
assert (
call_args["target"]
== f"https://us-south.ml.cloud.ibm.com/{endpoint_path}"
)

View file

@ -3185,3 +3185,358 @@ async def test_view_spend_logs_date_range_hashes_sk_api_key(client, monkeypatch)
assert where["api_key"] == "hashed::sk-raw-admin-token"
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
class _SpendScopeMockPrismaClient:
def __init__(self, get_data_returns=None, find_many_returns=None):
self._get_data_returns = (
get_data_returns if get_data_returns is not None else []
)
self._find_many_returns = (
find_many_returns if find_many_returns is not None else []
)
self.get_data_calls = []
self.find_many_calls = []
client = self
class _VerificationTokenTable:
async def find_many(self, where=None, order=None, include=None):
client.find_many_calls.append(
{"where": where, "order": order, "include": include}
)
return client._find_many_returns
class _DB:
def __init__(self):
self.litellm_verificationtoken = _VerificationTokenTable()
self.db = _DB()
async def get_data(self, table_name=None, query_type=None, **kwargs):
self.get_data_calls.append(
{"table_name": table_name, "query_type": query_type, **kwargs}
)
if query_type == "find_unique":
return self._get_data_returns[0] if self._get_data_returns else None
return self._get_data_returns
@pytest.mark.asyncio
async def test_spend_key_fn_proxy_admin_returns_all_keys(client, monkeypatch):
"""Admins keep their existing full-table view of /spend/keys."""
mock_keys = [
{"token": "hashed-a", "user_id": "alice", "spend": 10.0},
{"token": "hashed-b", "user_id": "bob", "spend": 5.0},
]
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=mock_keys)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"
)
try:
response = client.get(
"/spend/keys", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
# Admin path: goes through get_data (full table), never the scoped find_many
assert len(mock_prisma.get_data_calls) == 1
assert mock_prisma.get_data_calls[0]["table_name"] == "key"
assert mock_prisma.get_data_calls[0]["query_type"] == "find_all"
assert mock_prisma.find_many_calls == []
assert response.json() == mock_keys
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_key_fn_proxy_admin_view_only_returns_all_keys(client, monkeypatch):
"""View-only admins are still admins for this endpoint."""
mock_keys = [{"token": "hashed-a", "user_id": "alice"}]
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=mock_keys)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, user_id="admin_viewer"
)
try:
response = client.get(
"/spend/keys", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
assert mock_prisma.find_many_calls == []
assert len(mock_prisma.get_data_calls) == 1
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"role",
[LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY],
)
async def test_spend_key_fn_internal_user_scoped_to_own_keys(client, monkeypatch, role):
"""Both internal-user roles must only see keys they own."""
caller_owned_keys = [
{"token": "hashed-mine-1", "user_id": "alice", "spend": 2.0},
{"token": "hashed-mine-2", "user_id": "alice", "spend": 1.0},
]
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=caller_owned_keys)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=role, user_id="alice"
)
try:
response = client.get(
"/spend/keys", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
# Non-admin path goes through the same get_data helper as admin,
# but with a user_id scope so only the caller's rows come back.
assert mock_prisma.find_many_calls == []
assert len(mock_prisma.get_data_calls) == 1
call = mock_prisma.get_data_calls[0]
assert call["table_name"] == "key"
assert call["query_type"] == "find_all"
assert call["user_id"] == "alice"
assert response.json() == caller_owned_keys
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_key_fn_internal_user_without_user_id_returns_empty(
client, monkeypatch
):
"""
A non-admin key with no user_id has no tenant scope. Returning the full
table would re-introduce the leak; return an empty list instead.
"""
mock_prisma = _SpendScopeMockPrismaClient(
get_data_returns=[{"token": "do-not-leak"}],
find_many_returns=[{"token": "do-not-leak"}],
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id=None
)
try:
response = client.get(
"/spend/keys", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
assert response.json() == []
assert mock_prisma.get_data_calls == []
assert mock_prisma.find_many_calls == []
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_user_fn_proxy_admin_returns_all_users_without_user_id(
client, monkeypatch
):
"""Admins keep their existing full-table view of /spend/users."""
mock_users = [
{"user_id": "alice", "user_email": "alice@example.com", "spend": 1.0},
{"user_id": "bob", "user_email": "bob@example.com", "spend": 2.0},
]
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=mock_users)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"
)
try:
response = client.get(
"/spend/users", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
assert len(mock_prisma.get_data_calls) == 1
assert mock_prisma.get_data_calls[0]["table_name"] == "user"
assert mock_prisma.get_data_calls[0]["query_type"] == "find_all"
assert response.json() == mock_users
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_user_fn_proxy_admin_can_query_specific_user_id(
client, monkeypatch
):
"""Admins can still target a specific user_id."""
mock_user = {
"user_id": "carol",
"user_email": "carol@example.com",
"spend": 7.0,
}
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=[mock_user])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"
)
try:
response = client.get(
"/spend/users",
params={"user_id": "carol"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert len(mock_prisma.get_data_calls) == 1
assert mock_prisma.get_data_calls[0]["query_type"] == "find_unique"
assert mock_prisma.get_data_calls[0]["user_id"] == "carol"
assert response.json() == [mock_user]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"role",
[LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY],
)
async def test_spend_user_fn_internal_user_scoped_without_user_id(
client, monkeypatch, role
):
"""No user_id supplied -> must query the caller's own row, not the table."""
own_row = {"user_id": "alice", "user_email": "alice@example.com", "spend": 3.0}
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=[own_row])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=role, user_id="alice"
)
try:
response = client.get(
"/spend/users", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
assert len(mock_prisma.get_data_calls) == 1
assert mock_prisma.get_data_calls[0]["query_type"] == "find_unique"
assert mock_prisma.get_data_calls[0]["user_id"] == "alice"
assert response.json() == [own_row]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_user_fn_internal_user_supplying_other_user_id_returns_403(
client, monkeypatch
):
"""
An internal user passing user_id=victim must be rejected outright, not
silently rewritten. A 403 makes the attempt observable in logs.
"""
leaked_victim_row = {
"user_id": "victim",
"user_email": "victim@example.com",
"spend": 999.0,
}
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=[leaked_victim_row])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice"
)
try:
response = client.get(
"/spend/users",
params={"user_id": "victim"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 403
assert mock_prisma.get_data_calls == []
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_user_fn_internal_user_supplying_own_user_id_is_allowed(
client, monkeypatch
):
"""
Passing your own user_id explicitly is fine — the 403 only fires when
the supplied id differs from the caller's.
"""
own_row = {"user_id": "alice", "user_email": "alice@example.com", "spend": 3.0}
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=[own_row])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice"
)
try:
response = client.get(
"/spend/users",
params={"user_id": "alice"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert len(mock_prisma.get_data_calls) == 1
assert mock_prisma.get_data_calls[0]["query_type"] == "find_unique"
assert mock_prisma.get_data_calls[0]["user_id"] == "alice"
assert response.json() == [own_row]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_user_fn_internal_user_without_user_id_returns_empty(
client, monkeypatch
):
"""
A non-admin key with no user_id has no tenant scope -> return empty,
never the full table. Same defensive contract as /spend/keys.
"""
mock_prisma = _SpendScopeMockPrismaClient(
get_data_returns=[{"user_id": "do-not-leak"}]
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, user_id=None
)
try:
response = client.get(
"/spend/users", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
assert response.json() == []
assert mock_prisma.get_data_calls == []
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_user_fn_strips_password_field(client, monkeypatch):
"""
Existing password-redaction behavior must be preserved on the scoped
path so we don't regress a separate disclosure when adding the fix.
"""
own_row = {
"user_id": "alice",
"user_email": "alice@example.com",
"password": "hashed-password-must-not-leak",
"spend": 1.0,
}
mock_prisma = _SpendScopeMockPrismaClient(get_data_returns=[own_row])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice"
)
try:
response = client.get(
"/spend/users", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 200
body = response.json()
assert len(body) == 1
assert "password" not in body[0]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)

View file

@ -5165,6 +5165,110 @@ async def test_async_data_generator_passes_through_google_native_sse_bytes():
assert yielded_text[-1] == "data: [DONE]\n\n"
@pytest.mark.asyncio
async def test_async_data_generator_google_genai_stream_omits_openai_done():
"""
google-genai SDK streamGenerateContent?alt=sse must not receive data: [DONE].
"""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {
"model": "gemini-2.0-flash",
"_litellm_skip_openai_stream_done": True,
}
gemini_event = (
b'data: {"candidates": [{"content": {"parts": [{"text": "Hi"}]}}]}\n\n'
)
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
yield gemini_event
async def aclose(self):
pass
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
yielded_data = []
async for data in async_data_generator(
mock_response, mock_user_api_key_dict, mock_request_data
):
yielded_data.append(data)
yielded_text = [
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in yielded_data
]
assert yielded_text == [gemini_event.decode("utf-8")]
assert "[DONE]" not in "".join(yielded_text)
@pytest.mark.asyncio
async def test_async_data_generator_google_genai_stream_forwards_error_without_done():
"""Stream errors must still reach the client when OpenAI [DONE] is skipped."""
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
error_sse = 'data: {"error": {"message": "stream failed"}}\n\n'
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
mock_request_data = {
"model": "gemini-2.0-flash",
"_litellm_skip_openai_stream_done": True,
}
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
yield error_sse
async def aclose(self):
pass
mock_response = MockStream()
mock_response.aclose = AsyncMock()
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
yielded_data = []
async for data in async_data_generator(
mock_response, mock_user_api_key_dict, mock_request_data
):
yielded_data.append(data)
yielded_text = [
chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk
for chunk in yielded_data
]
assert yielded_text == [error_sse]
assert "[DONE]" not in "".join(yielded_text)
@pytest.mark.asyncio
async def test_async_data_generator_cleanup_on_normal_completion():
"""

View file

@ -0,0 +1,32 @@
# tests/test_litellm/proxy/test__types.py
from litellm.proxy._types import LiteLLM_TeamMembership
def test_team_membership_budget_table_optional_no_crash():
"""
Regression test for #28689
Pydantic v2: Optional[T] without default = required field.
When budget_id is null, DB join returns no litellm_budget_table key.
model_validate must NOT raise 'Field required'.
"""
data = {
"user_id": "test-user",
"team_id": "test-team",
"budget_id": None,
# litellm_budget_table intentionally absent (as DB join returns when budget_id is null)
}
result = LiteLLM_TeamMembership.model_validate(data)
assert result.litellm_budget_table is None
def test_team_membership_budget_table_present_still_works():
"""When budget_id exists, litellm_budget_table should still be populated."""
data = {
"user_id": "test-user",
"team_id": "test-team",
"budget_id": "some-budget-id",
"litellm_budget_table": None,
}
result = LiteLLM_TeamMembership.model_validate(data)
assert result.litellm_budget_table is None

View file

@ -0,0 +1,47 @@
"""
Validate that AWS GovCloud (Bedrock us-gov-*) Haiku 4.5 entries carry
the 1-hour cache write tier.
AWS Bedrock GovCloud pricing applies a +20% premium over global
Anthropic rates. Global Haiku 4.5 1h cache write is $2.00/MTok; us-gov
is therefore $2.40/MTok — exactly 1.6x the 5-minute rate of $1.50/MTok.
Source: https://aws.amazon.com/bedrock/pricing/
"""
import json
import os
import pytest
@pytest.fixture(scope="module")
def model_data():
json_path = os.path.join(
os.path.dirname(__file__), "../../model_prices_and_context_window.json"
)
with open(json_path) as f:
return json.load(f)
HAIKU_USGOV_KEYS = [
"bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0",
"bedrock/us-gov-west-1/anthropic.claude-haiku-4-5-20251001-v1:0",
]
@pytest.mark.parametrize("model_key", HAIKU_USGOV_KEYS)
def test_usgov_haiku_4_5_1hr_cache_write(model_data, model_key):
assert model_key in model_data, f"Missing model entry: {model_key}"
info = model_data[model_key]
assert (
info["cache_creation_input_token_cost"] == 1.5e-06
), f"{model_key}: 5m cache write should be $1.50/MTok"
assert (
info["cache_creation_input_token_cost_above_1hr"] == 2.4e-06
), f"{model_key}: 1h cache write should be $2.40/MTok"
ratio = (
info["cache_creation_input_token_cost_above_1hr"]
/ info["cache_creation_input_token_cost"]
)
assert abs(ratio - 1.6) < 1e-9, f"{model_key}: 1h/5m ratio is {ratio}, expected 1.6"

View file

@ -0,0 +1,132 @@
"""
Validate AWS GovCloud (Bedrock us-gov-*) Anthropic pricing entries.
AWS Bedrock pricing in GovCloud carries a +20% premium over the global
Anthropic prices (not the +10% commercial-US premium). Until 2026-05-22
these entries silently mirrored commercial US, undercharging customers
by ~9%.
Source: https://aws.amazon.com/bedrock/pricing/
Sonnet 4.5 in us-gov-* (per million tokens):
input = $3.60
output = $18.00
cache write 5m = $4.50
cache write 1h = $7.20
cache read = $0.36
Reference: https://github.com/BerriAI/litellm/issues/27120
"""
import json
import os
import pytest
@pytest.fixture(scope="module")
def model_data():
json_path = os.path.join(
os.path.dirname(__file__), "../../model_prices_and_context_window.json"
)
with open(json_path) as f:
return json.load(f)
SONNET_4_5_USGOV_KEYS = [
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0",
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0",
"bedrock/us-gov-east-1/claude-sonnet-4-5-20250929-v1:0",
"bedrock/us-gov-west-1/claude-sonnet-4-5-20250929-v1:0",
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0",
]
@pytest.mark.parametrize("model_key", SONNET_4_5_USGOV_KEYS)
def test_usgov_sonnet_4_5_pricing(model_data, model_key):
"""Each us-gov sonnet-4-5 entry must carry the +20%-over-global rates
that AWS publishes on the GovCloud pricing page.
"""
assert model_key in model_data, f"Missing model entry: {model_key}"
info = model_data[model_key]
assert info["input_cost_per_token"] == 3.6e-06, (
f"{model_key}: input_cost_per_token should be $3.60/MTok "
f"(got {info['input_cost_per_token']})"
)
assert (
info["output_cost_per_token"] == 1.8e-05
), f"{model_key}: output_cost_per_token should be $18.00/MTok"
assert (
info["cache_creation_input_token_cost"] == 4.5e-06
), f"{model_key}: 5m cache write should be $4.50/MTok"
assert (
info["cache_creation_input_token_cost_above_1hr"] == 7.2e-06
), f"{model_key}: 1h cache write should be $7.20/MTok"
assert (
info["cache_read_input_token_cost"] == 3.6e-07
), f"{model_key}: cache read should be $0.36/MTok"
def test_usgov_carries_20_percent_premium_over_global(model_data):
"""The us-gov rates must equal 1.2x the global anthropic.* rates,
matching AWS's documented GovCloud uplift.
"""
global_key = "anthropic.claude-sonnet-4-5-20250929-v1:0"
usgov_key = "bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0"
global_info = model_data[global_key]
usgov_info = model_data[usgov_key]
for field in (
"input_cost_per_token",
"output_cost_per_token",
"cache_creation_input_token_cost",
"cache_creation_input_token_cost_above_1hr",
"cache_read_input_token_cost",
):
ratio = usgov_info[field] / global_info[field]
assert (
abs(ratio - 1.2) < 1e-9
), f"{field}: us-gov / global ratio is {ratio}, expected 1.2"
# The us-gov.anthropic.* cross-region inference profile is the only us-gov
# entry that carries the 1M-context `_above_200k_tokens` pricing tier — the
# bedrock/us-gov-{east,west}-1/ entries are capped at 200k tokens.
USGOV_CROSS_REGION_KEY = "us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0"
EXPECTED_USGOV_ABOVE_200K = {
"input_cost_per_token_above_200k_tokens": 7.2e-06,
"output_cost_per_token_above_200k_tokens": 2.7e-05,
"cache_creation_input_token_cost_above_200k_tokens": 9.0e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.44e-05,
"cache_read_input_token_cost_above_200k_tokens": 7.2e-07,
}
@pytest.mark.parametrize("field,expected", EXPECTED_USGOV_ABOVE_200K.items())
def test_usgov_cross_region_above_200k_carries_gov_premium(model_data, field, expected):
"""The `_above_200k_tokens` tier on the us-gov cross-region inference
profile must also carry the +20% GovCloud uplift. The original PR
corrected the base rates but left the 200k-tier fields at the +10%
commercial-US rates, undercharging long-context requests.
"""
info = model_data[USGOV_CROSS_REGION_KEY]
assert field in info, f"{USGOV_CROSS_REGION_KEY}: missing field {field}"
assert (
info[field] == expected
), f"{USGOV_CROSS_REGION_KEY}: {field} should be {expected} (got {info[field]})"
def test_usgov_cross_region_above_200k_ratio_to_global(model_data):
"""Cross-check via the property-based invariant: every `_above_200k_tokens`
field on the us-gov cross-region profile must equal 1.2x the global
anthropic.* rate, the same GovCloud uplift the base tier carries.
"""
global_key = "anthropic.claude-sonnet-4-5-20250929-v1:0"
global_info = model_data[global_key]
usgov_info = model_data[USGOV_CROSS_REGION_KEY]
for field in EXPECTED_USGOV_ABOVE_200K:
ratio = usgov_info[field] / global_info[field]
assert (
abs(ratio - 1.2) < 1e-9
), f"{field}: us-gov / global ratio is {ratio}, expected 1.2"

Some files were not shown because too many files have changed in this diff Show more