mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_watsonx_orchestrate_agents
This commit is contained in:
commit
a931b54e59
118 changed files with 13482 additions and 632 deletions
25
.github/workflows/create-release-branch.yml
vendored
25
.github/workflows/create-release-branch.yml
vendored
|
|
@ -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}`);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -106,6 +106,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
# Health & ops
|
||||
"/health",
|
||||
"/metrics",
|
||||
"/watsonx"
|
||||
)
|
||||
|
||||
GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", "")}
|
||||
|
|
|
|||
|
|
@ -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", "")}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
litellm/llms/watsonx/passthrough/__init__.py
Normal file
0
litellm/llms/watsonx/passthrough/__init__.py
Normal file
69
litellm/llms/watsonx/passthrough/transformation.py
Normal file
69
litellm/llms/watsonx/passthrough/transformation.py
Normal 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)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
115
litellm/utils.py
115
litellm/utils.py
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
43
tests/test_litellm/litellm_core_utils/test_fallback_utils.py
Normal file
43
tests/test_litellm/litellm_core_utils/test_fallback_utils.py
Normal 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",
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]}}
|
||||
|
|
@ -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"}}
|
||||
|
|
@ -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"}}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)")
|
||||
|
|
@ -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"}},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 {})
|
||||
|
|
@ -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"
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
@ -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
|
||||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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")
|
||||
|
|
@ -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={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
32
tests/test_litellm/test__types.py
Normal file
32
tests/test_litellm/test__types.py
Normal 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
|
||||
47
tests/test_litellm/test_bedrock_usgov_haiku_1hr_cache.py
Normal file
47
tests/test_litellm/test_bedrock_usgov_haiku_1hr_cache.py
Normal 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"
|
||||
132
tests/test_litellm/test_bedrock_usgov_pricing.py
Normal file
132
tests/test_litellm/test_bedrock_usgov_pricing.py
Normal 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
Loading…
Add table
Reference in a new issue