mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge remote-tracking branch 'origin/main' into litellm-ocr-new-mapping
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> # Conflicts: # litellm-rust/crates/core/src/ocr/client.rs # litellm-rust/crates/core/src/ocr/document.rs # litellm-rust/crates/core/src/ocr/error.rs # litellm-rust/crates/core/src/ocr/lifecycle.rs # litellm-rust/crates/core/src/ocr/mod.rs # litellm-rust/crates/core/src/ocr/types.rs # litellm-rust/crates/core/src/ocr/wire.rs # litellm-rust/crates/core/tests/ocr.rs # litellm-rust/crates/python-bridge/src/routes/ocr/document.rs # litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs # litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs # litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs # litellm-rust/crates/python-bridge/src/routes/ocr/project.rs # litellm/proxy/ocr_endpoints/endpoints.py # litellm/rust_bridge/_native.pyi # tests/test_litellm/ocr/test_ocr_file_input.py # tests/test_litellm_rust/ocr/test_requests.py
This commit is contained in:
commit
f8e5092a34
128 changed files with 8941 additions and 815 deletions
|
|
@ -0,0 +1,5 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN IF NOT EXISTS "total_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0;
|
||||
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN IF NOT EXISTS "total_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0;
|
||||
|
|
@ -426,6 +426,7 @@ model LiteLLM_VerificationToken {
|
|||
key_alias String?
|
||||
soft_budget_cooldown Boolean @default(false) // key-level state on if budget alerts need to be cooled down
|
||||
spend Float @default(0.0)
|
||||
total_spend Float @default(0.0)
|
||||
expires DateTime?
|
||||
models String[]
|
||||
aliases Json @default("{}")
|
||||
|
|
@ -528,6 +529,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
key_alias String?
|
||||
soft_budget_cooldown Boolean @default(false)
|
||||
spend Float @default(0.0)
|
||||
total_spend Float @default(0.0)
|
||||
expires DateTime?
|
||||
models String[]
|
||||
aliases Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
|
|||
"router_general_settings",
|
||||
"ignore_invalid_deployments",
|
||||
"fallback_access_check",
|
||||
"fallback_budget_check",
|
||||
"auto_router_capability_limit",
|
||||
}
|
||||
)
|
||||
|
|
@ -53,6 +54,7 @@ S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64
|
|||
S3_PREFIX_DIGEST_CHARS: Final = 16
|
||||
# s3 allows 2048 bytes of combined metadata headers, which Content-Disposition counts against
|
||||
MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES: Final = 1024
|
||||
S3_LOG_PROMPTS_ONLY_ENV_VAR: Final = "S3_LOG_PROMPTS_ONLY"
|
||||
MAX_FILE_LIST_LIMIT: Final = 10000
|
||||
DEFAULT_SQS_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10))
|
||||
DEFAULT_NUM_WORKERS_LITELLM_PROXY: Final = int(os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1))
|
||||
|
|
@ -1499,6 +1501,7 @@ OUTPUT_TOKEN_CEILING_PARAMS: Final = frozenset({"max_tokens", "max_completion_to
|
|||
CLIENT_OUTPUT_CEILING_METADATA_KEY: Final = "_client_output_ceiling"
|
||||
CONSUMED_REQUEST_TAGS_METADATA_KEY: Final = "_consumed_request_tags"
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY: Final = "_routing_request_tags"
|
||||
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY: Final = "_litellm_router_usage_counted_tokens"
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin"
|
||||
SESSION_ID_GENERATED_METADATA_KEY: Final = "litellm_session_id_generated"
|
||||
SESSION_ID_OMITTED_METADATA_KEY: Final = "litellm_session_id_omitted"
|
||||
|
|
|
|||
|
|
@ -446,6 +446,12 @@
|
|||
"ui_name": "S3 Path Prefix",
|
||||
"description": "Path prefix within the bucket for organizing logs",
|
||||
"required": false
|
||||
},
|
||||
"s3_log_prompts_only": {
|
||||
"type": "boolean",
|
||||
"ui_name": "Log Prompts Only",
|
||||
"description": "Log request messages to S3 but drop the model response from each logged object",
|
||||
"required": false
|
||||
}
|
||||
},
|
||||
"description": "S3 Bucket (AWS) Logging Integration"
|
||||
|
|
|
|||
|
|
@ -118,6 +118,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
alias_map: Final = {
|
||||
"langfuse_otel": "langfuse",
|
||||
"s3_v2": "s3",
|
||||
}
|
||||
lookup_name: Final = alias_map.get(normalized_name, normalized_name)
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from collections.abc import Callable, Iterable, Mapping
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -20,7 +21,10 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
|
|||
OTELSemconvCategory,
|
||||
parse_semconv_opt_in,
|
||||
)
|
||||
from litellm.integrations.otel.mappers.utils import drop_none
|
||||
from litellm.integrations.otel.model.baggage import promoted_metadata
|
||||
from litellm.integrations.otel.model.db_endpoint import db_span_attributes
|
||||
from litellm.integrations.otel.model.metadata import flatten_metadata
|
||||
from litellm.integrations.otel.model.semconv import Metric
|
||||
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -205,6 +209,20 @@ def _resolve_metric_attribute_filter(
|
|||
)
|
||||
|
||||
|
||||
def _provider_label(custom_llm_provider: object) -> str | None:
|
||||
"""The provider label for one call's metrics and events, or None when the
|
||||
call carries no provider.
|
||||
|
||||
Every attribute set drops None before export, so the label is simply absent
|
||||
in that case: the OTLP encoder rejects a None attribute value outright, and a
|
||||
placeholder would mint a permanent metric series that no operator can act
|
||||
on. Mirrors the v2 integration's ``_provider_attributes``.
|
||||
"""
|
||||
if not isinstance(custom_llm_provider, str) or not custom_llm_provider:
|
||||
return None
|
||||
return custom_llm_provider
|
||||
|
||||
|
||||
def _normalize_team_metadata_keys(value: str | Iterable[object] | None) -> list[str]:
|
||||
"""Coerce a team-metadata allowlist from a list or comma-separated string.
|
||||
|
||||
|
|
@ -288,6 +306,7 @@ class OpenTelemetryConfig:
|
|||
# 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)
|
||||
baggage_metadata_keys: list[str] = field(default_factory=list)
|
||||
# Prometheus-style include/exclude control over which attributes are stamped
|
||||
# on emitted metrics, to cap metric cardinality.
|
||||
attributes: OTELMetricAttributeFilter | None = None
|
||||
|
|
@ -314,6 +333,9 @@ class OpenTelemetryConfig:
|
|||
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"))
|
||||
self.baggage_metadata_keys = _normalize_team_metadata_keys(
|
||||
self.baggage_metadata_keys
|
||||
) or _normalize_team_metadata_keys(os.getenv("LITELLM_OTEL_BAGGAGE_METADATA_KEYS"))
|
||||
|
||||
@classmethod
|
||||
def from_env(cls):
|
||||
|
|
@ -366,11 +388,14 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
**kwargs,
|
||||
):
|
||||
team_metadata_keys_override: Final = kwargs.pop("baggage_team_metadata_keys", None)
|
||||
metadata_keys_override: Final = kwargs.pop("baggage_metadata_keys", None)
|
||||
metric_attributes_override: Final = kwargs.pop("attributes", 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)
|
||||
if metadata_keys_override is not None:
|
||||
config.baggage_metadata_keys = _normalize_team_metadata_keys(metadata_keys_override)
|
||||
if metric_attributes_override is not None:
|
||||
config.attributes = _build_metric_attribute_filter(metric_attributes_override)
|
||||
|
||||
|
|
@ -1542,6 +1567,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
if team_metadata:
|
||||
self.safe_set_attribute(span=span, key=TEAM_METADATA_ATTRIBUTE, value=team_metadata)
|
||||
|
||||
if self.config.baggage_metadata_keys:
|
||||
flat_metadata: Final = MappingProxyType(dict(flatten_metadata(metadata)))
|
||||
for key, value in promoted_metadata(flat_metadata, tuple(self.config.baggage_metadata_keys)).items():
|
||||
self.safe_set_attribute(span=span, key=key, value=value)
|
||||
|
||||
model_group: Final = standard_logging_payload.get("model_group")
|
||||
if model_group:
|
||||
self.safe_set_attribute(span=span, key=MODEL_GROUP_ATTRIBUTE, value=model_group)
|
||||
|
|
@ -1601,19 +1631,22 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
) = _resolve_metric_attribute_filter(attributes)
|
||||
self._metric_attr_filter_resolved = True
|
||||
|
||||
def _filter_metric_attributes(self, attrs: dict[str, str]) -> dict[str, str]:
|
||||
def _filter_metric_attributes(self, attrs: Mapping[str, str | None]) -> dict[str, str]:
|
||||
if not self._metric_attr_filter_resolved:
|
||||
self._ensure_metric_attribute_filter()
|
||||
return {k: v for k, v in attrs.items() if v is not None and self._metric_attribute_allowed(k)}
|
||||
|
||||
def _metric_attribute_allowed(self, key: str) -> bool:
|
||||
if self._metric_attr_include is not None:
|
||||
return {k: v for k, v in attrs.items() if k in self._metric_attr_include}
|
||||
return key in self._metric_attr_include
|
||||
if self._metric_attr_exclude is not None:
|
||||
return {k: v for k, v in attrs.items() if k not in self._metric_attr_exclude}
|
||||
return attrs
|
||||
return key not in self._metric_attr_exclude
|
||||
return True
|
||||
|
||||
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
|
||||
duration_s: Final = (end_time - start_time).total_seconds()
|
||||
params: Final = kwargs.get("litellm_params") or {}
|
||||
provider: Final = params.get("custom_llm_provider", "Unknown")
|
||||
provider: Final = _provider_label(params.get("custom_llm_provider"))
|
||||
|
||||
common_attrs = {
|
||||
"gen_ai.operation.name": (
|
||||
|
|
@ -1857,7 +1890,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
otel_logger: Final = self._logger_provider.get_logger(LITELLM_LOGGER_NAME)
|
||||
|
||||
parent_ctx: Final = span.get_span_context()
|
||||
provider: Final = (kwargs.get("litellm_params") or {}).get("custom_llm_provider", "Unknown")
|
||||
provider: Final = _provider_label((kwargs.get("litellm_params") or {}).get("custom_llm_provider"))
|
||||
|
||||
if self._gen_ai_semconv_latest_experimental:
|
||||
self._emit_inference_details_event(
|
||||
|
|
@ -1894,7 +1927,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
severity_number=SeverityNumber.INFO,
|
||||
severity_text="INFO",
|
||||
body=body,
|
||||
attributes=attrs,
|
||||
attributes=drop_none(attrs),
|
||||
)
|
||||
otel_logger.emit(log_record)
|
||||
|
||||
|
|
@ -1926,7 +1959,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
severity_number=SeverityNumber.INFO,
|
||||
severity_text="INFO",
|
||||
body=body,
|
||||
attributes=attrs,
|
||||
attributes=drop_none(attrs),
|
||||
)
|
||||
otel_logger.emit(log_record)
|
||||
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from datetime import datetime
|
|||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.integrations.otel.mappers.utils import drop_none
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -195,13 +196,16 @@ class OTELGenAISemconvMixin:
|
|||
if value:
|
||||
self.safe_set_attribute(span=span, key=semconv_key, value=value)
|
||||
|
||||
def _build_inference_details_attrs(self, kwargs: dict, response_obj: dict, provider: str) -> dict[str, str]:
|
||||
def _build_inference_details_attrs(
|
||||
self, kwargs: dict, response_obj: dict, provider: str | None
|
||||
) -> dict[str, str | None]:
|
||||
"""Build the attribute payload for the inference-details event.
|
||||
|
||||
Always includes provider/operation; input/output messages are added
|
||||
Always includes operation and provider (None when the call carries none,
|
||||
dropped before the event is emitted); input/output messages are added
|
||||
only when content capture is enabled and non-empty. Mixin-internal.
|
||||
"""
|
||||
attrs: Final[dict[str, str]] = {
|
||||
attrs: Final[dict[str, str | None]] = {
|
||||
"event_name": _INFERENCE_DETAILS_EVENT_NAME,
|
||||
"gen_ai.provider.name": provider,
|
||||
"gen_ai.operation.name": self._gen_ai_operation_name(kwargs),
|
||||
|
|
@ -221,7 +225,7 @@ class OTELGenAISemconvMixin:
|
|||
self,
|
||||
kwargs: dict,
|
||||
response_obj: dict,
|
||||
provider: str,
|
||||
provider: str | None,
|
||||
otel_logger,
|
||||
parent_ctx,
|
||||
) -> None:
|
||||
|
|
@ -239,6 +243,6 @@ class OTELGenAISemconvMixin:
|
|||
severity_number=SeverityNumber.INFO,
|
||||
severity_text="INFO",
|
||||
body=None,
|
||||
attributes=self._build_inference_details_attrs(kwargs, response_obj, provider),
|
||||
attributes=drop_none(self._build_inference_details_attrs(kwargs, response_obj, provider)),
|
||||
)
|
||||
otel_logger.emit(log_record)
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from litellm.integrations.otel.model.metadata import (
|
|||
LLMCallEvent,
|
||||
RequestIdentity,
|
||||
auth_metadata,
|
||||
metadata_from_request_data,
|
||||
model_from_request_data,
|
||||
)
|
||||
from litellm.integrations.otel.model.payloads import (
|
||||
|
|
@ -679,7 +680,12 @@ class OpenTelemetryV2(CustomLogger):
|
|||
# / errors are the FastAPI instrumentor's job, so we don't touch it here.
|
||||
# ====================================================================== #
|
||||
|
||||
def seed_request_identity(self, user_api_key_dict: object, model: str | None = None) -> None:
|
||||
def seed_request_identity(
|
||||
self,
|
||||
user_api_key_dict: object,
|
||||
model: str | None = None,
|
||||
request_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Attach request-identity Baggage to the current context + server span.
|
||||
|
||||
Seeding identity into Baggage makes **every** span emitted afterwards for
|
||||
|
|
@ -691,7 +697,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
isn't determined yet, which is correct.
|
||||
"""
|
||||
try:
|
||||
identity: Final = RequestIdentity.from_user_api_key_auth(user_api_key_dict)
|
||||
identity: Final = RequestIdentity.from_user_api_key_auth(user_api_key_dict, request_metadata)
|
||||
bag: Final = promoted_baggage(
|
||||
identity,
|
||||
model,
|
||||
|
|
@ -743,6 +749,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
self.seed_request_identity(
|
||||
user_api_key_dict,
|
||||
model=model_from_request_data(data),
|
||||
request_metadata=metadata_from_request_data(data),
|
||||
)
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -15,9 +15,10 @@ never promoted whole.
|
|||
|
||||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.integrations.otel.model.metadata import RequestIdentity
|
||||
from litellm.integrations.otel.model.metadata import REQUESTER_METADATA_PATH, RequestIdentity
|
||||
from litellm.integrations.otel.model.semconv import GenAI, LiteLLM
|
||||
|
||||
# Attribute key -> value extractor over (identity, request_model,
|
||||
|
|
@ -79,17 +80,23 @@ def promoted_baggage(
|
|||
``team_metadata_keys`` selects sub-keys of the team's metadata to promote
|
||||
under ``litellm.team.metadata``. Empty values are dropped.
|
||||
"""
|
||||
out: Final[dict[str, str]] = {}
|
||||
for key, extract in _PROMOTABLE.items():
|
||||
if key in promoted_keys:
|
||||
value = extract(identity, request_model, team_metadata_keys)
|
||||
if value:
|
||||
out[key] = value
|
||||
for meta_key in metadata_keys:
|
||||
value = identity.metadata.get(meta_key)
|
||||
if value:
|
||||
out[f"{LiteLLM.METADATA_PREFIX}{meta_key}"] = value
|
||||
return out
|
||||
identity_values: Final = {
|
||||
key: value
|
||||
for key, extract in _PROMOTABLE.items()
|
||||
if key in promoted_keys and (value := extract(identity, request_model, team_metadata_keys))
|
||||
}
|
||||
return {**identity_values, **promoted_metadata(identity.metadata, metadata_keys)}
|
||||
|
||||
|
||||
def promoted_metadata(metadata: Mapping[str, str], metadata_keys: tuple[str, ...]) -> Mapping[str, str]:
|
||||
"""Allowlisted entries of a flattened metadata mapping under ``litellm.metadata.*``."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
f"{LiteLLM.METADATA_PREFIX}{meta_key.removeprefix(REQUESTER_METADATA_PATH)}": value
|
||||
for meta_key in metadata_keys
|
||||
if (value := metadata.get(meta_key))
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _filtered_team_metadata_json(
|
||||
|
|
|
|||
|
|
@ -210,7 +210,10 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
validation_alias=AliasChoices("baggage_metadata_keys", "LITELLM_OTEL_BAGGAGE_METADATA_KEYS"),
|
||||
description=(
|
||||
"Metadata sub-keys promoted under the ``litellm.metadata.*`` "
|
||||
"namespace. Configure via the ``LITELLM_OTEL_BAGGAGE_METADATA_KEYS`` "
|
||||
"namespace. A dotted path such as ``requester_metadata.trace_id`` "
|
||||
"reads the caller's nested ``metadata.trace_id`` and is promoted as "
|
||||
"``litellm.metadata.trace_id``; other dotted keys keep their full path. "
|
||||
"Configure via the ``LITELLM_OTEL_BAGGAGE_METADATA_KEYS`` "
|
||||
"env var (comma-separated) or "
|
||||
"``callback_settings.otel.baggage_metadata_keys`` in config.yaml."
|
||||
),
|
||||
|
|
|
|||
|
|
@ -49,6 +49,8 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
LANGFUSE_TRACE_NAME_HEADER: Final = "langfuse_trace_name"
|
||||
REQUESTER_METADATA_KEY: Final = "requester_metadata"
|
||||
REQUESTER_METADATA_PATH: Final = f"{REQUESTER_METADATA_KEY}."
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -78,7 +80,7 @@ class RequestIdentity:
|
|||
model, not just the user-facing one.
|
||||
"""
|
||||
raw_meta: Final = cast(Mapping[str, object], payload.get("metadata") or {})
|
||||
metadata = {key: str(value) for key, value in raw_meta.items() if isinstance(value, (str, bool, int, float))}
|
||||
metadata: Final = MappingProxyType(dict(flatten_metadata(raw_meta)))
|
||||
return cls(
|
||||
call_id=as_str(payload.get("litellm_call_id")) or as_str(payload.get("id")),
|
||||
# StandardLoggingMetadata's canonical key is ``user_api_key_team_id``;
|
||||
|
|
@ -95,7 +97,9 @@ class RequestIdentity:
|
|||
)
|
||||
|
||||
@classmethod
|
||||
def from_user_api_key_auth(cls, auth: object) -> RequestIdentity:
|
||||
def from_user_api_key_auth(
|
||||
cls, auth: object, request_metadata: Mapping[str, object] | None = None
|
||||
) -> RequestIdentity:
|
||||
"""Identity from a ``UserAPIKeyAuth`` (duck-typed to keep this module
|
||||
free of a proxy import).
|
||||
|
||||
|
|
@ -103,11 +107,13 @@ class RequestIdentity:
|
|||
guardrail, or service span is created — so the whole request's spans
|
||||
inherit identity, not just the LLM-call span. Metadata sub-keys use the
|
||||
``user_api_key_*`` names that ``baggage.DEFAULT_BAGGAGE_METADATA_KEYS``
|
||||
promotes.
|
||||
promotes; ``request_metadata`` (the caller's ``requester_metadata``
|
||||
snapshot) is flattened to dotted keys so ``requester_metadata.<key>``
|
||||
resolves too.
|
||||
"""
|
||||
get: Final = lambda name: getattr(auth, name, None) # noqa: E731
|
||||
metadata: Final = {
|
||||
meta_key: str(value)
|
||||
auth_meta: Final = tuple(
|
||||
(meta_key, str(value))
|
||||
for meta_key, attr in (
|
||||
("user_api_key_user_id", "user_id"),
|
||||
("user_api_key_org_id", "org_id"),
|
||||
|
|
@ -115,7 +121,9 @@ class RequestIdentity:
|
|||
("user_api_key_end_user_id", "end_user_id"),
|
||||
)
|
||||
if (value := get(attr))
|
||||
}
|
||||
)
|
||||
request_meta: Final = flatten_metadata(request_metadata) if request_metadata is not None else ()
|
||||
metadata: Final = MappingProxyType(dict((*request_meta, *auth_meta)))
|
||||
return cls(
|
||||
team_id=as_str(get("team_id")),
|
||||
team_alias=as_str(get("team_alias")),
|
||||
|
|
@ -351,6 +359,35 @@ def model_from_request_data(data: object) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def metadata_from_request_data(data: object) -> Mapping[str, object] | None:
|
||||
"""The caller's ``requester_metadata`` snapshot from a pre-call ``data`` dict, keyed under its wrapper.
|
||||
|
||||
The proxy stores it under ``metadata`` or ``litellm_metadata`` depending on the route;
|
||||
the proxy-owned siblings (``user_api_key_*``, ``requester_ip_address``) are not read.
|
||||
"""
|
||||
top: Final = _as_str_mapping(data)
|
||||
if top is None:
|
||||
return None
|
||||
snapshots: Final = tuple(
|
||||
snapshot
|
||||
for name in ("metadata", "litellm_metadata")
|
||||
if (nested := _as_str_mapping(top.get(name))) is not None
|
||||
and (snapshot := _as_str_mapping(nested.get(REQUESTER_METADATA_KEY))) is not None
|
||||
)
|
||||
return MappingProxyType({REQUESTER_METADATA_KEY: snapshots[0]}) if snapshots else None
|
||||
|
||||
|
||||
def flatten_metadata(raw: Mapping[str, object]) -> Iterator[tuple[str, str]]:
|
||||
"""Scalar leaves of a nested metadata mapping, keyed by their dotted path."""
|
||||
stack: Final = list(tuple(raw.items())[::-1]) # mutable-ok: iterative worklist keeps the walk off the call stack
|
||||
while stack:
|
||||
key, value = stack.pop()
|
||||
if (nested := _as_str_mapping(value)) is not None:
|
||||
stack.extend(tuple((f"{key}.{sub_key}", sub_value) for sub_key, sub_value in nested.items())[::-1])
|
||||
elif isinstance(value, (str, bool, int, float)):
|
||||
yield key, str(value)
|
||||
|
||||
|
||||
def resolve_provider_model(payload: StandardLoggingPayload) -> str | None:
|
||||
"""The model litellm dispatched to the provider, from the payload.
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from datetime import datetime, timedelta
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
|
|
@ -36,6 +37,7 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
from litellm.litellm_core_utils.service_tier_utils import (
|
||||
get_service_tier_from_standard_logging_payload,
|
||||
)
|
||||
from litellm.models.end_user import LiteLLM_EndUserTable
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_DeletedVerificationToken,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -43,7 +45,9 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.repositories.base_repository import BaseRepository
|
||||
from litellm.repositories.budget_repository import BudgetRepository
|
||||
from litellm.repositories.organization_repository import OrganizationRepository
|
||||
from litellm.repositories.table_repositories import EndUserRepository
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
|
@ -66,13 +70,26 @@ from litellm.types.utils import (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from prisma.types import (
|
||||
LiteLLM_BudgetTableWhereUniqueInput,
|
||||
LiteLLM_EndUserTableInclude,
|
||||
LiteLLM_EndUserTableOrderByInput,
|
||||
)
|
||||
from prometheus_client import Gauge
|
||||
from prometheus_client.metrics import MetricWrapperBase
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.router import Router
|
||||
else:
|
||||
AsyncIOScheduler = Any
|
||||
|
||||
_IsNotNull = TypedDict("_IsNotNull", {"not": ReadOnly[None]})
|
||||
|
||||
|
||||
class _BudgetedCustomerFilter(TypedDict):
|
||||
budget_id: ReadOnly[_IsNotNull]
|
||||
|
||||
|
||||
_BudgetRowT: Final = TypeVar("_BudgetRowT")
|
||||
_TableRowT: Final = TypeVar("_TableRowT", bound=BaseModel)
|
||||
|
||||
|
|
@ -116,8 +133,8 @@ def _paginated_table(repository: BaseRepository[_TableRowT]) -> _PaginatedPrisma
|
|||
)
|
||||
|
||||
|
||||
class _OrgBudgetRow(Protocol):
|
||||
"""The budget columns joined onto an organization row."""
|
||||
class _JoinedBudgetRow(Protocol):
|
||||
"""The budget columns joined onto an organization or customer row."""
|
||||
|
||||
@property
|
||||
def max_budget(self) -> float | None: ...
|
||||
|
|
@ -126,6 +143,23 @@ class _OrgBudgetRow(Protocol):
|
|||
def budget_reset_at(self) -> datetime | None: ...
|
||||
|
||||
|
||||
class _CustomerBudgetRow(Protocol):
|
||||
"""The columns of a customer (end user) row that budget gauges read."""
|
||||
|
||||
@property
|
||||
def user_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def spend(self) -> float: ...
|
||||
|
||||
@property
|
||||
def litellm_budget_table(self) -> _JoinedBudgetRow | None: ...
|
||||
|
||||
|
||||
def _customer_budget_metrics_enabled() -> bool:
|
||||
return litellm.enable_end_user_cost_tracking_prometheus_only is True and not litellm.disable_end_user_cost_tracking
|
||||
|
||||
|
||||
class _ExcludedLabelMetric:
|
||||
"""Proxies a prometheus metric whose declared ``labelnames`` had globally
|
||||
excluded labels removed, dropping those labels from every ``labels(...)``
|
||||
|
|
@ -471,6 +505,24 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=self.get_labels_for_metric("litellm_user_budget_remaining_hours_metric"),
|
||||
)
|
||||
|
||||
self.litellm_remaining_customer_budget_metric = self._gauge_factory(
|
||||
"litellm_remaining_customer_budget_metric",
|
||||
"Remaining budget for customer (end user)",
|
||||
labelnames=self.get_labels_for_metric("litellm_remaining_customer_budget_metric"),
|
||||
)
|
||||
|
||||
self.litellm_customer_max_budget_metric = self._gauge_factory(
|
||||
"litellm_customer_max_budget_metric",
|
||||
"Maximum budget set for customer (end user)",
|
||||
labelnames=self.get_labels_for_metric("litellm_customer_max_budget_metric"),
|
||||
)
|
||||
|
||||
self.litellm_customer_budget_remaining_hours_metric = self._gauge_factory(
|
||||
"litellm_customer_budget_remaining_hours_metric",
|
||||
"Remaining hours for customer (end user) budget to be reset",
|
||||
labelnames=self.get_labels_for_metric("litellm_customer_budget_remaining_hours_metric"),
|
||||
)
|
||||
|
||||
########################################
|
||||
# LiteLLM Virtual API KEY metrics
|
||||
########################################
|
||||
|
|
@ -1334,7 +1386,7 @@ class PrometheusLogger(CustomLogger):
|
|||
self,
|
||||
metric: Any,
|
||||
metric_name: DEFINED_PROMETHEUS_METRICS,
|
||||
labels: dict[str, str | None],
|
||||
labels: Mapping[str, str | None],
|
||||
) -> None:
|
||||
"""
|
||||
Cap the cardinality of metrics that include the ``end_user`` label.
|
||||
|
|
@ -1501,6 +1553,7 @@ class PrometheusLogger(CustomLogger):
|
|||
response_cost=response_cost,
|
||||
user_id=user_id,
|
||||
user_api_key_org_id=user_api_key_org_id,
|
||||
end_user_id=end_user_id,
|
||||
)
|
||||
|
||||
# set proxy virtual key rpm/tpm metrics
|
||||
|
|
@ -1930,12 +1983,14 @@ class PrometheusLogger(CustomLogger):
|
|||
response_cost: float,
|
||||
user_id: str | None = None,
|
||||
user_api_key_org_id: str | None = None,
|
||||
end_user_id: str | None = None,
|
||||
):
|
||||
if (
|
||||
isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric)
|
||||
and isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric)
|
||||
and isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric)
|
||||
and isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric)
|
||||
and self._customer_budget_gauges_are_noop()
|
||||
):
|
||||
return
|
||||
|
||||
|
|
@ -1990,6 +2045,10 @@ class PrometheusLogger(CustomLogger):
|
|||
carried=OrgBudgetSnapshot.from_metadata(_metadata),
|
||||
org_alias=_org_alias if isinstance(_org_alias, str) else None,
|
||||
),
|
||||
self._set_customer_budget_metrics_after_api_request(
|
||||
end_user_id=end_user_id,
|
||||
response_cost=response_cost,
|
||||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
try:
|
||||
|
|
@ -2006,7 +2065,7 @@ class PrometheusLogger(CustomLogger):
|
|||
if isinstance(r, Exception):
|
||||
verbose_logger.debug(
|
||||
"[Non-Blocking] Prometheus: Budget metric lookup %s failed: %s",
|
||||
["key", "team", "user", "org"][i],
|
||||
("key", "team", "user", "org", "customer")[i],
|
||||
r,
|
||||
)
|
||||
|
||||
|
|
@ -3574,9 +3633,9 @@ class PrometheusLogger(CustomLogger):
|
|||
|
||||
async def _initialize_budget_metrics(
|
||||
self,
|
||||
data_fetch_function: Callable[..., Awaitable[tuple[list[_BudgetRowT], int | None]]],
|
||||
set_metrics_function: Callable[[list[_BudgetRowT]], Awaitable[None]],
|
||||
data_type: Literal["teams", "keys", "users", "orgs"],
|
||||
data_fetch_function: Callable[..., Awaitable[tuple[Sequence[_BudgetRowT], int | None]]],
|
||||
set_metrics_function: Callable[[Sequence[_BudgetRowT]], Awaitable[None]],
|
||||
data_type: Literal["teams", "keys", "users", "orgs", "customers"],
|
||||
):
|
||||
"""
|
||||
Generic method to initialize budget metrics for teams or API keys.
|
||||
|
|
@ -3735,6 +3794,49 @@ class PrometheusLogger(CustomLogger):
|
|||
data_type="orgs",
|
||||
)
|
||||
|
||||
async def _initialize_customer_budget_metrics(self):
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_logger.debug("Prometheus: skipping customer metrics initialization, DB not initialized")
|
||||
return
|
||||
|
||||
if self._customer_budget_gauges_are_noop():
|
||||
return
|
||||
|
||||
if not _customer_budget_metrics_enabled():
|
||||
verbose_logger.debug("Prometheus: skipping customer metrics initialization, end_user tracking disabled")
|
||||
return
|
||||
|
||||
default_budget: Final = await self._get_default_customer_budget(prisma_client)
|
||||
customers_table: Final = EndUserRepository(prisma_client).table
|
||||
with_persisted_budget: Final[_BudgetedCustomerFilter] = {"budget_id": {"not": None}}
|
||||
budgeted_customers: Final = None if default_budget is not None else with_persisted_budget
|
||||
by_user_id: Final[LiteLLM_EndUserTableOrderByInput] = {"user_id": "asc"}
|
||||
with_budget: Final[LiteLLM_EndUserTableInclude] = {"litellm_budget_table": True}
|
||||
|
||||
async def fetch_customers(page_size: int, page: int) -> tuple[Sequence[_CustomerBudgetRow], int | None]:
|
||||
skip: Final = (page - 1) * page_size
|
||||
customers: Final = await customers_table.find_many(
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
where=budgeted_customers,
|
||||
order=by_user_id,
|
||||
include=with_budget,
|
||||
)
|
||||
total_count: Final = await customers_table.count(where=budgeted_customers) if page == 1 else None
|
||||
return customers, total_count
|
||||
|
||||
async def set_customer_metrics(customers: Sequence[_CustomerBudgetRow]) -> None:
|
||||
for customer in customers:
|
||||
self._set_customer_budget_metrics_from_row(customer, default_budget=default_budget)
|
||||
|
||||
await self._initialize_budget_metrics(
|
||||
data_fetch_function=fetch_customers,
|
||||
set_metrics_function=set_customer_metrics,
|
||||
data_type="customers",
|
||||
)
|
||||
|
||||
async def initialize_remaining_budget_metrics(self):
|
||||
"""
|
||||
Handler for initializing remaining budget metrics for all teams to avoid metric discrepancies.
|
||||
|
|
@ -3765,11 +3867,12 @@ class PrometheusLogger(CustomLogger):
|
|||
"""
|
||||
Helper to initialize remaining budget metrics for all teams, API keys, and users.
|
||||
"""
|
||||
verbose_logger.debug("Emitting key, team, user, org budget metrics....")
|
||||
verbose_logger.debug("Emitting key, team, user, org, customer budget metrics....")
|
||||
await self._initialize_team_budget_metrics()
|
||||
await self._initialize_api_key_budget_metrics()
|
||||
await self._initialize_user_budget_metrics()
|
||||
await self._initialize_org_budget_metrics()
|
||||
await self._initialize_customer_budget_metrics()
|
||||
await self._initialize_user_and_team_count_metrics()
|
||||
|
||||
async def _initialize_user_and_team_count_metrics(self):
|
||||
|
|
@ -3805,27 +3908,27 @@ class PrometheusLogger(CustomLogger):
|
|||
verbose_logger.exception("Error initializing user/team count metrics: %s", e)
|
||||
|
||||
async def _set_key_list_budget_metrics(
|
||||
self, keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]
|
||||
self, keys: Sequence[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]
|
||||
) -> None:
|
||||
"""Helper function to set budget metrics for a list of keys"""
|
||||
for key in keys:
|
||||
if isinstance(key, UserAPIKeyAuth):
|
||||
self._set_key_budget_metrics(key)
|
||||
|
||||
async def _set_team_list_budget_metrics(self, teams: list[LiteLLM_TeamTable]):
|
||||
async def _set_team_list_budget_metrics(self, teams: Sequence[LiteLLM_TeamTable]):
|
||||
"""Helper function to set budget metrics for a list of teams"""
|
||||
for team in teams:
|
||||
self._set_team_budget_metrics(team)
|
||||
|
||||
async def _set_user_list_budget_metrics(self, users: list[LiteLLM_UserTable]):
|
||||
async def _set_user_list_budget_metrics(self, users: Sequence[LiteLLM_UserTable]):
|
||||
"""Helper function to set budget metrics for a list of users"""
|
||||
for user in users:
|
||||
self._set_user_budget_metrics(user)
|
||||
|
||||
async def _set_org_list_budget_metrics(self, orgs: list):
|
||||
async def _set_org_list_budget_metrics(self, orgs: Sequence):
|
||||
"""Helper function to set budget metrics for a list of orgs"""
|
||||
for org in orgs:
|
||||
budget_table: _OrgBudgetRow | None = getattr(org, "litellm_budget_table", None)
|
||||
budget_table: _JoinedBudgetRow | None = getattr(org, "litellm_budget_table", None)
|
||||
self._set_org_budget_metrics(
|
||||
org_id=org.organization_id or "",
|
||||
org_alias=org.organization_alias or "",
|
||||
|
|
@ -3834,6 +3937,19 @@ class PrometheusLogger(CustomLogger):
|
|||
budget_reset_at=(getattr(budget_table, "budget_reset_at", None) if budget_table else None),
|
||||
)
|
||||
|
||||
def _set_customer_budget_metrics_from_row(
|
||||
self, customer: _CustomerBudgetRow, default_budget: _JoinedBudgetRow | None
|
||||
):
|
||||
budget_table: Final = (
|
||||
customer.litellm_budget_table if customer.litellm_budget_table is not None else default_budget
|
||||
)
|
||||
self._set_customer_budget_metrics(
|
||||
end_user_id=customer.user_id,
|
||||
spend=customer.spend,
|
||||
max_budget=budget_table.max_budget if budget_table is not None else None,
|
||||
budget_reset_at=budget_table.budget_reset_at if budget_table is not None else None,
|
||||
)
|
||||
|
||||
async def _set_team_budget_metrics_after_api_request(
|
||||
self,
|
||||
user_api_team: str | None,
|
||||
|
|
@ -4083,6 +4199,98 @@ class PrometheusLogger(CustomLogger):
|
|||
self._get_remaining_hours_for_budget_reset(budget_reset_at=budget_reset_at)
|
||||
)
|
||||
|
||||
async def _set_customer_budget_metrics_after_api_request(
|
||||
self,
|
||||
end_user_id: str | None,
|
||||
response_cost: float,
|
||||
):
|
||||
if self._customer_budget_gauges_are_noop() or not _customer_budget_metrics_enabled():
|
||||
return
|
||||
|
||||
if not end_user_id:
|
||||
return
|
||||
|
||||
from litellm.proxy.common_utils.user_api_key_cache import end_user_cache_key
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
try:
|
||||
cached_customer: Final = await user_api_key_cache.async_get_cache(
|
||||
key=end_user_cache_key(end_user_id),
|
||||
model_type=LiteLLM_EndUserTable,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("[Non-Blocking] Prometheus: Error getting customer info: %s", e)
|
||||
return
|
||||
|
||||
if cached_customer is None:
|
||||
return
|
||||
|
||||
budget_table: Final = cached_customer.litellm_budget_table
|
||||
self._set_customer_budget_metrics(
|
||||
end_user_id=end_user_id,
|
||||
spend=cached_customer.spend + response_cost,
|
||||
max_budget=budget_table.max_budget if budget_table is not None else None,
|
||||
budget_reset_at=None,
|
||||
)
|
||||
|
||||
async def _get_default_customer_budget(self, prisma_client: PrismaClient) -> _JoinedBudgetRow | None:
|
||||
default_budget_id: Final = litellm.max_end_user_budget_id
|
||||
if default_budget_id is None:
|
||||
return None
|
||||
default_budget_key: Final[LiteLLM_BudgetTableWhereUniqueInput] = {"budget_id": default_budget_id}
|
||||
try:
|
||||
return await BudgetRepository(prisma_client).table.find_unique(where=default_budget_key)
|
||||
except Exception as e:
|
||||
verbose_logger.debug("[Non-Blocking] Prometheus: Error getting default customer budget: %s", e)
|
||||
return None
|
||||
|
||||
def _customer_budget_gauges_are_noop(self) -> bool:
|
||||
return (
|
||||
isinstance(self.litellm_remaining_customer_budget_metric, NoOpMetric)
|
||||
and isinstance(self.litellm_customer_max_budget_metric, NoOpMetric)
|
||||
and isinstance(self.litellm_customer_budget_remaining_hours_metric, NoOpMetric)
|
||||
)
|
||||
|
||||
def _set_customer_budget_metrics(
|
||||
self,
|
||||
end_user_id: str,
|
||||
spend: float,
|
||||
max_budget: float | None,
|
||||
budget_reset_at: datetime | None,
|
||||
):
|
||||
_labels: Final[dict[str, str | None]] = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_remaining_customer_budget_metric"),
|
||||
enum_values=UserAPIKeyLabelValues(end_user=end_user_id),
|
||||
)
|
||||
if _labels.get(UserAPIKeyLabelNames.END_USER.value) is None:
|
||||
return
|
||||
|
||||
self.litellm_remaining_customer_budget_metric.labels(**_labels).set(
|
||||
self._safe_get_remaining_budget(
|
||||
max_budget=max_budget,
|
||||
spend=spend,
|
||||
)
|
||||
)
|
||||
self._track_end_user_metric_series(
|
||||
self.litellm_remaining_customer_budget_metric, "litellm_remaining_customer_budget_metric", _labels
|
||||
)
|
||||
|
||||
if max_budget is not None:
|
||||
self.litellm_customer_max_budget_metric.labels(**_labels).set(max_budget)
|
||||
self._track_end_user_metric_series(
|
||||
self.litellm_customer_max_budget_metric, "litellm_customer_max_budget_metric", _labels
|
||||
)
|
||||
|
||||
if budget_reset_at is not None:
|
||||
self.litellm_customer_budget_remaining_hours_metric.labels(**_labels).set(
|
||||
self._get_remaining_hours_for_budget_reset(budget_reset_at=budget_reset_at)
|
||||
)
|
||||
self._track_end_user_metric_series(
|
||||
self.litellm_customer_budget_remaining_hours_metric,
|
||||
"litellm_customer_budget_remaining_hours_metric",
|
||||
_labels,
|
||||
)
|
||||
|
||||
def _set_key_budget_metrics(self, user_api_key_dict: UserAPIKeyAuth):
|
||||
"""
|
||||
Set virtual key budget metrics
|
||||
|
|
|
|||
|
|
@ -2,19 +2,42 @@
|
|||
# On success + failure, log events to Supabase
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final, cast
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import (
|
||||
MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES,
|
||||
MAX_S3_OBJECT_KEY_BYTES,
|
||||
S3_BOUNDED_OBJECT_KEY_HEAD_BYTES,
|
||||
S3_LOG_PROMPTS_ONLY_ENV_VAR,
|
||||
S3_PREFIX_DIGEST_CHARS,
|
||||
)
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
_S3_LOG_PROMPTS_ONLY: Final = TypeAdapter(bool)
|
||||
|
||||
|
||||
def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | None = None) -> bool:
|
||||
env: Final = os.environ if environ is None else environ
|
||||
raw: Final = env.get(S3_LOG_PROMPTS_ONLY_ENV_VAR) if configured is None else configured
|
||||
if raw is None or raw == "":
|
||||
return False
|
||||
try:
|
||||
return _S3_LOG_PROMPTS_ONLY.validate_python(raw.strip() if isinstance(raw, str) else raw)
|
||||
except ValidationError:
|
||||
verbose_logger.warning("s3 logging: s3_log_prompts_only=%r is not a boolean, logging prompts only", raw)
|
||||
return True
|
||||
|
||||
|
||||
def prompts_only_payload(payload: StandardLoggingPayload) -> StandardLoggingPayload:
|
||||
return {**payload, "response": None}
|
||||
|
||||
|
||||
class S3Logger:
|
||||
# Class variables or attributes
|
||||
|
|
@ -33,6 +56,7 @@ class S3Logger:
|
|||
s3_config=None,
|
||||
s3_server_side_encryption: str | None = None,
|
||||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
import boto3
|
||||
|
|
@ -41,29 +65,30 @@ class S3Logger:
|
|||
verbose_logger.debug("in init s3 logger - s3_callback_params %s", litellm.s3_callback_params)
|
||||
|
||||
s3_use_team_prefix = False
|
||||
params: Final = {
|
||||
key: litellm.get_secret(value) if isinstance(value, str) and value.startswith("os.environ/") else value
|
||||
for key, value in (litellm.s3_callback_params or {}).items()
|
||||
}
|
||||
|
||||
if litellm.s3_callback_params is not None:
|
||||
# read in .env variables - example os.environ/AWS_BUCKET_NAME
|
||||
for key, value in litellm.s3_callback_params.items():
|
||||
if isinstance(value, str) and value.startswith("os.environ/"):
|
||||
litellm.s3_callback_params[key] = litellm.get_secret(value)
|
||||
# now set s3 params from litellm.s3_logger_params
|
||||
s3_bucket_name = litellm.s3_callback_params.get("s3_bucket_name")
|
||||
s3_region_name = litellm.s3_callback_params.get("s3_region_name")
|
||||
s3_api_version = litellm.s3_callback_params.get("s3_api_version")
|
||||
s3_use_ssl = litellm.s3_callback_params.get("s3_use_ssl", True)
|
||||
s3_verify = litellm.s3_callback_params.get("s3_verify")
|
||||
s3_endpoint_url = litellm.s3_callback_params.get("s3_endpoint_url")
|
||||
s3_aws_access_key_id = litellm.s3_callback_params.get("s3_aws_access_key_id")
|
||||
s3_aws_secret_access_key = litellm.s3_callback_params.get("s3_aws_secret_access_key")
|
||||
s3_aws_session_token = litellm.s3_callback_params.get("s3_aws_session_token")
|
||||
s3_config = litellm.s3_callback_params.get("s3_config")
|
||||
s3_path = litellm.s3_callback_params.get("s3_path")
|
||||
s3_server_side_encryption = litellm.s3_callback_params.get("s3_server_side_encryption")
|
||||
s3_sse_kms_key_id = litellm.s3_callback_params.get("s3_sse_kms_key_id")
|
||||
# done reading litellm.s3_callback_params
|
||||
s3_use_team_prefix = bool(litellm.s3_callback_params.get("s3_use_team_prefix", False))
|
||||
s3_bucket_name = params.get("s3_bucket_name")
|
||||
s3_region_name = params.get("s3_region_name")
|
||||
s3_api_version = params.get("s3_api_version")
|
||||
s3_use_ssl = params.get("s3_use_ssl", True)
|
||||
s3_verify = params.get("s3_verify")
|
||||
s3_endpoint_url = params.get("s3_endpoint_url")
|
||||
s3_aws_access_key_id = params.get("s3_aws_access_key_id")
|
||||
s3_aws_secret_access_key = params.get("s3_aws_secret_access_key")
|
||||
s3_aws_session_token = params.get("s3_aws_session_token")
|
||||
s3_config = params.get("s3_config")
|
||||
s3_path = params.get("s3_path")
|
||||
s3_server_side_encryption = params.get("s3_server_side_encryption")
|
||||
s3_sse_kms_key_id = params.get("s3_sse_kms_key_id")
|
||||
s3_use_team_prefix = bool(params.get("s3_use_team_prefix", False))
|
||||
self.s3_use_team_prefix = s3_use_team_prefix
|
||||
self.s3_log_prompts_only: object = (
|
||||
params.get("s3_log_prompts_only") if s3_log_prompts_only is None else s3_log_prompts_only
|
||||
)
|
||||
self.bucket_name = s3_bucket_name
|
||||
self.s3_path = s3_path
|
||||
self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params(
|
||||
|
|
@ -144,7 +169,9 @@ class S3Logger:
|
|||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
payload_str: Final = safe_dumps(payload)
|
||||
payload_str: Final = safe_dumps(
|
||||
prompts_only_payload(payload) if resolve_s3_log_prompts_only(self.s3_log_prompts_only) else payload
|
||||
)
|
||||
|
||||
print_verbose(f"\ns3 Logger - Logging payload = {payload_str}")
|
||||
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_S
|
|||
from litellm.integrations.s3 import (
|
||||
get_s3_object_download_filename,
|
||||
get_s3_object_key,
|
||||
prompts_only_payload,
|
||||
resolve_s3_log_prompts_only,
|
||||
resolve_sse_params,
|
||||
)
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
|
|
@ -68,6 +70,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_use_virtual_hosted_style: bool = False,
|
||||
s3_server_side_encryption: str | None = None,
|
||||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
s3_callback_params_override: dict | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -108,6 +111,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_use_virtual_hosted_style=s3_use_virtual_hosted_style,
|
||||
s3_server_side_encryption=s3_server_side_encryption,
|
||||
s3_sse_kms_key_id=s3_sse_kms_key_id,
|
||||
s3_log_prompts_only=s3_log_prompts_only,
|
||||
)
|
||||
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
|
||||
|
||||
|
|
@ -163,6 +167,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_use_virtual_hosted_style: bool = False,
|
||||
s3_server_side_encryption: str | None = None,
|
||||
s3_sse_kms_key_id: str | None = None,
|
||||
s3_log_prompts_only: bool | None = None,
|
||||
params_source: dict | None = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -212,6 +217,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style
|
||||
)
|
||||
|
||||
self.s3_log_prompts_only: object = (
|
||||
params.get("s3_log_prompts_only") if s3_log_prompts_only is None else s3_log_prompts_only
|
||||
)
|
||||
|
||||
self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params(
|
||||
params.get("s3_server_side_encryption") or s3_server_side_encryption,
|
||||
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
|
||||
|
|
@ -489,8 +498,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
s3_object_download_filename: Final = get_s3_object_download_filename(start_time, standard_logging_payload["id"])
|
||||
|
||||
payload: Final = (
|
||||
prompts_only_payload(standard_logging_payload)
|
||||
if resolve_s3_log_prompts_only(self.s3_log_prompts_only)
|
||||
else standard_logging_payload
|
||||
)
|
||||
return s3BatchLoggingElement(
|
||||
payload=dict(standard_logging_payload),
|
||||
payload=dict(payload),
|
||||
s3_object_key=s3_object_key,
|
||||
s3_object_download_filename=s3_object_download_filename,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1561,10 +1561,12 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
|
||||
stream_ended: Final = self._check_streaming_has_ended(responses_so_far)
|
||||
tool_use_fingerprints: Final = self._streamed_tool_use_fingerprints(responses_so_far)
|
||||
return StreamingScanKey(
|
||||
texts=(self.get_streaming_string_so_far(responses_so_far),),
|
||||
tool_calls=self._streamed_tool_use_fingerprints(responses_so_far) if stream_ended else (),
|
||||
tool_calls=tool_use_fingerprints if stream_ended else (),
|
||||
stream_ended=stream_ended,
|
||||
tool_calls_in_flight=bool(tool_use_fingerprints) and not stream_ended,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -632,6 +632,7 @@ class ModelResponseIterator:
|
|||
self.tool_name_reverse_map: dict[str, str] = tool_name_reverse_map or {}
|
||||
# Generate response ID once per stream to match OpenAI-compatible behavior
|
||||
self.response_id = _generate_id()
|
||||
self.served_model: str | None = None
|
||||
|
||||
# Track if we're currently streaming a response_format tool
|
||||
self.is_response_format_tool: bool = False
|
||||
|
|
@ -1067,6 +1068,9 @@ class ModelResponseIterator:
|
|||
}
|
||||
"""
|
||||
message_start_block: Final = MessageStartBlock(**chunk)
|
||||
start_message: Final = message_start_block["message"]
|
||||
if "model" in start_message:
|
||||
self.served_model = start_message["model"]
|
||||
if "usage" in message_start_block["message"]:
|
||||
usage = self._handle_usage(anthropic_usage_chunk=message_start_block["message"]["usage"])
|
||||
elif type_chunk == "error":
|
||||
|
|
@ -1098,6 +1102,7 @@ class ModelResponseIterator:
|
|||
],
|
||||
usage=usage,
|
||||
id=self.response_id,
|
||||
model=self.served_model,
|
||||
)
|
||||
|
||||
return returned_chunk
|
||||
|
|
|
|||
|
|
@ -78,6 +78,7 @@ _PROPAGATED_METADATA_KEYS: Final = (
|
|||
"user_api_key_end_user_id",
|
||||
"user_api_end_user_max_budget",
|
||||
"user_api_key_model_max_budget",
|
||||
"user_api_key_team_model_max_budget",
|
||||
"user_api_key_user_model_max_budget",
|
||||
"user_api_key_end_user_model_max_budget",
|
||||
"litellm_call_id",
|
||||
|
|
@ -395,9 +396,9 @@ async def _check_summary_model_budget(
|
|||
``user_api_key_auth`` runs for the client-requested model. Returns True outside the proxy or when no
|
||||
per-model budget is configured.
|
||||
|
||||
All three scopes are checked because the summary's spend is charged to all
|
||||
three: this file propagates the key, user and end-user budgets into the
|
||||
subrequest's metadata, so enforcing only two of them would let compaction
|
||||
Every scope is checked because the summary's spend is charged to every
|
||||
scope: this file propagates the key, team, user and end-user budgets into the
|
||||
subrequest's metadata, so skipping one of them would let compaction
|
||||
increment a counter it can never be refused by.
|
||||
"""
|
||||
if user_api_key_auth is None:
|
||||
|
|
@ -444,6 +445,26 @@ async def _check_summary_model_budget(
|
|||
)
|
||||
return False
|
||||
|
||||
team_model_max_budget: Final = user_api_key_auth.team_model_max_budget
|
||||
team_id: Final = user_api_key_auth.team_id
|
||||
if isinstance(team_model_max_budget, dict) and team_model_max_budget and team_id is not None:
|
||||
try:
|
||||
await model_max_budget_limiter.is_team_within_model_budget(
|
||||
team_id=team_id,
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=model_max_budget if isinstance(model_max_budget, dict) else None,
|
||||
model=summary_model,
|
||||
)
|
||||
except litellm.BudgetExceededError:
|
||||
return False
|
||||
except Exception as e: # noqa: BLE001 # a budget gate denies on any failure, as the other scopes do
|
||||
verbose_logger.warning(
|
||||
"compact_20260112: unexpected error during team model-budget check for summary_model=%s; denying: %s",
|
||||
summary_model,
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
||||
end_user_model_max_budget: Final[dict[str, object] | None] = getattr(
|
||||
user_api_key_auth, "end_user_model_max_budget", None
|
||||
)
|
||||
|
|
|
|||
|
|
@ -40,11 +40,15 @@ class StreamingScanKey:
|
|||
"""What a streaming guardrail round would hand to ``apply_guardrail``. Two keys
|
||||
compare equal when the round would scan the same content again; ``stream_ended``
|
||||
stays out of the comparison and only says whether the handler is on its
|
||||
end-of-stream path, where an empty payload is still scanned today."""
|
||||
end-of-stream path, where an empty payload is still scanned today.
|
||||
``tool_calls_in_flight`` also stays out of the comparison: it flags that tool
|
||||
calls have streamed which this round cannot scan yet, so a buffered window
|
||||
holding them must stay withheld until the end-of-stream scan covers them."""
|
||||
|
||||
texts: tuple[str, ...]
|
||||
tool_calls: tuple[str, ...] = ()
|
||||
stream_ended: bool = field(default=False, compare=False)
|
||||
tool_calls_in_flight: bool = field(default=False, compare=False)
|
||||
|
||||
@property
|
||||
def has_nothing_to_scan(self) -> bool:
|
||||
|
|
|
|||
|
|
@ -1902,7 +1902,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
return None
|
||||
tokens_5m: Final = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "5m")
|
||||
tokens_1h: Final = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "1h")
|
||||
if tokens_5m + tokens_1h != usage.get("cacheWriteInputTokens", 0):
|
||||
if tokens_5m + tokens_1h != AmazonConverseConfig._cache_write_count(usage):
|
||||
return None
|
||||
return CacheCreationTokenDetails(
|
||||
ephemeral_5m_input_tokens=tokens_5m,
|
||||
|
|
@ -1933,6 +1933,15 @@ class AmazonConverseConfig(BaseConfig):
|
|||
return int(value)
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
def _cache_read_count(usage_object: Mapping[str, object]) -> int:
|
||||
"""Converse reports ``cacheReadInputTokens``; InvokeModel reports ``cacheReadInputTokenCount``."""
|
||||
return AmazonConverseConfig._usage_count(usage_object, "cacheReadInputTokens", "cacheReadInputTokenCount")
|
||||
|
||||
@staticmethod
|
||||
def _cache_write_count(usage_object: Mapping[str, object]) -> int:
|
||||
return AmazonConverseConfig._usage_count(usage_object, "cacheWriteInputTokens", "cacheWriteInputTokenCount")
|
||||
|
||||
def usage_from_batch_output(self, usage_object: Mapping[str, object]) -> Usage:
|
||||
"""Read a Converse-shaped usage block out of a batch output line.
|
||||
|
||||
|
|
@ -1942,8 +1951,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
"""
|
||||
input_tokens: Final = self._usage_count(usage_object, "inputTokens")
|
||||
output_tokens: Final = self._usage_count(usage_object, "outputTokens")
|
||||
cache_read: Final = self._usage_count(usage_object, "cacheReadInputTokens", "cacheReadInputTokenCount")
|
||||
cache_write: Final = self._usage_count(usage_object, "cacheWriteInputTokens", "cacheWriteInputTokenCount")
|
||||
cache_read: Final = self._cache_read_count(usage_object)
|
||||
cache_write: Final = self._cache_write_count(usage_object)
|
||||
return self.transform_usage(
|
||||
ConverseTokenUsageBlock(
|
||||
inputTokens=input_tokens,
|
||||
|
|
@ -1963,19 +1972,12 @@ class AmazonConverseConfig(BaseConfig):
|
|||
thinking_ran: bool = False,
|
||||
provider_reasoning_tokens: int | None = None,
|
||||
) -> Usage:
|
||||
input_tokens = usage["inputTokens"]
|
||||
raw_input_tokens: Final = usage["inputTokens"]
|
||||
output_tokens: Final = usage["outputTokens"]
|
||||
total_tokens: Final = usage["totalTokens"]
|
||||
cache_creation_input_tokens: int = 0
|
||||
cache_read_input_tokens: int = 0
|
||||
|
||||
raw_input_tokens: Final = input_tokens # capture before inflation
|
||||
if "cacheReadInputTokens" in usage:
|
||||
cache_read_input_tokens = usage["cacheReadInputTokens"]
|
||||
input_tokens += cache_read_input_tokens
|
||||
if "cacheWriteInputTokens" in usage:
|
||||
cache_creation_input_tokens = usage["cacheWriteInputTokens"]
|
||||
input_tokens += cache_creation_input_tokens
|
||||
cache_read_input_tokens: Final = self._cache_read_count(usage)
|
||||
cache_creation_input_tokens: Final = self._cache_write_count(usage)
|
||||
input_tokens: Final = raw_input_tokens + cache_read_input_tokens + cache_creation_input_tokens
|
||||
total_tokens: Final = usage.get("totalTokens", input_tokens + output_tokens)
|
||||
|
||||
prompt_tokens_details: Final = PromptTokensDetailsWrapper(
|
||||
cached_tokens=cache_read_input_tokens,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from collections.abc import AsyncIterator, Iterator
|
|||
from typing import Final, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
|
|
@ -51,6 +52,15 @@ bedrock_tool_name_mappings: Final[InMemoryCache] = InMemoryCache(max_size_in_mem
|
|||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
|
||||
converse_config: Final = AmazonConverseConfig()
|
||||
NOVA_INVOKE_STREAM_EVENT_TYPES: Final = (
|
||||
"messageStart",
|
||||
"contentBlockStart",
|
||||
"contentBlockDelta",
|
||||
"contentBlockStop",
|
||||
"messageStop",
|
||||
"metadata",
|
||||
)
|
||||
NOVA_INVOKE_STREAM_EVENT_PAYLOAD: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
class AmazonCohereChatConfig:
|
||||
|
|
@ -601,14 +611,12 @@ class AWSEventStreamDecoder:
|
|||
if thinking_blocks:
|
||||
self._thinking_ran = True
|
||||
|
||||
carries_message_content: Final = any(
|
||||
key in chunk_data for key in ("start", "delta", "contentBlockIndex", "stopReason", "trace")
|
||||
trace: Final = chunk_data.get("trace")
|
||||
carries_message_content: Final = bool(trace) or any(
|
||||
key in chunk_data for key in ("start", "delta", "contentBlockIndex", "stopReason")
|
||||
)
|
||||
|
||||
model_response_provider_specific_fields: Final = {}
|
||||
if "trace" in chunk_data:
|
||||
trace: Final = chunk_data.get("trace")
|
||||
model_response_provider_specific_fields["trace"] = trace
|
||||
model_response_provider_specific_fields: Final = {"trace": trace} if trace else {}
|
||||
response: Final = ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
|
|
@ -654,10 +662,10 @@ class AWSEventStreamDecoder:
|
|||
):
|
||||
return self.converse_chunk_parser(chunk_data=chunk_data)
|
||||
######### /bedrock/invoke nova mappings ###############
|
||||
elif "contentBlockDelta" in chunk_data:
|
||||
# when using /bedrock/invoke/nova, the chunk_data is nested under "contentBlockDelta"
|
||||
_chunk_data: Final = chunk_data.get("contentBlockDelta", {})
|
||||
return self.converse_chunk_parser(chunk_data=_chunk_data)
|
||||
elif nova_event_type := next((key for key in NOVA_INVOKE_STREAM_EVENT_TYPES if key in chunk_data), None):
|
||||
return self.converse_chunk_parser(
|
||||
chunk_data=NOVA_INVOKE_STREAM_EVENT_PAYLOAD.validate_python(chunk_data[nova_event_type])
|
||||
)
|
||||
######## bedrock.mistral mappings ###############
|
||||
elif "outputs" in chunk_data:
|
||||
if len(chunk_data["outputs"]) == 1 and chunk_data["outputs"][0].get("text", None) is not None:
|
||||
|
|
|
|||
|
|
@ -6,12 +6,21 @@ Inherits from `AmazonConverseConfig`
|
|||
Nova + Invoke API Tutorial: https://docs.aws.amazon.com/nova/latest/userguide/using-invoke-api.html
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from functools import reduce
|
||||
from typing import TYPE_CHECKING, Final, TypeVar
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.llms.bedrock import BedrockInvokeNovaRequest
|
||||
from litellm.types.llms.bedrock import (
|
||||
BedrockInvokeNovaRequest,
|
||||
CachePointBlock,
|
||||
ContentBlock,
|
||||
MessageBlock,
|
||||
SystemContentBlock,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -21,6 +30,50 @@ from .base_invoke_transformation import AmazonInvokeConfig
|
|||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
_CachePointCarrier = TypeVar("_CachePointCarrier", SystemContentBlock, ContentBlock)
|
||||
_INJECTION_POINTS: Final = TypeAdapter(tuple[Mapping[str, object], ...])
|
||||
|
||||
|
||||
def _without_tool_config_injection_points(optional_params: Mapping[str, object]) -> dict[str, object]:
|
||||
"""InvokeModel has no tool caching, and a ``tool_config`` point the Converse transform
|
||||
placed would credit the gateway for a cachePoint this request cannot carry.
|
||||
"""
|
||||
raw_points: Final = optional_params.get("cache_control_injection_points")
|
||||
if raw_points is None:
|
||||
return dict(optional_params)
|
||||
try:
|
||||
points = _INJECTION_POINTS.validate_python(raw_points)
|
||||
except ValidationError:
|
||||
return dict(optional_params)
|
||||
return {
|
||||
**optional_params,
|
||||
"cache_control_injection_points": [point for point in points if point.get("location") != "tool_config"],
|
||||
}
|
||||
|
||||
|
||||
def _system_block_with_cache_point(block: SystemContentBlock, cache_point: CachePointBlock) -> SystemContentBlock:
|
||||
return {**block, "cachePoint": cache_point}
|
||||
|
||||
|
||||
def _content_block_with_cache_point(block: ContentBlock, cache_point: CachePointBlock) -> ContentBlock:
|
||||
return {**block, "cachePoint": cache_point}
|
||||
|
||||
|
||||
def _inline_block_cache_points(
|
||||
blocks: Sequence[_CachePointCarrier],
|
||||
with_cache_point: Callable[[_CachePointCarrier, CachePointBlock], _CachePointCarrier],
|
||||
) -> list[_CachePointCarrier]:
|
||||
def attach(inlined: tuple[_CachePointCarrier, ...], block: _CachePointCarrier) -> tuple[_CachePointCarrier, ...]:
|
||||
cache_point: Final = block.get("cachePoint")
|
||||
if cache_point is None or len(block) != 1:
|
||||
return (*inlined, block)
|
||||
anchor: Final = next((index for index in reversed(range(len(inlined))) if "text" in inlined[index]), None)
|
||||
if anchor is None:
|
||||
return inlined
|
||||
return (*inlined[:anchor], with_cache_point(inlined[anchor], cache_point), *inlined[anchor + 1 :])
|
||||
|
||||
return list(reduce(attach, blocks, ()))
|
||||
|
||||
|
||||
class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
||||
"""
|
||||
|
|
@ -46,7 +99,7 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
|||
self,
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
|
|
@ -54,11 +107,13 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
|||
self,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
optional_params=_without_tool_config_injection_points(optional_params),
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
_bedrock_invoke_nova_request: Final = BedrockInvokeNovaRequest(**_transformed_nova_request)
|
||||
_bedrock_invoke_nova_request: Final = self._inline_cache_points(
|
||||
BedrockInvokeNovaRequest(**_transformed_nova_request)
|
||||
)
|
||||
self._remove_empty_system_messages(_bedrock_invoke_nova_request)
|
||||
bedrock_invoke_nova_request: Final = self._filter_allowed_fields(_bedrock_invoke_nova_request)
|
||||
return bedrock_invoke_nova_request
|
||||
|
|
@ -92,6 +147,24 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
|||
json_mode,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _inline_cache_points(request: BedrockInvokeNovaRequest) -> BedrockInvokeNovaRequest:
|
||||
"""InvokeModel takes ``cachePoint`` as a key of the text block it caches: it rejects the
|
||||
standalone ``{"cachePoint": ...}`` blocks Converse accepts and the key on image, toolUse,
|
||||
and toolResult blocks, so a point behind one of those moves back to the last text block.
|
||||
"""
|
||||
return {
|
||||
**request,
|
||||
"system": _inline_block_cache_points(request.get("system", []), _system_block_with_cache_point),
|
||||
"messages": [
|
||||
MessageBlock(
|
||||
role=message["role"],
|
||||
content=_inline_block_cache_points(message["content"], _content_block_with_cache_point),
|
||||
)
|
||||
for message in request.get("messages", [])
|
||||
],
|
||||
}
|
||||
|
||||
def _filter_allowed_fields(self, bedrock_invoke_nova_request: BedrockInvokeNovaRequest) -> dict:
|
||||
"""
|
||||
Filter out fields that are not allowed in the `BedrockInvokeNovaRequest` dataclass.
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import ssl
|
|||
import sys
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
|
||||
from http.cookiejar import CookieJar, DefaultCookiePolicy
|
||||
from io import BytesIO
|
||||
|
|
@ -185,6 +186,33 @@ def _handler_may_close_client(client_refcount: int, owns_client: bool) -> bool:
|
|||
return owns_client and client_refcount <= _CLIENT_REFCOUNT_WHEN_HANDLER_IS_SOLE_REFERRER
|
||||
|
||||
|
||||
def _drop_streaming_anchor(_handler: object) -> None:
|
||||
"""Release a handler anchored to a streaming response. See ``_anchor_handler_to``.
|
||||
|
||||
The work is the reference held until this point, so there is nothing to do here.
|
||||
"""
|
||||
|
||||
|
||||
def _anchor_handler_to(response: httpx.Response, handler: object) -> None:
|
||||
"""Keep the handler alive for as long as a streaming response can still read.
|
||||
|
||||
A body still arriving reads through the handler's connection pool, and closing
|
||||
the client tears that pool down. The refcount ``_handler_may_close_client``
|
||||
reads cannot see that body: the reference graph runs response -> stream ->
|
||||
connection and stops there, so a client carrying one looks exactly like an
|
||||
unreferenced client, and the finalizer closes it mid-body.
|
||||
|
||||
``weakref.finalize`` holds the handler in its own registry rather than on the
|
||||
response, which matters twice. The handler stays out of the response's
|
||||
reference cycle, so it is finalized by refcount once the anchor drops and can
|
||||
still schedule an async close, instead of being finalized inside a cyclic
|
||||
collection that reaps its aiohttp session in the same pass. And a handler
|
||||
serving several streams collects only once every one of them is done, because
|
||||
each anchor holds it separately.
|
||||
"""
|
||||
weakref.finalize(response, _drop_streaming_anchor, handler)
|
||||
|
||||
|
||||
def blocked_cookie_jar() -> CookieJar:
|
||||
"""A jar that stores no response cookie and sends none, for httpx clients.
|
||||
|
||||
|
|
@ -778,6 +806,8 @@ class AsyncHTTPHandler:
|
|||
content=request_content,
|
||||
)
|
||||
response: Final = await self.client.send(req, stream=stream)
|
||||
if stream:
|
||||
_anchor_handler_to(response, self)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except (httpx.RemoteProtocolError, httpx.ConnectError):
|
||||
|
|
@ -982,6 +1012,8 @@ class AsyncHTTPHandler:
|
|||
content=request_content,
|
||||
)
|
||||
response: Final = await self.client.send(req, stream=stream)
|
||||
if stream:
|
||||
_anchor_handler_to(response, self)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except (httpx.RemoteProtocolError, httpx.ConnectError):
|
||||
|
|
@ -1451,6 +1483,8 @@ class HTTPHandler:
|
|||
content=request_content,
|
||||
)
|
||||
response: Final = self.client.send(req, stream=stream)
|
||||
if stream:
|
||||
_anchor_handler_to(response, self)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except httpx.TimeoutException:
|
||||
|
|
@ -1501,6 +1535,8 @@ class HTTPHandler:
|
|||
content=request_content,
|
||||
)
|
||||
response: Final = self.client.send(req, stream=stream)
|
||||
if stream:
|
||||
_anchor_handler_to(response, self)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except httpx.TimeoutException:
|
||||
|
|
@ -1551,6 +1587,8 @@ class HTTPHandler:
|
|||
content=request_content,
|
||||
)
|
||||
response: Final = self.client.send(req, stream=stream)
|
||||
if stream:
|
||||
_anchor_handler_to(response, self)
|
||||
return response
|
||||
except httpx.TimeoutException:
|
||||
raise litellm.Timeout(
|
||||
|
|
@ -1600,6 +1638,8 @@ class HTTPHandler:
|
|||
content=request_content,
|
||||
)
|
||||
response: Final = self.client.send(req, stream=stream)
|
||||
if stream:
|
||||
_anchor_handler_to(response, self)
|
||||
response.raise_for_status()
|
||||
return response
|
||||
except httpx.TimeoutException:
|
||||
|
|
|
|||
|
|
@ -49,6 +49,17 @@ if TYPE_CHECKING:
|
|||
import tiktoken
|
||||
|
||||
|
||||
def _map_reasoning_effort(value: object) -> object:
|
||||
effort: Final[object] = cast(Mapping[str, object], value).get("effort") if isinstance(value, Mapping) else value
|
||||
if effort is True:
|
||||
return "medium"
|
||||
if effort is False:
|
||||
return "none"
|
||||
if effort == "auto":
|
||||
return None
|
||||
return effort
|
||||
|
||||
|
||||
def _extract_fireworks_hidden_params(payload: dict) -> dict:
|
||||
"""
|
||||
Collect Fireworks-specific response fields (perf_metrics, prompt_token_ids,
|
||||
|
|
@ -327,12 +338,9 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
|
|||
elif param == "max_completion_tokens":
|
||||
optional_params["max_tokens"] = value
|
||||
elif param == "reasoning_effort":
|
||||
if value is True:
|
||||
optional_params["reasoning_effort"] = "medium"
|
||||
elif value is False:
|
||||
optional_params["reasoning_effort"] = "none"
|
||||
elif value != "auto":
|
||||
optional_params["reasoning_effort"] = value
|
||||
effort = _map_reasoning_effort(value)
|
||||
if effort is not None:
|
||||
optional_params["reasoning_effort"] = effort
|
||||
elif param in supported_openai_params:
|
||||
if value is not None:
|
||||
optional_params[param] = value
|
||||
|
|
|
|||
|
|
@ -792,10 +792,12 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None:
|
||||
chunks: Final = tuple(chunk for chunk in responses_so_far if isinstance(chunk, ModelResponseStream))
|
||||
stream_ended: Final = self._first_choice_has_finished(responses_so_far)
|
||||
tool_call_fingerprints: Final = self._streamed_tool_call_fingerprints(responses_so_far)
|
||||
return StreamingScanKey(
|
||||
texts=tuple(self._combine_streaming_texts(chunks).values()),
|
||||
tool_calls=self._streamed_tool_call_fingerprints(responses_so_far) if stream_ended else (),
|
||||
tool_calls=tool_call_fingerprints if stream_ended else (),
|
||||
stream_ended=stream_ended,
|
||||
tool_calls_in_flight=bool(tool_call_fingerprints) and not stream_ended,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -804,7 +806,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
stream_item_fingerprint(tool_call)
|
||||
for chunk in responses_so_far
|
||||
for choice in _stream_chunk_choices(chunk)
|
||||
for tool_call in stream_item_items(stream_item_field(choice, "delta"), "tool_calls")
|
||||
for tool_call in _streamed_delta_tool_calls(stream_item_field(choice, "delta"))
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1342,6 +1344,12 @@ def _stream_chunk_choices(item: object) -> Sequence[object]:
|
|||
return ()
|
||||
|
||||
|
||||
def _streamed_delta_tool_calls(delta: object) -> tuple[object, ...]:
|
||||
function_call: Final = stream_item_field(delta, "function_call")
|
||||
legacy: Final = () if function_call is None else (function_call,)
|
||||
return stream_item_items(delta, "tool_calls") + legacy
|
||||
|
||||
|
||||
def _blocked_stream_identity(
|
||||
exc: "ModifyResponseException", responses_so_far: Sequence[object]
|
||||
) -> tuple[str, int, str]:
|
||||
|
|
|
|||
|
|
@ -1175,11 +1175,22 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
last_event_type: Final = stream_item_field(last_event, "type")
|
||||
if last_event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE.value:
|
||||
return None
|
||||
if last_event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED.value:
|
||||
if last_event_type in _TERMINAL_ENVELOPE_EVENT_TYPES:
|
||||
return self._completed_response_scan_key(stream_item_field(last_event, "response"))
|
||||
return StreamingScanKey(
|
||||
texts=(self.get_streaming_string_so_far(responses_so_far),),
|
||||
stream_ended=self._check_streaming_has_ended(responses_so_far),
|
||||
tool_calls_in_flight=self._has_streamed_tool_call_events(responses_so_far),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _has_streamed_tool_call_events(responses_so_far: Sequence[object]) -> bool:
|
||||
return any(
|
||||
stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_EVENT_TYPES
|
||||
or (
|
||||
stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES
|
||||
and stream_item_field(stream_item_field(event, "item"), "type") in _TOOL_CALL_ITEM_TYPES
|
||||
)
|
||||
for event in responses_so_far
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -18,6 +18,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
|
|||
key_name: str | None = None
|
||||
key_alias: str | None = None
|
||||
spend: float = 0.0
|
||||
total_spend: float = 0.0
|
||||
max_budget: float | None = None
|
||||
expires: str | datetime | None = None
|
||||
models: list = []
|
||||
|
|
@ -69,6 +70,7 @@ class LiteLLM_DeletedVerificationToken(LiteLLM_VerificationToken):
|
|||
"""Audit record for deleted keys; mirrors the token plus deletion metadata."""
|
||||
|
||||
id: str | None = None
|
||||
organization_id: str | None = None
|
||||
deleted_at: datetime | None = None
|
||||
deleted_by: str | None = None
|
||||
deleted_by_api_key: str | None = None
|
||||
|
|
|
|||
|
|
@ -2004,6 +2004,13 @@ RouterSettingsDict = Annotated[
|
|||
class NewTeamRequest(TeamBase):
|
||||
router_settings: RouterSettingsDict | None = None
|
||||
model_aliases: dict | None = None
|
||||
model_max_budget: GenericBudgetConfigType | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Max budget per model for every key on the team, overridable per key "
|
||||
"(e.g. {'gpt-4o': {'max_budget': 10, 'budget_duration': '1d'}})"
|
||||
),
|
||||
)
|
||||
tags: list | None = None
|
||||
guardrails: list[str] | None = None
|
||||
policies: list[str] | None = None
|
||||
|
|
@ -2105,6 +2112,13 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
access_group_ids: list[str] | None = None
|
||||
budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows
|
||||
default_team_member_models: list[str] | None = None # default allowed_models seeded onto new team members
|
||||
model_max_budget: GenericBudgetConfigType | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Max budget per model for every key on the team, overridable per key "
|
||||
"(e.g. {'gpt-4o': {'max_budget': 10, 'budget_duration': '1d'}})"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class PatchTeamRequest(UpdateTeamRequest):
|
||||
|
|
@ -3032,6 +3046,7 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken):
|
|||
team_tpd_limit: int | None = None
|
||||
team_max_budget: float | None = None
|
||||
team_soft_budget: float | None = None
|
||||
team_model_max_budget: dict[str, object] | None = None
|
||||
team_models: list = []
|
||||
team_blocked: bool = False
|
||||
soft_budget: float | None = None
|
||||
|
|
@ -3710,6 +3725,7 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
|
|||
"AWS_ACCESS_KEY_ID",
|
||||
"AWS_SECRET_ACCESS_KEY",
|
||||
"AWS_REGION_NAME",
|
||||
"S3_LOG_PROMPTS_ONLY",
|
||||
],
|
||||
)
|
||||
|
||||
|
|
@ -4462,6 +4478,7 @@ class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):
|
|||
# Parent org's model ceiling, reported only to callers who can manage the team.
|
||||
# None = no org or not a manager; [] or ["all-proxy-models"] = no ceiling.
|
||||
organization_models: list[str] | None = None
|
||||
model_max_budget_usage: Mapping[str, Mapping[str, object]] | None = None
|
||||
|
||||
|
||||
class TeamInfoResponseObject(TypedDict):
|
||||
|
|
|
|||
166
litellm/proxy/auth/fallback_budget.py
Normal file
166
litellm/proxy/auth/fallback_budget.py
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
"""
|
||||
Enforce the caller's budget against router fallback targets.
|
||||
|
||||
Budget is checked once, during auth, against the *requested* model group. A zero-cost group takes
|
||||
`_is_model_cost_zero`'s bypass and waives every budget check; the router then picks a fallback
|
||||
target after auth, inside `run_async_fallback`, and nothing re-checks budget on the group that
|
||||
actually bills. So a free model with a paid fallback spends without a gate.
|
||||
|
||||
This predicate is injected into the router to re-check budget for each fallback target before it is
|
||||
attempted, mirroring `fallback_model_access.py`. It deliberately leaves the primary attempt alone:
|
||||
a zero-cost model is never blocked by budget, and only the paid fallback is refused. On by default;
|
||||
set `general_settings.enforce_fallback_budget: false` to restore the unguarded behaviour.
|
||||
|
||||
Scope: the key's and the user's `max_budget`. Not covered yet, and each needs a read-only evaluation
|
||||
path before it can be: team, team-member, end-user, org, global and per-model budgets, whose
|
||||
auth-path functions enforce rather than report (they raise), so reusing them would fire threshold
|
||||
alerts and take spend reservations for a target that is then skipped; and the key's rolling
|
||||
`budget_limits` windows, whose accumulated spend lives only in per-window counters
|
||||
(`spend:key:{token}:window:{budget_duration}`), so enforcing them means more counter reads on the
|
||||
fallback path rather than reusing state auth already loaded.
|
||||
|
||||
Two known limitations of that narrow scope, both shared with `fallback_model_access.py`:
|
||||
|
||||
* This reads the spend counter, it does not reserve against it. Requests already in flight all
|
||||
observe the same pre-billing figure, so a cap can be crossed by roughly the number of concurrent
|
||||
fallbacks times their cost. Auth-time enforcement avoids this by pre-filling the counter through
|
||||
`reserve_budget_for_request`, which the zero-cost bypass skips. Turning the soft cap into a hard
|
||||
one means reserving per fallback attempt and reconciling on completion.
|
||||
* A request that reaches the router without `metadata["user_api_key_auth"]` is not restricted.
|
||||
Only `add_litellm_data_to_request` populates that key, so endpoints that assemble metadata by
|
||||
hand (for example `/queue/chat/completions`) fall through as unauthenticated.
|
||||
"""
|
||||
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_is_model_cost_zero, # pyright: ignore[reportPrivateUsage] # the zero-cost predicate the auth-time budget checks use; no public equivalent
|
||||
)
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
class _RequestMetadata(BaseModel):
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None
|
||||
|
||||
|
||||
class _FallbackBudgetSettings(BaseModel):
|
||||
enforce_fallback_budget: bool = True
|
||||
|
||||
|
||||
def _token_in_metadata(metadata: object) -> UserAPIKeyAuth | None:
|
||||
try:
|
||||
return _RequestMetadata.model_validate(metadata).user_api_key_auth
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _user_api_key_auth_from_request(request_kwargs: Mapping[str, object]) -> UserAPIKeyAuth | None:
|
||||
return next(
|
||||
(
|
||||
token
|
||||
for field in ("metadata", "litellm_metadata")
|
||||
if (token := _token_in_metadata(request_kwargs.get(field))) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def _enforced_by_general_settings() -> bool:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
return _FallbackBudgetSettings.model_validate(general_settings).enforce_fallback_budget
|
||||
|
||||
|
||||
def _applies_user_budget_to_team_keys() -> bool:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
return general_settings.get("apply_user_budget_to_team_keys") is True
|
||||
|
||||
|
||||
async def _counter_spend(counter_key: str, fallback_spend: float, max_budget: float) -> float:
|
||||
"""
|
||||
Read a spend counter the same way the auth-time budget checks do.
|
||||
|
||||
`max_budget` is not advisory: it makes `get_current_spend` re-check the counter against the
|
||||
authoritative recorded spend before admitting. A counter restored from an older Redis snapshot
|
||||
reads as a hit rather than a clean miss, so without this the reseed path never runs and a
|
||||
stale-low counter would keep admitting paid fallbacks past the cap.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import get_current_spend
|
||||
|
||||
return await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=fallback_spend,
|
||||
max_budget=max_budget,
|
||||
)
|
||||
|
||||
|
||||
async def is_token_within_budget_for_model(*, model: str, valid_token: UserAPIKeyAuth, llm_router: Router) -> bool:
|
||||
"""
|
||||
True when the key and the user behind it can still pay for `model`.
|
||||
|
||||
A zero-cost fallback target is always allowed: refusing it would deny a request on spend some
|
||||
other model accrued, which is the same reasoning behind the auth-time bypass.
|
||||
"""
|
||||
if _is_model_cost_zero(model=model, llm_router=llm_router):
|
||||
return True
|
||||
|
||||
key_budget: Final = valid_token.max_budget
|
||||
if key_budget is not None and valid_token.token is not None:
|
||||
key_spend: Final = await _counter_spend(
|
||||
counter_key=f"spend:key:{valid_token.token}",
|
||||
fallback_spend=valid_token.spend or 0.0,
|
||||
max_budget=key_budget,
|
||||
)
|
||||
if key_spend >= key_budget:
|
||||
return False
|
||||
|
||||
# Mirrors `_PROXY_MaxBudgetLimiter`: a team key does not carry the key owner's personal budget
|
||||
# unless the proxy opts in, so the personal cap must not gate the fallback either.
|
||||
user_budget: Final = valid_token.user_max_budget
|
||||
if (
|
||||
user_budget is not None
|
||||
and valid_token.user_id is not None
|
||||
and (valid_token.team_id is None or _applies_user_budget_to_team_keys())
|
||||
):
|
||||
user_spend: Final = await _counter_spend(
|
||||
counter_key=f"spend:user:{valid_token.user_id}",
|
||||
fallback_spend=valid_token.user_spend or 0.0,
|
||||
max_budget=user_budget,
|
||||
)
|
||||
if user_spend >= user_budget:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouterFallbackBudgetCheck:
|
||||
"""
|
||||
`FallbackBudgetCheck` for the proxy's router: while `is_enforced()` is true, a paid fallback
|
||||
target is attempted only when the caller is still within budget. Requests that carry no key
|
||||
(for example internal health checks) are not restricted.
|
||||
"""
|
||||
|
||||
is_enforced: Callable[[], bool]
|
||||
|
||||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: Router) -> bool:
|
||||
if not self.is_enforced():
|
||||
return True
|
||||
valid_token: Final = _user_api_key_auth_from_request(request_kwargs)
|
||||
if valid_token is None:
|
||||
return True
|
||||
try:
|
||||
return await is_token_within_budget_for_model(model=model, valid_token=valid_token, llm_router=llm_router)
|
||||
except Exception as e: # noqa: BLE001 # fail closed: a spend lookup failure must not bill the caller
|
||||
verbose_proxy_logger.warning("Skipping fallback to model=%s: budget lookup failed: %s", model, e)
|
||||
return False
|
||||
|
||||
|
||||
router_fallback_budget_check: Final = RouterFallbackBudgetCheck(is_enforced=_enforced_by_general_settings)
|
||||
|
|
@ -59,6 +59,7 @@ class TeamGrants(TypedDict, total=False):
|
|||
team_tpd_limit: ReadOnly[int | None]
|
||||
team_max_budget: ReadOnly[float | None]
|
||||
team_soft_budget: ReadOnly[float | None]
|
||||
team_model_max_budget: ReadOnly[dict[str, object] | None]
|
||||
team_spend: ReadOnly[float | None]
|
||||
team_models: ReadOnly[Sequence[str]]
|
||||
team_blocked: ReadOnly[bool]
|
||||
|
|
@ -101,6 +102,7 @@ def team_grants(
|
|||
team_tpd_limit=team_object.tpd_limit,
|
||||
team_max_budget=team_object.max_budget,
|
||||
team_soft_budget=team_object.soft_budget,
|
||||
team_model_max_budget=team_object.model_max_budget,
|
||||
team_spend=team_object.spend,
|
||||
team_models=tuple(team_object.models),
|
||||
team_blocked=team_object.blocked,
|
||||
|
|
|
|||
|
|
@ -304,6 +304,16 @@ class _UserModelBudgetLimiter(Protocol):
|
|||
) -> bool: ...
|
||||
|
||||
|
||||
class _TeamModelBudgetLimiter(Protocol):
|
||||
async def is_team_within_model_budget(
|
||||
self,
|
||||
team_id: str,
|
||||
team_model_max_budget: Mapping[str, object],
|
||||
key_model_max_budget: Mapping[str, object] | None,
|
||||
model: str,
|
||||
) -> bool: ...
|
||||
|
||||
|
||||
class _TokenTeamModels(Protocol):
|
||||
@property
|
||||
def team_models(self) -> list[str]: ...
|
||||
|
|
@ -374,6 +384,25 @@ async def _check_user_model_budget(
|
|||
)
|
||||
|
||||
|
||||
async def _check_team_model_budget(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
model_max_budget_limiter: _TeamModelBudgetLimiter,
|
||||
models: list[str],
|
||||
) -> None:
|
||||
"""Enforce the team's `model_max_budget` for every requested model the key does not override."""
|
||||
team_model_max_budget: Final = valid_token.team_model_max_budget
|
||||
if valid_token.team_id is None or not team_model_max_budget:
|
||||
return
|
||||
key_model_max_budget: Final[Mapping[str, object] | None] = valid_token.model_max_budget
|
||||
for model_name in models:
|
||||
await model_max_budget_limiter.is_team_within_model_budget(
|
||||
team_id=valid_token.team_id,
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
model=model_name,
|
||||
)
|
||||
|
||||
|
||||
async def _check_key_model_budget_with_fallback(
|
||||
valid_token: UserAPIKeyAuth,
|
||||
model_max_budget_limiter: _KeyModelBudgetLimiter,
|
||||
|
|
@ -2376,6 +2405,7 @@ async def _user_api_key_auth_builder(
|
|||
team_id=valid_token.team_id,
|
||||
max_budget=valid_token.team_max_budget,
|
||||
soft_budget=valid_token.team_soft_budget,
|
||||
model_max_budget=valid_token.team_model_max_budget,
|
||||
spend=valid_token.team_spend,
|
||||
tpm_limit=valid_token.team_tpm_limit,
|
||||
rpm_limit=valid_token.team_rpm_limit,
|
||||
|
|
@ -2530,6 +2560,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached
|
|||
team_id=valid_token.team_id,
|
||||
max_budget=valid_token.team_max_budget,
|
||||
soft_budget=valid_token.team_soft_budget,
|
||||
model_max_budget=valid_token.team_model_max_budget,
|
||||
spend=valid_token.team_spend,
|
||||
tpm_limit=valid_token.team_tpm_limit,
|
||||
rpm_limit=valid_token.team_rpm_limit,
|
||||
|
|
@ -2571,6 +2602,13 @@ def _token_can_vouch_for_team(valid_token: UserAPIKeyAuth, lookup_error: BaseExc
|
|||
return PrismaDBExceptionHandler.should_allow_request_on_db_unavailable()
|
||||
|
||||
|
||||
def is_no_auth_dev_mode(master_key: str | None, general_settings: Mapping[str, object]) -> bool:
|
||||
return master_key is None and not any(
|
||||
general_settings.get(flag, False)
|
||||
for flag in ("enable_jwt_auth", "enable_oauth2_auth", "enable_oauth2_proxy_auth")
|
||||
)
|
||||
|
||||
|
||||
@tracer.wrap()
|
||||
async def _run_centralized_common_checks(
|
||||
user_api_key_auth_obj: UserAPIKeyAuth,
|
||||
|
|
@ -2599,6 +2637,7 @@ async def _run_centralized_common_checks(
|
|||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
master_key,
|
||||
model_max_budget_limiter,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
|
|
@ -2630,11 +2669,7 @@ async def _run_centralized_common_checks(
|
|||
# Running common_checks would block every admin route on these
|
||||
# deployments where that was previously not the contract. If any
|
||||
# authn is enabled (JWT, OAuth2, OAuth2-proxy), authz must run.
|
||||
if master_key is None and not (
|
||||
general_settings.get("enable_jwt_auth", False)
|
||||
or general_settings.get("enable_oauth2_auth", False)
|
||||
or general_settings.get("enable_oauth2_proxy_auth", False)
|
||||
):
|
||||
if is_no_auth_dev_mode(master_key, general_settings):
|
||||
return
|
||||
|
||||
if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False):
|
||||
|
|
@ -2871,6 +2906,21 @@ async def _run_centralized_common_checks(
|
|||
finally:
|
||||
release_spend_counter_batch()
|
||||
|
||||
if not skip_budget_checks:
|
||||
await _check_team_model_budget(
|
||||
valid_token=user_api_key_auth_obj,
|
||||
model_max_budget_limiter=model_max_budget_limiter,
|
||||
models=_get_model_names_for_budget_checks(
|
||||
model=_get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=user_api_key_auth_obj.team_id,
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
await _reserve_budget_after_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request=request,
|
||||
|
|
|
|||
|
|
@ -78,6 +78,7 @@ async def create_missing_views(db: SupportsRawQueries) -> None:
|
|||
v.*,
|
||||
t.spend AS team_spend,
|
||||
t.max_budget AS team_max_budget,
|
||||
t.model_max_budget AS team_model_max_budget,
|
||||
t.tpm_limit AS team_tpm_limit,
|
||||
t.rpm_limit AS team_rpm_limit,
|
||||
t.tpd_limit AS team_tpd_limit,
|
||||
|
|
|
|||
|
|
@ -18,6 +18,8 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
|
||||
from urllib.parse import quote, unquote
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import RedisCache
|
||||
|
|
@ -109,6 +111,10 @@ def _batch_cost_row_to_write(payload: SpendLogsPayload, disable_spend_logs: bool
|
|||
return MappingProxyType({field: value for field, value in payload.items() if field in _BATCH_COST_CLAIM_FIELDS})
|
||||
|
||||
|
||||
class _SpendIncrement(TypedDict):
|
||||
increment: ReadOnly[float]
|
||||
|
||||
|
||||
class _SpendBatch(Protocol):
|
||||
litellm_usertable: BatchTable
|
||||
litellm_verificationtoken: BatchTable
|
||||
|
|
@ -1615,10 +1621,12 @@ class DBSpendUpdateWriter:
|
|||
async with transaction.batch_() as batcher:
|
||||
# Sort by token for consistent lock ordering across pods to prevent deadlocks.
|
||||
for token, response_cost in sorted(key_list_transactions.items()):
|
||||
spend_increment: _SpendIncrement = {"increment": response_cost}
|
||||
batcher.litellm_verificationtoken.update_many( # 'update_many' prevents error from being raised if no row exists
|
||||
where={"token": token},
|
||||
data={
|
||||
"spend": {"increment": response_cost},
|
||||
"spend": spend_increment,
|
||||
"total_spend": spend_increment,
|
||||
"last_active": datetime.now(timezone.utc),
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.types.guardrails import (
|
|||
ApplyGuardrailResponse,
|
||||
BaseLitellmParams,
|
||||
BedrockGuardrailConfigModel,
|
||||
BedrockGuardrailStreamingParams,
|
||||
Guardrail,
|
||||
GuardrailEventHooks,
|
||||
GuardrailInfoResponse,
|
||||
|
|
@ -1959,7 +1960,10 @@ async def get_provider_specific_params():
|
|||
```
|
||||
"""
|
||||
# Get fields from the models
|
||||
bedrock_fields: Final = _get_fields_from_model(BedrockGuardrailConfigModel)
|
||||
bedrock_fields: Final = {
|
||||
**_get_fields_from_model(BedrockGuardrailConfigModel),
|
||||
**_get_fields_from_model(BedrockGuardrailStreamingParams),
|
||||
}
|
||||
presidio_fields: Final = _get_fields_from_model(PresidioPresidioConfigModelUserInterface)
|
||||
lakera_v2_fields: Final = _get_fields_from_model(LakeraV2GuardrailConfigModel)
|
||||
tool_permission_fields: Final = _get_fields_from_model(ToolPermissionGuardrailConfigModel)
|
||||
|
|
|
|||
|
|
@ -248,6 +248,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
streaming_buffer_until_moderated: bool | None = None,
|
||||
streaming_sampling_rate: int | None = None,
|
||||
streaming_end_of_stream_only: bool | None = None,
|
||||
streaming_buffer_release_on_scan: bool | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
|
||||
|
|
@ -258,6 +259,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
"streaming_buffer_until_moderated": streaming_buffer_until_moderated,
|
||||
"streaming_sampling_rate": streaming_sampling_rate,
|
||||
"streaming_end_of_stream_only": streaming_end_of_stream_only,
|
||||
"streaming_buffer_release_on_scan": streaming_buffer_release_on_scan,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
|
@ -321,13 +323,18 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
self.streaming_buffer_until_moderated = streaming_params.streaming_buffer_until_moderated
|
||||
self.streaming_sampling_rate = streaming_params.streaming_sampling_rate
|
||||
self.streaming_end_of_stream_only = streaming_params.streaming_end_of_stream_only
|
||||
self.streaming_buffer_release_on_scan = streaming_params.streaming_buffer_release_on_scan
|
||||
|
||||
def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
|
||||
super().update_in_memory_litellm_params(litellm_params)
|
||||
self._set_streaming_params(BedrockGuardrailStreamingParams.from_extras(litellm_params.model_extra))
|
||||
|
||||
def _streams_incrementally(self) -> bool:
|
||||
return not self.streaming_buffer_until_moderated and not self.mask_response_content
|
||||
if self.mask_response_content:
|
||||
return False
|
||||
if not self.streaming_buffer_until_moderated:
|
||||
return True
|
||||
return self.streaming_buffer_release_on_scan and not self.streaming_end_of_stream_only
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
fail_on_error=litellm_params.fail_on_error,
|
||||
streaming_buffer_until_moderated=streaming_params.streaming_buffer_until_moderated,
|
||||
streaming_buffer_release_on_scan=streaming_params.streaming_buffer_release_on_scan,
|
||||
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
|
||||
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -260,6 +260,8 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
fail_on_error: bool | None = True,
|
||||
streaming_buffer_until_moderated: bool | None = None,
|
||||
streaming_buffer_release_on_scan: bool | None = None,
|
||||
streaming_end_of_stream_only: bool | None = None,
|
||||
streaming_sampling_rate: int | None = None,
|
||||
async_handler: AsyncHTTPHandler | None = None,
|
||||
|
|
@ -287,6 +289,8 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
CrowdStrikeAIDRGuardrailConfigModelOptionalParams(
|
||||
streaming_end_of_stream_only=streaming_end_of_stream_only,
|
||||
streaming_sampling_rate=streaming_sampling_rate,
|
||||
streaming_buffer_until_moderated=streaming_buffer_until_moderated,
|
||||
streaming_buffer_release_on_scan=streaming_buffer_release_on_scan,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -310,6 +314,8 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
)
|
||||
|
||||
def _set_streaming_params(self, streaming_params: CrowdStrikeAIDRGuardrailConfigModelOptionalParams) -> None:
|
||||
self.streaming_buffer_until_moderated: bool = streaming_params.streaming_buffer_until_moderated or False
|
||||
self.streaming_buffer_release_on_scan: bool = streaming_params.streaming_buffer_release_on_scan or False
|
||||
self.streaming_end_of_stream_only: bool = streaming_params.streaming_end_of_stream_only or False
|
||||
self.streaming_sampling_rate: int = streaming_params.streaming_sampling_rate or 5
|
||||
|
||||
|
|
|
|||
|
|
@ -956,6 +956,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
buffer_until_moderated: bool = _streaming_flag(
|
||||
"streaming_buffer_until_moderated", buffer_until_moderated_default
|
||||
)
|
||||
release_on_scan: Final[bool] = _streaming_flag("streaming_buffer_release_on_scan", False)
|
||||
|
||||
if (
|
||||
buffer_until_moderated
|
||||
|
|
@ -970,9 +971,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
)
|
||||
buffer_until_moderated = False
|
||||
|
||||
# Buffering can only moderate the assembled response, so it always
|
||||
# defers to end-of-stream.
|
||||
if buffer_until_moderated:
|
||||
if buffer_until_moderated and not release_on_scan:
|
||||
end_of_stream_only = True
|
||||
|
||||
if guardrail_to_apply is None:
|
||||
|
|
@ -1026,12 +1025,14 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
chunk_counter = 0
|
||||
responses_so_far: Final[list[object]] = []
|
||||
responses_yielded: Final[list[object]] = []
|
||||
withheld_items: Final[list[object]] = [] # mutable-ok: streaming window must be released incrementally
|
||||
pending_end_of_stream_items: Final[list[object]] = []
|
||||
# Whether any real response chunk has been forwarded to the client.
|
||||
# Drives how a block terminates the stream: continue the in-progress
|
||||
# message (True) vs emit a standalone block message (False, buffered).
|
||||
chunks_yielded = False
|
||||
last_scan_key: StreamingScanKey | None = None # rebind-ok: replaced after every scan round
|
||||
tool_calls_in_flight = False # rebind-ok: tracks the latest scan key's unscanned tool calls
|
||||
|
||||
async for item in response:
|
||||
chunk_counter += 1
|
||||
|
|
@ -1069,21 +1070,37 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
else:
|
||||
withheld_items.append(item)
|
||||
continue
|
||||
|
||||
# Process chunk based on sampling rate
|
||||
if buffer_until_moderated:
|
||||
withheld_items.append(item)
|
||||
if chunk_counter % sampling_rate == 0:
|
||||
endpoint_translation = mappings[CallTypes(call_type)]()
|
||||
scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far)
|
||||
if scan_key is not None:
|
||||
tool_calls_in_flight = scan_key.tool_calls_in_flight
|
||||
hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight)
|
||||
if _is_redundant_scan(scan_key, last_scan_key):
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round",
|
||||
chunk_counter,
|
||||
guardrail_to_apply.guardrail_name,
|
||||
)
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
if buffer_until_moderated:
|
||||
if hold_window:
|
||||
continue
|
||||
for withheld_item in withheld_items:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(withheld_item)
|
||||
yield withheld_item
|
||||
withheld_items.clear()
|
||||
else:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
continue
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -1093,13 +1110,9 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
guardrail_to_apply.guardrail_name,
|
||||
)
|
||||
|
||||
# Deep-copy the current chunk before guardrail processing.
|
||||
# process_output_streaming_response modifies responses_so_far
|
||||
# in-place: it puts the combined guardrailed text in the first
|
||||
# chunk and clears all subsequent chunks to "". Without this
|
||||
# copy, yielding processed_items[-1] would yield an empty
|
||||
# string, permanently losing this chunk's content.
|
||||
original_item = copy.deepcopy(item)
|
||||
original_items = (
|
||||
tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),)
|
||||
)
|
||||
|
||||
try:
|
||||
await endpoint_translation.process_output_streaming_response(
|
||||
|
|
@ -1144,13 +1157,24 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
return
|
||||
if scan_key is not None:
|
||||
last_scan_key = scan_key
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(original_item)
|
||||
yield original_item
|
||||
if hold_window:
|
||||
verbose_proxy_logger.debug(
|
||||
"Holding %s buffered chunks for guardrail %s: this round could not scan the whole window",
|
||||
len(withheld_items),
|
||||
guardrail_to_apply.guardrail_name,
|
||||
)
|
||||
withheld_items[:] = original_items
|
||||
continue
|
||||
for original_item in original_items:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(original_item)
|
||||
yield original_item
|
||||
withheld_items.clear()
|
||||
else:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
if not buffer_until_moderated:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
|
||||
# Stream has ended - do final processing with all collected chunks
|
||||
if call_type is not None and CallTypes(call_type) in mappings:
|
||||
|
|
@ -1162,14 +1186,13 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
endpoint_translation = mappings[CallTypes(call_type)]()
|
||||
|
||||
# When buffering, snapshot the original chunks before moderation.
|
||||
# A shallow copy suffices: end-of-stream
|
||||
# process_output_streaming_response builds a separate assembled
|
||||
# response (it does not mutate the individual chunks in place), and
|
||||
# the chunks themselves are replayed verbatim -- so we only need to
|
||||
# preserve the list, not clone every chunk (deepcopy would double
|
||||
# peak memory for large responses).
|
||||
buffered_items: Final = list(responses_so_far) if buffer_until_moderated else None
|
||||
buffered_items: Final = (
|
||||
tuple(copy.deepcopy(withheld_items))
|
||||
if buffer_until_moderated and release_on_scan and not end_of_stream_only
|
||||
else tuple(withheld_items)
|
||||
if buffer_until_moderated
|
||||
else None
|
||||
)
|
||||
end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far)
|
||||
if _is_redundant_scan(end_scan_key, last_scan_key):
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail):
|
|||
streaming_buffer_until_moderated=streaming_params.streaming_buffer_until_moderated,
|
||||
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
|
||||
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
|
||||
streaming_buffer_release_on_scan=streaming_params.streaming_buffer_release_on_scan,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback)
|
||||
return _bedrock_callback
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.utils import _hash_token_if_needed
|
||||
from litellm.secret_managers.base_secret_manager import BaseSecretManager
|
||||
|
||||
# NOTE: This is the prefix for all virtual keys stored in AWS Secrets Manager
|
||||
LITELLM_PREFIX_STORED_VIRTUAL_KEYS: Final = "litellm/"
|
||||
|
|
@ -100,6 +101,7 @@ class KeyManagementEventHooks:
|
|||
Post /key/update processing hook
|
||||
|
||||
Handles the following:
|
||||
- Renaming the key's secret in the secret manager when the alias changes
|
||||
- Storing Audit Logs for key update
|
||||
"""
|
||||
from litellm.proxy.management_helpers.audit_logs import (
|
||||
|
|
@ -109,6 +111,16 @@ class KeyManagementEventHooks:
|
|||
)
|
||||
from litellm.proxy.proxy_server import litellm_proxy_admin_name
|
||||
|
||||
if data.key_alias is not None and data.key_alias != existing_key_row.key_alias:
|
||||
try:
|
||||
await KeyManagementEventHooks._rename_virtual_key_in_secret_manager(
|
||||
current_secret_name=existing_key_row.key_alias or f"virtual-key-{existing_key_row.token}",
|
||||
new_secret_name=data.key_alias,
|
||||
team_id=existing_key_row.team_id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("Failed to rename virtual key in secret manager: %s", e)
|
||||
|
||||
if is_audit_logging_enabled():
|
||||
updated_fields: Final = {
|
||||
**data.model_dump(exclude_none=True),
|
||||
|
|
@ -153,10 +165,11 @@ class KeyManagementEventHooks:
|
|||
from litellm.proxy.proxy_server import litellm_proxy_admin_name
|
||||
|
||||
# Store the generated key in the secret manager - non-blocking, independent operation
|
||||
if data is not None and response.token_id is not None:
|
||||
if response.token_id is not None:
|
||||
try:
|
||||
initial_secret_name: Final = existing_key_row.key_alias or f"virtual-key-{existing_key_row.token}"
|
||||
new_secret_name: Final = response.key_alias or data.key_alias or initial_secret_name
|
||||
requested_alias: Final = data.key_alias if data is not None else None
|
||||
new_secret_name: Final = response.key_alias or requested_alias or initial_secret_name
|
||||
verbose_proxy_logger.info(
|
||||
"Updating secret in secret manager: secret_name=%s",
|
||||
new_secret_name,
|
||||
|
|
@ -305,21 +318,66 @@ class KeyManagementEventHooks:
|
|||
new_secret_value: New value of the virtual key (example: sk-1234)
|
||||
team_id: Optional team ID to get team-specific secret manager settings
|
||||
"""
|
||||
if litellm._key_management_settings is not None:
|
||||
if litellm._key_management_settings.store_virtual_keys is True:
|
||||
from litellm.secret_managers.base_secret_manager import (
|
||||
BaseSecretManager,
|
||||
)
|
||||
secret_manager: Final = KeyManagementEventHooks._stored_virtual_key_secret_manager()
|
||||
if secret_manager is None:
|
||||
return
|
||||
optional_params: Final = await KeyManagementEventHooks._get_secret_manager_optional_params(team_id)
|
||||
await secret_manager.async_rotate_secret(
|
||||
current_secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name),
|
||||
new_secret_name=KeyManagementEventHooks._get_secret_name(new_secret_name),
|
||||
new_secret_value=new_secret_value,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
# store the key in the secret manager
|
||||
if isinstance(litellm.secret_manager_client, BaseSecretManager):
|
||||
optional_params: Final = await KeyManagementEventHooks._get_secret_manager_optional_params(team_id)
|
||||
await litellm.secret_manager_client.async_rotate_secret(
|
||||
current_secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name),
|
||||
new_secret_name=KeyManagementEventHooks._get_secret_name(new_secret_name),
|
||||
new_secret_value=new_secret_value,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
@staticmethod
|
||||
def _stored_virtual_key_secret_manager() -> BaseSecretManager | None:
|
||||
"""
|
||||
The secret manager client that stores virtual keys, or None when virtual keys are not stored in one
|
||||
"""
|
||||
if litellm._key_management_settings is None or litellm._key_management_settings.store_virtual_keys is not True:
|
||||
return None
|
||||
if not isinstance(litellm.secret_manager_client, BaseSecretManager):
|
||||
return None
|
||||
return litellm.secret_manager_client
|
||||
|
||||
@staticmethod
|
||||
async def _rename_virtual_key_in_secret_manager(
|
||||
current_secret_name: str,
|
||||
new_secret_name: str,
|
||||
team_id: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Move a virtual key to a new secret name, keeping its current value
|
||||
|
||||
Args:
|
||||
current_secret_name: Current name of the virtual key
|
||||
new_secret_name: New name of the virtual key
|
||||
team_id: Optional team ID to get team-specific secret manager settings
|
||||
"""
|
||||
secret_manager: Final = KeyManagementEventHooks._stored_virtual_key_secret_manager()
|
||||
if secret_manager is None:
|
||||
return
|
||||
optional_params: Final = await KeyManagementEventHooks._get_secret_manager_optional_params(team_id)
|
||||
current_secret_value: Final = await secret_manager.async_read_secret(
|
||||
secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name),
|
||||
optional_params=optional_params,
|
||||
)
|
||||
if current_secret_value is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Secret %s not found in secret manager, skipping rename to %s", current_secret_name, new_secret_name
|
||||
)
|
||||
return
|
||||
verbose_proxy_logger.info(
|
||||
"Renaming secret in secret manager: current_secret_name=%s new_secret_name=%s",
|
||||
current_secret_name,
|
||||
new_secret_name,
|
||||
)
|
||||
await secret_manager.async_rotate_secret(
|
||||
current_secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name),
|
||||
new_secret_name=KeyManagementEventHooks._get_secret_name(new_secret_name),
|
||||
new_secret_value=current_secret_value,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_secret_name(secret_name: str) -> str:
|
||||
|
|
|
|||
|
|
@ -19,12 +19,14 @@ from litellm.types.utils import BudgetConfig, StandardLoggingPayload
|
|||
VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX: Final = "virtual_key_spend"
|
||||
END_USER_SPEND_CACHE_KEY_PREFIX: Final = "end_user_model_spend"
|
||||
USER_SPEND_CACHE_KEY_PREFIX: Final = "user_model_spend"
|
||||
TEAM_SPEND_CACHE_KEY_PREFIX: Final = "team_model_spend"
|
||||
|
||||
_SPEND_CACHE_KEY_PREFIXES: Final = MappingProxyType(
|
||||
{
|
||||
Litellm_EntityType.KEY: VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX,
|
||||
Litellm_EntityType.USER: USER_SPEND_CACHE_KEY_PREFIX,
|
||||
Litellm_EntityType.END_USER: END_USER_SPEND_CACHE_KEY_PREFIX,
|
||||
Litellm_EntityType.TEAM: TEAM_SPEND_CACHE_KEY_PREFIX,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -37,6 +39,7 @@ _BUDGET_START_TIME_KEY_PREFIXES: Final = MappingProxyType(
|
|||
Litellm_EntityType.KEY: "virtual_key_budget_start_time",
|
||||
Litellm_EntityType.USER: "user_model_budget_start_time",
|
||||
Litellm_EntityType.END_USER: "end_user_budget_start_time",
|
||||
Litellm_EntityType.TEAM: "team_model_budget_start_time",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -139,6 +142,18 @@ def resolve_model_budget(model: str, model_max_budget: Mapping[str, object]) ->
|
|||
return None
|
||||
|
||||
|
||||
def team_model_budget_applies(model: str, key_model_max_budget: Mapping[str, object] | None) -> bool:
|
||||
"""A key entry that spend-gates `model` overrides the team cap: it is then gated on and billed to the key alone."""
|
||||
if not key_model_max_budget:
|
||||
return True
|
||||
resolved: Final = resolve_model_budget(model=model, model_max_budget=key_model_max_budget)
|
||||
return resolved is None or not _spend_gated(resolved.budget_config)
|
||||
|
||||
|
||||
def _spend_gated(budget_config: BudgetConfig) -> bool:
|
||||
return budget_config.max_budget is not None and budget_config.max_budget >= 0
|
||||
|
||||
|
||||
def _budget_model_candidates(model: str) -> tuple[str, ...]:
|
||||
"""Names a budget may be configured under for a request on `model`, most specific first.
|
||||
|
||||
|
|
@ -346,6 +361,30 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
exceeded_message=f"LiteLLM End User: {end_user_id}, exceeded budget for model={model}",
|
||||
)
|
||||
|
||||
async def is_team_within_model_budget(
|
||||
self,
|
||||
team_id: str,
|
||||
team_model_max_budget: Mapping[str, object],
|
||||
key_model_max_budget: Mapping[str, object] | None,
|
||||
model: str,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if the team is within the model budget, unless the key's own
|
||||
`model_max_budget` overrides it for `model`
|
||||
|
||||
Raises:
|
||||
BudgetExceededError: If the team has exceeded the model budget
|
||||
"""
|
||||
if not team_model_budget_applies(model=model, key_model_max_budget=key_model_max_budget):
|
||||
return True
|
||||
return await self._is_entity_within_model_budget(
|
||||
entity_type=Litellm_EntityType.TEAM,
|
||||
entity_id=team_id,
|
||||
model_max_budget=team_model_max_budget,
|
||||
model=model,
|
||||
exceeded_message=f"LiteLLM Team: {team_id}, exceeded budget for model={model}",
|
||||
)
|
||||
|
||||
async def _is_entity_within_model_budget(
|
||||
self,
|
||||
entity_type: Litellm_EntityType,
|
||||
|
|
@ -456,11 +495,26 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
return
|
||||
|
||||
response_cost: Final[float] = standard_logging_payload.get("response_cost", 0)
|
||||
key_model_max_budget: Final = _metadata.get("user_api_key_model_max_budget")
|
||||
entity_budgets: Final = (
|
||||
(
|
||||
Litellm_EntityType.KEY,
|
||||
payload_metadata.get("user_api_key_hash"),
|
||||
_metadata.get("user_api_key_model_max_budget"),
|
||||
key_model_max_budget,
|
||||
),
|
||||
(
|
||||
Litellm_EntityType.TEAM,
|
||||
payload_metadata.get("user_api_key_team_id"),
|
||||
(
|
||||
_metadata.get("user_api_key_team_model_max_budget")
|
||||
if team_model_budget_applies(
|
||||
model=model,
|
||||
key_model_max_budget=(
|
||||
key_model_max_budget if isinstance(key_model_max_budget, Mapping) else None
|
||||
),
|
||||
)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
Litellm_EntityType.USER,
|
||||
|
|
@ -478,7 +532,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
|
|||
if not resolved_budgets:
|
||||
verbose_proxy_logger.debug(
|
||||
"Not running _PROXY_VirtualKeyModelMaxBudgetLimiter.async_log_success_event: "
|
||||
"no key, user or end-user model_max_budget covers model=%s",
|
||||
"no key, team, user or end-user model_max_budget covers model=%s",
|
||||
model,
|
||||
)
|
||||
return
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.constants import (
|
|||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
OTEL_SERVICE_NAME_METADATA_KEYS,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
SESSION_ID_GENERATED_METADATA_KEY,
|
||||
|
|
@ -369,7 +370,13 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg
|
|||
# and read by spend logs as fact; a client value has no legitimate meaning and no
|
||||
# key or team setting keeps it, so the strip is never gated.
|
||||
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset(
|
||||
{"attempted_fallbacks", "original_model_group", "request_retry_count", CLIENT_OUTPUT_CEILING_METADATA_KEY}
|
||||
{
|
||||
"attempted_fallbacks",
|
||||
"original_model_group",
|
||||
"request_retry_count",
|
||||
CLIENT_OUTPUT_CEILING_METADATA_KEY,
|
||||
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY,
|
||||
}
|
||||
)
|
||||
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"
|
||||
|
||||
|
|
@ -2327,6 +2334,7 @@ async def add_litellm_data_to_request(
|
|||
# Team spend, budget - used by prometheus.py
|
||||
data[_metadata_variable_name]["user_api_key_team_max_budget"] = user_api_key_dict.team_max_budget
|
||||
data[_metadata_variable_name]["user_api_key_team_spend"] = user_api_key_dict.team_spend
|
||||
data[_metadata_variable_name]["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget
|
||||
data[_metadata_variable_name]["user_api_key_request_route"] = user_api_key_dict.request_route
|
||||
|
||||
# API Key spend, budget - used by prometheus.py
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ def validate_budget_duration(budget_duration: str | None, status_code: int = 400
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import (
|
||||
CommonProxyErrors,
|
||||
KeyRequestBase,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
|
|
@ -73,12 +74,62 @@ from litellm.proxy._types import ( # noqa: F401 re-exported
|
|||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.utils import _premium_user_check
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.types.utils import BudgetConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import NewProjectRequest, UpdateProjectRequest
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
|
||||
def validate_team_model_max_budget(
|
||||
model_max_budget: Mapping[str, BudgetConfig] | None,
|
||||
premium_user: bool,
|
||||
) -> None:
|
||||
"""Reject a team `model_max_budget` the limiter could not enforce (no duration, bad cap, tpm/rpm limits)."""
|
||||
if not model_max_budget:
|
||||
return
|
||||
if premium_user is not True:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": f"Setting model_max_budget on a team is an enterprise feature. {CommonProxyErrors.not_premium_user.value}"
|
||||
},
|
||||
)
|
||||
for model_name, budget_config in model_max_budget.items():
|
||||
if not model_name.strip():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "model_max_budget keys must be non-empty model names"},
|
||||
)
|
||||
max_budget = budget_config.max_budget
|
||||
if max_budget is None or not math.isfinite(max_budget) or max_budget < 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
f"model_max_budget[{model_name!r}].max_budget must be a non-negative finite number. "
|
||||
f"Received: {max_budget}"
|
||||
)
|
||||
},
|
||||
)
|
||||
if budget_config.budget_duration is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"model_max_budget[{model_name!r}] requires a budget_duration, e.g. '1d' or '30d'"},
|
||||
)
|
||||
validate_budget_duration(budget_config.budget_duration)
|
||||
if budget_config.tpm_limit is not None or budget_config.rpm_limit is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
f"model_max_budget[{model_name!r}] tpm_limit/rpm_limit are not enforced on a team; "
|
||||
"set per-model rate limits on the key instead"
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def require_caller_user_id_for_non_admin(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> str:
|
||||
|
|
|
|||
|
|
@ -4166,7 +4166,10 @@ async def info_key_fn(
|
|||
|
||||
Returns:
|
||||
- key: str - The key that was looked up, echoed back as it was passed in
|
||||
- info: dict - The key's row, minus the hashed token
|
||||
- info: dict - The key's row, minus the hashed token. Deleted keys are served from the
|
||||
LiteLLM_DeletedVerificationToken archive and carry deleted_at / deleted_by
|
||||
- status: "active" | "expired" | "revoked" | "deleted" - Derived from blocked, expires and
|
||||
whether the row came from the archive
|
||||
- key_alias: str | None - User-friendly key alias
|
||||
- spend: float - Amount spent by the key. When budget_duration is set this covers only the
|
||||
current budget window, not the key's lifetime
|
||||
|
|
@ -4220,10 +4223,15 @@ async def info_key_fn(
|
|||
hashed_key: str | None = key
|
||||
if key is not None:
|
||||
hashed_key = _hash_token_if_needed(token=key)
|
||||
key_info = await _prisma_table(VerificationTokenRepository(prisma_client)).find_unique(
|
||||
live_key_info: Final = await _prisma_table(VerificationTokenRepository(prisma_client)).find_unique(
|
||||
where={"token": hashed_key},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
key_info: Final = (
|
||||
live_key_info
|
||||
if live_key_info is not None
|
||||
else await _find_deleted_key_info(prisma_client=prisma_client, hashed_key=hashed_key)
|
||||
)
|
||||
if key_info is None:
|
||||
raise ProxyException(
|
||||
message="Key not found in database",
|
||||
|
|
@ -4231,7 +4239,6 @@ async def info_key_fn(
|
|||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
if (
|
||||
await _can_user_query_key_info(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -4245,38 +4252,46 @@ async def info_key_fn(
|
|||
detail=f"You are not allowed to access this key's info. Your role={user_api_key_dict.user_role}",
|
||||
)
|
||||
## REMOVE HASHED TOKEN INFO BEFORE RETURNING ##
|
||||
try:
|
||||
key_info = key_info.model_dump()
|
||||
except Exception:
|
||||
# if using pydantic v1
|
||||
key_info = key_info.dict() # pyright: ignore[reportDeprecated] # deliberate pydantic v1 fallback
|
||||
key_token_hash: Final[str | None] = key_info.pop("token")
|
||||
key_info_dict: Final = key_info.model_dump()
|
||||
key_token_hash: Final[str | None] = key_info_dict.pop("token")
|
||||
key_info_dict["status"] = (
|
||||
"deleted" if live_key_info is None else _derive_key_status(key_info_dict, now=datetime.now(timezone.utc))
|
||||
)
|
||||
|
||||
model_max_budget = key_info.get("model_max_budget") or {}
|
||||
budget_table: Final = key_info.get("litellm_budget_table") or {}
|
||||
model_max_budget = key_info_dict.get("model_max_budget") or {}
|
||||
budget_table: Final = key_info_dict.get("litellm_budget_table") or {}
|
||||
if not model_max_budget and isinstance(budget_table, dict):
|
||||
model_max_budget = budget_table.get("model_max_budget") or {}
|
||||
if model_max_budget and key_token_hash:
|
||||
key_info["model_max_budget_usage"] = await _build_model_max_budget_usage(
|
||||
key_info_dict["model_max_budget_usage"] = await _build_model_max_budget_usage(
|
||||
api_key_hash=key_token_hash,
|
||||
model_max_budget=model_max_budget,
|
||||
user_api_key_cache=model_max_budget_limiter.dual_cache,
|
||||
)
|
||||
budget_limits_usage: Final = await _build_budget_limits_usage(
|
||||
budget_limits=key_info.get("budget_limits"),
|
||||
budget_limits=key_info_dict.get("budget_limits"),
|
||||
api_key_hash=key_token_hash,
|
||||
)
|
||||
if budget_limits_usage is not None:
|
||||
key_info["budget_limits_usage"] = budget_limits_usage
|
||||
key_info_dict["budget_limits_usage"] = budget_limits_usage
|
||||
|
||||
# Attach object_permission if object_permission_id is set
|
||||
key_info = await attach_object_permission_to_dict(key_info, prisma_client)
|
||||
|
||||
return {"key": key, "info": key_info}
|
||||
return {"key": key, "info": await attach_object_permission_to_dict(key_info_dict, prisma_client)}
|
||||
except Exception as e:
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
|
||||
async def _find_deleted_key_info(
|
||||
prisma_client: PrismaClient, hashed_key: str | None
|
||||
) -> LiteLLM_DeletedVerificationToken | None:
|
||||
archived_row: Final = await _deleted_verification_token_table(prisma_client).find_first(
|
||||
where={"token": hashed_key},
|
||||
order={"deleted_at": "desc"},
|
||||
)
|
||||
if archived_row is None:
|
||||
return None
|
||||
return LiteLLM_DeletedVerificationToken.model_validate(archived_row.model_dump())
|
||||
|
||||
|
||||
def _check_model_access_group(models: list[str] | None, llm_router: Router | None, premium_user: bool) -> Literal[True]:
|
||||
"""
|
||||
if is_model_access_group is True + is_wildcard_route is True, check if user is a premium user
|
||||
|
|
@ -6216,6 +6231,24 @@ async def get_member_team_ids(
|
|||
|
||||
VALID_EXPIRES_FILTER_VALUES: Final = frozenset({"active", "expired"})
|
||||
|
||||
KeyStatus = Literal["active", "expired", "revoked", "deleted"]
|
||||
VALID_STATUS_FILTER_VALUES: Final[frozenset[KeyStatus]] = frozenset({"active", "expired", "revoked", "deleted"})
|
||||
|
||||
|
||||
class _KeyStatusSource(BaseModel):
|
||||
blocked: bool | None = None
|
||||
expires: datetime | None = None
|
||||
|
||||
|
||||
def _derive_key_status(row: Mapping[str, object], now: datetime) -> KeyStatus:
|
||||
source: Final = _KeyStatusSource.model_validate(row)
|
||||
if source.blocked is True:
|
||||
return "revoked"
|
||||
if source.expires is None:
|
||||
return "active"
|
||||
expires_utc: Final = source.expires if source.expires.tzinfo else source.expires.replace(tzinfo=timezone.utc)
|
||||
return "expired" if expires_utc < now else "active"
|
||||
|
||||
|
||||
@router.get(
|
||||
"/key/list",
|
||||
|
|
@ -6252,7 +6285,10 @@ async def list_keys(
|
|||
),
|
||||
sort_order: str = Query(default="desc", description="Sort order ('asc' or 'desc')"),
|
||||
expand: list[str] | None = Query(None, description="Expand related objects (e.g. 'user')"),
|
||||
status: str | None = Query(None, description="Filter by status (e.g. 'deleted')"),
|
||||
status: str | None = Query(
|
||||
None,
|
||||
description="Filter by status: 'active' (not blocked, not expired), 'expired' (not blocked, past expiry), 'revoked' (blocked) or 'deleted' (archived keys). Omit to return live keys regardless of status.",
|
||||
),
|
||||
project_id: str | None = Query(None, description="Filter keys by project ID"),
|
||||
access_group_id: str | None = Query(None, description="Filter keys by access group ID"),
|
||||
agent_id: str | None = Query(None, description="Filter keys by agent ID"),
|
||||
|
|
@ -6270,7 +6306,9 @@ async def list_keys(
|
|||
|
||||
Parameters:
|
||||
expand: Optional[List[str]] - Expand related objects (e.g. 'user' to include user information)
|
||||
status: Optional[str] - Filter by status. Currently supports "deleted" to query deleted keys.
|
||||
status: Optional[str] - Filter by status: "active", "expired", "revoked" (blocked) or "deleted".
|
||||
"deleted" reads the LiteLLM_DeletedVerificationToken archive; the other values partition the
|
||||
live key table, so every live key matches exactly one of them.
|
||||
|
||||
Returns:
|
||||
{
|
||||
|
|
@ -6292,11 +6330,10 @@ async def list_keys(
|
|||
verbose_proxy_logger.error("Database not connected")
|
||||
raise Exception("Database not connected")
|
||||
|
||||
# Validate status parameter
|
||||
if status is not None and status != "deleted":
|
||||
if status is not None and status not in VALID_STATUS_FILTER_VALUES:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": "Invalid status value. Currently only 'deleted' is supported."},
|
||||
detail={"error": "Invalid status value. Supported: 'active', 'expired', 'revoked', 'deleted'."},
|
||||
)
|
||||
|
||||
if isinstance(expires, str) and expires not in VALID_EXPIRES_FILTER_VALUES:
|
||||
|
|
@ -6608,6 +6645,18 @@ def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str,
|
|||
return {"OR": [{"expires": None}, {"expires": {"gte": now}}]}
|
||||
|
||||
|
||||
def _not_blocked_where_clause() -> dict[str, object]:
|
||||
return {"OR": [{"blocked": None}, {"blocked": False}]}
|
||||
|
||||
|
||||
def _build_status_where_clause(status_filter: str | None, now: datetime) -> dict[str, object] | None:
|
||||
if status_filter == "revoked":
|
||||
return {"blocked": True}
|
||||
if status_filter in ("expired", "active"):
|
||||
return {"AND": [_not_blocked_where_clause(), _build_expires_where_clause(status_filter, now)]}
|
||||
return None
|
||||
|
||||
|
||||
def _build_key_search_where(search: str) -> KeySearchWhere:
|
||||
search_where: Final[KeySearchWhere] = {
|
||||
"OR": (
|
||||
|
|
@ -6635,6 +6684,7 @@ def _build_key_filter_conditions(
|
|||
use_key_alias_substring_matching: bool = False,
|
||||
expires_filter: str | None = None,
|
||||
search: str | None = None,
|
||||
status_filter: str | None = None,
|
||||
) -> Mapping[str, object]:
|
||||
"""Build filter conditions for key listing.
|
||||
|
||||
|
|
@ -6724,6 +6774,8 @@ def _build_key_filter_conditions(
|
|||
|
||||
# Apply team_id, project_id and access_group_id as global AND filters so they
|
||||
# narrow results across all visibility conditions (own keys, team keys, etc.)
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
status_where: Final = _build_status_where_clause(status_filter, now)
|
||||
global_filters: Final[tuple[Mapping[str, object], ...]] = (
|
||||
*(
|
||||
(
|
||||
|
|
@ -6741,10 +6793,11 @@ def _build_key_filter_conditions(
|
|||
*(({"access_group_ids": {"hasSome": [access_group_id]}},) if access_group_id else ()),
|
||||
*(({"agent_id": agent_id},) if agent_id and isinstance(agent_id, str) else ()),
|
||||
*(
|
||||
(_build_expires_where_clause(expires_filter, datetime.now(timezone.utc)),)
|
||||
(_build_expires_where_clause(expires_filter, now),)
|
||||
if expires_filter is not None and expires_filter in VALID_EXPIRES_FILTER_VALUES
|
||||
else ()
|
||||
),
|
||||
*((status_where,) if status_where is not None else ()),
|
||||
)
|
||||
combined_where: Final[Mapping[str, object]] = {"AND": [where, *global_filters]} if global_filters else where
|
||||
verbose_proxy_logger.debug("Filter conditions: %s", combined_where)
|
||||
|
|
@ -6817,6 +6870,7 @@ async def _list_key_helper(
|
|||
use_key_alias_substring_matching=use_key_alias_substring_matching,
|
||||
expires_filter=expires_filter,
|
||||
search=search,
|
||||
status_filter=status,
|
||||
)
|
||||
|
||||
# Calculate skip for pagination
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ from typing import TYPE_CHECKING, Annotated, Final, NamedTuple, NoReturn, Protoc
|
|||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
|
||||
from pydantic import BaseModel, JsonValue
|
||||
from pydantic import BaseModel, JsonValue, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
|
|
@ -38,6 +38,7 @@ from litellm.proxy._types import (
|
|||
DeleteTeamRequest,
|
||||
LiteLLM_AuditLogs,
|
||||
LiteLLM_DeletedTeamTable,
|
||||
Litellm_EntityType,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
LiteLLM_ModelTable,
|
||||
|
|
@ -95,6 +96,10 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.hooks.model_max_budget_limiter import (
|
||||
build_model_max_budget_usage,
|
||||
resolve_model_budget,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import (
|
||||
get_daily_activity_aggregated,
|
||||
)
|
||||
|
|
@ -108,6 +113,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
_upsert_budget_and_membership,
|
||||
_user_has_admin_view,
|
||||
validate_budget_duration,
|
||||
validate_team_model_max_budget,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
add_member_to_organization,
|
||||
|
|
@ -177,6 +183,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
TeamUserSpendRow,
|
||||
UpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
from litellm.types.utils import BudgetConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
|
|
@ -1170,6 +1177,62 @@ def _check_team_budget_update_authority(
|
|||
)
|
||||
|
||||
|
||||
def _existing_model_cap(raw_budget_config: object) -> BudgetConfig | None:
|
||||
try:
|
||||
return BudgetConfig.model_validate(raw_budget_config)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _check_team_model_budget_update_authority(
|
||||
data: UpdateTeamRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
existing_model_max_budget: Mapping[str, object] | None,
|
||||
) -> None:
|
||||
"""Like `_check_team_budget_update_authority`: only a proxy admin may raise, re-window or drop a per-model cap."""
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
|
||||
return
|
||||
if "model_max_budget" not in data.model_fields_set or not existing_model_max_budget:
|
||||
return
|
||||
requested: Final[Mapping[str, BudgetConfig]] = data.model_max_budget or {}
|
||||
for model_name, raw_existing in existing_model_max_budget.items():
|
||||
existing = _existing_model_cap(raw_existing)
|
||||
if existing is None or existing.max_budget is None or model_name in requested:
|
||||
continue
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": (
|
||||
f"Only a proxy admin can remove a team's model_max_budget for {model_name!r}. "
|
||||
f"Current max_budget={existing.max_budget}."
|
||||
)
|
||||
},
|
||||
)
|
||||
for model_name, proposed in requested.items():
|
||||
governing = resolve_model_budget(model=model_name, model_max_budget=existing_model_max_budget)
|
||||
if governing is None:
|
||||
continue
|
||||
cap = governing.budget_config
|
||||
if cap.max_budget is None:
|
||||
continue
|
||||
if (
|
||||
proposed.max_budget is None
|
||||
or proposed.max_budget > cap.max_budget
|
||||
or proposed.budget_duration != cap.budget_duration
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": (
|
||||
f"Only a proxy admin can raise a team's model_max_budget for {model_name!r} or change its "
|
||||
f"budget_duration. Current max_budget={cap.max_budget} per {cap.budget_duration} "
|
||||
f"(entry {governing.budget_model!r}), requested={proposed.max_budget} per "
|
||||
f"{proposed.budget_duration}."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _should_auto_add_team_creator(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
general_settings: Mapping[str, object],
|
||||
|
|
@ -1230,6 +1293,7 @@ async def new_team(
|
|||
- prompts: Optional[List[str]] - List of prompts that the team is allowed to use.
|
||||
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
|
||||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- model_max_budget: Optional[dict] - Per-model max budget every key on the team inherits unless the key sets its own for that model. Example: {"gpt-4o": {"max_budget": 10, "budget_duration": "1d"}}
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies)
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
|
|
@ -1291,6 +1355,7 @@ async def new_team(
|
|||
general_settings,
|
||||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
|
@ -1321,6 +1386,7 @@ async def new_team(
|
|||
|
||||
validate_budget_duration(data.budget_duration)
|
||||
validate_budget_duration(data.team_member_budget_duration)
|
||||
validate_team_model_max_budget(model_max_budget=data.model_max_budget, premium_user=premium_user)
|
||||
|
||||
if data.soft_budget is not None:
|
||||
if data.max_budget is not None:
|
||||
|
|
@ -1980,6 +2046,7 @@ async def update_team(
|
|||
- tags: Optional[List[str]] - Tags for [tracking spend](https://litellm.vercel.app/docs/proxy/enterprise#tracking-spend-for-custom-tags) and/or doing [tag-based routing](https://litellm.vercel.app/docs/proxy/tag_routing).
|
||||
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
|
||||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- model_max_budget: Optional[dict] - Per-model max budget every key on the team inherits unless the key sets its own for that model. Example: {"gpt-4o": {"max_budget": 10, "budget_duration": "1d"}}
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies)
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
|
|
@ -2031,6 +2098,7 @@ async def update_team(
|
|||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
|
|
@ -2069,6 +2137,7 @@ async def update_team(
|
|||
|
||||
validate_budget_duration(data.budget_duration)
|
||||
validate_budget_duration(data.team_member_budget_duration)
|
||||
validate_team_model_max_budget(model_max_budget=data.model_max_budget, premium_user=premium_user)
|
||||
|
||||
existing_team_row = await _raw_team_db(TeamRepository(prisma_client)).find_unique(
|
||||
where={"team_id": data.team_id}
|
||||
|
|
@ -2204,8 +2273,15 @@ async def update_team(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
existing_team_max_budget=existing_team_row.max_budget,
|
||||
)
|
||||
_check_team_model_budget_update_authority(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
existing_model_max_budget=existing_team_row.model_max_budget,
|
||||
)
|
||||
|
||||
updated_kv = data.json(exclude_unset=True)
|
||||
if "model_max_budget" in updated_kv and updated_kv["model_max_budget"] is None:
|
||||
updated_kv["model_max_budget"] = {}
|
||||
|
||||
# Drop server-owned metadata keys from caller input so they can only
|
||||
# be written by the same code path that creates the underlying rows.
|
||||
|
|
@ -4473,7 +4549,7 @@ async def team_info(
|
|||
```
|
||||
"""
|
||||
from litellm.proxy._types import TeamInfoResponseObjectTeamTable
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import model_max_budget_limiter, prisma_client
|
||||
|
||||
try:
|
||||
if prisma_client is None:
|
||||
|
|
@ -4573,6 +4649,12 @@ async def team_info(
|
|||
update={ # mutable-ok: pydantic update payload
|
||||
"members_with_roles": hydrated_members,
|
||||
"organization_models": organization_models,
|
||||
"model_max_budget_usage": await build_model_max_budget_usage(
|
||||
entity_type=Litellm_EntityType.TEAM,
|
||||
entity_id=team_id,
|
||||
model_max_budget=resolved_team_info.model_max_budget,
|
||||
cache=model_max_budget_limiter.dual_cache,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.proxy.auth.handle_jwt import JWTHandler
|
|||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_get_bearer_token,
|
||||
is_no_auth_dev_mode,
|
||||
user_api_key_auth,
|
||||
user_api_key_auth_websocket,
|
||||
)
|
||||
|
|
@ -709,8 +710,7 @@ async def anthropic_proxy_route(
|
|||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=str(updated_url),
|
||||
custom_headers=auth_header if auth_header is not None else {},
|
||||
_forward_headers=True,
|
||||
custom_headers=_upstream_headers_for_anthropic_route(request, user_api_key_dict, auth_header),
|
||||
is_streaming_request=is_streaming_request,
|
||||
) # dynamically construct pass-through endpoint based on incoming path
|
||||
received_value: Final = await endpoint_func(
|
||||
|
|
@ -1989,6 +1989,19 @@ _HEADERS_NEVER_FORWARDED_TO_VERTEX: Final = frozenset({"content-length", "host"}
|
|||
SpecialHeaders.litellm_credential_header_names() - _VERTEX_UPSTREAM_CREDENTIAL_HEADERS
|
||||
)
|
||||
|
||||
_CREDENTIALLESS_ANTHROPIC_MISSING_CREDENTIAL_DETAIL: Final = (
|
||||
"No Anthropic credential is configured on this proxy and the request carried no upstream "
|
||||
"Anthropic credential. The LiteLLM virtual key is not forwarded to Anthropic. Configure an "
|
||||
"Anthropic credential (ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN, or a model with "
|
||||
"use_in_pass_through: true), or send your own Anthropic API key in the x-api-key header or "
|
||||
"your own Anthropic OAuth token in the Authorization header."
|
||||
)
|
||||
|
||||
_ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS: Final = frozenset({"authorization", "x-api-key"})
|
||||
_HEADERS_NEVER_FORWARDED_TO_ANTHROPIC: Final = frozenset({"content-length", "host", "accept-encoding"}) | (
|
||||
SpecialHeaders.litellm_credential_header_names() - _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS
|
||||
)
|
||||
|
||||
|
||||
_MAPPED_ROUTE_CALLER_KEY_HEADER: Final = "litellm_user_api_key"
|
||||
|
||||
|
|
@ -2026,8 +2039,11 @@ def _is_authenticated_caller_jwt(value: str, jwt_claims: Mapping[str, object]) -
|
|||
|
||||
|
||||
def _is_authenticated_caller_secret(value: str, user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
"""Whether a header value is the master key, the JWT that authenticated, or the key stored as ``api_key``."""
|
||||
from litellm.proxy.proxy_server import master_key
|
||||
"""Whether a header value is the master key, the JWT that authenticated, or the key stored as ``api_key``.
|
||||
|
||||
A proxy in no-auth dev mode without custom auth authenticated nothing, so none of the caller's values is one.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import general_settings, master_key, user_custom_auth
|
||||
|
||||
normalized: Final = _normalize_credential_value(value)
|
||||
if master_key is not None and hmac.compare_digest(normalized.encode(), master_key.encode()):
|
||||
|
|
@ -2035,35 +2051,54 @@ def _is_authenticated_caller_secret(value: str, user_api_key_dict: UserAPIKeyAut
|
|||
jwt_claims: Final = user_api_key_dict.jwt_claims
|
||||
if jwt_claims and _is_authenticated_caller_jwt(normalized, jwt_claims):
|
||||
return True
|
||||
if is_no_auth_dev_mode(master_key, general_settings) and user_custom_auth is None:
|
||||
return False
|
||||
authenticated_key: Final = user_api_key_dict.api_key
|
||||
if authenticated_key is None:
|
||||
return False
|
||||
if master_key is None and not normalized.startswith("sk-"):
|
||||
return False
|
||||
stored_representation: Final = UserAPIKeyAuth._safe_hash_litellm_api_key(normalized) # pyright: ignore[reportPrivateUsage] # the exact transform auth applied when it stored api_key
|
||||
return hmac.compare_digest(stored_representation.encode(), authenticated_key.encode())
|
||||
|
||||
|
||||
def _caller_headers_without_litellm_secrets(
|
||||
request: Request, user_api_key_dict: UserAPIKeyAuth, never_forwarded: frozenset[str]
|
||||
) -> Mapping[str, str]:
|
||||
incoming: Final = _safe_get_request_headers(request)
|
||||
dropped_by_name: Final = never_forwarded.union(
|
||||
(_MAPPED_ROUTE_CALLER_KEY_HEADER, *_operator_configured_caller_key_header_names())
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
name: value
|
||||
for name, value in incoming.items()
|
||||
if name not in dropped_by_name and not _is_authenticated_caller_secret(value, user_api_key_dict)
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _forwarded_headers_for_credentialless_vertex_passthrough(
|
||||
request: Request, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> Mapping[str, str]:
|
||||
"""Caller headers to forward on the bring-your-own-credentials Vertex branch, minus LiteLLM secrets."""
|
||||
incoming: Final = _safe_get_request_headers(request)
|
||||
never_forwarded: Final = _HEADERS_NEVER_FORWARDED_TO_VERTEX.union(
|
||||
(_MAPPED_ROUTE_CALLER_KEY_HEADER, *_operator_configured_caller_key_header_names())
|
||||
forwarded: Final = _caller_headers_without_litellm_secrets(
|
||||
request, user_api_key_dict, _HEADERS_NEVER_FORWARDED_TO_VERTEX
|
||||
)
|
||||
forwarded: Final = MappingProxyType(
|
||||
{
|
||||
name: value
|
||||
for name, value in incoming.items()
|
||||
if name not in never_forwarded and not _is_authenticated_caller_secret(value, user_api_key_dict)
|
||||
}
|
||||
)
|
||||
if "authorization" not in forwarded and "x-goog-api-key" not in forwarded:
|
||||
if _VERTEX_UPSTREAM_CREDENTIAL_HEADERS.isdisjoint(forwarded):
|
||||
raise HTTPException(status_code=401, detail=_CREDENTIALLESS_VERTEX_MISSING_CREDENTIAL_DETAIL)
|
||||
return forwarded
|
||||
|
||||
|
||||
def _upstream_headers_for_anthropic_route(
|
||||
request: Request, user_api_key_dict: UserAPIKeyAuth, proxy_auth_header: Mapping[str, str] | None
|
||||
) -> Mapping[str, str]:
|
||||
caller_headers: Final = _caller_headers_without_litellm_secrets(
|
||||
request, user_api_key_dict, _HEADERS_NEVER_FORWARDED_TO_ANTHROPIC
|
||||
)
|
||||
if proxy_auth_header is None and _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS.isdisjoint(caller_headers):
|
||||
raise HTTPException(status_code=401, detail=_CREDENTIALLESS_ANTHROPIC_MISSING_CREDENTIAL_DETAIL)
|
||||
return MappingProxyType({**caller_headers, **(proxy_auth_header or {})})
|
||||
|
||||
|
||||
async def _prepare_vertex_auth_headers(
|
||||
request: Request,
|
||||
vertex_credentials: VertexPassThroughCredentials | None,
|
||||
|
|
|
|||
|
|
@ -609,6 +609,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
# merely shares the name.
|
||||
if not request_dispatched_to_pass_through_endpoint(request):
|
||||
_metadata["user_api_key_model_max_budget"] = user_api_key_dict.model_max_budget
|
||||
_metadata["user_api_key_team_model_max_budget"] = user_api_key_dict.team_model_max_budget
|
||||
_metadata["user_api_key_user_model_max_budget"] = user_api_key_dict.user_model_max_budget
|
||||
_metadata["user_api_key_end_user_model_max_budget"] = user_api_key_dict.end_user_model_max_budget
|
||||
_metadata.update(
|
||||
|
|
|
|||
|
|
@ -323,6 +323,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
log_once_if_budget_reservation_disabled,
|
||||
warn_once_if_custom_auth_skips_common_checks,
|
||||
)
|
||||
from litellm.proxy.auth.fallback_budget import router_fallback_budget_check
|
||||
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_REMEDY, LicenseCheck
|
||||
|
|
@ -6161,6 +6162,7 @@ class ProxyConfig:
|
|||
),
|
||||
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
fallback_budget_check=router_fallback_budget_check,
|
||||
auto_router_capability_limit=_license_check.auto_router_capability_limit,
|
||||
)
|
||||
|
||||
|
|
@ -6622,6 +6624,7 @@ class ProxyConfig:
|
|||
search_tools=search_tools,
|
||||
ignore_invalid_deployments=True,
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
fallback_budget_check=router_fallback_budget_check,
|
||||
auto_router_capability_limit=_license_check.auto_router_capability_limit,
|
||||
)
|
||||
verbose_proxy_logger.debug("updated llm_router: %s", llm_router)
|
||||
|
|
|
|||
|
|
@ -426,6 +426,7 @@ model LiteLLM_VerificationToken {
|
|||
key_alias String?
|
||||
soft_budget_cooldown Boolean @default(false) // key-level state on if budget alerts need to be cooled down
|
||||
spend Float @default(0.0)
|
||||
total_spend Float @default(0.0)
|
||||
expires DateTime?
|
||||
models String[]
|
||||
aliases Json @default("{}")
|
||||
|
|
@ -528,6 +529,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
key_alias String?
|
||||
soft_budget_cooldown Boolean @default(false)
|
||||
spend Float @default(0.0)
|
||||
total_spend Float @default(0.0)
|
||||
expires DateTime?
|
||||
models String[]
|
||||
aliases Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ def carry_team_and_user_budget_state(
|
|||
budget_reset_at=team_object.budget_reset_at,
|
||||
max_budget=team_object.max_budget,
|
||||
)
|
||||
valid_token.team_model_max_budget = team_object.model_max_budget # rebind-ok: caller keeps this object
|
||||
if user_object is not None:
|
||||
valid_token.user_budget_snapshot = UserBudgetSnapshot( # rebind-ok: same object the caller keeps using
|
||||
budget_reset_at=user_object.budget_reset_at,
|
||||
|
|
|
|||
|
|
@ -4340,6 +4340,7 @@ class PrismaClient:
|
|||
v.*,
|
||||
t.spend AS team_spend,
|
||||
t.max_budget AS team_max_budget,
|
||||
t.model_max_budget AS team_model_max_budget,
|
||||
t.tpm_limit AS team_tpm_limit,
|
||||
t.rpm_limit AS team_rpm_limit,
|
||||
t.tpd_limit AS team_tpd_limit
|
||||
|
|
@ -4779,6 +4780,7 @@ class PrismaClient:
|
|||
t.spend AS team_spend,
|
||||
t.max_budget AS team_max_budget,
|
||||
t.soft_budget AS team_soft_budget,
|
||||
t.model_max_budget AS team_model_max_budget,
|
||||
t.tpm_limit AS team_tpm_limit,
|
||||
t.rpm_limit AS team_rpm_limit,
|
||||
t.tpd_limit AS team_tpd_limit,
|
||||
|
|
|
|||
|
|
@ -63,6 +63,7 @@ from litellm.constants import (
|
|||
DEFAULT_MAX_LRU_CACHE_SIZE,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
OUTPUT_TOKEN_CEILING_PARAMS,
|
||||
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY,
|
||||
ROUTING_REQUEST_TAGS_METADATA_KEY,
|
||||
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
|
||||
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
|
||||
|
|
@ -132,7 +133,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
get_hidden_params_dict,
|
||||
prepare_response_for_header_attachment,
|
||||
replace_complexity_router_headers,
|
||||
response_in_flight_token_count,
|
||||
response_total_token_count,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
AUTO_ROUTER_MODEL_PREFIX,
|
||||
|
|
@ -215,6 +216,8 @@ from litellm.router_utils.reasoning_effort_capability import (
|
|||
resolve_supported_reasoning_efforts,
|
||||
)
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import (
|
||||
find_deployment_metadata,
|
||||
get_counted_usage_tokens,
|
||||
increment_deployment_failures_for_current_minute,
|
||||
increment_deployment_successes_for_current_minute,
|
||||
)
|
||||
|
|
@ -240,6 +243,7 @@ from litellm.types.router import (
|
|||
DeploymentModelListingInfo,
|
||||
DeploymentTypedDict,
|
||||
FallbackAccessCheck,
|
||||
FallbackBudgetCheck,
|
||||
GuardrailTypedDict,
|
||||
LiteLLM_Params,
|
||||
MockRouterTestingParams,
|
||||
|
|
@ -777,6 +781,7 @@ class Router:
|
|||
background_health_check_model_groups: Sequence[str] | None = None,
|
||||
enable_weighted_failover: bool = False,
|
||||
fallback_access_check: FallbackAccessCheck | None = None,
|
||||
fallback_budget_check: FallbackBudgetCheck | None = None,
|
||||
auto_router_capability_limit: AutoRouterCapabilityLimit | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -815,6 +820,7 @@ class Router:
|
|||
ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error.
|
||||
enable_weighted_failover (bool): When True and the routing strategy is "simple-shuffle", a retryable failure on one deployment causes the request to re-pick (weighted) across the other deployments in the same model group before any cross-group fallback runs. Bounded by `max_fallbacks`. Async-only: currently honored by `router.acompletion()` and other async entrypoints. The sync `router.completion()` path falls back to the regular fallback flow. Defaults to False.
|
||||
fallback_access_check (Optional[FallbackAccessCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects is skipped. Defaults to None (every configured fallback is attempted).
|
||||
fallback_budget_check (Optional[FallbackBudgetCheck]): Awaited before each cross-model-group fallback attempt on the async path; a fallback target it rejects as over budget is skipped. Defaults to None (budget is not re-checked on fallback).
|
||||
Returns:
|
||||
Router: An instance of the litellm.Router class.
|
||||
|
||||
|
|
@ -856,6 +862,7 @@ class Router:
|
|||
self.ignore_invalid_deployments = ignore_invalid_deployments
|
||||
self.auto_router_capability_limit = auto_router_capability_limit
|
||||
self.fallback_access_check: Final = fallback_access_check
|
||||
self.fallback_budget_check: Final = fallback_budget_check
|
||||
self.debug_level = debug_level
|
||||
self.enable_pre_call_checks = enable_pre_call_checks
|
||||
self.enable_tag_filtering = enable_tag_filtering
|
||||
|
|
@ -7937,6 +7944,7 @@ class Router:
|
|||
response = original_function(*args, **kwargs)
|
||||
if coroutine_checker.is_async_callable(response) or inspect.isawaitable(response):
|
||||
response = await response
|
||||
await self.increment_deployment_usage_for_response(response=response, request_kwargs=kwargs)
|
||||
## PROCESS RESPONSE HEADERS
|
||||
response = await self.set_response_headers(response=response, model_group=model_group, request_kwargs=kwargs)
|
||||
|
||||
|
|
@ -8153,8 +8161,6 @@ class Router:
|
|||
"""
|
||||
Track remaining tpm/rpm quota for model in model_list
|
||||
"""
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
||||
try:
|
||||
# WS session wrappers fire with result=None; per-turn costs tracked by inner calls.
|
||||
if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"):
|
||||
|
|
@ -8162,114 +8168,135 @@ class Router:
|
|||
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
|
||||
if standard_logging_object is None:
|
||||
raise ValueError("standard_logging_object is None")
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
else:
|
||||
deployment_name: Final = kwargs["litellm_params"]["metadata"].get(
|
||||
"deployment", None
|
||||
) # stable name - works for wildcard routes as well
|
||||
# Get model_group and id from kwargs like the sync version does
|
||||
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
|
||||
model_info: Final = kwargs["litellm_params"].get("model_info", {}) or {}
|
||||
id = model_info.get("id", None)
|
||||
if model_group is None or id is None:
|
||||
return
|
||||
elif isinstance(id, int):
|
||||
id = str(id)
|
||||
litellm_params: Final = kwargs["litellm_params"]
|
||||
metadata: Final = litellm_params.get("metadata")
|
||||
if metadata is None:
|
||||
return
|
||||
model_group: Final = metadata.get("model_group", None)
|
||||
model_info: Final = litellm_params.get("model_info", {}) or {}
|
||||
deployment_id: Final = model_info.get("id", None)
|
||||
if model_group is None or deployment_id is None or self.get_deployment(model_id=str(deployment_id)) is None:
|
||||
return
|
||||
|
||||
## get deployment info
|
||||
deployment_info: Final = self.get_deployment(model_id=id)
|
||||
# Always track deployment successes for cooldown logic, regardless of TPM/RPM limits
|
||||
increment_deployment_successes_for_current_minute(
|
||||
litellm_router_instance=self,
|
||||
deployment_id=str(deployment_id),
|
||||
)
|
||||
|
||||
if deployment_info is None:
|
||||
return
|
||||
else:
|
||||
deployment_model_info: Final = self.get_router_model_info(
|
||||
deployment=deployment_info,
|
||||
received_model_name=model_group,
|
||||
)
|
||||
# get tpm/rpm from deployment info
|
||||
tpm: Final = deployment_info.get("tpm", None)
|
||||
rpm: Final = deployment_info.get("rpm", None)
|
||||
|
||||
## check tpm/rpm in litellm_params
|
||||
tpm_litellm_params: Final = deployment_info.litellm_params.tpm
|
||||
rpm_litellm_params: Final = deployment_info.litellm_params.rpm
|
||||
|
||||
## check tpm/rpm in model_info
|
||||
tpm_model_info: Final = deployment_model_info.get("tpm", None)
|
||||
rpm_model_info: Final = deployment_model_info.get("rpm", None)
|
||||
|
||||
# Always track deployment successes for cooldown logic, regardless of TPM/RPM limits
|
||||
increment_deployment_successes_for_current_minute(
|
||||
litellm_router_instance=self,
|
||||
deployment_id=id,
|
||||
)
|
||||
|
||||
deployment_dict = deployment_info if isinstance(deployment_info, dict) else deployment_info.model_dump()
|
||||
has_io_token_limits: Final = deployment_has_io_token_limits(deployment_dict)
|
||||
|
||||
## Nothing to track only when neither tpm/rpm nor itpm/otpm limits are
|
||||
## set. IO deployments still record TPM/RPM usage here so TPM-aware
|
||||
## routing strategies see their real load in mixed model groups; their
|
||||
## itpm/otpm enforcement runs separately in ModelRateLimitingCheck.
|
||||
if (
|
||||
tpm is None
|
||||
and rpm is None
|
||||
and tpm_litellm_params is None
|
||||
and rpm_litellm_params is None
|
||||
and tpm_model_info is None
|
||||
and rpm_model_info is None
|
||||
and not has_io_token_limits
|
||||
):
|
||||
return
|
||||
|
||||
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(kwargs)
|
||||
total_tokens: Final[float] = standard_logging_object.get("total_tokens", 0)
|
||||
|
||||
# ------------
|
||||
# Setup values
|
||||
# ------------
|
||||
dt: Final = get_utc_datetime()
|
||||
current_minute: Final = dt.strftime("%H-%M") # use the same timezone regardless of system clock
|
||||
|
||||
tpm_key = RouterCacheEnum.TPM.value.format(id=id, current_minute=current_minute, model=deployment_name)
|
||||
# ------------
|
||||
# Update usage
|
||||
# ------------
|
||||
# update cache
|
||||
pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = []
|
||||
|
||||
## TPM
|
||||
pipeline_operations.append(
|
||||
RedisPipelineIncrementOperation(
|
||||
key=tpm_key,
|
||||
increment_value=total_tokens,
|
||||
ttl=RoutingArgs.ttl.value,
|
||||
)
|
||||
)
|
||||
|
||||
## RPM
|
||||
rpm_key = RouterCacheEnum.RPM.value.format(id=id, current_minute=current_minute, model=deployment_name)
|
||||
pipeline_operations.append(
|
||||
RedisPipelineIncrementOperation(
|
||||
key=rpm_key,
|
||||
increment_value=1,
|
||||
ttl=RoutingArgs.ttl.value,
|
||||
)
|
||||
)
|
||||
|
||||
await self.cache.async_increment_cache_pipeline(
|
||||
increment_list=pipeline_operations,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
return tpm_key
|
||||
total_tokens: Final[float] = standard_logging_object.get("total_tokens", 0)
|
||||
counted_tokens: Final = get_counted_usage_tokens(litellm_params)
|
||||
deployment_name: Final = metadata.get("deployment", None)
|
||||
return await self._increment_deployment_usage(
|
||||
deployment_id=str(deployment_id),
|
||||
deployment_name=deployment_name if isinstance(deployment_name, str) else None,
|
||||
model_group=model_group,
|
||||
total_tokens=total_tokens if counted_tokens is None else max(0, total_tokens - counted_tokens),
|
||||
rpm_increment=1 if counted_tokens is None else 0,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_router_logger.debug(
|
||||
"litellm.router.Router::deployment_callback_on_success(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
async def increment_deployment_usage_for_response(
|
||||
self,
|
||||
response: object,
|
||||
request_kwargs: dict[str, object],
|
||||
) -> None:
|
||||
if response is None:
|
||||
return
|
||||
try:
|
||||
deployment_metadata: Final = find_deployment_metadata(request_kwargs)
|
||||
model_group: Final = request_kwargs.get("model")
|
||||
if deployment_metadata is None or not isinstance(model_group, str):
|
||||
return
|
||||
model_info: Final = deployment_metadata["model_info"]
|
||||
deployment_id: Final = model_info.get("id") if isinstance(model_info, dict) else None
|
||||
if deployment_id is None:
|
||||
return
|
||||
total_tokens: Final = response_total_token_count(response)
|
||||
deployment_name: Final = deployment_metadata.get("deployment")
|
||||
deployment_metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] = total_tokens
|
||||
try:
|
||||
await self._increment_deployment_usage(
|
||||
deployment_id=str(deployment_id),
|
||||
deployment_name=deployment_name if isinstance(deployment_name, str) else None,
|
||||
model_group=model_group,
|
||||
total_tokens=total_tokens,
|
||||
rpm_increment=1,
|
||||
parent_otel_span=_get_parent_otel_span_from_kwargs(request_kwargs),
|
||||
)
|
||||
except Exception:
|
||||
deployment_metadata.pop(ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY, None)
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_router_logger.debug(
|
||||
"litellm.router.Router::increment_deployment_usage_for_response(): Exception occured - %s", e
|
||||
)
|
||||
|
||||
async def _increment_deployment_usage(
|
||||
self,
|
||||
*,
|
||||
deployment_id: str,
|
||||
deployment_name: str | None,
|
||||
model_group: str,
|
||||
total_tokens: float,
|
||||
rpm_increment: int,
|
||||
parent_otel_span: Span | None,
|
||||
) -> str | None:
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
||||
deployment_info: Final = self.get_deployment(model_id=deployment_id)
|
||||
if deployment_info is None:
|
||||
return None
|
||||
deployment_model_info: Final = self.get_router_model_info(
|
||||
deployment=deployment_info,
|
||||
received_model_name=model_group,
|
||||
)
|
||||
configured_limits: Final = (
|
||||
deployment_info.get("tpm", None),
|
||||
deployment_info.get("rpm", None),
|
||||
deployment_info.litellm_params.tpm,
|
||||
deployment_info.litellm_params.rpm,
|
||||
deployment_model_info.get("tpm", None),
|
||||
deployment_model_info.get("rpm", None),
|
||||
)
|
||||
## Nothing to track only when neither tpm/rpm nor itpm/otpm limits are
|
||||
## set. IO deployments still record TPM/RPM usage here so TPM-aware
|
||||
## routing strategies see their real load in mixed model groups; their
|
||||
## itpm/otpm enforcement runs separately in ModelRateLimitingCheck.
|
||||
if all(limit is None for limit in configured_limits) and not deployment_has_io_token_limits(
|
||||
deployment_info.model_dump()
|
||||
):
|
||||
return None
|
||||
if total_tokens <= 0 and rpm_increment <= 0:
|
||||
return None
|
||||
|
||||
current_minute: Final = get_utc_datetime().strftime("%H-%M") # use the same timezone regardless of system clock
|
||||
tpm_key: Final = RouterCacheEnum.TPM.value.format(
|
||||
id=deployment_id, current_minute=current_minute, model=deployment_name
|
||||
)
|
||||
rpm_key: Final = RouterCacheEnum.RPM.value.format(
|
||||
id=deployment_id, current_minute=current_minute, model=deployment_name
|
||||
)
|
||||
pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [
|
||||
RedisPipelineIncrementOperation(key=key, increment_value=increment_value, ttl=RoutingArgs.ttl.value)
|
||||
for key, increment_value in ((tpm_key, total_tokens), (rpm_key, rpm_increment))
|
||||
]
|
||||
post_increment_values: Final = await self.cache.async_increment_cache_pipeline(
|
||||
increment_list=pipeline_operations,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
if post_increment_values is not None and self.cache.redis_cache is not None:
|
||||
for operation, value in zip(pipeline_operations, post_increment_values):
|
||||
await self.cache.async_set_cache(
|
||||
operation["key"], int(value), local_only=True, ttl=RoutingArgs.ttl.value
|
||||
)
|
||||
return tpm_key
|
||||
|
||||
def sync_deployment_callback_on_success(
|
||||
self,
|
||||
kwargs, # kwargs to completion
|
||||
|
|
@ -11205,15 +11232,7 @@ class Router:
|
|||
|
||||
if model_group is not None:
|
||||
remaining_usage: Final = await self.get_remaining_model_group_usage(model_group)
|
||||
# get_remaining_model_group_usage reads the router's TPM/RPM counter,
|
||||
# which is incremented post-response by deployment_callback_on_success.
|
||||
# Replay the in-flight increment for TPM/RPM only (LIT-2719); ITPM/OTPM
|
||||
# counters are incremented at reservation time and must not be adjusted.
|
||||
apply_remaining_usage_headers(
|
||||
additional_headers,
|
||||
remaining_usage,
|
||||
response_in_flight_token_count(response),
|
||||
)
|
||||
apply_remaining_usage_headers(additional_headers, remaining_usage)
|
||||
return response
|
||||
|
||||
def _build_model_name_index(self, model_list: list) -> None:
|
||||
|
|
|
|||
|
|
@ -151,7 +151,7 @@ def apply_quality_router_decision_headers(
|
|||
additional_headers[header] = str(decision[field])
|
||||
|
||||
|
||||
def response_in_flight_token_count(response: object) -> int:
|
||||
def response_total_token_count(response: object) -> int:
|
||||
usage: Final = response.get("usage") if isinstance(response, dict) else getattr(response, "usage", None)
|
||||
if usage is None:
|
||||
return 0
|
||||
|
|
@ -166,15 +166,10 @@ def response_in_flight_token_count(response: object) -> int:
|
|||
def apply_remaining_usage_headers(
|
||||
additional_headers: dict[str, object],
|
||||
remaining_usage: dict[str, int],
|
||||
in_flight_tokens: int,
|
||||
) -> None:
|
||||
in_flight_delta: Final = {
|
||||
"x-ratelimit-remaining-tokens": in_flight_tokens,
|
||||
"x-ratelimit-remaining-requests": 1,
|
||||
}
|
||||
for header, value in remaining_usage.items():
|
||||
if value is not None and header not in additional_headers:
|
||||
additional_headers[header] = value - in_flight_delta.get(header, 0)
|
||||
additional_headers[header] = value
|
||||
|
||||
|
||||
def _normalize_hidden_params(hidden_params: object) -> dict[str, object]:
|
||||
|
|
|
|||
|
|
@ -421,6 +421,25 @@ async def _is_fallback_target_authorized(
|
|||
return False
|
||||
|
||||
|
||||
async def _is_fallback_target_within_budget(
|
||||
litellm_router: LitellmRouter,
|
||||
fallback_entry: str | Mapping[str, object],
|
||||
original_model_group: str,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> bool:
|
||||
budget_check: Final = litellm_router.fallback_budget_check
|
||||
target: Final = _get_fallback_target_model_group(fallback_entry)
|
||||
if budget_check is None or target is None or target == original_model_group:
|
||||
return True
|
||||
if await budget_check(model=target, request_kwargs=kwargs, llm_router=litellm_router):
|
||||
return True
|
||||
verbose_router_logger.info(
|
||||
"Skipping fallback to model_group = %s: caller is over budget",
|
||||
mask_sensitive_structure(fallback_entry),
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def references_provider_scoped_resource(kwargs: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
True when a file, batch, or fine-tuning job operation names an id that only exists
|
||||
|
|
@ -528,6 +547,8 @@ async def run_async_fallback(
|
|||
continue
|
||||
if not await _is_fallback_target_authorized(litellm_router, mg, original_model_group, kwargs):
|
||||
continue
|
||||
if not await _is_fallback_target_within_budget(litellm_router, mg, original_model_group, kwargs):
|
||||
continue
|
||||
attempt_key = fallback_attempt_key(mg)
|
||||
if attempt_key is not None:
|
||||
if attempt_key in attempted:
|
||||
|
|
|
|||
|
|
@ -9,8 +9,11 @@ get_deployment_failures_for_current_minute
|
|||
get_deployment_successes_for_current_minute
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm.constants import ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router as _Router
|
||||
|
||||
|
|
@ -18,6 +21,26 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
LitellmRouter = Any
|
||||
|
||||
_METADATA_CHANNELS: Final = ("litellm_metadata", "metadata")
|
||||
|
||||
|
||||
def find_deployment_metadata(kwargs: Mapping[str, object]) -> dict[str, object] | None:
|
||||
buckets: Final = (kwargs.get(channel) for channel in _METADATA_CHANNELS)
|
||||
return next((bucket for bucket in buckets if isinstance(bucket, dict) and "model_info" in bucket), None)
|
||||
|
||||
|
||||
def get_counted_usage_tokens(litellm_params: Mapping[str, object]) -> int | None:
|
||||
buckets: Final = (litellm_params.get(channel) for channel in _METADATA_CHANNELS)
|
||||
counted: Final = next(
|
||||
(
|
||||
bucket[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY]
|
||||
for bucket in buckets
|
||||
if isinstance(bucket, dict) and ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY in bucket
|
||||
),
|
||||
None,
|
||||
)
|
||||
return counted if isinstance(counted, int) and not isinstance(counted, bool) else None
|
||||
|
||||
|
||||
def increment_deployment_successes_for_current_minute(
|
||||
litellm_router_instance: LitellmRouter,
|
||||
|
|
|
|||
|
|
@ -682,6 +682,12 @@ class BedrockGuardrailStreamingParams(BaseModel):
|
|||
"and the scan result lands in guardrail_information; a flagged response still ends the "
|
||||
"stream with a block message (disable_exception_on_block=true) or an error frame.",
|
||||
)
|
||||
streaming_buffer_release_on_scan: bool = Field(
|
||||
default=False,
|
||||
description="When buffering, scan the accumulated response every streaming_sampling_rate chunks "
|
||||
"and release the withheld chunks once the scan passes, instead of holding everything to end of stream. "
|
||||
"Flagged content is never released. Ignored when streaming_end_of_stream_only is true.",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_extras(cls, extras: Mapping[str, object] | None) -> "BedrockGuardrailStreamingParams":
|
||||
|
|
|
|||
|
|
@ -262,6 +262,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_remaining_user_budget_metric",
|
||||
"litellm_user_max_budget_metric",
|
||||
"litellm_user_budget_remaining_hours_metric",
|
||||
"litellm_remaining_customer_budget_metric",
|
||||
"litellm_customer_max_budget_metric",
|
||||
"litellm_customer_budget_remaining_hours_metric",
|
||||
"litellm_deployment_state",
|
||||
"litellm_deployment_failure_responses",
|
||||
"litellm_deployment_total_requests",
|
||||
|
|
@ -733,6 +736,12 @@ class PrometheusMetricLabels:
|
|||
|
||||
litellm_user_budget_remaining_hours_metric = litellm_remaining_user_budget_metric
|
||||
|
||||
litellm_remaining_customer_budget_metric = (UserAPIKeyLabelNames.END_USER.value,)
|
||||
|
||||
litellm_customer_max_budget_metric = litellm_remaining_customer_budget_metric
|
||||
|
||||
litellm_customer_budget_remaining_hours_metric = litellm_remaining_customer_budget_metric
|
||||
|
||||
litellm_remaining_api_key_requests_for_model = [
|
||||
UserAPIKeyLabelNames.API_KEY_HASH.value,
|
||||
UserAPIKeyLabelNames.API_KEY_ALIAS.value,
|
||||
|
|
|
|||
|
|
@ -231,7 +231,7 @@ class CacheDetailBlock(TypedDict):
|
|||
class ConverseTokenUsageBlock(TypedDict, total=False):
|
||||
inputTokens: Required[ReadOnly[int]]
|
||||
outputTokens: Required[ReadOnly[int]]
|
||||
totalTokens: Required[ReadOnly[int]]
|
||||
totalTokens: ReadOnly[int]
|
||||
cacheReadInputTokenCount: ReadOnly[int]
|
||||
cacheReadInputTokens: ReadOnly[int]
|
||||
cacheWriteInputTokenCount: ReadOnly[int]
|
||||
|
|
|
|||
|
|
@ -4,6 +4,14 @@ from .base import GuardrailConfigModel
|
|||
|
||||
|
||||
class CrowdStrikeAIDRGuardrailConfigModelOptionalParams(BaseModel):
|
||||
streaming_buffer_until_moderated: bool | None = Field(
|
||||
default=None,
|
||||
description="When True, withhold streamed chunks until moderation passes. Defaults to False when unset.",
|
||||
)
|
||||
streaming_buffer_release_on_scan: bool | None = Field(
|
||||
default=None,
|
||||
description="When buffering, release withheld chunks after each passing scan. Defaults to False when unset.",
|
||||
)
|
||||
streaming_end_of_stream_only: bool | None = Field(
|
||||
default=None,
|
||||
description="If False (default when unset), post_call scans the accumulated streamed response every "
|
||||
|
|
|
|||
|
|
@ -963,6 +963,19 @@ class FallbackAccessCheck(Protocol):
|
|||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
|
||||
|
||||
|
||||
class FallbackBudgetCheck(Protocol):
|
||||
"""
|
||||
Decides whether the caller behind `request_kwargs` is still within budget for fallback `model`.
|
||||
|
||||
Budget is enforced once during auth, against the *requested* model group. A fallback target is
|
||||
chosen later, inside the router, so a zero-cost group that falls back to a priced one bills
|
||||
without any budget gate. The router runs this before every cross-model-group fallback attempt
|
||||
and skips targets it rejects, leaving the free attempt itself untouched.
|
||||
"""
|
||||
|
||||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
|
||||
|
||||
|
||||
class AutoRouterCapabilityLimit(Protocol):
|
||||
"""
|
||||
Resolves how many complexity routers may claim each licensed capability right now; None means unlimited.
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -426,6 +426,7 @@ model LiteLLM_VerificationToken {
|
|||
key_alias String?
|
||||
soft_budget_cooldown Boolean @default(false) // key-level state on if budget alerts need to be cooled down
|
||||
spend Float @default(0.0)
|
||||
total_spend Float @default(0.0)
|
||||
expires DateTime?
|
||||
models String[]
|
||||
aliases Json @default("{}")
|
||||
|
|
@ -528,6 +529,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
key_alias String?
|
||||
soft_budget_cooldown Boolean @default(false)
|
||||
spend Float @default(0.0)
|
||||
total_spend Float @default(0.0)
|
||||
expires DateTime?
|
||||
models String[]
|
||||
aliases Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -15,6 +15,14 @@ all read the whole table and all pass. That is deliberate: a rule wide enough to
|
|||
reach them fires on most ordinary migrations, and a marker everyone adds by reflex
|
||||
stops carrying information. The outage this was written for was a backfill.
|
||||
|
||||
The one schema change banned outright is `ADD COLUMN ... DEFAULT` on a table in
|
||||
`REQUEST_LOG_TABLES`, the tables that hold a row per request. Postgres 11 stores such
|
||||
a default as metadata and touches no rows, but Postgres 10, which is supported,
|
||||
rewrites the whole heap and rebuilds every index under an `ACCESS EXCLUSIVE` lock,
|
||||
which on a spend-log-sized table is the same outage as a backfill. Every other table
|
||||
is small enough that the rewrite is not worth a rule, and a column added to a log
|
||||
table without a default is still free on every version.
|
||||
|
||||
Flagged, per statement, by its leading keyword:
|
||||
|
||||
UPDATE rewrites every matching row, and `WHERE` does not bound the scan
|
||||
|
|
@ -32,6 +40,10 @@ Flagged, per statement, by its leading keyword:
|
|||
against the part of the statement holding it, so a writable CTE
|
||||
bounded by its own `VALUES` list is not handed the query the statement
|
||||
ends with as the rows it copies
|
||||
ALTER only `ALTER TABLE` on a request-log table, and only when one of its
|
||||
actions adds a column with a `DEFAULT`. An `ALTER COLUMN ... SET
|
||||
DEFAULT` written after the column exists changes metadata alone, so it
|
||||
passes, as does an `ADD CONSTRAINT`
|
||||
|
||||
Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a
|
||||
statement's leading keyword, so they pass.
|
||||
|
|
@ -85,7 +97,7 @@ would let one written for a `DO` block silence a rewrite added to that block lat
|
|||
|
||||
`GRANDFATHERED` freezes the violations that predate this check. Prisma records a
|
||||
checksum for every applied migration and this repo treats applied files as
|
||||
immutable, so those two cannot take an inline marker. The set is closed; a new
|
||||
immutable, so those files cannot take an inline marker. The set is closed; a new
|
||||
migration belongs nowhere in it.
|
||||
"""
|
||||
|
||||
|
|
@ -102,11 +114,15 @@ MIGRATIONS_DIR = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" / "
|
|||
|
||||
GRANDFATHERED = frozenset(
|
||||
{
|
||||
"20250425182129_add_session_id",
|
||||
"20260817000000_shadow_eval_multi_key",
|
||||
"20260818000000_add_spend_log_timestamps",
|
||||
"20260818224500_add_shadow_eval_stopped_by",
|
||||
}
|
||||
)
|
||||
|
||||
REQUEST_LOG_TABLES = frozenset({"LiteLLM_SpendLogs", "LiteLLM_ErrorLogs"})
|
||||
|
||||
MARKER = re.compile(r"--[ \t]*data-migration-ok:[ \t]*(\S.*?)[ \t]*$", re.MULTILINE)
|
||||
DOLLAR_TAG = re.compile(r"\$(?:[A-Za-z_][A-Za-z0-9_]*)?\$")
|
||||
FIRST_WORD = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
|
||||
|
|
@ -128,6 +144,8 @@ DEFINES_A_ROUTINE = re.compile(
|
|||
)
|
||||
QUALIFIED_NAME = r"(?:\"[^\"]*\"|[A-Za-z_][A-Za-z0-9_$]*)"
|
||||
ROUTINE_NAME = re.compile(rf"\s*(?:{QUALIFIED_NAME}\s*\.\s*)?({QUALIFIED_NAME})")
|
||||
TABLE_NAME = ROUTINE_NAME
|
||||
ALTERS_A_TABLE = re.compile(r"\bALTER\s+TABLE\b(?:\s+IF\s+EXISTS)?(?:\s+ONLY)?", re.IGNORECASE)
|
||||
OPENS_A_CALL = re.compile(r"\s*\(")
|
||||
NAMES_AN_INDEX = re.compile(r"\bCREATE\b.+\bINDEX\b", re.IGNORECASE | re.DOTALL)
|
||||
INTRODUCES_A_RELATION = frozenset({"TABLE", "INTO", "REFERENCES", "EXISTS", "COPY"})
|
||||
|
|
@ -185,6 +203,10 @@ statement with the bound spelled out:
|
|||
|
||||
-- data-migration-ok: <what bounds this>
|
||||
UPDATE ...
|
||||
|
||||
On Postgres 10 an `ADD COLUMN ... DEFAULT` on a request-log table rewrites the table
|
||||
too. Add the column nullable with no default, then set the default in a separate
|
||||
`ALTER COLUMN ... SET DEFAULT`, which never touches existing rows.
|
||||
"""
|
||||
|
||||
|
||||
|
|
@ -537,6 +559,51 @@ def row_source_in(text: str) -> str | None:
|
|||
return next((word for word in ("SELECT", "TABLE") if contains(text, word)), None)
|
||||
|
||||
|
||||
def rewrites_a_log_table(clause: str, region: str, base: int) -> str | None:
|
||||
"""The keyword to report when an `ALTER TABLE` adds a defaulted column to a request-log
|
||||
table, which Postgres 10 answers by rewriting the whole table. The table is read from the
|
||||
region rather than the masked clause, since masking blanks the quoted name in place, after
|
||||
stepping over any comment sitting between `TABLE` and the name, which masking blanked as
|
||||
well. Each action of the statement is read on its own so that a `SET DEFAULT` on one column
|
||||
does not stand in for a default on a column another action adds."""
|
||||
altered = ALTERS_A_TABLE.search(clause)
|
||||
if altered is None:
|
||||
return None
|
||||
named = TABLE_NAME.match(region, skip_comments(region, base + altered.end()))
|
||||
if named is None or named.group(1).strip('"') not in REQUEST_LOG_TABLES:
|
||||
return None
|
||||
actions = strip_parens(clause[named.end() - base :]).split(",")
|
||||
if not any(adds_a_defaulted_column(action) for action in actions):
|
||||
return None
|
||||
return f"ADD COLUMN ... DEFAULT on {named.group(1)}"
|
||||
|
||||
|
||||
def skip_comments(sql: str, start: int) -> int:
|
||||
index = start
|
||||
while index < len(sql):
|
||||
pair = sql[index : index + 2]
|
||||
if pair == "--":
|
||||
stop = sql.find("\n", index)
|
||||
index = len(sql) if stop == -1 else stop
|
||||
elif pair == "/*":
|
||||
index = skip_block_comment(sql, index)
|
||||
elif sql[index].isspace():
|
||||
index += 1
|
||||
else:
|
||||
return index
|
||||
return index
|
||||
|
||||
|
||||
def adds_a_defaulted_column(action: str) -> bool:
|
||||
"""Whether an `ALTER TABLE` action is an `ADD COLUMN` carrying a column default. A `DEFAULT`
|
||||
right after `SET` is the referential action of an inline foreign key, which fills nothing
|
||||
in, so it does not count."""
|
||||
words = tuple(word.group().upper() for word in FIRST_WORD.finditer(action))
|
||||
if words[:1] != ("ADD",) or words[1:2] == ("CONSTRAINT",):
|
||||
return False
|
||||
return any(word == "DEFAULT" and previous != "SET" for previous, word in zip(words, words[1:]))
|
||||
|
||||
|
||||
def hands_off_sql(statement: str, executed: frozenset[str]) -> bool:
|
||||
"""Whether a statement gives the server a string literal to run as SQL. `EXECUTE` runs one
|
||||
outright, and so does `DO`, whose body is a string wherever it is not dollar-quoted. An
|
||||
|
|
@ -724,9 +791,12 @@ def scan_region(
|
|||
)
|
||||
|
||||
keyword = offending_keyword(clause)
|
||||
if keyword is None or exempt:
|
||||
if exempt:
|
||||
continue
|
||||
yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), keyword)
|
||||
found = keyword or rewrites_a_log_table(clause, region, base)
|
||||
if found is None:
|
||||
continue
|
||||
yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), found)
|
||||
|
||||
for body in bodies:
|
||||
if not runs_when_applied(masked, region, bodies, runnable, identifiers, body):
|
||||
|
|
|
|||
|
|
@ -75,6 +75,9 @@ _VCR_INCOMPATIBLE_FILES = frozenset(
|
|||
"test_router_caching.py",
|
||||
# Hits the local fake OpenAI endpoint on 127.0.0.1; nothing to record.
|
||||
"test_fake_openai_endpoint.py",
|
||||
# Needs the real connection pool a collected handler tears down; vcrpy
|
||||
# patches the transport that pool lives in.
|
||||
"test_handler_gc_does_not_close_client.py",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
315
tests/local_testing/test_handler_gc_does_not_close_client.py
Normal file
315
tests/local_testing/test_handler_gc_does_not_close_client.py
Normal file
|
|
@ -0,0 +1,315 @@
|
|||
"""
|
||||
Collecting an HTTP handler must not abort a response that is still on the wire.
|
||||
|
||||
``HTTPHandler`` and ``AsyncHTTPHandler`` close their client from ``__del__``.
|
||||
Closing a client tears down the connection pool, which aborts every response
|
||||
still streaming through it. ``_handler_may_close_client`` already withholds the
|
||||
close from a client someone else holds, but a streaming response holds the
|
||||
connection it is reading from and never the client, so the refcount it reads
|
||||
says "sole referrer" for exactly the client that is busiest. The handler is
|
||||
routinely collectable at that moment: a provider's streaming call returns the
|
||||
response and drops the handler, and ``get_async_httpx_client`` caches handlers
|
||||
behind a one-hour TTL and then lets them go.
|
||||
|
||||
The fix anchors the handler to the streaming response, so these tests turn on
|
||||
*when* the handler is collected rather than on whether it is: pinned while the
|
||||
body can still arrive, released once the caller is done with the response.
|
||||
|
||||
Nothing here re-tests the shapes ``_handler_may_close_client`` covers -- a
|
||||
borrowed ``handler.client``, a caller-supplied client, an evicted-but-held
|
||||
client. Those are pinned in ``tests/test_litellm/llms/custom_httpx/
|
||||
test_http_handler.py``. What is uncovered there is the in-flight response, so no
|
||||
test here may keep the client in a local: that inflates the very refcount under
|
||||
test, and the test then passes on a broken handler. They hold weak references
|
||||
instead, which the refcount does not count.
|
||||
|
||||
These live here rather than under ``tests/test_litellm/`` because they need a
|
||||
real connection pool: a mocked transport goes on yielding chunks after its
|
||||
client is closed, so the very teardown under test is what a mock cannot
|
||||
reproduce. The server is a hermetic, credential-free ``ThreadingHTTPServer`` on
|
||||
an ephemeral loopback port, and needs no network access beyond it.
|
||||
|
||||
Related: https://github.com/BerriAI/litellm/issues/24929
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
FRAME_COUNT = 6
|
||||
# Generous: the server emits all frames in ~0.3s. A client whose pool was torn
|
||||
# down mid-stream can stall silently instead of raising, so reads are bounded.
|
||||
READ_TIMEOUT_SECONDS = 15.0
|
||||
RELEASE_TIMEOUT_SECONDS = 3.0
|
||||
|
||||
BOTH_TRANSPORTS = pytest.mark.parametrize("disable_aiohttp_transport", [False, True], ids=["aiohttp", "httpcore"])
|
||||
|
||||
STILL_PINNED = "the handler was released while its response could still read"
|
||||
NOT_RELEASED = "the handler outlived the response that was holding it"
|
||||
|
||||
|
||||
class _ChunkedSSEServer:
|
||||
"""In-process HTTP/1.1 server that answers every request with chunked SSE frames."""
|
||||
|
||||
def __init__(self, frame_count: int = FRAME_COUNT, frame_delay: float = 0.05) -> None:
|
||||
self.frame_count = frame_count
|
||||
self.frame_delay = frame_delay
|
||||
parent = self
|
||||
|
||||
class _Handler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def _stream(self):
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "text/event-stream")
|
||||
self.send_header("Transfer-Encoding", "chunked")
|
||||
self.end_headers()
|
||||
try:
|
||||
for index in range(parent.frame_count):
|
||||
frame = f"data: frame-{index}\n\n".encode()
|
||||
self.wfile.write(b"%x\r\n" % len(frame) + frame + b"\r\n")
|
||||
self.wfile.flush()
|
||||
time.sleep(parent.frame_delay)
|
||||
self.wfile.write(b"0\r\n\r\n")
|
||||
self.wfile.flush()
|
||||
except (BrokenPipeError, ConnectionResetError):
|
||||
pass
|
||||
|
||||
do_GET = _stream
|
||||
do_POST = _stream
|
||||
|
||||
def log_message(self, *args):
|
||||
pass
|
||||
|
||||
self._server = ThreadingHTTPServer(("127.0.0.1", 0), _Handler)
|
||||
self.url = f"http://127.0.0.1:{self._server.server_address[1]}/stream"
|
||||
|
||||
def __enter__(self):
|
||||
threading.Thread(target=self._server.serve_forever, daemon=True).start()
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc_info):
|
||||
self._server.shutdown()
|
||||
self._server.server_close()
|
||||
|
||||
|
||||
def _select_transport(monkeypatch, disable_aiohttp_transport: bool) -> None:
|
||||
monkeypatch.delenv("DISABLE_AIOHTTP_TRANSPORT", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", disable_aiohttp_transport)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", False)
|
||||
|
||||
|
||||
async def _read_frames(response: httpx.Response) -> int:
|
||||
"""Count SSE frames, collecting garbage between chunks so a finalizer has every chance to fire.
|
||||
|
||||
The body is joined before counting: a chunk boundary can fall inside the
|
||||
marker, which a per-chunk count would miss.
|
||||
"""
|
||||
chunks = []
|
||||
async for chunk in response.aiter_bytes():
|
||||
chunks.append(chunk)
|
||||
gc.collect()
|
||||
return b"".join(chunks).count(b"data: frame-")
|
||||
|
||||
|
||||
async def _wait_until(is_done, failure: str) -> None:
|
||||
deadline = time.monotonic() + RELEASE_TIMEOUT_SECONDS
|
||||
while time.monotonic() < deadline:
|
||||
if is_done():
|
||||
return
|
||||
await asyncio.sleep(0.05)
|
||||
pytest.fail(failure)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@BOTH_TRANSPORTS
|
||||
async def test_async_stream_survives_handler_collection(monkeypatch, disable_aiohttp_transport):
|
||||
"""A response still streaming keeps working after its handler goes out of scope.
|
||||
|
||||
The caller holds the response and nothing else, which is what a provider's
|
||||
streaming path is left with once ``post(..., stream=True)`` has returned.
|
||||
"""
|
||||
_select_transport(monkeypatch, disable_aiohttp_transport)
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = AsyncHTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
response = await handler.post(server.url, stream=True)
|
||||
|
||||
ref = weakref.ref(handler)
|
||||
del handler
|
||||
gc.collect()
|
||||
await asyncio.sleep(0) # let any close the finalizer scheduled run
|
||||
|
||||
assert ref() is not None, STILL_PINNED
|
||||
assert await asyncio.wait_for(_read_frames(response), timeout=READ_TIMEOUT_SECONDS) == FRAME_COUNT
|
||||
|
||||
del response
|
||||
gc.collect()
|
||||
assert ref() is None, NOT_RELEASED
|
||||
|
||||
|
||||
def test_sync_stream_survives_handler_collection(monkeypatch):
|
||||
"""The sync handler closes inline from its finalizer, so a stream must hold it off.
|
||||
|
||||
litellm/main.py builds a sync handler only for non-streaming calls, commented
|
||||
"Keep this here, otherwise, the httpx.client closes and streaming is
|
||||
impossible" -- a workaround for this finalizer rather than a fix for it.
|
||||
"""
|
||||
monkeypatch.setattr(litellm, "force_ipv4", False)
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = HTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
response = handler.post(server.url, stream=True)
|
||||
|
||||
ref = weakref.ref(handler)
|
||||
del handler
|
||||
gc.collect()
|
||||
assert ref() is not None, STILL_PINNED
|
||||
|
||||
# Joined before counting, as in ``_read_frames``.
|
||||
chunks = []
|
||||
for chunk in response.iter_bytes():
|
||||
chunks.append(chunk)
|
||||
gc.collect()
|
||||
assert b"".join(chunks).count(b"data: frame-") == FRAME_COUNT
|
||||
|
||||
del response
|
||||
gc.collect()
|
||||
assert ref() is None, NOT_RELEASED
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@BOTH_TRANSPORTS
|
||||
async def test_an_abandoned_stream_still_releases_its_handler(monkeypatch, disable_aiohttp_transport):
|
||||
"""A caller that drops a stream unread must not pin the handler for good.
|
||||
|
||||
Tying the handler to the response's own lifetime is what bounds this. No
|
||||
deadline, and no poll of the connection's state, can tell an abandoned body
|
||||
from one the upstream is merely slow to finish: httpx leaves the connection
|
||||
checked out until the response is read or closed, and a legitimate stream is
|
||||
bounded only by how long the upstream keeps sending.
|
||||
"""
|
||||
_select_transport(monkeypatch, disable_aiohttp_transport)
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = AsyncHTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
client_ref = weakref.ref(handler.client)
|
||||
response = await handler.post(server.url, stream=True)
|
||||
|
||||
ref = weakref.ref(handler)
|
||||
del handler, response
|
||||
gc.collect()
|
||||
|
||||
assert ref() is None, NOT_RELEASED
|
||||
await _wait_until(
|
||||
lambda: client_ref() is None or client_ref().is_closed,
|
||||
"the client outlived the abandoned stream without being closed",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@BOTH_TRANSPORTS
|
||||
async def test_the_pool_is_released_once_the_stream_it_carried_ends(monkeypatch, disable_aiohttp_transport):
|
||||
"""Holding the finalizer off must defer the close, not drop it.
|
||||
|
||||
Otherwise a collected handler leaks its pool for every streaming request it
|
||||
was carrying, and on aiohttp warns "Unclosed client session" when the
|
||||
collector eventually takes it. The pool and the session are children of the
|
||||
client, so keeping one here does not inflate the refcount the finalizer
|
||||
reads, the way keeping the client would.
|
||||
"""
|
||||
_select_transport(monkeypatch, disable_aiohttp_transport)
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = AsyncHTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
transport = handler.client._transport
|
||||
if disable_aiohttp_transport:
|
||||
pool = transport._pool
|
||||
|
||||
def is_released() -> bool:
|
||||
return pool.connections == []
|
||||
else:
|
||||
session = transport._get_valid_client_session()
|
||||
|
||||
def is_released() -> bool:
|
||||
return session.closed
|
||||
|
||||
response = await handler.post(server.url, stream=True)
|
||||
|
||||
del handler, transport
|
||||
gc.collect()
|
||||
assert not is_released(), "the pool was torn down while it was still carrying a body"
|
||||
|
||||
assert await asyncio.wait_for(_read_frames(response), timeout=READ_TIMEOUT_SECONDS) == FRAME_COUNT
|
||||
del response
|
||||
gc.collect()
|
||||
|
||||
await _wait_until(is_released, "the pool outlived the stream it carried, unclosed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@BOTH_TRANSPORTS
|
||||
async def test_a_non_streaming_response_does_not_pin_its_handler(monkeypatch, disable_aiohttp_transport):
|
||||
"""Only a body that can still arrive holds the handler.
|
||||
|
||||
A non-streaming response has been read in full by the time ``post`` returns,
|
||||
so pinning the handler to it would delay every client close behind whatever
|
||||
the caller goes on to do with the response.
|
||||
"""
|
||||
_select_transport(monkeypatch, disable_aiohttp_transport)
|
||||
|
||||
with _ChunkedSSEServer(frame_count=1, frame_delay=0.0) as server:
|
||||
handler = AsyncHTTPHandler(timeout=httpx.Timeout(10.0, connect=5.0))
|
||||
response = await handler.post(server.url)
|
||||
assert response.status_code == 200
|
||||
|
||||
ref = weakref.ref(handler)
|
||||
del handler
|
||||
gc.collect()
|
||||
|
||||
assert ref() is None, "a fully-read response pinned its handler"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@BOTH_TRANSPORTS
|
||||
async def test_cached_handler_eviction_does_not_abort_an_in_flight_stream(monkeypatch, disable_aiohttp_transport):
|
||||
"""Evicting a cached handler mid-stream leaves the stream alone.
|
||||
|
||||
``get_async_httpx_client`` caches handlers for an hour. When that TTL
|
||||
expires the cache drops the only reference to a handler whose client is
|
||||
still streaming -- the production shape of #24929.
|
||||
"""
|
||||
_select_transport(monkeypatch, disable_aiohttp_transport)
|
||||
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
|
||||
|
||||
with _ChunkedSSEServer() as server:
|
||||
handler = get_async_httpx_client(llm_provider=LlmProviders.OPENAI)
|
||||
response = await handler.post(server.url, stream=True)
|
||||
|
||||
# An hour passes: the TTL expires and the cache lets the handler go.
|
||||
ref = weakref.ref(handler)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
del handler
|
||||
gc.collect()
|
||||
|
||||
assert ref() is not None, STILL_PINNED
|
||||
assert await asyncio.wait_for(_read_frames(response), timeout=READ_TIMEOUT_SECONDS) == FRAME_COUNT
|
||||
|
||||
del response
|
||||
gc.collect()
|
||||
assert ref() is None, NOT_RELEASED
|
||||
|
|
@ -587,6 +587,8 @@ def _success_kwargs(
|
|||
response_cost=0.5,
|
||||
key_hash=None,
|
||||
key_model_max_budget=None,
|
||||
team_id=None,
|
||||
team_model_max_budget=None,
|
||||
user_id=None,
|
||||
user_model_max_budget=None,
|
||||
end_user_id=None,
|
||||
|
|
@ -600,6 +602,7 @@ def _success_kwargs(
|
|||
"end_user": end_user_id,
|
||||
"metadata": {
|
||||
"user_api_key_hash": key_hash,
|
||||
"user_api_key_team_id": team_id,
|
||||
"user_api_key_user_id": user_id,
|
||||
"user_api_key_end_user_id": end_user_id,
|
||||
},
|
||||
|
|
@ -607,6 +610,7 @@ def _success_kwargs(
|
|||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key_model_max_budget": key_model_max_budget,
|
||||
"user_api_key_team_model_max_budget": team_model_max_budget,
|
||||
"user_api_key_user_model_max_budget": user_model_max_budget,
|
||||
"user_api_key_end_user_model_max_budget": end_user_model_max_budget,
|
||||
},
|
||||
|
|
@ -1417,3 +1421,266 @@ async def test_spend_logged_on_one_replica_is_enforced_and_reported_on_another()
|
|||
replica_c = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache(redis_cache=shared_redis))
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await replica_c.is_key_within_model_budget(user_api_key, "gpt-4")
|
||||
|
||||
|
||||
def _log_success(limiter, **kwargs):
|
||||
return limiter.async_log_success_event(
|
||||
_success_kwargs(**kwargs), response_obj=None, start_time=None, end_time=None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_model",
|
||||
["gpt-4", "openai/gpt-4"],
|
||||
ids=["bare_model", "provider_prefixed_model"],
|
||||
)
|
||||
async def test_team_model_budget_is_shared_by_every_key_without_an_override(request_model):
|
||||
"""
|
||||
Two keys on the same team, neither carrying a matching key-level entry,
|
||||
charge one team counter and are both refused once it is spent.
|
||||
"""
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}}
|
||||
check = lambda: limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=None,
|
||||
model=request_model,
|
||||
)
|
||||
|
||||
assert await check() is True
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group=request_model,
|
||||
response_cost=0.6,
|
||||
key_hash="vk-a",
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
assert await check() is True
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group=request_model,
|
||||
response_cost=0.6,
|
||||
key_hash="vk-b",
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:gpt-4:1d") == pytest.approx(1.2)
|
||||
with pytest.raises(litellm.BudgetExceededError) as exc:
|
||||
await check()
|
||||
assert exc.value.entity_type == Litellm_EntityType.TEAM.value
|
||||
assert await build_model_max_budget_usage(
|
||||
entity_type=Litellm_EntityType.TEAM,
|
||||
entity_id="team-1",
|
||||
model_max_budget=team_model_max_budget,
|
||||
cache=dual_cache,
|
||||
) == {"gpt-4": {"current_spend": pytest.approx(1.2), "budget_limit": 1.0, "time_period": "1d"}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_override_replaces_the_team_cap_for_that_model():
|
||||
"""
|
||||
A key with its own entry for the model is gated on the key counter alone:
|
||||
the exhausted team counter does not block it, and its spend never lands on
|
||||
the team counter.
|
||||
"""
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}}
|
||||
key_model_max_budget = {"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}}
|
||||
await dual_cache.async_set_cache(key="team_model_spend:team-1:gpt-4:1d", value=9.0)
|
||||
|
||||
assert (
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
model="openai/gpt-4",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group="openai/gpt-4",
|
||||
response_cost=2.0,
|
||||
key_hash="vk-override",
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:gpt-4:1d") == 9.0
|
||||
assert await dual_cache.async_get_cache(key="virtual_key_spend:vk-override:gpt-4:1d") == 2.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_entry_for_another_model_does_not_lift_the_team_cap():
|
||||
"""A key override only covers the model it names; other models stay on the team counter."""
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}}
|
||||
key_model_max_budget = {"claude-3": {"budget_limit": 5.0, "time_period": "1d"}}
|
||||
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group="gpt-4",
|
||||
response_cost=1.5,
|
||||
key_hash="vk-other",
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:gpt-4:1d") == 1.5
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
model="gpt-4",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_budget_leaves_unconfigured_models_alone():
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {"gpt-4": {"budget_limit": 0.0, "time_period": "1d"}}
|
||||
|
||||
assert (
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=None,
|
||||
model="claude-3",
|
||||
)
|
||||
is True
|
||||
)
|
||||
with patch.object(limiter, "_increment_spend_for_key", new_callable=AsyncMock) as mock_increment:
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group="claude-3",
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
mock_increment.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_counters_are_isolated_by_team_model_and_window():
|
||||
"""Same model on two teams, and two models with different windows on one team, never share a counter."""
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {
|
||||
"gpt-4": {"budget_limit": 10.0, "time_period": "1d"},
|
||||
"claude-3": {"budget_limit": 10.0, "time_period": "30d"},
|
||||
}
|
||||
|
||||
for team_id, model in (("team-1", "gpt-4"), ("team-2", "gpt-4"), ("team-1", "claude-3")):
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group=model,
|
||||
response_cost=1.0,
|
||||
team_id=team_id,
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:gpt-4:1d") == 1.0
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-2:gpt-4:1d") == 1.0
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:claude-3:30d") == 1.0
|
||||
assert await dual_cache.async_get_cache(key="team_model_budget_start_time:team-1:claude-3:30d") is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_malformed_team_entry_is_skipped_and_its_sibling_still_enforced():
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache())
|
||||
team_model_max_budget = {
|
||||
"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"},
|
||||
"claude-3": {"budget_limit": 0.0, "time_period": "1d"},
|
||||
}
|
||||
|
||||
assert (
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=None,
|
||||
model="gpt-4",
|
||||
)
|
||||
is True
|
||||
)
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=None,
|
||||
model="claude-3",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_malformed_key_entry_does_not_count_as_an_override():
|
||||
"""A key entry the limiter cannot enforce must not also switch the team cap off."""
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}}
|
||||
key_model_max_budget = {"gpt-4": {"budget_limit": "not-a-number", "time_period": "1d"}}
|
||||
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group="gpt-4",
|
||||
response_cost=1.5,
|
||||
key_hash="vk-bad",
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:gpt-4:1d") == 1.5
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
model="gpt-4",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"key_entry",
|
||||
[
|
||||
{"time_period": "1d", "tpm_limit": 100},
|
||||
{"time_period": "1d", "rpm_limit": 10},
|
||||
{"budget_limit": -1.0, "time_period": "1d"},
|
||||
],
|
||||
)
|
||||
async def test_key_entry_without_a_spend_cap_does_not_lift_the_team_cap(key_entry):
|
||||
"""A key row that only rate-limits the model, or has no enforceable cap, leaves the team cap in force."""
|
||||
dual_cache = DualCache()
|
||||
limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=dual_cache)
|
||||
team_model_max_budget = {"gpt-4": {"budget_limit": 1.0, "time_period": "1d"}}
|
||||
key_model_max_budget = {"gpt-4": key_entry}
|
||||
|
||||
await _log_success(
|
||||
limiter,
|
||||
model_group="openai/gpt-4",
|
||||
response_cost=1.5,
|
||||
key_hash="vk-rate-limited",
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
)
|
||||
|
||||
assert await dual_cache.async_get_cache(key="team_model_spend:team-1:gpt-4:1d") == 1.5
|
||||
with pytest.raises(litellm.BudgetExceededError):
|
||||
await limiter.is_team_within_model_budget(
|
||||
team_id="team-1",
|
||||
team_model_max_budget=team_model_max_budget,
|
||||
key_model_max_budget=key_model_max_budget,
|
||||
model="openai/gpt-4",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import traceback
|
||||
|
|
@ -10,9 +11,14 @@ import pytest
|
|||
import litellm
|
||||
from unittest.mock import patch, MagicMock, AsyncMock
|
||||
from create_mock_standard_logging_payload import create_standard_logging_payload
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.types.utils import ModelResponse, StandardLoggingPayload
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
|
||||
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo
|
||||
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
|
||||
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS, ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -928,18 +934,7 @@ async def test_set_response_headers(model_list):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_response_headers_subtracts_in_flight_delta(model_list):
|
||||
"""
|
||||
LIT-2719: router-derived `x-ratelimit-remaining-*` headers must be
|
||||
post-decrement (match OpenAI/Anthropic vendor semantics) so the proxy's
|
||||
HTTP response headers and the prometheus gauges that read them stay
|
||||
comparable across providers.
|
||||
|
||||
Router's TPM/RPM counter is incremented post-response by
|
||||
`deployment_callback_on_success`, so `get_remaining_model_group_usage`
|
||||
sees pre-decrement values. `set_response_headers` must replay the
|
||||
in-flight increment before writing the headers.
|
||||
"""
|
||||
async def test_set_response_headers_passes_through_post_increment_counters(model_list):
|
||||
from pydantic import BaseModel
|
||||
|
||||
class _Usage(BaseModel):
|
||||
|
|
@ -952,49 +947,10 @@ async def test_set_response_headers_subtracts_in_flight_delta(model_list):
|
|||
router = Router(model_list=model_list)
|
||||
router.get_remaining_model_group_usage = AsyncMock(
|
||||
return_value={
|
||||
"x-ratelimit-remaining-tokens": 1000,
|
||||
"x-ratelimit-remaining-tokens": 958,
|
||||
"x-ratelimit-limit-tokens": 1000,
|
||||
"x-ratelimit-remaining-requests": 100,
|
||||
"x-ratelimit-remaining-requests": 99,
|
||||
"x-ratelimit-limit-requests": 100,
|
||||
}
|
||||
)
|
||||
|
||||
resp = _Resp()
|
||||
resp._hidden_params = {}
|
||||
await router.set_response_headers(response=resp, model_group="gpt-3.5-turbo")
|
||||
|
||||
headers = resp._hidden_params["additional_headers"]
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 958
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
# Limit headers pass through unmodified.
|
||||
assert headers["x-ratelimit-limit-tokens"] == 1000
|
||||
assert headers["x-ratelimit-limit-requests"] == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_response_headers_in_flight_delta_only_adjusts_tpm_rpm(model_list):
|
||||
"""
|
||||
The in-flight replay applies only to the post-incremented TPM/RPM counters
|
||||
(`x-ratelimit-remaining-tokens` / `-requests`). The ITPM/OTPM counters are
|
||||
incremented at reservation time (pre-call), so the input/output token
|
||||
headers already reflect this request and must pass through untouched.
|
||||
"""
|
||||
from pydantic import BaseModel
|
||||
|
||||
class _Usage(BaseModel):
|
||||
total_tokens: int = 30
|
||||
prompt_tokens: int = 20
|
||||
completion_tokens: int = 10
|
||||
|
||||
class _Resp(BaseModel):
|
||||
usage: _Usage = _Usage()
|
||||
_hidden_params: dict = {}
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
router.get_remaining_model_group_usage = AsyncMock(
|
||||
return_value={
|
||||
"x-ratelimit-remaining-tokens": 1000,
|
||||
"x-ratelimit-remaining-requests": 100,
|
||||
"x-ratelimit-remaining-input-tokens": 1000,
|
||||
"x-ratelimit-remaining-output-tokens": 500,
|
||||
}
|
||||
|
|
@ -1005,14 +961,336 @@ async def test_set_response_headers_in_flight_delta_only_adjusts_tpm_rpm(model_l
|
|||
await router.set_response_headers(response=resp, model_group="gpt-3.5-turbo")
|
||||
|
||||
headers = resp._hidden_params["additional_headers"]
|
||||
# TPM/RPM headers replay the in-flight increment...
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 970
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 958
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
# ...but the reservation-based input/output headers pass through unchanged.
|
||||
assert headers["x-ratelimit-limit-tokens"] == 1000
|
||||
assert headers["x-ratelimit-limit-requests"] == 100
|
||||
assert headers["x-ratelimit-remaining-input-tokens"] == 1000
|
||||
assert headers["x-ratelimit-remaining-output-tokens"] == 500
|
||||
|
||||
|
||||
def _rpm_tpm_router(model_id: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake", "tpm": 1000, "rpm": 100},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _ratelimit_headers(response: ModelResponse | CustomStreamWrapper) -> dict[str, int]:
|
||||
return {k: v for k, v in response._hidden_params["additional_headers"].items() if k.startswith("x-ratelimit-")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_headers_read_post_increment_counter_and_count_once():
|
||||
router = _rpm_tpm_router("lit-3058-async")
|
||||
|
||||
response = await router.acompletion(
|
||||
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong"
|
||||
)
|
||||
total_tokens = response.usage.total_tokens
|
||||
assert total_tokens > 0
|
||||
|
||||
headers = _ratelimit_headers(response)
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_wildcard_route_headers_and_counter_use_resolved_deployment_name():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {"model": "openai/*", "api_key": "sk-fake", "tpm": 1000, "rpm": 100},
|
||||
"model_info": {"id": "lit-3058-wildcard"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
response = await router.acompletion(
|
||||
model="openai/gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong"
|
||||
)
|
||||
total_tokens = response.usage.total_tokens
|
||||
|
||||
headers = _ratelimit_headers(response)
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000 - total_tokens
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("openai/gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_stream_counts_request_before_headers_and_tokens_once_on_completion():
|
||||
router = _rpm_tpm_router("lit-3058-stream")
|
||||
|
||||
stream = await router.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="pong pong pong",
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
headers = _ratelimit_headers(stream)
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (0, 1)
|
||||
|
||||
chunks = [chunk async for chunk in stream]
|
||||
total_tokens = chunks[-1].usage.total_tokens
|
||||
assert total_tokens > 0
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (total_tokens, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_callback_on_success_adds_only_uncounted_tokens():
|
||||
import time
|
||||
|
||||
router = _rpm_tpm_router("lit-3058-callback")
|
||||
standard_logging_payload = create_standard_logging_payload()
|
||||
standard_logging_payload["total_tokens"] = 100
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"deployment": "gpt-5-mini",
|
||||
"model_group": "gpt-5-mini",
|
||||
ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY: 60,
|
||||
},
|
||||
"model_info": {"id": "lit-3058-callback"},
|
||||
},
|
||||
"standard_logging_object": standard_logging_payload,
|
||||
}
|
||||
|
||||
tpm_key = await router.deployment_callback_on_success(
|
||||
kwargs=kwargs,
|
||||
completion_response=litellm.ModelResponse(model="gpt-5-mini", usage={"total_tokens": 100}),
|
||||
start_time=time.time(),
|
||||
end_time=time.time(),
|
||||
)
|
||||
|
||||
assert tpm_key is not None
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (40, 0)
|
||||
|
||||
|
||||
class _GatedIncrementCache(DualCache):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(in_memory_cache=InMemoryCache())
|
||||
self.first_increment_started = asyncio.Event()
|
||||
self.release_first_increment = asyncio.Event()
|
||||
self.increment_calls = 0
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
increment_list: list[RedisPipelineIncrementOperation],
|
||||
local_only: bool = False,
|
||||
parent_otel_span: object = None,
|
||||
**kwargs: object,
|
||||
) -> list[float] | None:
|
||||
self.increment_calls += 1
|
||||
if self.increment_calls == 1:
|
||||
self.first_increment_started.set()
|
||||
await self.release_first_increment.wait()
|
||||
return await super().async_increment_cache_pipeline(
|
||||
increment_list, local_only=local_only, parent_otel_span=parent_otel_span, **kwargs
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_callback_running_during_pre_header_increment_does_not_double_count():
|
||||
router = _rpm_tpm_router("lit-3058-race")
|
||||
cache = _GatedIncrementCache()
|
||||
router.cache = cache
|
||||
|
||||
request = asyncio.ensure_future(
|
||||
router.acompletion(model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong")
|
||||
)
|
||||
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5)
|
||||
for _ in range(50):
|
||||
if get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
assert get_deployment_successes_for_current_minute(router, "lit-3058-race") == 1
|
||||
assert cache.increment_calls == 1
|
||||
|
||||
cache.release_first_increment.set()
|
||||
response = await request
|
||||
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (response.usage.total_tokens, 1)
|
||||
|
||||
|
||||
class _UnavailableIncrementCache(DualCache):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(in_memory_cache=InMemoryCache())
|
||||
self.first_increment_started = asyncio.Event()
|
||||
self.release_first_increment = asyncio.Event()
|
||||
self.increment_calls = 0
|
||||
|
||||
async def async_increment_cache_pipeline(
|
||||
self,
|
||||
increment_list: list[RedisPipelineIncrementOperation],
|
||||
local_only: bool = False,
|
||||
parent_otel_span: object = None,
|
||||
**kwargs: object,
|
||||
) -> list[float] | None:
|
||||
self.increment_calls += 1
|
||||
if self.increment_calls == 1:
|
||||
self.first_increment_started.set()
|
||||
await self.release_first_increment.wait()
|
||||
raise RuntimeError("cache unavailable")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_observing_stamp_before_pre_header_increment_fails_leaves_no_stamp_behind():
|
||||
router = _rpm_tpm_router("lit-3058-fail")
|
||||
cache = _UnavailableIncrementCache()
|
||||
router.cache = cache
|
||||
metadata: dict[str, object] = {}
|
||||
|
||||
request = asyncio.ensure_future(
|
||||
router.acompletion(
|
||||
model="gpt-5-mini", messages=[{"role": "user", "content": "hi"}], mock_response="pong", metadata=metadata
|
||||
)
|
||||
)
|
||||
await asyncio.wait_for(cache.first_increment_started.wait(), timeout=5)
|
||||
assert metadata[ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY] == 30
|
||||
for _ in range(50):
|
||||
if get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
assert get_deployment_successes_for_current_minute(router, "lit-3058-fail") == 1
|
||||
assert cache.increment_calls == 1
|
||||
|
||||
cache.release_first_increment.set()
|
||||
response = await request
|
||||
|
||||
assert response.usage.total_tokens == 30
|
||||
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in metadata
|
||||
assert _ratelimit_headers(response)["x-ratelimit-remaining-requests"] == 100
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_deployment_usage_for_response_skips_session_wrappers():
|
||||
router = _rpm_tpm_router("lit-3058-ws")
|
||||
request_kwargs = {
|
||||
"model": "gpt-5-mini",
|
||||
"litellm_metadata": {"model_group": "gpt-5-mini", "model_info": {"id": "lit-3058-ws"}},
|
||||
}
|
||||
|
||||
await router.increment_deployment_usage_for_response(response=None, request_kwargs=request_kwargs)
|
||||
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (None, None)
|
||||
assert ROUTER_USAGE_COUNTED_TOKENS_METADATA_KEY not in request_kwargs["litellm_metadata"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_deployment_usage_writes_only_positive_deltas_for_limited_deployments():
|
||||
router = _rpm_tpm_router("lit-3058-delta")
|
||||
unlimited = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5-mini",
|
||||
"litellm_params": {"model": "gpt-5-mini", "api_key": "sk-fake"},
|
||||
"model_info": {"id": "lit-3058-unlimited"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
tpm_key = await router._increment_deployment_usage(
|
||||
deployment_id="lit-3058-delta",
|
||||
deployment_name="gpt-5-mini",
|
||||
model_group="gpt-5-mini",
|
||||
total_tokens=25,
|
||||
rpm_increment=1,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
assert tpm_key is not None
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (25, 1)
|
||||
|
||||
assert (
|
||||
await router._increment_deployment_usage(
|
||||
deployment_id="lit-3058-delta",
|
||||
deployment_name="gpt-5-mini",
|
||||
model_group="gpt-5-mini",
|
||||
total_tokens=0,
|
||||
rpm_increment=0,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert await router.get_model_group_usage("gpt-5-mini") == (25, 1)
|
||||
|
||||
assert (
|
||||
await unlimited._increment_deployment_usage(
|
||||
deployment_id="lit-3058-unlimited",
|
||||
deployment_name="gpt-5-mini",
|
||||
model_group="gpt-5-mini",
|
||||
total_tokens=25,
|
||||
rpm_increment=1,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert await unlimited.get_model_group_usage("gpt-5-mini") == (None, None)
|
||||
|
||||
|
||||
def _shared_redis_stub(store: dict) -> MagicMock:
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
async def increment_pipeline(increment_list, **kwargs):
|
||||
for op in increment_list:
|
||||
store[op["key"]] = store.get(op["key"], 0.0) + op["increment_value"]
|
||||
return [store[op["key"]] for op in increment_list]
|
||||
|
||||
async def batch_get(keys, **kwargs):
|
||||
return {key: store.get(key) for key in keys}
|
||||
|
||||
redis_stub = MagicMock(spec=RedisCache)
|
||||
redis_stub.async_increment_pipeline = increment_pipeline
|
||||
redis_stub.async_batch_get_cache = batch_get
|
||||
return redis_stub
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_headers_on_fresh_worker_reflect_shared_redis_usage():
|
||||
store: dict = {}
|
||||
worker_a = _rpm_tpm_router("lit-3058-workers")
|
||||
worker_b = _rpm_tpm_router("lit-3058-workers")
|
||||
worker_a.cache = DualCache(redis_cache=_shared_redis_stub(store), in_memory_cache=InMemoryCache())
|
||||
worker_b.cache = DualCache(redis_cache=_shared_redis_stub(store), in_memory_cache=InMemoryCache())
|
||||
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
tokens_on_a = 0
|
||||
for _ in range(3):
|
||||
response = await worker_a.acompletion(model="gpt-5-mini", messages=messages, mock_response="pong")
|
||||
tokens_on_a += response.usage.total_tokens
|
||||
|
||||
response = await worker_b.acompletion(model="gpt-5-mini", messages=messages, mock_response="pong")
|
||||
headers = _ratelimit_headers(response)
|
||||
assert headers["x-ratelimit-remaining-requests"] == 96
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000 - tokens_on_a - response.usage.total_tokens
|
||||
|
||||
counted_tokens = tokens_on_a + response.usage.total_tokens
|
||||
for _ in range(2):
|
||||
response = await worker_a.acompletion(model="gpt-5-mini", messages=messages, mock_response="pong")
|
||||
counted_tokens += response.usage.total_tokens
|
||||
|
||||
stream = await worker_b.acompletion(model="gpt-5-mini", messages=messages, mock_response="pong", stream=True)
|
||||
stream_headers = _ratelimit_headers(stream)
|
||||
assert stream_headers["x-ratelimit-remaining-requests"] == 93
|
||||
assert stream_headers["x-ratelimit-remaining-tokens"] == 1000 - counted_tokens
|
||||
assert [chunk async for chunk in stream]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_model_group_io_token_usage_sums_across_deployments():
|
||||
"""
|
||||
|
|
@ -1154,8 +1432,8 @@ async def test_set_response_headers_native_input_token_header_does_not_suppress_
|
|||
await router.set_response_headers(response=resp, model_group="gpt-3.5-turbo")
|
||||
|
||||
headers = resp._hidden_params["additional_headers"]
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 958
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000
|
||||
assert headers["x-ratelimit-remaining-requests"] == 100
|
||||
# the provider's native header is left untouched
|
||||
assert headers["x-ratelimit-remaining-input-tokens"] == 5
|
||||
|
||||
|
|
@ -1187,7 +1465,7 @@ async def test_set_response_headers_native_token_header_does_not_suppress_io_hea
|
|||
|
||||
headers = resp._hidden_params["additional_headers"]
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 5
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert headers["x-ratelimit-remaining-requests"] == 100
|
||||
assert headers["x-ratelimit-remaining-input-tokens"] == 900
|
||||
assert headers["x-ratelimit-remaining-output-tokens"] == 450
|
||||
|
||||
|
|
@ -1196,8 +1474,7 @@ async def test_set_response_headers_native_token_header_does_not_suppress_io_hea
|
|||
async def test_set_response_headers_handles_missing_usage(model_list):
|
||||
"""
|
||||
Streaming chunks and some response shapes may lack a `usage` attribute or
|
||||
populated `total_tokens`. The in-flight subtraction must default to 0
|
||||
tokens (still subtract 1 from requests) and never raise.
|
||||
populated `total_tokens`. Header composition must not depend on usage and never raise.
|
||||
"""
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -1218,7 +1495,7 @@ async def test_set_response_headers_handles_missing_usage(model_list):
|
|||
|
||||
headers = resp._hidden_params["additional_headers"]
|
||||
assert headers["x-ratelimit-remaining-tokens"] == 1000
|
||||
assert headers["x-ratelimit-remaining-requests"] == 99
|
||||
assert headers["x-ratelimit-remaining-requests"] == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -168,6 +168,46 @@ def test_allowlisted_metadata_subkey_promoted_blob_excluded():
|
|||
assert all("private_note" not in k for k in span.attributes)
|
||||
|
||||
|
||||
def test_nested_metadata_key_promoted_under_caller_path():
|
||||
"""A dotted allowlist entry reads the nested caller metadata the proxy stores
|
||||
under ``requester_metadata`` and lands on the LLM-call span under the caller's
|
||||
own path (``litellm.metadata.trace_id``, ``litellm.metadata.nested.deep``);
|
||||
a pre-existing flat dotted key keeps its full name, and unlisted siblings and
|
||||
the blob stay out."""
|
||||
engine, exporter = _engine_and_exporter()
|
||||
payload = _payload()
|
||||
payload["metadata"]["a.b"] = "flat"
|
||||
payload["metadata"]["requester_metadata"] = {
|
||||
"trace_id": "abc",
|
||||
"attempt": 0,
|
||||
"empty": "",
|
||||
"nested": {"deep": "x", "skipped": "y"},
|
||||
}
|
||||
data = LLMCallSpanData.from_standard_logging_payload(payload)
|
||||
bag = promoted_baggage(
|
||||
data.identity,
|
||||
data.request_model,
|
||||
BAGGAGE_PROMOTED_KEYS,
|
||||
metadata_keys=(
|
||||
"requester_metadata.trace_id",
|
||||
"requester_metadata.attempt",
|
||||
"requester_metadata.empty",
|
||||
"requester_metadata.nested.deep",
|
||||
"a.b",
|
||||
),
|
||||
)
|
||||
engine.emit(SpanRole.LLM_CALL, data, ctx_mod.set_request_baggage(bag))
|
||||
(span,) = exporter.get_finished_spans()
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}trace_id"] == "abc"
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}attempt"] == "0"
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}nested.deep"] == "x"
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}a.b"] == "flat"
|
||||
assert f"{LiteLLM.METADATA_PREFIX}empty" not in span.attributes
|
||||
assert f"{LiteLLM.METADATA_PREFIX}deep" not in span.attributes
|
||||
assert f"{LiteLLM.METADATA_PREFIX}nested.skipped" not in span.attributes
|
||||
assert not any(k.startswith(f"{LiteLLM.METADATA_PREFIX}requester_metadata") for k in span.attributes)
|
||||
|
||||
|
||||
def test_http_attributes_never_promoted():
|
||||
"""Even if http.* is present in baggage, the processor must not stamp it on
|
||||
child spans (it belongs on the SERVER span only)."""
|
||||
|
|
|
|||
|
|
@ -1623,17 +1623,22 @@ def test_provider_model_and_team_metadata_on_real_boundary_flow():
|
|||
def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans():
|
||||
"""The pre-call hook seeds identity Baggage in the request context so the
|
||||
server span (stamped directly) AND later child spans (service here, via the
|
||||
Baggage processor) carry identity — not just the LLM-call span."""
|
||||
Baggage processor) carry identity — not just the LLM-call span. Only the
|
||||
caller's ``requester_metadata`` is read from the request dict, so a proxy-owned
|
||||
sibling such as ``requester_ip_address`` is not stamped from here even though
|
||||
the default allowlist names it, and an unlisted caller key is not promoted."""
|
||||
logger, exporter = _logger()
|
||||
server = logger._emitter.start_span(
|
||||
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
)
|
||||
data = {
|
||||
"model": "gpt-4o",
|
||||
"metadata": {"requester_ip_address": "127.0.0.1", "requester_metadata": {"trace_id": "abc"}},
|
||||
}
|
||||
|
||||
async def _flow():
|
||||
# pre-call seeds baggage + stamps the active server span
|
||||
await logger.async_pre_call_hook(
|
||||
_Auth(), None, {"model": "gpt-4o"}, "completion"
|
||||
)
|
||||
await logger.async_pre_call_hook(_Auth(), None, data, "completion")
|
||||
# a later service call (same task) must inherit the identity
|
||||
await logger.async_service_success_hook(
|
||||
payload=_ServicePayload("redis", "set"), parent_otel_span=server
|
||||
|
|
@ -1653,6 +1658,46 @@ def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans():
|
|||
srv.attributes[LiteLLM.TEAM_ID] == "t1"
|
||||
) # stamped directly on the server span
|
||||
assert srv.attributes[f"{LiteLLM.METADATA_PREFIX}user_api_key_user_id"] == "u1"
|
||||
assert not any(
|
||||
k in (f"{LiteLLM.METADATA_PREFIX}requester_ip_address", f"{LiteLLM.METADATA_PREFIX}trace_id")
|
||||
for s in (redis, srv)
|
||||
for k in s.attributes
|
||||
)
|
||||
|
||||
|
||||
def test_pre_call_hook_promotes_nested_request_metadata_key():
|
||||
"""``baggage_metadata_keys: [requester_metadata.trace_id]`` reads the caller's
|
||||
``metadata.trace_id`` (snapshotted by the proxy under ``requester_metadata``)
|
||||
and stamps ``litellm.metadata.trace_id`` on the server, LLM-call and service
|
||||
spans of the request; unlisted siblings are not promoted."""
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory", baggage_metadata_keys=["requester_metadata.trace_id"])
|
||||
exporter = InMemorySpanExporter()
|
||||
logger = OpenTelemetryV2(config=cfg, tracer_provider=providers.build_tracer_provider(cfg, exporter=exporter))
|
||||
server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
|
||||
data = {"model": "gpt-4o", "metadata": {"requester_metadata": {"trace_id": "abc", "nested": {"deep": "x"}}}}
|
||||
kwargs = _kwargs()
|
||||
|
||||
async def _flow():
|
||||
await logger.async_pre_call_hook(_Auth(), None, data, "completion")
|
||||
logger.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs)
|
||||
await logger.async_log_success_event(kwargs, None, None, None)
|
||||
await logger.async_service_success_hook(payload=_ServicePayload("redis", "set"), parent_otel_span=server)
|
||||
|
||||
with trace.use_span(server, end_on_exit=False):
|
||||
asyncio.run(_flow())
|
||||
server.end()
|
||||
|
||||
spans = {s.name: s for s in exporter.get_finished_spans()}
|
||||
key = f"{LiteLLM.METADATA_PREFIX}trace_id"
|
||||
assert spans[LITELLM_PROXY_REQUEST_SPAN_NAME].attributes[key] == "abc"
|
||||
assert spans["chat gpt-4o"].attributes[key] == "abc"
|
||||
assert spans["redis set"].attributes[key] == "abc"
|
||||
assert data == {"model": "gpt-4o", "metadata": {"requester_metadata": {"trace_id": "abc", "nested": {"deep": "x"}}}}
|
||||
assert not any(
|
||||
k.startswith(f"{LiteLLM.METADATA_PREFIX}requester_metadata") or k.endswith("deep")
|
||||
for s in spans.values()
|
||||
for k in s.attributes
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
|
|||
|
|
@ -15,10 +15,11 @@ from unittest.mock import MagicMock, patch
|
|||
|
||||
# Adds the grandparent directory to sys.path to allow importing project modules
|
||||
from opentelemetry import trace
|
||||
from opentelemetry.sdk._logs import LogData
|
||||
from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
|
||||
from opentelemetry.sdk._logs.export import InMemoryLogExporter, SimpleLogRecordProcessor
|
||||
from opentelemetry.sdk.metrics import MeterProvider
|
||||
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
|
||||
from opentelemetry.sdk.metrics.export import InMemoryMetricReader, MetricsData
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
|
||||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||
|
|
@ -5581,6 +5582,38 @@ class TestOpenTelemetryInferenceIdentityAttributes(unittest.TestCase):
|
|||
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
|
||||
assert "http.route" not in self._attr(span, exp)
|
||||
|
||||
def test_nested_metadata_key_promoted_under_caller_path(self):
|
||||
"""``baggage_metadata_keys: [requester_metadata.trace_id]`` stamps the
|
||||
caller's nested metadata value as ``litellm.metadata.trace_id`` and a deeper
|
||||
path keeps its dotted name; unlisted siblings stay inside the
|
||||
``metadata.requester_metadata`` blob."""
|
||||
otel = OpenTelemetry(
|
||||
config=OpenTelemetryConfig(
|
||||
baggage_metadata_keys=["requester_metadata.trace_id", "requester_metadata.nested.deep"]
|
||||
)
|
||||
)
|
||||
kwargs = self._kwargs()
|
||||
kwargs["standard_logging_object"]["metadata"]["requester_metadata"] = {
|
||||
"trace_id": "abc",
|
||||
"nested": {"deep": "x", "skipped": "y"},
|
||||
}
|
||||
span, exp = self._span()
|
||||
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
|
||||
attrs = self._attr(span, exp)
|
||||
assert attrs["litellm.metadata.trace_id"] == "abc"
|
||||
assert attrs["litellm.metadata.nested.deep"] == "x"
|
||||
assert "litellm.metadata.deep" not in attrs
|
||||
assert "litellm.metadata.nested.skipped" not in attrs
|
||||
assert not any(k.startswith("litellm.metadata.requester_metadata") for k in attrs)
|
||||
|
||||
def test_metadata_keys_default_to_none_promoted(self):
|
||||
otel = OpenTelemetry()
|
||||
kwargs = self._kwargs()
|
||||
kwargs["standard_logging_object"]["metadata"]["requester_metadata"] = {"trace_id": "abc"}
|
||||
span, exp = self._span()
|
||||
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
|
||||
assert not any(k.startswith("litellm.metadata.") for k in self._attr(span, exp))
|
||||
|
||||
def test_team_metadata_json_helper(self):
|
||||
keys = ["a", "b"]
|
||||
assert OpenTelemetry._team_metadata_json(None, keys) is None
|
||||
|
|
@ -5631,6 +5664,11 @@ class TestOpenTelemetryTeamMetadataKeysConfig(unittest.TestCase):
|
|||
cfg = OpenTelemetryConfig(baggage_team_metadata_keys=["from_arg"])
|
||||
assert cfg.baggage_team_metadata_keys == ["from_arg"]
|
||||
|
||||
def test_metadata_keys_from_kwargs_and_env(self):
|
||||
with patch.dict("os.environ", {"LITELLM_OTEL_BAGGAGE_METADATA_KEYS": "requester_metadata.trace_id, a.b"}):
|
||||
assert OpenTelemetryConfig().baggage_metadata_keys == ["requester_metadata.trace_id", "a.b"]
|
||||
assert OpenTelemetry(baggage_metadata_keys="x.y").config.baggage_metadata_keys == ["x.y"]
|
||||
|
||||
|
||||
class TestOpenTelemetryMetricAttributeFiltering(unittest.TestCase):
|
||||
"""LIT-3600: include/exclude control over which attributes are stamped on
|
||||
|
|
@ -5884,13 +5922,11 @@ class TestOpenTelemetryMetricAttributeFiltering(unittest.TestCase):
|
|||
}
|
||||
)
|
||||
|
||||
def test_no_filter_returns_attrs_object_unchanged(self):
|
||||
"""The no-config path is a hot-path no-op: it returns the same dict
|
||||
object, so default emission pays zero copy cost. Locking identity makes
|
||||
a future refactor that always copies/filters trip here."""
|
||||
def test_no_filter_keeps_every_attribute(self):
|
||||
"""The no-config path drops nothing: every attribute the caller set reaches the meter."""
|
||||
otel = OpenTelemetry(config=OpenTelemetryConfig(exporter="console"))
|
||||
attrs = {"gen_ai.request.model": "m", "hidden_params": "{}"}
|
||||
self.assertIs(otel._filter_metric_attributes(attrs), attrs)
|
||||
self.assertEqual(otel._filter_metric_attributes(attrs), attrs)
|
||||
|
||||
def test_token_type_discriminator_rejected_from_either_list(self):
|
||||
"""gen_ai.token.type is a structural discriminator stamped onto the
|
||||
|
|
@ -6031,6 +6067,118 @@ class TestOTELServiceTierAttributes(unittest.TestCase):
|
|||
self.assertEqual(attributes[self.RESPONSE_KEY], "tier-added-by-provider-later")
|
||||
|
||||
|
||||
class TestOpenTelemetryProviderlessCallAttributes(unittest.TestCase):
|
||||
"""Regression for the OTLP exporter rejecting a None gen_ai.system or gen_ai.request.model
|
||||
attribute on every export cycle."""
|
||||
|
||||
HERE = os.path.dirname(__file__)
|
||||
POLL_INTERVAL = 0.05
|
||||
POLL_TIMEOUT = 2.0
|
||||
|
||||
def _providerless_kwargs(self) -> tuple[dict[str, object], dict[str, object]]:
|
||||
with open(os.path.join(self.HERE, "open_telemetry", "data", "captured_kwargs.json")) as f:
|
||||
kwargs = json.load(f)
|
||||
with open(os.path.join(self.HERE, "open_telemetry", "data", "captured_response.json")) as f:
|
||||
response_obj = json.load(f)
|
||||
kwargs["litellm_params"]["custom_llm_provider"] = None
|
||||
return kwargs, response_obj
|
||||
|
||||
def _modelless_kwargs(self) -> tuple[dict[str, object], dict[str, object]]:
|
||||
kwargs, response_obj = self._providerless_kwargs()
|
||||
kwargs["model"] = None
|
||||
return kwargs, response_obj
|
||||
|
||||
def _recorded_metrics(self, kwargs: dict[str, object], response_obj: dict[str, object]) -> MetricsData | None:
|
||||
metric_reader = InMemoryMetricReader()
|
||||
meter_provider = MeterProvider(metric_readers=[metric_reader])
|
||||
tracer_provider = TracerProvider()
|
||||
tracer_provider.add_span_processor(SimpleSpanProcessor(InMemorySpanExporter()))
|
||||
otel = OpenTelemetry(
|
||||
config=OpenTelemetryConfig(exporter="console", enable_metrics=True),
|
||||
tracer_provider=tracer_provider,
|
||||
meter_provider=meter_provider,
|
||||
)
|
||||
otel.tracer = tracer_provider.get_tracer(__name__)
|
||||
|
||||
start = datetime.utcnow()
|
||||
otel._handle_success(kwargs, response_obj, start, start + timedelta(seconds=1))
|
||||
|
||||
deadline = time.time() + self.POLL_TIMEOUT
|
||||
while time.time() < deadline:
|
||||
data = metric_reader.get_metrics_data()
|
||||
if data and getattr(data, "resource_metrics", None):
|
||||
return data
|
||||
time.sleep(self.POLL_INTERVAL)
|
||||
return None
|
||||
|
||||
def _emitted_log_records(self, semconv_opt_in: str) -> tuple[LogData, ...]:
|
||||
log_exporter = InMemoryLogExporter()
|
||||
logger_provider = OTLoggerProvider()
|
||||
logger_provider.add_log_record_processor(SimpleLogRecordProcessor(log_exporter))
|
||||
with patch.dict(os.environ, {"OTEL_SEMCONV_STABILITY_OPT_IN": semconv_opt_in}):
|
||||
handler = OpenTelemetry(
|
||||
config=OpenTelemetryConfig(exporter="console", enable_events=True),
|
||||
logger_provider=logger_provider,
|
||||
)
|
||||
handler.message_logging = True
|
||||
|
||||
kwargs, response_obj = self._providerless_kwargs()
|
||||
span = handler.tracer.start_span("test")
|
||||
with self.assertNoLogs("opentelemetry.attributes", level="WARNING"):
|
||||
handler._emit_semantic_logs(kwargs, response_obj, span)
|
||||
span.end()
|
||||
handler._logger_provider.force_flush(2000)
|
||||
return log_exporter.get_finished_logs()
|
||||
|
||||
def _assert_every_attribute_encodes(self, attrs: dict[str, object]) -> None:
|
||||
from opentelemetry.exporter.otlp.proto.common._internal import _encode_attributes
|
||||
|
||||
self.assertEqual(len(_encode_attributes(attrs) or []), len(attrs))
|
||||
|
||||
def _recorded_data_points(self, kwargs: dict[str, object], response_obj: dict[str, object]) -> list[object]:
|
||||
data = self._recorded_metrics(kwargs, response_obj)
|
||||
self.assertIsNotNone(data, "no metrics were recorded")
|
||||
data_points = [
|
||||
dp
|
||||
for rm in data.resource_metrics
|
||||
for sm in rm.scope_metrics
|
||||
for m in sm.metrics
|
||||
for dp in m.data.data_points
|
||||
]
|
||||
self.assertTrue(data_points, "no metric data points were recorded")
|
||||
return data_points
|
||||
|
||||
def test_metrics_are_encodable_and_carry_no_provider_label(self):
|
||||
kwargs, response_obj = self._providerless_kwargs()
|
||||
for dp in self._recorded_data_points(kwargs, response_obj):
|
||||
self.assertNotIn("gen_ai.system", dp.attributes)
|
||||
self.assertEqual(dp.attributes["gen_ai.request.model"], kwargs["model"])
|
||||
self._assert_every_attribute_encodes(dict(dp.attributes))
|
||||
|
||||
def test_metrics_are_encodable_and_carry_no_model_label_when_the_call_has_none(self):
|
||||
for dp in self._recorded_data_points(*self._modelless_kwargs()):
|
||||
self.assertNotIn("gen_ai.request.model", dp.attributes)
|
||||
self._assert_every_attribute_encodes(dict(dp.attributes))
|
||||
|
||||
def test_legacy_content_events_are_encodable_and_carry_no_provider_label(self):
|
||||
logs = self._emitted_log_records("")
|
||||
self.assertTrue(logs, "no content events were emitted")
|
||||
for log in logs:
|
||||
attrs = dict(log.log_record.attributes or {})
|
||||
self.assertNotIn("gen_ai.system", attrs)
|
||||
self.assertNotIn(None, attrs.values())
|
||||
self._assert_every_attribute_encodes(attrs)
|
||||
|
||||
def test_inference_details_event_is_encodable_and_carries_no_provider_label(self):
|
||||
logs = self._emitted_log_records("gen_ai_latest_experimental")
|
||||
self.assertEqual(len(logs), 1)
|
||||
attrs = dict(logs[0].log_record.attributes or {})
|
||||
self.assertEqual(attrs["event_name"], "gen_ai.client.inference.operation.details")
|
||||
self.assertNotIn("gen_ai.provider.name", attrs)
|
||||
self.assertNotIn(None, attrs.values())
|
||||
self._assert_every_attribute_encodes(attrs)
|
||||
|
||||
|
||||
class TestDynamicTracerProviderCache(unittest.TestCase):
|
||||
"""Every credential-scoped TracerProvider that owns its exporter also owns a
|
||||
BatchSpanProcessor worker thread that only stops on shutdown, so the cache holding them
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from datetime import datetime, timedelta, timezone
|
||||
from time import monotonic
|
||||
|
||||
import pytest
|
||||
|
|
@ -179,3 +180,50 @@ def test_prometheus_end_user_not_tracked_by_default():
|
|||
|
||||
prometheus_labels = prometheus_label_factory(labels, label_values)
|
||||
assert prometheus_labels["end_user"] is None
|
||||
|
||||
|
||||
def test_prometheus_customer_budget_series_are_capped_per_metric(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_end_user_cost_tracking_prometheus_only", True)
|
||||
monkeypatch.setattr(litellm, "prometheus_end_user_metrics_max_series_per_metric", 2)
|
||||
monkeypatch.setattr(litellm, "prometheus_end_user_metrics_ttl_seconds", None)
|
||||
logger = PrometheusLogger()
|
||||
|
||||
for index in range(5):
|
||||
logger._set_customer_budget_metrics(
|
||||
end_user_id=f"customer-{index}",
|
||||
spend=1.0,
|
||||
max_budget=10.0,
|
||||
budget_reset_at=None,
|
||||
)
|
||||
|
||||
assert set(logger.litellm_remaining_customer_budget_metric._metrics) == {("customer-3",), ("customer-4",)}
|
||||
assert set(logger.litellm_customer_max_budget_metric._metrics) == {("customer-3",), ("customer-4",)}
|
||||
|
||||
|
||||
def test_prometheus_customer_budget_series_expire_by_ttl(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_end_user_cost_tracking_prometheus_only", True)
|
||||
monkeypatch.setattr(litellm, "prometheus_end_user_metrics_max_series_per_metric", None)
|
||||
monkeypatch.setattr(litellm, "prometheus_end_user_metrics_ttl_seconds", 10.0)
|
||||
monkeypatch.setattr(litellm, "prometheus_end_user_metrics_cleanup_interval_seconds", 0.0)
|
||||
logger = PrometheusLogger()
|
||||
|
||||
current_time = [monotonic()]
|
||||
monkeypatch.setattr(bounded_prometheus_series_tracker.time, "monotonic", lambda: current_time[0])
|
||||
logger._set_customer_budget_metrics(
|
||||
end_user_id="customer-with-removed-budget",
|
||||
spend=1.0,
|
||||
max_budget=10.0,
|
||||
budget_reset_at=datetime.now(timezone.utc) + timedelta(hours=1),
|
||||
)
|
||||
|
||||
current_time[0] += 11.0
|
||||
logger._set_customer_budget_metrics(
|
||||
end_user_id="customer-still-budgeted",
|
||||
spend=1.0,
|
||||
max_budget=10.0,
|
||||
budget_reset_at=datetime.now(timezone.utc) + timedelta(hours=1),
|
||||
)
|
||||
|
||||
assert set(logger.litellm_remaining_customer_budget_metric._metrics) == {("customer-still-budgeted",)}
|
||||
assert set(logger.litellm_customer_max_budget_metric._metrics) == {("customer-still-budgeted",)}
|
||||
assert set(logger.litellm_customer_budget_remaining_hours_metric._metrics) == {("customer-still-budgeted",)}
|
||||
|
|
|
|||
|
|
@ -923,6 +923,438 @@ async def test_initialize_org_budget_metrics(prometheus_logger):
|
|||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def customer_metrics_enabled(monkeypatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "enable_end_user_cost_tracking_prometheus_only", True)
|
||||
monkeypatch.setattr(litellm, "disable_end_user_cost_tracking", False)
|
||||
monkeypatch.setattr(litellm, "max_end_user_budget_id", None)
|
||||
|
||||
|
||||
def _customer_sample(metric_name: str, end_user_id: str):
|
||||
return REGISTRY.get_sample_value(metric_name, {"end_user": end_user_id})
|
||||
|
||||
|
||||
def _mock_customer_row(user_id: str, spend: float, max_budget: float | None, budget_reset_at):
|
||||
budget_mock = MagicMock()
|
||||
budget_mock.max_budget = max_budget
|
||||
budget_mock.budget_reset_at = budget_reset_at
|
||||
row = MagicMock()
|
||||
row.user_id = user_id
|
||||
row.spend = spend
|
||||
row.litellm_budget_table = budget_mock
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"spend, max_budget, expected_remaining",
|
||||
[(125.0, 500.0, 375.0), (500.0, 500.0, 0.0), (0.0, 500.0, 500.0)],
|
||||
)
|
||||
def test_set_customer_budget_metrics_emits_remaining_and_max_budget(
|
||||
prometheus_logger, customer_metrics_enabled, spend, max_budget, expected_remaining
|
||||
):
|
||||
prometheus_logger._set_customer_budget_metrics(
|
||||
end_user_id="cust-1",
|
||||
spend=spend,
|
||||
max_budget=max_budget,
|
||||
budget_reset_at=None,
|
||||
)
|
||||
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-1") == pytest.approx(
|
||||
expected_remaining
|
||||
)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-1") == pytest.approx(max_budget)
|
||||
assert _customer_sample("litellm_customer_budget_remaining_hours_metric", "cust-1") is None
|
||||
|
||||
|
||||
def test_set_customer_budget_metrics_remaining_hours(prometheus_logger, customer_metrics_enabled):
|
||||
reset_at = datetime(2099, 1, 1, tzinfo=timezone.utc)
|
||||
prometheus_logger._set_customer_budget_metrics(
|
||||
end_user_id="cust-1",
|
||||
spend=1.0,
|
||||
max_budget=10.0,
|
||||
budget_reset_at=reset_at,
|
||||
)
|
||||
|
||||
expected_hours = (reset_at - datetime.now(timezone.utc)).total_seconds() / 3600
|
||||
assert _customer_sample("litellm_customer_budget_remaining_hours_metric", "cust-1") == pytest.approx(
|
||||
expected_hours, abs=0.1
|
||||
)
|
||||
|
||||
|
||||
def test_set_customer_budget_metrics_not_emitted_when_end_user_tracking_off(prometheus_logger, monkeypatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "enable_end_user_cost_tracking_prometheus_only", False)
|
||||
monkeypatch.setattr(litellm, "disable_end_user_cost_tracking", False)
|
||||
|
||||
prometheus_logger._set_customer_budget_metrics(
|
||||
end_user_id="cust-off",
|
||||
spend=1.0,
|
||||
max_budget=10.0,
|
||||
budget_reset_at=datetime(2099, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
|
||||
assert prometheus_logger.litellm_remaining_customer_budget_metric._metrics == {}
|
||||
assert prometheus_logger.litellm_customer_max_budget_metric._metrics == {}
|
||||
assert prometheus_logger.litellm_customer_budget_remaining_hours_metric._metrics == {}
|
||||
|
||||
|
||||
def test_set_customer_budget_metrics_without_budget_only_emits_remaining(prometheus_logger, customer_metrics_enabled):
|
||||
prometheus_logger._set_customer_budget_metrics(
|
||||
end_user_id="cust-free",
|
||||
spend=3.0,
|
||||
max_budget=None,
|
||||
budget_reset_at=None,
|
||||
)
|
||||
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-free") == float("inf")
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-free") is None
|
||||
assert _customer_sample("litellm_customer_budget_remaining_hours_metric", "cust-free") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_remaining_budget_metrics_emits_customer_gauges_from_cached_end_user(
|
||||
prometheus_logger, customer_metrics_enabled
|
||||
):
|
||||
import sys
|
||||
|
||||
from litellm.models.budget import LiteLLM_BudgetTable
|
||||
from litellm.models.end_user import LiteLLM_EndUserTable
|
||||
|
||||
end_user = LiteLLM_EndUserTable(
|
||||
user_id="cust-req",
|
||||
blocked=False,
|
||||
spend=300.0,
|
||||
budget_id="budget-1",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-1", max_budget=1000.0),
|
||||
)
|
||||
get_end_user_object = AsyncMock()
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = None
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock(return_value=end_user)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}),
|
||||
patch("litellm.proxy.auth.auth_checks.get_end_user_object", get_end_user_object), # test-quality-ok: [TQ008] assert the request path never reaches the DB-backed auth lookup
|
||||
):
|
||||
await prometheus_logger._increment_remaining_budget_metrics(
|
||||
user_api_team=None,
|
||||
user_api_team_alias=None,
|
||||
user_api_key=None,
|
||||
user_api_key_alias=None,
|
||||
litellm_params={"metadata": {}},
|
||||
response_cost=50.0,
|
||||
end_user_id="cust-req",
|
||||
)
|
||||
|
||||
get_end_user_object.assert_not_awaited()
|
||||
cache_read = mock_proxy_server.user_api_key_cache.async_get_cache
|
||||
cache_read.assert_awaited_once()
|
||||
assert cache_read.await_args.kwargs["key"] == "end_user_id:cust-req"
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-req") == pytest.approx(650.0)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-req") == pytest.approx(1000.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_customer_budget_metrics_after_api_request_uses_cached_default_budget(
|
||||
prometheus_logger, customer_metrics_enabled
|
||||
):
|
||||
import sys
|
||||
|
||||
from litellm.models.budget import LiteLLM_BudgetTable
|
||||
from litellm.models.end_user import LiteLLM_EndUserTable
|
||||
|
||||
end_user = LiteLLM_EndUserTable(
|
||||
user_id="cust-default",
|
||||
blocked=False,
|
||||
spend=0.5,
|
||||
budget_id=None,
|
||||
litellm_budget_table=LiteLLM_BudgetTable(budget_id="default-budget", max_budget=3.0),
|
||||
)
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock(return_value=end_user)
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._set_customer_budget_metrics_after_api_request(
|
||||
end_user_id="cust-default",
|
||||
response_cost=0.5,
|
||||
)
|
||||
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-default") == pytest.approx(2.0)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-default") == pytest.approx(3.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_customer_budget_metrics_after_api_request_without_budget_only_emits_remaining(
|
||||
prometheus_logger, customer_metrics_enabled
|
||||
):
|
||||
import sys
|
||||
|
||||
from litellm.models.end_user import LiteLLM_EndUserTable
|
||||
|
||||
end_user = LiteLLM_EndUserTable(user_id="cust-no-budget", blocked=False, spend=2.0, budget_id=None)
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock(return_value=end_user)
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._set_customer_budget_metrics_after_api_request(
|
||||
end_user_id="cust-no-budget",
|
||||
response_cost=1.0,
|
||||
)
|
||||
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-no-budget") == float("inf")
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-no-budget") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_customer_budget_metrics_after_api_request_skips_uncached_customer(
|
||||
prometheus_logger, customer_metrics_enabled
|
||||
):
|
||||
import sys
|
||||
|
||||
get_end_user_object = AsyncMock()
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = MagicMock()
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}),
|
||||
patch("litellm.proxy.auth.auth_checks.get_end_user_object", get_end_user_object), # test-quality-ok: [TQ008] assert a cache miss does not fall back to the DB-backed auth lookup
|
||||
):
|
||||
await prometheus_logger._set_customer_budget_metrics_after_api_request(
|
||||
end_user_id="cust-uncached",
|
||||
response_cost=1.0,
|
||||
)
|
||||
|
||||
get_end_user_object.assert_not_awaited()
|
||||
mock_proxy_server.prisma_client.assert_not_called()
|
||||
assert prometheus_logger.litellm_remaining_customer_budget_metric._metrics == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_customer_budget_metrics_after_api_request_without_end_user_is_noop(prometheus_logger):
|
||||
import sys
|
||||
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock()
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._set_customer_budget_metrics_after_api_request(
|
||||
end_user_id=None,
|
||||
response_cost=1.0,
|
||||
)
|
||||
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache.assert_not_awaited()
|
||||
assert prometheus_logger.litellm_remaining_customer_budget_metric._metrics == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_customer_budget_metrics_after_api_request_skips_cache_when_end_user_tracking_off(
|
||||
prometheus_logger, monkeypatch
|
||||
):
|
||||
import sys
|
||||
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "enable_end_user_cost_tracking_prometheus_only", False)
|
||||
monkeypatch.setattr(litellm, "disable_end_user_cost_tracking", False)
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock()
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._set_customer_budget_metrics_after_api_request(
|
||||
end_user_id="cust-off",
|
||||
response_cost=1.0,
|
||||
)
|
||||
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache.assert_not_awaited()
|
||||
assert prometheus_logger.litellm_remaining_customer_budget_metric._metrics == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_customer_budget_metrics_emits_gauges_for_budgeted_customers(
|
||||
prometheus_logger, customer_metrics_enabled
|
||||
):
|
||||
import sys
|
||||
|
||||
reset_at = datetime(2099, 1, 1, tzinfo=timezone.utc)
|
||||
rows = [
|
||||
_mock_customer_row("cust-a", 100.0, 500.0, None),
|
||||
_mock_customer_row("cust-b", 20.0, 50.0, reset_at),
|
||||
]
|
||||
find_many = AsyncMock(return_value=rows)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = find_many
|
||||
mock_prisma.db.litellm_endusertable.count = AsyncMock(return_value=len(rows))
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = mock_prisma
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._initialize_customer_budget_metrics()
|
||||
|
||||
assert find_many.await_args.kwargs["where"] == {"budget_id": {"not": None}}
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-a") == pytest.approx(400.0)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-a") == pytest.approx(500.0)
|
||||
assert _customer_sample("litellm_customer_budget_remaining_hours_metric", "cust-a") is None
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-b") == pytest.approx(30.0)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-b") == pytest.approx(50.0)
|
||||
assert _customer_sample("litellm_customer_budget_remaining_hours_metric", "cust-b") > 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"enable_prometheus_only, disable_end_user",
|
||||
[(False, False), (True, True)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_customer_budget_metrics_skips_when_end_user_tracking_off(
|
||||
prometheus_logger, monkeypatch, enable_prometheus_only, disable_end_user
|
||||
):
|
||||
import sys
|
||||
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "enable_end_user_cost_tracking_prometheus_only", enable_prometheus_only)
|
||||
monkeypatch.setattr(litellm, "disable_end_user_cost_tracking", disable_end_user)
|
||||
|
||||
find_many = AsyncMock(return_value=[_mock_customer_row("cust-a", 100.0, 500.0, None)])
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = find_many
|
||||
mock_prisma.db.litellm_endusertable.count = AsyncMock(return_value=1)
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = mock_prisma
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._initialize_customer_budget_metrics()
|
||||
|
||||
find_many.assert_not_awaited()
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-a") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_remaining_budget_metrics_includes_customers(prometheus_logger, customer_metrics_enabled):
|
||||
import sys
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(
|
||||
return_value=[_mock_customer_row("cust-startup", 5.0, 25.0, None)]
|
||||
)
|
||||
mock_prisma.db.litellm_endusertable.count = AsyncMock(return_value=1)
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = mock_prisma
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._initialize_remaining_budget_metrics()
|
||||
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-startup") == pytest.approx(20.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_customer_budget_metrics_counts_once_across_pages(prometheus_logger, customer_metrics_enabled):
|
||||
import sys
|
||||
|
||||
pages = [
|
||||
[_mock_customer_row(f"cust-{i}", 1.0, 10.0, None) for i in range(50)],
|
||||
[_mock_customer_row(f"cust-{i}", 1.0, 10.0, None) for i in range(50, 100)],
|
||||
[_mock_customer_row("cust-100", 1.0, 10.0, None)],
|
||||
]
|
||||
find_many = AsyncMock(side_effect=pages)
|
||||
count = AsyncMock(return_value=101)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = find_many
|
||||
mock_prisma.db.litellm_endusertable.count = count
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = mock_prisma
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._initialize_customer_budget_metrics()
|
||||
|
||||
assert find_many.await_count == 3
|
||||
count.assert_awaited_once()
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-100") == pytest.approx(9.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_customer_budget_metrics_applies_default_budget_to_unbudgeted_customers(
|
||||
prometheus_logger, customer_metrics_enabled, monkeypatch
|
||||
):
|
||||
import sys
|
||||
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "max_end_user_budget_id", "default-customer-budget")
|
||||
reset_at = datetime(2099, 1, 1, tzinfo=timezone.utc)
|
||||
default_budget = MagicMock()
|
||||
default_budget.max_budget = 10.0
|
||||
default_budget.budget_reset_at = reset_at
|
||||
explicit_row = _mock_customer_row("cust-explicit", 5.0, 100.0, None)
|
||||
default_row = _mock_customer_row("cust-default", 2.0, None, None)
|
||||
default_row.litellm_budget_table = None
|
||||
find_many = AsyncMock(return_value=[explicit_row, default_row])
|
||||
find_unique = AsyncMock(return_value=default_budget)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = find_many
|
||||
mock_prisma.db.litellm_endusertable.count = AsyncMock(return_value=2)
|
||||
mock_prisma.db.litellm_budgettable.find_unique = find_unique
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = mock_prisma
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await prometheus_logger._initialize_customer_budget_metrics()
|
||||
|
||||
assert find_unique.await_args.kwargs["where"] == {"budget_id": "default-customer-budget"}
|
||||
assert find_many.await_args.kwargs["where"] is None
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-explicit") == pytest.approx(95.0)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-explicit") == pytest.approx(100.0)
|
||||
assert _customer_sample("litellm_remaining_customer_budget_metric", "cust-default") == pytest.approx(8.0)
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-default") == pytest.approx(10.0)
|
||||
assert _customer_sample("litellm_customer_budget_remaining_hours_metric", "cust-default") > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_customer_max_budget_gauge_emitted_when_only_it_is_configured(customer_metrics_enabled, monkeypatch):
|
||||
import sys
|
||||
|
||||
import litellm
|
||||
from litellm.models.budget import LiteLLM_BudgetTable
|
||||
from litellm.models.end_user import LiteLLM_EndUserTable
|
||||
from litellm.types.integrations.prometheus import NoOpMetric
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"prometheus_metrics_config",
|
||||
[{"group": "customer-max-only", "metrics": ["litellm_customer_max_budget_metric"]}],
|
||||
)
|
||||
logger = PrometheusLogger()
|
||||
assert isinstance(logger.litellm_remaining_customer_budget_metric, NoOpMetric)
|
||||
assert not isinstance(logger.litellm_customer_max_budget_metric, NoOpMetric)
|
||||
|
||||
end_user = LiteLLM_EndUserTable(
|
||||
user_id="cust-max-only",
|
||||
blocked=False,
|
||||
spend=1.0,
|
||||
budget_id="budget-1",
|
||||
litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget-1", max_budget=40.0),
|
||||
)
|
||||
mock_proxy_server = MagicMock()
|
||||
mock_proxy_server.prisma_client = None
|
||||
mock_proxy_server.user_api_key_cache.async_get_cache = AsyncMock(return_value=end_user)
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}):
|
||||
await logger._increment_remaining_budget_metrics(
|
||||
user_api_team=None,
|
||||
user_api_team_alias=None,
|
||||
user_api_key=None,
|
||||
user_api_key_alias=None,
|
||||
litellm_params={"metadata": {}},
|
||||
response_cost=1.0,
|
||||
end_user_id="cust-max-only",
|
||||
)
|
||||
|
||||
assert _customer_sample("litellm_customer_max_budget_metric", "cust-max-only") == pytest.approx(40.0)
|
||||
|
||||
|
||||
def test_default_latency_buckets(prometheus_logger):
|
||||
"""PrometheusLogger uses the new reduced default latency buckets."""
|
||||
from litellm.types.integrations.prometheus import LATENCY_BUCKETS
|
||||
|
|
|
|||
|
|
@ -1,16 +1,24 @@
|
|||
import copy
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.constants import MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES, MAX_S3_OBJECT_KEY_BYTES
|
||||
from litellm.integrations.s3 import S3Logger
|
||||
from litellm.integrations.s3 import S3Logger, prompts_only_payload, resolve_s3_log_prompts_only
|
||||
|
||||
TEST_KMS_KEY_ARN = "arn:aws:kms:us-east-1:111122223333:key/test-key-id"
|
||||
TEST_MESSAGES = [{"role": "user", "content": "Reply with exactly the word PINEAPPLE."}]
|
||||
TEST_RESPONSE = {"choices": [{"message": {"role": "assistant", "content": "PINEAPPLE"}}]}
|
||||
|
||||
|
||||
def _standard_logging_payload(response_id: str = "chatcmpl-test-id") -> dict:
|
||||
return {
|
||||
"id": response_id,
|
||||
"messages": copy.deepcopy(TEST_MESSAGES),
|
||||
"response": copy.deepcopy(TEST_RESPONSE),
|
||||
"metadata": {"user_api_key_team_alias": None},
|
||||
}
|
||||
|
||||
|
|
@ -22,7 +30,9 @@ def _log_event_kwargs(response_id: str = "chatcmpl-test-id") -> dict:
|
|||
}
|
||||
|
||||
|
||||
def _run_log_event(callback_params: dict, response_id: str = "chatcmpl-test-id") -> MagicMock:
|
||||
def _run_log_event(
|
||||
callback_params: dict, response_id: str = "chatcmpl-test-id", log_kwargs: dict[str, object] | None = None
|
||||
) -> MagicMock:
|
||||
original = litellm.s3_callback_params
|
||||
litellm.s3_callback_params = callback_params
|
||||
try:
|
||||
|
|
@ -31,7 +41,7 @@ def _run_log_event(callback_params: dict, response_id: str = "chatcmpl-test-id")
|
|||
mock_boto3_client.return_value = mock_s3_client
|
||||
logger = S3Logger()
|
||||
logger.log_event(
|
||||
kwargs=_log_event_kwargs(response_id),
|
||||
kwargs=_log_event_kwargs(response_id) if log_kwargs is None else log_kwargs,
|
||||
response_obj={"id": response_id},
|
||||
start_time=datetime(2026, 7, 30, 12, 0, 0),
|
||||
end_time=datetime(2026, 7, 30, 12, 0, 1),
|
||||
|
|
@ -182,3 +192,123 @@ def test_put_object_keeps_the_configured_path_intact_when_only_the_id_has_to_shr
|
|||
key = mock_s3_client.put_object.call_args.kwargs["Key"]
|
||||
assert key.startswith(long_path + "/2026-07-30/")
|
||||
assert len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES
|
||||
|
||||
|
||||
def _uploaded_body(mock_s3_client: MagicMock) -> dict[str, object]:
|
||||
return json.loads(mock_s3_client.put_object.call_args.kwargs["Body"])
|
||||
|
||||
|
||||
def test_log_event_prompts_only_drops_response_and_keeps_messages(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||||
log_kwargs = _log_event_kwargs()
|
||||
original_payload = copy.deepcopy(log_kwargs["standard_logging_object"])
|
||||
|
||||
mock_s3_client = _run_log_event(
|
||||
{"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1", "s3_log_prompts_only": True},
|
||||
log_kwargs=log_kwargs,
|
||||
)
|
||||
|
||||
body = _uploaded_body(mock_s3_client)
|
||||
assert body["messages"] == TEST_MESSAGES
|
||||
assert body["response"] is None
|
||||
assert body["id"] == "chatcmpl-test-id"
|
||||
assert log_kwargs["standard_logging_object"] == original_payload
|
||||
|
||||
|
||||
def test_log_event_default_keeps_response(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||||
|
||||
mock_s3_client = _run_log_event({"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1"})
|
||||
|
||||
body = _uploaded_body(mock_s3_client)
|
||||
assert body["response"] == TEST_RESPONSE
|
||||
assert body["messages"] == TEST_MESSAGES
|
||||
|
||||
|
||||
def test_log_event_reads_prompts_only_env_var_at_log_time(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||||
original = litellm.s3_callback_params
|
||||
litellm.s3_callback_params = {"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1"}
|
||||
try:
|
||||
with patch("boto3.client") as mock_boto3_client:
|
||||
mock_s3_client = MagicMock()
|
||||
mock_boto3_client.return_value = mock_s3_client
|
||||
logger = S3Logger()
|
||||
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
|
||||
logger.log_event(
|
||||
kwargs=_log_event_kwargs(),
|
||||
response_obj={"id": "chatcmpl-test-id"},
|
||||
start_time=datetime(2026, 7, 30, 12, 0, 0),
|
||||
end_time=datetime(2026, 7, 30, 12, 0, 1),
|
||||
print_verbose=lambda *args, **kwargs: None,
|
||||
)
|
||||
finally:
|
||||
litellm.s3_callback_params = original
|
||||
|
||||
body = _uploaded_body(mock_s3_client)
|
||||
assert body["response"] is None
|
||||
assert body["messages"] == TEST_MESSAGES
|
||||
|
||||
|
||||
def test_log_event_explicit_false_param_beats_env_var(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
|
||||
|
||||
mock_s3_client = _run_log_event(
|
||||
{"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1", "s3_log_prompts_only": False}
|
||||
)
|
||||
|
||||
assert _uploaded_body(mock_s3_client)["response"] == TEST_RESPONSE
|
||||
|
||||
|
||||
def test_s3_logger_init_does_not_mutate_global_callback_params(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("MY_S3_BUCKET", "resolved-bucket")
|
||||
callback_params = {"s3_bucket_name": "os.environ/MY_S3_BUCKET", "s3_region_name": "us-east-1"}
|
||||
snapshot = copy.deepcopy(callback_params)
|
||||
original = litellm.s3_callback_params
|
||||
litellm.s3_callback_params = callback_params
|
||||
try:
|
||||
with patch("boto3.client"):
|
||||
logger = S3Logger()
|
||||
finally:
|
||||
litellm.s3_callback_params = original
|
||||
|
||||
assert logger.bucket_name == "resolved-bucket"
|
||||
assert callback_params == snapshot
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"configured,env_value,expected",
|
||||
[
|
||||
(True, None, True),
|
||||
(False, "true", False),
|
||||
("true", None, True),
|
||||
("False", "true", False),
|
||||
("1", None, True),
|
||||
("0", None, False),
|
||||
(" yes ", None, True),
|
||||
(None, None, False),
|
||||
(None, "true", True),
|
||||
(None, "false", False),
|
||||
(None, "", False),
|
||||
("", "true", False),
|
||||
],
|
||||
)
|
||||
def test_resolve_s3_log_prompts_only(configured: object, env_value: str | None, expected: bool):
|
||||
environ = {} if env_value is None else {"S3_LOG_PROMPTS_ONLY": env_value}
|
||||
assert resolve_s3_log_prompts_only(configured, environ) is expected
|
||||
|
||||
|
||||
def test_resolve_s3_log_prompts_only_unparseable_value_fails_toward_prompts_only():
|
||||
assert resolve_s3_log_prompts_only("enabled", {}) is True
|
||||
|
||||
|
||||
def test_prompts_only_payload_returns_copy_with_response_cleared():
|
||||
payload = _standard_logging_payload()
|
||||
snapshot = copy.deepcopy(payload)
|
||||
|
||||
stripped = prompts_only_payload(payload)
|
||||
|
||||
assert stripped["response"] is None
|
||||
assert stripped["messages"] == TEST_MESSAGES
|
||||
assert stripped is not payload
|
||||
assert payload == snapshot
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import re
|
||||
import sys
|
||||
import textwrap
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
|
@ -10,6 +13,7 @@ from unittest.mock import AsyncMock, MagicMock, call, patch
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
|
@ -2310,3 +2314,137 @@ def _s3_logger_for_region(region_name: str) -> S3Logger:
|
|||
)
|
||||
def test_build_object_url_uses_partition_dns_suffix(region_name: str, expected_url: str) -> None:
|
||||
assert _s3_logger_for_region(region_name)._build_object_url("2025-01-01/key.json") == expected_url
|
||||
|
||||
|
||||
def _prompts_only_logger(s3_log_prompts_only: bool | None = None) -> S3Logger:
|
||||
return S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1",
|
||||
s3_log_prompts_only=s3_log_prompts_only,
|
||||
)
|
||||
|
||||
|
||||
def _chat_payload() -> StandardLoggingPayload:
|
||||
return StandardLoggingPayload(
|
||||
id="chatcmpl-prompts-only",
|
||||
messages=[{"role": "user", "content": "Reply with exactly the word PINEAPPLE."}],
|
||||
response={"choices": [{"message": {"role": "assistant", "content": "PINEAPPLE"}}]},
|
||||
metadata={"user_api_key_team_alias": None},
|
||||
)
|
||||
|
||||
|
||||
async def _queued_body_via_async_upload(
|
||||
logger: S3Logger, log_event: Callable[..., Awaitable[None]]
|
||||
) -> dict[str, object]:
|
||||
payload = _chat_payload()
|
||||
original = copy.deepcopy(payload)
|
||||
await log_event(
|
||||
kwargs={"standard_logging_object": payload},
|
||||
response_obj=None,
|
||||
start_time=datetime(2026, 7, 30, 12, 0, 0),
|
||||
end_time=datetime(2026, 7, 30, 12, 0, 1),
|
||||
)
|
||||
assert payload == original, "the caller's standard_logging_object must not be mutated"
|
||||
(element,) = logger.log_queue
|
||||
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.raise_for_status = MagicMock()
|
||||
logger.async_httpx_client = AsyncMock()
|
||||
logger.async_httpx_client.put.return_value = response
|
||||
await logger.async_upload_data_to_s3(element)
|
||||
return json.loads(logger.async_httpx_client.put.call_args.kwargs["data"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("event_name", ["async_log_success_event", "async_log_failure_event"])
|
||||
async def test_prompts_only_drops_response_but_keeps_messages_in_uploaded_object(
|
||||
monkeypatch: pytest.MonkeyPatch, event_name: str
|
||||
):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_log_prompts_only": True})
|
||||
logger = _prompts_only_logger()
|
||||
|
||||
log_event: Callable[..., Awaitable[None]] = (
|
||||
logger.async_log_success_event if event_name == "async_log_success_event" else logger.async_log_failure_event
|
||||
)
|
||||
body = await _queued_body_via_async_upload(logger, log_event)
|
||||
|
||||
assert body["messages"] == _chat_payload()["messages"]
|
||||
assert body["response"] is None
|
||||
assert body["id"] == "chatcmpl-prompts-only"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompts_only_default_off_keeps_response_in_uploaded_object(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {})
|
||||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||||
logger = _prompts_only_logger()
|
||||
|
||||
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
|
||||
|
||||
assert body["response"] == _chat_payload()["response"]
|
||||
assert body["messages"] == _chat_payload()["messages"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompts_only_explicit_false_in_params_beats_env_var(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_log_prompts_only": False})
|
||||
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
|
||||
logger = _prompts_only_logger()
|
||||
|
||||
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
|
||||
|
||||
assert body["response"] == _chat_payload()["response"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompts_only_env_var_applies_when_param_unset(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {})
|
||||
logger = _prompts_only_logger()
|
||||
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
|
||||
|
||||
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
|
||||
|
||||
assert body["response"] is None
|
||||
assert body["messages"] == _chat_payload()["messages"]
|
||||
|
||||
|
||||
@respx.mock
|
||||
def test_prompts_only_constructor_kwarg_applies_to_sync_upload(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm
|
||||
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {})
|
||||
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
|
||||
logger = _prompts_only_logger(s3_log_prompts_only=True)
|
||||
payload = _chat_payload()
|
||||
|
||||
element = logger.create_s3_batch_logging_element(
|
||||
start_time=datetime(2026, 7, 30, 12, 0, 0),
|
||||
standard_logging_payload=payload,
|
||||
)
|
||||
assert element is not None
|
||||
assert payload["response"] == _chat_payload()["response"]
|
||||
|
||||
put_route = respx.put(url__regex=r"https://test-bucket\.s3\..*").mock(return_value=httpx.Response(200))
|
||||
logger.upload_data_to_s3(element)
|
||||
|
||||
body = json.loads(put_route.calls.last.request.content)
|
||||
assert body["response"] is None
|
||||
assert body["messages"] == _chat_payload()["messages"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("callback_name", ["s3", "s3_v2"])
|
||||
def test_prompts_only_toggle_is_exposed_to_admin_ui_for_both_s3_callbacks(callback_name: str):
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
assert "S3_LOG_PROMPTS_ONLY" in CustomLogger.get_callback_env_vars(callback_name)
|
||||
|
|
|
|||
|
|
@ -2482,7 +2482,10 @@ class TestAnthropicMessagesHandlerStreamingScanKey:
|
|||
open_key = handler.get_streaming_scan_key([self._text_delta("hi"), tool_use])
|
||||
ended_key = handler.get_streaming_scan_key([self._text_delta("hi"), tool_use, self._stop("tool_use")])
|
||||
assert open_key == StreamingScanKey(texts=("hi",))
|
||||
assert open_key.tool_calls_in_flight is True
|
||||
assert handler.get_streaming_scan_key([self._text_delta("hi")]).tool_calls_in_flight is False
|
||||
assert len(ended_key.tool_calls) == 1 and "get_weather" in ended_key.tool_calls[0]
|
||||
assert ended_key.tool_calls_in_flight is False
|
||||
assert ended_key != open_key
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2719,3 +2719,74 @@ class TestRustChatCompletionsHook:
|
|||
"model": "m",
|
||||
"messages": [],
|
||||
}
|
||||
|
||||
|
||||
def _served_model_stream_chunks(model: str | None) -> list[dict[str, object]]:
|
||||
return [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_served",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 1},
|
||||
**({"model": model} if model is not None else {}),
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "Hello"},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 2},
|
||||
},
|
||||
{"type": "message_stop"},
|
||||
]
|
||||
|
||||
|
||||
def test_message_start_model_is_carried_on_stream_chunks():
|
||||
iterator: Final = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
parsed: Final = [iterator.chunk_parser(chunk) for chunk in _served_model_stream_chunks("claude-served-1")]
|
||||
|
||||
assert all(chunk.model == "claude-served-1" for chunk in parsed)
|
||||
|
||||
|
||||
def test_message_start_without_model_leaves_chunk_model_unset():
|
||||
iterator: Final = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
parsed: Final = [iterator.chunk_parser(chunk) for chunk in _served_model_stream_chunks(None)]
|
||||
|
||||
assert all(chunk.model is None for chunk in parsed)
|
||||
|
||||
|
||||
def test_served_model_reaches_assembled_stream_through_custom_stream_wrapper():
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
served_model: Final = "claude-served-1"
|
||||
sse_lines: Final = [f"data: {json.dumps(chunk)}\n".encode() for chunk in _served_model_stream_chunks(served_model)]
|
||||
iterator: Final = ModelResponseIterator(iter(sse_lines), sync_stream=True)
|
||||
wrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=iter(iterator),
|
||||
model="anthropic/claude-requested",
|
||||
custom_llm_provider="anthropic",
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
chunks: Final = list(wrapper)
|
||||
|
||||
assert len(chunks) > 1
|
||||
for chunk in chunks[1:]:
|
||||
assert chunk._hidden_params["provider_response_model"] == served_model
|
||||
assembled: Final = litellm.stream_chunk_builder(chunks=list(chunks), messages=[{"role": "user", "content": "hi"}])
|
||||
assert assembled._hidden_params["provider_response_model"] == served_model
|
||||
|
|
|
|||
|
|
@ -1200,6 +1200,7 @@ def _fake_user_api_key_auth(
|
|||
team_models=None,
|
||||
team_id=None,
|
||||
model_max_budget=None,
|
||||
team_model_max_budget=None,
|
||||
end_user_model_max_budget=None,
|
||||
end_user_id=None,
|
||||
user_model_max_budget=None,
|
||||
|
|
@ -1220,6 +1221,7 @@ def _fake_user_api_key_auth(
|
|||
auth.team_id = team_id
|
||||
auth.team_model_aliases = None
|
||||
auth.model_max_budget = model_max_budget
|
||||
auth.team_model_max_budget = team_model_max_budget
|
||||
auth.end_user_model_max_budget = end_user_model_max_budget
|
||||
auth.end_user_id = end_user_id
|
||||
auth.user_model_max_budget = user_model_max_budget
|
||||
|
|
@ -1860,6 +1862,78 @@ async def test_summary_model_rate_limit_skipped_for_legacy_limiter():
|
|||
assert not result.applied_edits[0].get("error")
|
||||
|
||||
|
||||
async def test_summary_model_denied_when_team_over_model_budget():
|
||||
"""The team per-model budget gates the summary subrequest, whose spend is
|
||||
charged to the team counter via the propagated `user_api_key_team_model_max_budget`.
|
||||
The key's own `model_max_budget` is handed to the limiter so a key-level
|
||||
override keeps taking precedence over the team cap here as it does in auth."""
|
||||
import litellm
|
||||
|
||||
messages = _simple_messages()
|
||||
mock_call = AsyncMock(return_value=_make_mock_response("<summary>x</summary>"))
|
||||
key_budget = {"claude-opus-4-8": {"budget_limit": 1}}
|
||||
team_budget = {"claude-haiku-4-5": {"budget_limit": 5, "time_period": "1d"}}
|
||||
|
||||
auth = _fake_user_api_key_auth(
|
||||
key_models=["all-proxy-models"],
|
||||
model_max_budget=key_budget,
|
||||
team_model_max_budget=team_budget,
|
||||
team_id="team-over-budget",
|
||||
token="hashed-token",
|
||||
)
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.is_key_within_model_budget = AsyncMock(return_value=True)
|
||||
limiter.is_team_within_model_budget = AsyncMock(
|
||||
side_effect=litellm.BudgetExceededError(
|
||||
message="over budget", current_cost=10, max_budget=5
|
||||
)
|
||||
)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: apply_compact_20260112 reads the summary model setting as a module global, no seam
|
||||
"litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact._read_summary_model_setting",
|
||||
return_value="claude-haiku-4-5",
|
||||
),
|
||||
patch("litellm.token_counter", return_value=200_000), # test-quality-ok: forces the over-threshold branch
|
||||
patch( # test-quality-ok: the summary call is the observable that must NOT happen when the team is over budget
|
||||
"litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact._call_summary_model",
|
||||
mock_call,
|
||||
),
|
||||
patch( # test-quality-ok: the limiter is a proxy_server module global the editor imports, no injection seam
|
||||
"litellm.proxy.proxy_server.model_max_budget_limiter", limiter
|
||||
),
|
||||
):
|
||||
result = await apply_compact_20260112(
|
||||
model=MODEL,
|
||||
messages=messages,
|
||||
tools=None,
|
||||
system=None,
|
||||
edit_spec=_EDIT_SPEC_DEFAULT,
|
||||
user_api_key_auth=auth,
|
||||
)
|
||||
|
||||
mock_call.assert_not_awaited()
|
||||
assert result.applied_edits[0].get("error") == "summary_model_budget_exceeded"
|
||||
limiter.is_team_within_model_budget.assert_awaited_once_with(
|
||||
team_id="team-over-budget",
|
||||
team_model_max_budget=team_budget,
|
||||
key_model_max_budget=key_budget,
|
||||
model="claude-haiku-4-5",
|
||||
)
|
||||
import inspect
|
||||
|
||||
from litellm.proxy.hooks.model_max_budget_limiter import (
|
||||
_PROXY_VirtualKeyModelMaxBudgetLimiter,
|
||||
)
|
||||
|
||||
real_params = inspect.signature(
|
||||
_PROXY_VirtualKeyModelMaxBudgetLimiter.is_team_within_model_budget
|
||||
).parameters
|
||||
for kwarg in ("team_id", "team_model_max_budget", "key_model_max_budget", "model"):
|
||||
assert kwarg in real_params, f"compact.py passes {kwarg}=, which the limiter does not accept"
|
||||
|
||||
|
||||
async def test_scoped_budget_metadata_propagated_to_summary_call():
|
||||
"""The end-user/project scope identifiers and the end-user budget the post-call
|
||||
spend and rate-limit hooks key on are forwarded to the summary subrequest, and
|
||||
|
|
|
|||
|
|
@ -0,0 +1,95 @@
|
|||
import json
|
||||
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.amazon_nova_transformation import (
|
||||
AmazonInvokeNovaConfig,
|
||||
)
|
||||
from litellm.types.integrations.anthropic_cache_control_hook import GATEWAY_INJECTED_CACHE_METADATA_KEY
|
||||
|
||||
MODEL = "us.amazon.nova-pro-v1:0"
|
||||
EPHEMERAL = {"type": "ephemeral"}
|
||||
DEFAULT_CACHE_POINT = {"type": "default"}
|
||||
TOOLS = [{"type": "function", "function": {"name": "f", "parameters": {"type": "object", "properties": {}}}}]
|
||||
TOOL_CALL = {"id": "call_1", "type": "function", "function": {"name": "f", "arguments": "{}"}}
|
||||
PNG_DATA_URL = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
|
||||
|
||||
|
||||
def _transform_request(messages, optional_params, litellm_params=None):
|
||||
return AmazonInvokeNovaConfig().transform_request(
|
||||
model=MODEL,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params if litellm_params is not None else {},
|
||||
headers={},
|
||||
)
|
||||
|
||||
|
||||
def test_cache_points_are_inlined_into_the_block_they_cache(local_model_cost_map):
|
||||
"""InvokeModel rejects the standalone ``{"cachePoint": ...}`` block Converse emits
|
||||
(``#/system/1: required key [text] not found``); it wants ``cachePoint`` as a key of the
|
||||
block being cached."""
|
||||
request = _transform_request(
|
||||
messages=[
|
||||
{"role": "system", "content": [{"type": "text", "text": "long system prompt", "cache_control": EPHEMERAL}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "hello", "cache_control": EPHEMERAL}]},
|
||||
{"role": "assistant", "content": "hi there", "cache_control": EPHEMERAL},
|
||||
{"role": "user", "content": "again"},
|
||||
],
|
||||
optional_params={"max_tokens": 20},
|
||||
)
|
||||
assert request["system"] == [{"text": "long system prompt", "cachePoint": DEFAULT_CACHE_POINT}]
|
||||
assert [message["content"] for message in request["messages"]] == [
|
||||
[{"text": "hello", "cachePoint": DEFAULT_CACHE_POINT}],
|
||||
[{"text": "hi there", "cachePoint": DEFAULT_CACHE_POINT}],
|
||||
[{"text": "again"}],
|
||||
]
|
||||
|
||||
|
||||
def test_cache_point_behind_a_non_text_block_moves_back_to_the_last_text_block(local_model_cost_map):
|
||||
"""InvokeModel rejects ``cachePoint`` on image, toolUse, and toolResult blocks
|
||||
(``extraneous key [cachePoint] is not permitted``), so the point a user put on an image or a
|
||||
tool result lands on the closest text block before it, and a message with no text block at
|
||||
all sends no point rather than a request AWS refuses.
|
||||
"""
|
||||
request = _transform_request(
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "what is in this picture?"},
|
||||
{"type": "image_url", "image_url": {"url": PNG_DATA_URL}, "cache_control": EPHEMERAL},
|
||||
],
|
||||
},
|
||||
{"role": "assistant", "content": None, "tool_calls": [TOOL_CALL]},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "sunny", "cache_control": EPHEMERAL},
|
||||
],
|
||||
optional_params={"tools": TOOLS},
|
||||
)
|
||||
picture, image = request["messages"][0]["content"]
|
||||
assert picture == {"text": "what is in this picture?", "cachePoint": DEFAULT_CACHE_POINT}
|
||||
assert set(image) == {"image"}
|
||||
assert [set(block) for block in request["messages"][2]["content"]] == [{"toolResult"}]
|
||||
|
||||
|
||||
def test_cache_point_with_nothing_before_it_is_dropped():
|
||||
request = AmazonInvokeNovaConfig._inline_cache_points(
|
||||
{
|
||||
"system": [{"cachePoint": DEFAULT_CACHE_POINT}],
|
||||
"messages": [{"role": "user", "content": [{"cachePoint": DEFAULT_CACHE_POINT}, {"text": "hi"}]}],
|
||||
}
|
||||
)
|
||||
assert request["system"] == []
|
||||
assert request["messages"] == [{"role": "user", "content": [{"text": "hi"}]}]
|
||||
|
||||
|
||||
def test_tool_config_injection_point_is_neither_placed_nor_credited(local_model_cost_map):
|
||||
"""InvokeModel has no tool caching, so the point cannot land and the gateway must not be
|
||||
credited for it in spend attribution."""
|
||||
metadata = {"user_api_key": "sk-test"}
|
||||
request = _transform_request(
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"tools": TOOLS, "cache_control_injection_points": [{"location": "tool_config"}]},
|
||||
litellm_params={"metadata": metadata, "litellm_metadata": None, "model_info": {"id": "dep-bedrock"}},
|
||||
)
|
||||
assert [tool["toolSpec"]["name"] for tool in request["toolConfig"]["tools"]] == ["f"]
|
||||
assert "cachePoint" not in json.dumps(request)
|
||||
assert GATEWAY_INJECTED_CACHE_METADATA_KEY not in metadata
|
||||
|
|
@ -139,6 +139,118 @@ def test_bedrock_converse_1h_cache_write_billed_at_1h_rate(monkeypatch):
|
|||
assert completion_cost == pytest.approx(4 * model_info["output_cost_per_token"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"usage, expected_prompt_tokens, expected_cached_tokens, expected_cache_creation_tokens",
|
||||
[
|
||||
pytest.param(
|
||||
{
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 3,
|
||||
"totalTokens": 12270,
|
||||
"cacheReadInputTokenCount": 12262,
|
||||
"cacheWriteInputTokenCount": 0,
|
||||
},
|
||||
12267,
|
||||
12262,
|
||||
0,
|
||||
id="invoke-model-cache-read",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 3,
|
||||
"totalTokens": 12270,
|
||||
"cacheReadInputTokenCount": 0,
|
||||
"cacheWriteInputTokenCount": 12262,
|
||||
},
|
||||
12267,
|
||||
0,
|
||||
12262,
|
||||
id="invoke-model-cache-write",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 3,
|
||||
"cacheReadInputTokenCount": 12262,
|
||||
"cacheWriteInputTokenCount": 0,
|
||||
},
|
||||
12267,
|
||||
12262,
|
||||
0,
|
||||
id="invoke-model-streaming-metadata-without-totalTokens",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_transform_usage_reads_invoke_model_count_suffixed_cache_keys(
|
||||
usage, expected_prompt_tokens, expected_cached_tokens, expected_cache_creation_tokens
|
||||
):
|
||||
"""InvokeModel Nova reports ``cacheReadInputTokenCount`` and ``cacheWriteInputTokenCount``
|
||||
where Converse reports the un-suffixed keys, and ``inputTokens`` excludes both."""
|
||||
openai_usage = AmazonConverseConfig().transform_usage(ConverseTokenUsageBlock(**usage))
|
||||
assert openai_usage.prompt_tokens == expected_prompt_tokens
|
||||
assert openai_usage.prompt_tokens_details.cached_tokens == expected_cached_tokens
|
||||
assert openai_usage._cache_read_input_tokens == expected_cached_tokens
|
||||
assert openai_usage._cache_creation_input_tokens == expected_cache_creation_tokens
|
||||
assert openai_usage.completion_tokens == 3
|
||||
assert openai_usage.total_tokens == 12270
|
||||
|
||||
|
||||
def test_bedrock_invoke_nova_cache_read_billed_at_discounted_rate(monkeypatch):
|
||||
"""Nova cache reads are billed at the entry's discounted cache read rate; without a
|
||||
``cache_read_input_token_cost`` entry the cached tokens were billed at nothing."""
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
usage = ConverseTokenUsageBlock(
|
||||
**{
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 3,
|
||||
"totalTokens": 12270,
|
||||
"cacheReadInputTokenCount": 12262,
|
||||
"cacheWriteInputTokenCount": 0,
|
||||
}
|
||||
)
|
||||
openai_usage = AmazonConverseConfig().transform_usage(usage)
|
||||
model = "bedrock/invoke/us.amazon.nova-pro-v1:0"
|
||||
prompt_cost, completion_cost = litellm.cost_calculator.cost_per_token(model=model, usage_object=openai_usage)
|
||||
model_info = litellm.get_model_info(model=model)
|
||||
assert 0 < model_info["cache_read_input_token_cost"] < model_info["input_cost_per_token"]
|
||||
assert prompt_cost == pytest.approx(
|
||||
5 * model_info["input_cost_per_token"] + 12262 * model_info["cache_read_input_token_cost"]
|
||||
)
|
||||
assert prompt_cost > 5 * model_info["input_cost_per_token"]
|
||||
assert completion_cost == pytest.approx(3 * model_info["output_cost_per_token"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"amazon.nova-micro-v1:0",
|
||||
"amazon.nova-lite-v1:0",
|
||||
"amazon.nova-pro-v1:0",
|
||||
"us.amazon.nova-micro-v1:0",
|
||||
"us.amazon.nova-lite-v1:0",
|
||||
"us.amazon.nova-pro-v1:0",
|
||||
"eu.amazon.nova-micro-v1:0",
|
||||
"eu.amazon.nova-lite-v1:0",
|
||||
"eu.amazon.nova-pro-v1:0",
|
||||
"apac.amazon.nova-micro-v1:0",
|
||||
"apac.amazon.nova-lite-v1:0",
|
||||
"apac.amazon.nova-pro-v1:0",
|
||||
"bedrock/us-gov-west-1/amazon.nova-micro-v1:0",
|
||||
"bedrock/us-gov-west-1/amazon.nova-lite-v1:0",
|
||||
"bedrock/us-gov-west-1/amazon.nova-pro-v1:0",
|
||||
"bedrock/us-gov-east-1/amazon.nova-pro-v1:0",
|
||||
],
|
||||
)
|
||||
def test_nova_prompt_caching_models_price_cache_reads_below_the_input_rate(model, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
entry = litellm.model_cost[model]
|
||||
assert entry["supports_prompt_caching"] is True
|
||||
assert 0 < entry["cache_read_input_token_cost"] < entry["input_cost_per_token"]
|
||||
|
||||
|
||||
def test_transform_usage_with_reasoning_content():
|
||||
"""Test that completion_tokens_details correctly tracks reasoning vs text tokens."""
|
||||
usage = ConverseTokenUsageBlock(
|
||||
|
|
|
|||
|
|
@ -324,18 +324,18 @@ CONVERSE_METADATA_EVENT = {
|
|||
}
|
||||
|
||||
|
||||
def _converse_stream_wrapper(events):
|
||||
def _converse_stream_wrapper(events, model=CONVERSE_MODEL):
|
||||
async def bedrock_stream():
|
||||
decoder = AWSEventStreamDecoder(model=CONVERSE_MODEL)
|
||||
decoder = AWSEventStreamDecoder(model=model)
|
||||
for event in events:
|
||||
yield decoder._chunk_parser(chunk_data=event)
|
||||
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=bedrock_stream(),
|
||||
model=CONVERSE_MODEL,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=LiteLLMLoggingObj(
|
||||
model=CONVERSE_MODEL,
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="completion",
|
||||
|
|
@ -427,6 +427,46 @@ async def test_converse_stream_ends_on_finish_reason_chunk(events, expected_fini
|
|||
assert any(getattr(chunk, "usage", None) is not None for chunk in wrapper.chunks)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nova_invoke_stream_reports_bedrock_usage_and_finish_reason():
|
||||
"""InvokeModel Nova wraps every Converse event under its event-type key and reports usage
|
||||
without ``totalTokens``; the stream must end on Bedrock's finish reason and surface the
|
||||
cached tokens instead of a token-count estimate."""
|
||||
events = (
|
||||
{"messageStart": {"role": "assistant"}},
|
||||
{"contentBlockDelta": {"delta": {"text": "OK"}, "contentBlockIndex": 0}},
|
||||
{"contentBlockDelta": {"delta": {"text": "."}, "contentBlockIndex": 0}},
|
||||
{"contentBlockStop": {"contentBlockIndex": 0}},
|
||||
{"messageStop": {"stopReason": "end_turn"}},
|
||||
{
|
||||
"metadata": {
|
||||
"usage": {
|
||||
"inputTokens": 5,
|
||||
"outputTokens": 3,
|
||||
"cacheReadInputTokenCount": 12262,
|
||||
"cacheWriteInputTokenCount": 0,
|
||||
},
|
||||
"metrics": {},
|
||||
"trace": {},
|
||||
}
|
||||
},
|
||||
)
|
||||
wrapper = _converse_stream_wrapper(events, model="bedrock/invoke/us.amazon.nova-pro-v1:0")
|
||||
|
||||
chunks = [chunk async for chunk in wrapper]
|
||||
|
||||
assert "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == "OK."
|
||||
finish_reasons = [choice.finish_reason for chunk in chunks for choice in chunk.choices if choice.finish_reason]
|
||||
assert finish_reasons == ["stop"]
|
||||
assert chunks[-1].choices[0].finish_reason == "stop"
|
||||
usages = [chunk.usage for chunk in wrapper.chunks if getattr(chunk, "usage", None) is not None]
|
||||
assert len(usages) == 1
|
||||
assert usages[0].prompt_tokens == 12267
|
||||
assert usages[0].prompt_tokens_details.cached_tokens == 12262
|
||||
assert usages[0].completion_tokens == 3
|
||||
assert usages[0].total_tokens == 12270
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_converse_stream_still_emits_guardrail_trace_after_finish_reason():
|
||||
"""Guardrail metadata events carry a trace payload alongside usage; that chunk must still reach the caller
|
||||
|
|
|
|||
|
|
@ -5,12 +5,15 @@ from typing import NamedTuple
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Message,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -152,3 +155,55 @@ def test_bedrock_gpt_5_6_offers_tools_and_reasoning_effort_but_not_thinking(prof
|
|||
assert "reasoning_effort" in supported
|
||||
assert "thinking" not in supported
|
||||
assert "output_config" not in supported
|
||||
|
||||
|
||||
# Cache-read prices are the `*-cache-read-input-tokens` usagetype rows of the AWS Price List API, us-east-1,
|
||||
# https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json on 2026-09-15
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected_cache_read",
|
||||
[
|
||||
("amazon.nova-lite-v1:0", 1.5e-8),
|
||||
("us.amazon.nova-lite-v1:0", 1.5e-8),
|
||||
("amazon.nova-micro-v1:0", 8.75e-9),
|
||||
("us.amazon.nova-micro-v1:0", 8.75e-9),
|
||||
("amazon.nova-pro-v1:0", 2e-7),
|
||||
("us.amazon.nova-pro-v1:0", 2e-7),
|
||||
("us.amazon.nova-premier-v1:0", 6.25e-7),
|
||||
],
|
||||
)
|
||||
def test_bedrock_nova_cache_read_prices(
|
||||
model, expected_cache_read, local_model_cost_map
|
||||
):
|
||||
model_info = litellm.model_cost[model]
|
||||
assert model_info["cache_read_input_token_cost"] == expected_cache_read
|
||||
usage = Usage(
|
||||
prompt_tokens=1_000,
|
||||
completion_tokens=100,
|
||||
total_tokens=1_100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=400),
|
||||
)
|
||||
response = _bedrock_response(model, usage)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=response,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
expected_cost = (
|
||||
600 * model_info["input_cost_per_token"]
|
||||
+ 400 * expected_cache_read
|
||||
+ 100 * model_info["output_cost_per_token"]
|
||||
)
|
||||
assert cost == pytest.approx(expected_cost)
|
||||
|
||||
uncached_usage = Usage(
|
||||
prompt_tokens=1_000,
|
||||
completion_tokens=100,
|
||||
total_tokens=1_100,
|
||||
)
|
||||
uncached_cost = completion_cost(
|
||||
completion_response=_bedrock_response(model, uncached_usage),
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
assert cost < uncached_cost
|
||||
|
|
|
|||
|
|
@ -1025,6 +1025,86 @@ def test_handed_out_sync_client_pool_survives_handler_collection(keepalive_serve
|
|||
consumer_client.close()
|
||||
|
||||
|
||||
def _mock_transport() -> httpx.MockTransport:
|
||||
"""Answers anything with a short body, left unread when the caller asked to stream."""
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, request=request, content=b"ab")
|
||||
|
||||
return httpx.MockTransport(respond)
|
||||
|
||||
|
||||
RELEASED_TOO_EARLY = "the handler was released while its response could still read"
|
||||
NEVER_RELEASED = "the handler outlived the response that was holding it"
|
||||
|
||||
# Every method that can hand back a body the caller has not read yet, which is
|
||||
# every one that passes stream= down to send(). Parametrized so a method added
|
||||
# later is covered here rather than being the one that forgets to anchor.
|
||||
ASYNC_STREAMING_SENDS = ["post", "delete"]
|
||||
SYNC_STREAMING_SENDS = ["post", "patch", "put", "delete"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("method", ASYNC_STREAMING_SENDS)
|
||||
async def test_a_streaming_response_holds_its_handler_until_it_is_released(method):
|
||||
"""The finalizer must not run while a body this handler issued can still arrive.
|
||||
|
||||
``_handler_may_close_client`` cannot see that body: it holds the connection it
|
||||
reads from and never the client. Anchoring the handler to the response is what
|
||||
withholds the close, and releasing the anchor is what still delivers one.
|
||||
"""
|
||||
handler = AsyncHTTPHandler()
|
||||
handler.client._transport = _mock_transport()
|
||||
ref = weakref.ref(handler)
|
||||
response = await getattr(handler, method)("https://example.invalid/stream", stream=True)
|
||||
|
||||
del handler
|
||||
gc.collect()
|
||||
assert ref() is not None, RELEASED_TOO_EARLY
|
||||
|
||||
assert await response.aread() == b"ab"
|
||||
del response
|
||||
gc.collect()
|
||||
assert ref() is None, NEVER_RELEASED
|
||||
|
||||
|
||||
@pytest.mark.parametrize("method", SYNC_STREAMING_SENDS)
|
||||
def test_a_sync_streaming_response_holds_its_handler_until_it_is_released(method):
|
||||
"""The sync finalizer closes inline, so the same anchor has to hold it off."""
|
||||
handler = HTTPHandler()
|
||||
handler.client._transport = _mock_transport()
|
||||
ref = weakref.ref(handler)
|
||||
response = getattr(handler, method)("https://example.invalid/stream", stream=True)
|
||||
|
||||
del handler
|
||||
gc.collect()
|
||||
assert ref() is not None, RELEASED_TOO_EARLY
|
||||
|
||||
assert response.read() == b"ab"
|
||||
del response
|
||||
gc.collect()
|
||||
assert ref() is None, NEVER_RELEASED
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_fully_read_response_does_not_hold_its_handler():
|
||||
"""A non-streaming response is complete when ``post`` returns, so it anchors nothing.
|
||||
|
||||
Otherwise every client close would wait on whatever the caller does next with
|
||||
a response it has already read.
|
||||
"""
|
||||
handler = AsyncHTTPHandler()
|
||||
handler.client._transport = _mock_transport()
|
||||
ref = weakref.ref(handler)
|
||||
response = await handler.post("https://example.invalid/whole")
|
||||
assert response.content == b"ab"
|
||||
|
||||
del handler
|
||||
gc.collect()
|
||||
|
||||
assert ref() is None, "a fully-read response pinned its handler"
|
||||
|
||||
|
||||
def test_sync_close_leaves_caller_supplied_client_open():
|
||||
supplied = httpx.Client()
|
||||
handler = HTTPHandler(client=supplied)
|
||||
|
|
|
|||
|
|
@ -1189,6 +1189,28 @@ def test_reasoning_effort_integer_passthrough():
|
|||
assert isinstance(result["reasoning_effort"], int)
|
||||
|
||||
|
||||
def test_reasoning_effort_dict_from_anthropic_adapter_flattened_to_effort_string():
|
||||
config = FireworksAIConfig()
|
||||
result = config.map_openai_params(
|
||||
{"reasoning_effort": {"effort": "medium", "summary": "detailed"}},
|
||||
{},
|
||||
_REASONING_MODEL,
|
||||
drop_params=False,
|
||||
)
|
||||
assert result["reasoning_effort"] == "medium"
|
||||
|
||||
|
||||
def test_reasoning_effort_dict_without_effort_key_dropped():
|
||||
config = FireworksAIConfig()
|
||||
result = config.map_openai_params(
|
||||
{"reasoning_effort": {"summary": "detailed"}},
|
||||
{},
|
||||
_REASONING_MODEL,
|
||||
drop_params=False,
|
||||
)
|
||||
assert "reasoning_effort" not in result
|
||||
|
||||
|
||||
def test_reasoning_effort_auto_dropped_to_model_default():
|
||||
config = FireworksAIConfig()
|
||||
result = config.map_openai_params(
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.llms.gemini.cost_calculator import (
|
||||
cost_per_google_maps_grounding_request,
|
||||
cost_per_web_search_request,
|
||||
|
|
@ -18,6 +17,7 @@ from litellm.types.utils import (
|
|||
ImageResponse,
|
||||
ImageUsage,
|
||||
ImageUsageInputTokensDetails,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
|
@ -452,6 +452,42 @@ def test_map_traffic_type_to_service_tier(
|
|||
)
|
||||
|
||||
|
||||
# Alias targets are the `modelVersion` returned by
|
||||
# POST https://generativelanguage.googleapis.com/v1beta/models/<alias>:generateContent on 2026-09-15
|
||||
@pytest.mark.parametrize(
|
||||
"alias,target",
|
||||
[
|
||||
("gemini/gemini-flash-latest", "gemini/gemini-3.8-flash"),
|
||||
("gemini/gemini-flash-lite-latest", "gemini/gemini-3.5-flash-lite"),
|
||||
("gemini/gemini-pro-latest", "gemini/gemini-3.1-pro-preview"),
|
||||
],
|
||||
)
|
||||
def test_latest_aliases_cost_the_same_as_their_current_target(
|
||||
monkeypatch, alias, target
|
||||
):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=1_000,
|
||||
completion_tokens=500,
|
||||
total_tokens=1_500,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=400),
|
||||
)
|
||||
|
||||
def cost_of(model: str) -> float:
|
||||
return completion_cost(
|
||||
completion_response=ModelResponse(model=model, usage=usage),
|
||||
model=model,
|
||||
custom_llm_provider="gemini",
|
||||
)
|
||||
|
||||
alias_cost = cost_of(alias)
|
||||
target_cost = cost_of(target)
|
||||
assert alias_cost == pytest.approx(target_cost)
|
||||
assert alias_cost > 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"prefixed,bare",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -2206,10 +2206,35 @@ class TestStreamingScanKey:
|
|||
[self._chunk("hi"), tool_chunk, self._chunk(None, finish_reason="stop")]
|
||||
)
|
||||
assert open_key == StreamingScanKey(texts=("hi",))
|
||||
assert open_key.tool_calls_in_flight is True
|
||||
assert handler.get_streaming_scan_key([self._chunk("hi")]).tool_calls_in_flight is False
|
||||
assert ended_key.texts == ("hi",)
|
||||
assert len(ended_key.tool_calls) == 1 and "get_weather" in ended_key.tool_calls[0]
|
||||
assert ended_key.tool_calls_in_flight is False
|
||||
assert ended_key != open_key
|
||||
|
||||
def test_legacy_function_call_delta_is_held_like_a_tool_call(self):
|
||||
from litellm.types.utils import Delta, FunctionCall, ModelResponseStream, StreamingChoices
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
function_chunk = ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=None, function_call=FunctionCall(name="run_shell", arguments='{"cmd": "rm"}')),
|
||||
finish_reason=None,
|
||||
)
|
||||
]
|
||||
)
|
||||
open_key = handler.get_streaming_scan_key([self._chunk("hi"), function_chunk])
|
||||
ended_key = handler.get_streaming_scan_key(
|
||||
[self._chunk("hi"), function_chunk, self._chunk(None, finish_reason="function_call")]
|
||||
)
|
||||
assert open_key.tool_calls_in_flight is True
|
||||
assert open_key.tool_calls == ()
|
||||
assert len(ended_key.tool_calls) == 1 and "run_shell" in ended_key.tool_calls[0]
|
||||
assert ended_key.tool_calls_in_flight is False
|
||||
|
||||
def test_text_after_the_first_choice_finishes_still_changes_the_key(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
first_done = [self._chunk("a", index=0), self._chunk("b", finish_reason="stop", index=0)]
|
||||
|
|
|
|||
|
|
@ -3211,3 +3211,42 @@ class TestOpenAIResponsesHandlerStreamingScanKey:
|
|||
def test_output_item_done_round_is_never_deduped(self):
|
||||
done = {"type": "response.output_item.done", "sequence_number": 1, "item": {"type": "function_call"}}
|
||||
assert OpenAIResponsesHandler().get_streaming_scan_key([self._delta(0, "hi"), done]) is None
|
||||
|
||||
@pytest.mark.parametrize("terminal_type", ["response.incomplete", "response.failed"])
|
||||
def test_non_completed_terminal_envelopes_key_their_output_items(self, terminal_type):
|
||||
handler = OpenAIResponsesHandler()
|
||||
arguments_delta = {
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"sequence_number": 1,
|
||||
"item_id": "fc_1",
|
||||
"delta": '{"city":',
|
||||
}
|
||||
function_call = {"type": "function_call", "call_id": "call_1", "name": "get_weather", "arguments": '{"city":'}
|
||||
terminal = {"type": terminal_type, "sequence_number": 2, "response": {"id": "resp_1", "output": [function_call]}}
|
||||
mid_stream_key = handler.get_streaming_scan_key([arguments_delta])
|
||||
ended_key = handler.get_streaming_scan_key([arguments_delta, terminal])
|
||||
assert ended_key.stream_ended is True
|
||||
assert ended_key.tool_calls_in_flight is False
|
||||
assert len(ended_key.tool_calls) == 1
|
||||
assert ended_key != mid_stream_key
|
||||
|
||||
def test_streamed_tool_call_events_flag_tool_calls_in_flight_until_the_stream_ends(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
added = {
|
||||
"type": "response.output_item.added",
|
||||
"sequence_number": 1,
|
||||
"item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "get_weather"},
|
||||
}
|
||||
arguments_delta = {
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"sequence_number": 2,
|
||||
"item_id": "fc_1",
|
||||
"delta": '{"city":',
|
||||
}
|
||||
function_call = {"type": "function_call", "call_id": "call_1", "name": "get_weather", "arguments": "{}"}
|
||||
assert handler.get_streaming_scan_key([self._delta(0, "hi")]).tool_calls_in_flight is False
|
||||
assert handler.get_streaming_scan_key([self._delta(0, "hi"), added]).tool_calls_in_flight is True
|
||||
assert handler.get_streaming_scan_key([self._delta(0, "hi"), arguments_delta]).tool_calls_in_flight is True
|
||||
ended_key = handler.get_streaming_scan_key([self._delta(0, "hi"), added, self._completed(3, [function_call])])
|
||||
assert ended_key.tool_calls_in_flight is False
|
||||
assert len(ended_key.tool_calls) == 1
|
||||
|
|
|
|||
|
|
@ -362,6 +362,14 @@ class TestVerificationToken:
|
|||
assert deleted.deleted_at is not None
|
||||
assert deleted.token == "t1"
|
||||
|
||||
def test_total_spend_is_carried_separately_from_resettable_spend(self):
|
||||
token = LiteLLM_VerificationToken(token="t1", spend=0.0, total_spend=12.5)
|
||||
assert token.model_dump()["total_spend"] == 12.5
|
||||
assert token.model_dump()["spend"] == 0.0
|
||||
|
||||
deleted = LiteLLM_DeletedVerificationToken.model_validate({**token.model_dump(), "deleted_by": "admin"})
|
||||
assert deleted.total_spend == 12.5
|
||||
|
||||
|
||||
class TestConfigTable:
|
||||
def test_config_creation(self):
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ import orjson
|
|||
import pytest
|
||||
from starlette.datastructures import FormData
|
||||
|
||||
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
|
||||
from litellm.ocr.legacy import convert_file_document_to_url_document, get_mime_type
|
||||
|
||||
|
||||
class TestGetMimeType:
|
||||
|
|
@ -487,9 +487,9 @@ class TestProxySecurityGuard:
|
|||
async def test_proxy_upload_stops_reading_at_size_limit() -> None:
|
||||
from starlette.datastructures import UploadFile
|
||||
|
||||
from litellm.proxy.ocr_endpoints.endpoints import _parse_multipart_form
|
||||
from litellm.proxy.ocr_endpoints.endpoints import _MAX_FILE_BYTES, _parse_multipart_form
|
||||
|
||||
limit: Final = 50 * 1024 * 1024
|
||||
limit: Final = _MAX_FILE_BYTES
|
||||
with tempfile.TemporaryFile() as stream:
|
||||
stream.truncate(limit * 2)
|
||||
upload: Final = UploadFile(file=stream, filename="large.pdf")
|
||||
|
|
|
|||
202
tests/test_litellm/proxy/auth/test_fallback_budget.py
Normal file
202
tests/test_litellm/proxy/auth/test_fallback_budget.py
Normal file
|
|
@ -0,0 +1,202 @@
|
|||
import pytest
|
||||
|
||||
from litellm import Router
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.fallback_budget import (
|
||||
RouterFallbackBudgetCheck,
|
||||
is_token_within_budget_for_model,
|
||||
router_fallback_budget_check,
|
||||
)
|
||||
|
||||
FREE_MODEL = {
|
||||
"model_name": "free-model",
|
||||
"litellm_params": {
|
||||
"model": "ollama/llama2",
|
||||
"api_base": "http://localhost:11434",
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
},
|
||||
"model_info": {
|
||||
"id": "free-model-id",
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
},
|
||||
}
|
||||
|
||||
PAID_MODEL = {
|
||||
"model_name": "paid-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "k"},
|
||||
"model_info": {"id": "paid-model-id"},
|
||||
}
|
||||
|
||||
|
||||
def _router() -> Router:
|
||||
return Router(model_list=[FREE_MODEL, PAID_MODEL], fallbacks=[{"free-model": ["paid-model"]}])
|
||||
|
||||
|
||||
def _token(**overrides) -> UserAPIKeyAuth:
|
||||
fields = {
|
||||
"api_key": "hashed",
|
||||
"token": "hashed",
|
||||
"spend": 0.0,
|
||||
"max_budget": None,
|
||||
"user_id": "u1",
|
||||
"user_spend": 0.0,
|
||||
"user_max_budget": None,
|
||||
}
|
||||
fields.update(overrides)
|
||||
return UserAPIKeyAuth(**fields)
|
||||
|
||||
|
||||
ENFORCED = RouterFallbackBudgetCheck(is_enforced=lambda: True)
|
||||
NOT_ENFORCED = RouterFallbackBudgetCheck(is_enforced=lambda: False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paid_target_allowed_when_under_budget():
|
||||
token = _token(spend=1.0, max_budget=50.0, user_spend=1.0, user_max_budget=50.0)
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paid_target_refused_when_over_key_budget():
|
||||
token = _token(spend=100.0, max_budget=50.0)
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paid_target_refused_when_over_user_budget():
|
||||
token = _token(user_spend=1900.0, user_max_budget=50.0)
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_zero_cost_target_allowed_even_when_over_budget():
|
||||
"""Refusing a free target would deny a request on spend some other model accrued."""
|
||||
token = _token(user_spend=1900.0, user_max_budget=50.0)
|
||||
assert await is_token_within_budget_for_model(model="free-model", valid_token=token, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_budget_configured_is_always_within_budget():
|
||||
token = _token(spend=9999.0, user_spend=9999.0)
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_key_does_not_inherit_personal_budget_by_default(monkeypatch):
|
||||
"""Mirrors _PROXY_MaxBudgetLimiter: a team key ignores the owner's personal cap."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {}, raising=False)
|
||||
token = _token(team_id="t1", user_spend=1900.0, user_max_budget=50.0)
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_key_inherits_personal_budget_when_opted_in(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"apply_user_budget_to_team_keys": True}, raising=False)
|
||||
token = _token(team_id="t1", user_spend=1900.0, user_max_budget=50.0)
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_is_a_no_op_while_not_enforced():
|
||||
request = {"metadata": {"user_api_key_auth": _token(user_spend=1900.0, user_max_budget=50.0)}}
|
||||
assert await NOT_ENFORCED(model="paid-model", request_kwargs=request, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_without_a_key_is_unrestricted():
|
||||
assert await ENFORCED(model="paid-model", request_kwargs={}, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("metadata_field", ["metadata", "litellm_metadata"])
|
||||
async def test_enforced_check_reads_the_key_from_request_metadata(metadata_field: str):
|
||||
over = {metadata_field: {"user_api_key_auth": _token(user_spend=1900.0, user_max_budget=50.0)}}
|
||||
under = {metadata_field: {"user_api_key_auth": _token(user_spend=1.0, user_max_budget=50.0)}}
|
||||
|
||||
assert await ENFORCED(model="paid-model", request_kwargs=over, llm_router=_router()) is False
|
||||
assert await ENFORCED(model="paid-model", request_kwargs=under, llm_router=_router()) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_stale_low_counter_still_refuses_a_paid_target(monkeypatch):
|
||||
"""
|
||||
The counter can read low (e.g. restored from an older Redis snapshot). Passing the budget makes
|
||||
`get_current_spend` verify against authoritative spend instead of trusting that read, so the
|
||||
paid target is still refused.
|
||||
"""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
seen: list[dict] = []
|
||||
|
||||
async def _stale_counter(**kwargs):
|
||||
seen.append(kwargs)
|
||||
# a stale-low counter read; the authoritative spend is what the budget must be judged on
|
||||
return 0.0 if kwargs.get("max_budget") is None else kwargs["fallback_spend"]
|
||||
|
||||
monkeypatch.setattr(proxy_server, "get_current_spend", _stale_counter, raising=False)
|
||||
token = _token(user_spend=1900.0, user_max_budget=50.0)
|
||||
|
||||
assert await is_token_within_budget_for_model(model="paid-model", valid_token=token, llm_router=_router()) is False
|
||||
assert [call["max_budget"] for call in seen] == [50.0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_fails_closed_when_the_spend_lookup_breaks(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
async def _boom(**kwargs):
|
||||
raise RuntimeError("spend counter unavailable")
|
||||
|
||||
monkeypatch.setattr(proxy_server, "get_current_spend", _boom, raising=False)
|
||||
request = {"metadata": {"user_api_key_auth": _token(user_spend=1.0, user_max_budget=50.0)}}
|
||||
|
||||
assert await ENFORCED(model="paid-model", request_kwargs=request, llm_router=_router()) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_skips_the_paid_fallback_target_when_over_budget():
|
||||
from litellm.router_utils.fallback_event_handlers import _is_fallback_target_within_budget
|
||||
|
||||
router = Router(
|
||||
model_list=[FREE_MODEL, PAID_MODEL],
|
||||
fallbacks=[{"free-model": ["paid-model"]}],
|
||||
fallback_budget_check=ENFORCED,
|
||||
)
|
||||
over = {"metadata": {"user_api_key_auth": _token(user_spend=1900.0, user_max_budget=50.0)}}
|
||||
under = {"metadata": {"user_api_key_auth": _token(user_spend=1.0, user_max_budget=50.0)}}
|
||||
|
||||
assert await _is_fallback_target_within_budget(router, "paid-model", "free-model", over) is False
|
||||
assert await _is_fallback_target_within_budget(router, "paid-model", "free-model", under) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_without_a_budget_check_attempts_every_fallback():
|
||||
from litellm.router_utils.fallback_event_handlers import _is_fallback_target_within_budget
|
||||
|
||||
router = _router() # fallback_budget_check defaults to None
|
||||
over = {"metadata": {"user_api_key_auth": _token(user_spend=1900.0, user_max_budget=50.0)}}
|
||||
|
||||
assert await _is_fallback_target_within_budget(router, "paid-model", "free-model", over) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enforcement_is_on_by_default_and_opt_out_restores_the_leak(monkeypatch):
|
||||
"""
|
||||
Leaving the paid fallback unguarded is the budget bypass this module exists to close, so an
|
||||
unconfigured proxy has to enforce. `enforce_fallback_budget: false` is the deliberate opt-out.
|
||||
"""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
over = {"metadata": {"user_api_key_auth": _token(user_spend=1900.0, user_max_budget=50.0)}}
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {}, raising=False)
|
||||
assert await router_fallback_budget_check(model="paid-model", request_kwargs=over, llm_router=_router()) is False
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"enforce_fallback_budget": False}, raising=False)
|
||||
assert await router_fallback_budget_check(model="paid-model", request_kwargs=over, llm_router=_router()) is True
|
||||
|
|
@ -31,6 +31,7 @@ def _full_team(model_aliases=ALIASES) -> LiteLLM_TeamTable:
|
|||
max_budget=50.0,
|
||||
soft_budget=25.0,
|
||||
spend=12.5,
|
||||
model_max_budget={"gpt-4o": {"max_budget": 5.0, "budget_duration": "1d"}},
|
||||
models=["gpt-4o", "gpt-4o-mini"],
|
||||
blocked=True,
|
||||
metadata={"tier": "gold"},
|
||||
|
|
@ -72,6 +73,7 @@ def test_team_grants_cover_every_team_field_the_key_path_gets():
|
|||
assert token.team_max_budget == 50.0
|
||||
assert token.team_soft_budget == 25.0
|
||||
assert token.team_spend == 12.5
|
||||
assert token.team_model_max_budget == {"gpt-4o": {"max_budget": 5.0, "budget_duration": "1d"}}
|
||||
assert token.team_models == ["gpt-4o", "gpt-4o-mini"]
|
||||
assert token.team_blocked is True
|
||||
assert token.team_metadata == {"tier": "gold"}
|
||||
|
|
|
|||
|
|
@ -4374,6 +4374,74 @@ async def test_centralized_common_checks_carries_team_and_user_budget_state_on_t
|
|||
}
|
||||
|
||||
|
||||
class _RecordingTeamModelBudgetLimiter:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
async def is_team_within_model_budget(self, team_id, team_model_max_budget, key_model_max_budget, model):
|
||||
self.calls.append((team_id, dict(team_model_max_budget), key_model_max_budget, model))
|
||||
return True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_centralized_common_checks_enforces_team_model_max_budget_from_the_resolved_team():
|
||||
"""The team's model_max_budget is enforced at the single authz gate, off the
|
||||
team object auth resolved (not the possibly stale token copy), and the key's
|
||||
own model_max_budget is handed to the limiter so a matching key entry can
|
||||
override the team cap."""
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
|
||||
team_caps = {"gpt-4o": {"max_budget": 5.0, "budget_duration": "1d"}}
|
||||
key_caps = {"claude-sonnet-4-6": {"max_budget": 1.0, "budget_duration": "1d"}}
|
||||
token = UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
token="hashed",
|
||||
team_id="t1",
|
||||
team_model_max_budget={"gpt-4o": {"max_budget": 999.0, "budget_duration": "30d"}},
|
||||
model_max_budget=key_caps,
|
||||
)
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key="team_id:t1",
|
||||
value=LiteLLM_TeamTableCachedObj(team_id="t1", model_max_budget=team_caps),
|
||||
)
|
||||
limiter = _RecordingTeamModelBudgetLimiter()
|
||||
attrs = {
|
||||
**_proxy_attrs_for_centralized_checks(user_custom_auth=None),
|
||||
"prisma_client": MagicMock(),
|
||||
"user_api_key_cache": user_api_key_cache,
|
||||
"model_max_budget_limiter": limiter,
|
||||
}
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
with (
|
||||
patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock), # test-quality-ok: stubs the sibling check so only the team model-budget gate is under test
|
||||
patch( # test-quality-ok: stubs the budget reservation so only the team model-budget gate is under test
|
||||
"litellm.proxy.auth.user_api_key_auth._reserve_budget_after_common_checks",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
assert limiter.calls == [("t1", team_caps, key_caps, "gpt-4o")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_centralized_common_checks_skipped_for_custom_auth_without_flag():
|
||||
"""Existing RPS guarantee: custom-auth deployments without
|
||||
|
|
|
|||
|
|
@ -291,6 +291,23 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client):
|
|||
assert set(write["data"].keys()) == {"spend", "budget_reset_at"}
|
||||
|
||||
|
||||
def test_reset_budget_for_key_leaves_lifetime_total_spend_alone(reset_budget_job, mock_prisma_client):
|
||||
"""A period reset zeroes spend but must neither write nor touch the lifetime total_spend."""
|
||||
now = datetime.now(timezone.utc)
|
||||
key = LiteLLM_VerificationToken(
|
||||
token="tok-key-1", spend=100.0, total_spend=340.0, budget_duration="30d", budget_reset_at=now
|
||||
)
|
||||
mock_prisma_client.data["key"] = [key]
|
||||
|
||||
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
|
||||
|
||||
(write,) = _batch_writes(mock_prisma_client, "key")
|
||||
assert write["data"]["spend"] == {"decrement": 100.0}
|
||||
assert "total_spend" not in write["data"]
|
||||
assert key.spend == 0.0
|
||||
assert key.total_spend == 340.0
|
||||
|
||||
|
||||
def test_reset_budget_for_key_honors_injected_reset_time(mock_prisma_client, mock_proxy_logging):
|
||||
"""Injected BudgetResetSettings drives the written reset time end to end (DI, no globals).
|
||||
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ async def test_create_views_creates_view_on_does_not_exist():
|
|||
mock_db.execute_raw.assert_called_once()
|
||||
created_sql = mock_db.execute_raw.call_args[0][0]
|
||||
assert 'CREATE VIEW "LiteLLM_VerificationTokenView"' in created_sql
|
||||
assert "t.model_max_budget AS team_model_max_budget" in created_sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1658,6 +1658,57 @@ async def test_commit_key_spend_updates_includes_last_active():
|
|||
assert before_call <= last_active <= after_call
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_spend_updates_to_db_increments_key_total_spend_alongside_spend():
|
||||
"""
|
||||
The key table write must increment the lifetime total_spend by the same amount as the
|
||||
resettable spend, in the same update so the two cannot drift.
|
||||
"""
|
||||
db_writer = DBSpendUpdateWriter()
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_batcher.litellm_verificationtoken = MagicMock()
|
||||
mock_batcher.litellm_verificationtoken.update_many = MagicMock()
|
||||
|
||||
mock_transaction = AsyncMock()
|
||||
mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction)
|
||||
mock_transaction.__aexit__ = AsyncMock(return_value=False)
|
||||
mock_transaction.batch_ = MagicMock(
|
||||
return_value=AsyncMock(
|
||||
__aenter__=AsyncMock(return_value=mock_batcher),
|
||||
__aexit__=AsyncMock(return_value=False),
|
||||
)
|
||||
)
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction)
|
||||
|
||||
db_spend_update_transactions = {
|
||||
"user_list_transactions": {},
|
||||
"end_user_list_transactions": {},
|
||||
"key_list_transactions": {"hashed_token_abc": 0.05, "hashed_token_def": 1.25},
|
||||
"team_list_transactions": {},
|
||||
"team_member_list_transactions": {},
|
||||
"org_list_transactions": {},
|
||||
"tag_list_transactions": {},
|
||||
"agent_list_transactions": {},
|
||||
}
|
||||
|
||||
await db_writer._commit_spend_updates_to_db(
|
||||
prisma_client=mock_prisma_client,
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
db_spend_update_transactions=db_spend_update_transactions,
|
||||
)
|
||||
|
||||
calls = mock_batcher.litellm_verificationtoken.update_many.call_args_list
|
||||
assert [c.kwargs["where"] for c in calls] == [{"token": "hashed_token_abc"}, {"token": "hashed_token_def"}]
|
||||
for call, expected_cost in zip(calls, (0.05, 1.25)):
|
||||
assert call.kwargs["data"]["spend"] == {"increment": expected_cost}
|
||||
assert call.kwargs["data"]["total_spend"] == call.kwargs["data"]["spend"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_database_creates_single_task():
|
||||
"""
|
||||
|
|
@ -2813,7 +2864,7 @@ async def test_commit_spend_updates_to_db_does_not_stamp_key_settings_updated_at
|
|||
mock_batcher.litellm_verificationtoken.update_many.assert_called_once()
|
||||
call_kwargs = mock_batcher.litellm_verificationtoken.update_many.call_args[1]
|
||||
assert call_kwargs["where"] == {"token": token}
|
||||
assert set(call_kwargs["data"]) == {"spend", "last_active"}
|
||||
assert set(call_kwargs["data"]) == {"spend", "total_spend", "last_active"}
|
||||
assert call_kwargs["data"]["spend"] == {"increment": response_cost}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -5603,6 +5603,7 @@ def test_initialize_bedrock_wires_streaming_flags():
|
|||
streaming_buffer_until_moderated=False,
|
||||
streaming_sampling_rate=3,
|
||||
streaming_end_of_stream_only=True,
|
||||
streaming_buffer_release_on_scan=True,
|
||||
),
|
||||
{"guardrail_name": "bedrock-streaming"},
|
||||
)
|
||||
|
|
@ -5616,9 +5617,11 @@ def test_initialize_bedrock_wires_streaming_flags():
|
|||
assert configured.streaming_buffer_until_moderated is False
|
||||
assert configured.streaming_sampling_rate == 3
|
||||
assert configured.streaming_end_of_stream_only is True
|
||||
assert configured.streaming_buffer_release_on_scan is True
|
||||
assert defaulted.streaming_buffer_until_moderated is True
|
||||
assert defaulted.streaming_sampling_rate == 5
|
||||
assert defaulted.streaming_end_of_stream_only is False
|
||||
assert defaulted.streaming_buffer_release_on_scan is False
|
||||
|
||||
|
||||
def test_initialize_bedrock_rejects_non_positive_sampling_rate():
|
||||
|
|
@ -5721,6 +5724,44 @@ async def test_buffered_default_hook_scans_before_any_chunk():
|
|||
assert len([e for e in events if e != "scan"]) >= 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_buffered_release_on_scan_hook_releases_each_window_after_its_scan():
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-release-on-scan",
|
||||
guardrailIdentifier="test-id",
|
||||
guardrailVersion="DRAFT",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
default_on=True,
|
||||
streaming_buffer_release_on_scan=True,
|
||||
streaming_sampling_rate=1,
|
||||
)
|
||||
|
||||
assert guardrail._streams_incrementally() is True
|
||||
events = await _run_streaming_hook_recording_order(guardrail)
|
||||
|
||||
assert events == ["scan", ("chunk", "Hello"), "scan", ("chunk", " world"), ("chunk", "")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_buffered_release_on_scan_defers_to_end_of_stream_only():
|
||||
guardrail = BedrockGuardrail(
|
||||
guardrail_name="bedrock-release-on-scan-end-only",
|
||||
guardrailIdentifier="test-id",
|
||||
guardrailVersion="DRAFT",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
default_on=True,
|
||||
streaming_buffer_release_on_scan=True,
|
||||
streaming_end_of_stream_only=True,
|
||||
streaming_sampling_rate=1,
|
||||
)
|
||||
|
||||
assert guardrail._streams_incrementally() is False
|
||||
events = await _run_streaming_hook_recording_order(guardrail)
|
||||
|
||||
assert events.count("scan") == 1
|
||||
assert events[0] == "scan"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_masking_keeps_buffered_path_even_when_unbuffered_configured():
|
||||
guardrail = BedrockGuardrail(
|
||||
|
|
|
|||
|
|
@ -1622,10 +1622,23 @@ def test_initialize_guardrail_rejects_unsupported_mode_instead_of_running_other_
|
|||
def test_initialize_guardrail_defaults_streaming_params() -> None:
|
||||
handler = _initialize_from_config(mode="post_call")
|
||||
|
||||
assert handler.streaming_buffer_until_moderated is False
|
||||
assert handler.streaming_buffer_release_on_scan is False
|
||||
assert handler.streaming_end_of_stream_only is False
|
||||
assert handler.streaming_sampling_rate == 5
|
||||
|
||||
|
||||
def test_initialize_guardrail_forwards_buffer_streaming_params() -> None:
|
||||
handler = _initialize_from_config(
|
||||
mode="post_call",
|
||||
streaming_buffer_until_moderated=True,
|
||||
streaming_buffer_release_on_scan=True,
|
||||
)
|
||||
|
||||
assert handler.streaming_buffer_until_moderated is True
|
||||
assert handler.streaming_buffer_release_on_scan is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"configured",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -1291,8 +1291,9 @@ class TestToolPermissionGuardrailAnthropicMessages:
|
|||
async def test_rewrite_mode_keeps_the_stream_identity_it_had_before_the_shared_helper(self):
|
||||
"""Well-formed SSE must round-trip exactly as it did before the helpers were shared.
|
||||
|
||||
The shared module can stamp the upstream message id and model onto the assembled response
|
||||
for callers that ask for it; this path never did, and a client reads those bytes.
|
||||
The shared module can stamp the upstream message id onto the assembled response for
|
||||
callers that ask for it; this path never did, and a client reads those bytes. The model,
|
||||
though, is now the upstream's, matching what the untouched passthrough shows clients.
|
||||
"""
|
||||
with patch.object(self.rewriting, "should_run_guardrail", return_value=True):
|
||||
out = await self._drain(self.rewriting, self._sse_chunks("Read"))
|
||||
|
|
@ -1304,7 +1305,7 @@ class TestToolPermissionGuardrailAnthropicMessages:
|
|||
if line.startswith("data: ") and json.loads(line[6:]).get("type") == "message_start"
|
||||
)["message"]
|
||||
assert message_start["id"].startswith("chatcmpl-"), "the rewritten stream must not adopt the upstream message id"
|
||||
assert message_start["model"] == "unknown-model", "the rewritten stream must not adopt the upstream model"
|
||||
assert message_start["model"] == "claude-sonnet-4-5", "the rewritten stream reports the model the upstream served"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_start_without_a_dict_message_fails_closed(self):
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ released unchanged after moderation passes.
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, List, Literal, Optional
|
||||
from typing import Any, AsyncGenerator, List, Literal, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -19,14 +19,25 @@ from litellm.integrations.custom_guardrail import (
|
|||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import StreamingScanKey
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
_is_redundant_scan,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
Function,
|
||||
FunctionCall,
|
||||
GenericGuardrailAPIInputs,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
BLOCK_MESSAGE = "Blocked by policy: this response was withheld."
|
||||
ORIGINAL_MARKER = "ORIGINAL-SECRET-ANSWER"
|
||||
TOOL_ARGUMENTS_MARKER = "TOOL-ARGS-SECRET"
|
||||
|
||||
|
||||
class _BlockingGuardrail(CustomGuardrail):
|
||||
|
|
@ -60,6 +71,85 @@ class _PassingGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
|
||||
class _CountingPassingGuardrail(_PassingGuardrail):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.scan_count = 0
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.scan_count += 1
|
||||
return inputs
|
||||
|
||||
|
||||
class _ToolCallRecordingGuardrail(_CountingPassingGuardrail):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.tool_call_scan_indexes: List[int] = []
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.scan_count += 1
|
||||
if inputs.get("tool_calls"):
|
||||
self.tool_call_scan_indexes.append(self.scan_count)
|
||||
return inputs
|
||||
|
||||
|
||||
class _SecondScanBlockingGuardrail(_CountingPassingGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.scan_count += 1
|
||||
if self.scan_count == 2:
|
||||
raise ModifyResponseException(
|
||||
message=BLOCK_MESSAGE,
|
||||
model="gpt-4",
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
return inputs
|
||||
|
||||
|
||||
class _MarkerBlockingGuardrail(_CountingPassingGuardrail):
|
||||
"""Blocks as soon as the inspected input field (texts or tool_calls) carries the marker."""
|
||||
|
||||
def __init__(self, *args, marker: str, field: Literal["texts", "tool_calls"] = "texts", **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.marker = marker
|
||||
self.field = field
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.scan_count += 1
|
||||
if self.marker in json.dumps(inputs.get(self.field, [])):
|
||||
raise ModifyResponseException(
|
||||
message=BLOCK_MESSAGE,
|
||||
model="gpt-4o",
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
return inputs
|
||||
|
||||
|
||||
def _sse_event(event_type: str, data: dict) -> bytes:
|
||||
return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode()
|
||||
|
||||
|
|
@ -115,6 +205,212 @@ def _decode(chunks: List[Any]) -> str:
|
|||
return "".join(c.decode() if isinstance(c, bytes) else str(c) for c in chunks)
|
||||
|
||||
|
||||
def _chat_chunk(content: str = "", finish_reason: str | None = None) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-windowed",
|
||||
created=1724900000,
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(role="assistant", content=content),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _tool_call_chunk(
|
||||
arguments: str, finish_reason: str | None = None, legacy_function_call: bool = False
|
||||
) -> ModelResponseStream:
|
||||
delta = (
|
||||
Delta(role="assistant", content=None, function_call=FunctionCall(name="run_shell", arguments=arguments))
|
||||
if legacy_function_call
|
||||
else Delta(
|
||||
role="assistant",
|
||||
content=None,
|
||||
tool_calls=[
|
||||
ChatCompletionDeltaToolCall(
|
||||
id="call_1",
|
||||
type="function",
|
||||
index=0,
|
||||
function=Function(name="run_shell", arguments=arguments),
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
return ModelResponseStream(
|
||||
id="chatcmpl-windowed",
|
||||
created=1724900000,
|
||||
model="gpt-4",
|
||||
choices=[StreamingChoices(index=0, delta=delta, finish_reason=finish_reason)],
|
||||
)
|
||||
|
||||
|
||||
async def _windowed_chat_stream(
|
||||
yielded_count: List[int],
|
||||
collected: List[Any],
|
||||
content_chunks: List[str],
|
||||
tool_argument_chunks: List[str] | None = None,
|
||||
legacy_function_call: bool = False,
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
for content in content_chunks:
|
||||
yielded_count.append(len(collected))
|
||||
yield _chat_chunk(content)
|
||||
for arguments in tool_argument_chunks or []:
|
||||
yielded_count.append(len(collected))
|
||||
yield _tool_call_chunk(arguments, legacy_function_call=legacy_function_call)
|
||||
yielded_count.append(len(collected))
|
||||
yield _chat_chunk(finish_reason="tool_calls" if tool_argument_chunks else "stop")
|
||||
|
||||
|
||||
def _tool_argument_text(chunks: List[Any]) -> str:
|
||||
return "".join(
|
||||
tool_call.function.arguments or ""
|
||||
for chunk in chunks
|
||||
if isinstance(chunk, ModelResponseStream)
|
||||
for choice in chunk.choices
|
||||
for tool_call in choice.delta.tool_calls or []
|
||||
)
|
||||
|
||||
|
||||
def _function_call_argument_text(chunks: list[Any]) -> str:
|
||||
return "".join(
|
||||
choice.delta.function_call.arguments or ""
|
||||
for chunk in chunks
|
||||
if isinstance(chunk, ModelResponseStream)
|
||||
for choice in chunk.choices
|
||||
if choice.delta.function_call is not None
|
||||
)
|
||||
|
||||
|
||||
async def _run_windowed(
|
||||
guardrail: CustomGuardrail,
|
||||
content_chunks: List[str],
|
||||
end_of_stream_only: bool = False,
|
||||
tool_argument_chunks: List[str] | None = None,
|
||||
legacy_function_call: bool = False,
|
||||
) -> tuple[List[Any], List[int]]:
|
||||
guardrail.streaming_buffer_until_moderated = True
|
||||
guardrail.streaming_buffer_release_on_scan = True
|
||||
guardrail.streaming_end_of_stream_only = end_of_stream_only
|
||||
guardrail.streaming_sampling_rate = 2
|
||||
unified = UnifiedLLMGuardrails()
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/v1/chat/completions")
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": [guardrail.guardrail_name]},
|
||||
}
|
||||
collected: List[Any] = []
|
||||
yielded_count: List[int] = []
|
||||
async for chunk in unified.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=_windowed_chat_stream(
|
||||
yielded_count, collected, content_chunks, tool_argument_chunks, legacy_function_call
|
||||
),
|
||||
request_data=request_data,
|
||||
):
|
||||
collected.append(chunk)
|
||||
return collected, yielded_count
|
||||
|
||||
|
||||
def _responses_message_stream_events(text_chunks: List[str]) -> List[dict]:
|
||||
message = {"type": "message", "id": "msg_1", "status": "completed", "role": "assistant"}
|
||||
content = [{"type": "output_text", "text": "".join(text_chunks), "annotations": []}]
|
||||
return [
|
||||
{"type": "response.output_item.added", "output_index": 0, "item": {**message, "content": []}},
|
||||
*(
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_1",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": text,
|
||||
}
|
||||
for text in text_chunks
|
||||
),
|
||||
{"type": "response.output_item.done", "output_index": 0, "item": {**message, "content": content}},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_1",
|
||||
"model": "gpt-4o",
|
||||
"status": "completed",
|
||||
"output": [{**message, "content": content}],
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _responses_truncated_function_call_events(text: str, argument_chunks: List[str]) -> List[dict]:
|
||||
message = {"type": "message", "id": "msg_1", "status": "completed", "role": "assistant"}
|
||||
content = [{"type": "output_text", "text": text, "annotations": []}]
|
||||
function_call = {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "run_shell"}
|
||||
return [
|
||||
{"type": "response.output_item.added", "output_index": 0, "item": {**message, "content": []}},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_1",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": text,
|
||||
},
|
||||
{"type": "response.output_item.added", "output_index": 1, "item": {**function_call, "arguments": ""}},
|
||||
*(
|
||||
{"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 1, "delta": arguments}
|
||||
for arguments in argument_chunks
|
||||
),
|
||||
{
|
||||
"type": "response.incomplete",
|
||||
"response": {
|
||||
"id": "resp_1",
|
||||
"model": "gpt-4o",
|
||||
"status": "incomplete",
|
||||
"output": [
|
||||
{**message, "content": content},
|
||||
{**function_call, "arguments": "".join(argument_chunks), "status": "incomplete"},
|
||||
],
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
async def _replay(events: List[dict]) -> AsyncGenerator[dict, None]:
|
||||
for event in events:
|
||||
yield event
|
||||
|
||||
|
||||
async def _run_windowed_responses(guardrail: CustomGuardrail, events: List[dict]) -> str:
|
||||
guardrail.streaming_buffer_until_moderated = True
|
||||
guardrail.streaming_buffer_release_on_scan = True
|
||||
guardrail.streaming_sampling_rate = 2
|
||||
unified = UnifiedLLMGuardrails()
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test", request_route="/v1/responses")
|
||||
request_data = {
|
||||
"input": "hi",
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": [guardrail.guardrail_name]},
|
||||
}
|
||||
collected: List[Any] = []
|
||||
async for chunk in unified.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=_replay(events),
|
||||
request_data=request_data,
|
||||
):
|
||||
collected.append(chunk)
|
||||
return json.dumps([chunk if isinstance(chunk, dict) else str(chunk) for chunk in collected])
|
||||
|
||||
|
||||
def _chat_text(chunks: List[Any]) -> str:
|
||||
return "".join(
|
||||
choice.delta.content or ""
|
||||
for chunk in chunks
|
||||
if isinstance(chunk, ModelResponseStream)
|
||||
for choice in chunk.choices
|
||||
)
|
||||
|
||||
|
||||
async def _run(guardrail: CustomGuardrail) -> str:
|
||||
# Rubrik's real config: end-of-stream-only moderation. Without buffering
|
||||
# this releases every chunk before moderation runs (content leaks on
|
||||
|
|
@ -159,6 +455,110 @@ async def test_buffered_clean_releases_all_content():
|
|||
assert BLOCK_MESSAGE not in raw
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_buffer_releases_after_each_passing_scan():
|
||||
guardrail = _CountingPassingGuardrail(guardrail_name="windowed-pass", event_hook="post_call")
|
||||
content_chunks = ["one ", "two ", "three ", "four ", "five ", "six "]
|
||||
|
||||
collected, yielded_count = await _run_windowed(guardrail, content_chunks)
|
||||
|
||||
assert yielded_count[2] >= 2
|
||||
assert yielded_count == [0, 0, 2, 2, 4, 4, 6]
|
||||
assert _chat_text(collected) == "".join(content_chunks)
|
||||
assert guardrail.scan_count > 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_buffer_drops_blocked_window():
|
||||
guardrail = _SecondScanBlockingGuardrail(guardrail_name="windowed-block", event_hook="post_call")
|
||||
content_chunks = ["one ", "two ", "MARKER ", "four ", "five ", "six "]
|
||||
|
||||
collected, _ = await _run_windowed(guardrail, content_chunks)
|
||||
raw = _decode(collected)
|
||||
|
||||
assert _chat_text(collected) == "one two "
|
||||
assert "MARKER" not in raw
|
||||
assert BLOCK_MESSAGE in raw
|
||||
assert '"error"' not in raw
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_buffer_holds_tool_call_windows_until_end_of_stream_scan():
|
||||
guardrail = _ToolCallRecordingGuardrail(guardrail_name="windowed-tools", event_hook="post_call")
|
||||
content_chunks = ["one ", "two ", "three "]
|
||||
tool_argument_chunks = ['{"cmd": "', TOOL_ARGUMENTS_MARKER, '"}']
|
||||
|
||||
collected, yielded_count = await _run_windowed(guardrail, content_chunks, tool_argument_chunks=tool_argument_chunks)
|
||||
|
||||
assert yielded_count == [0, 0, 2, 2, 2, 2, 2]
|
||||
assert _chat_text(collected) == "".join(content_chunks)
|
||||
assert _tool_argument_text(collected) == "".join(tool_argument_chunks)
|
||||
assert guardrail.tool_call_scan_indexes == [guardrail.scan_count]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_buffer_holds_legacy_function_call_windows_until_end_of_stream():
|
||||
guardrail = _PassingGuardrail(guardrail_name="windowed-functions", event_hook="post_call")
|
||||
content_chunks = ["one ", "two ", "three "]
|
||||
function_argument_chunks = ['{"cmd": "', TOOL_ARGUMENTS_MARKER, '"}']
|
||||
|
||||
collected, yielded_count = await _run_windowed(
|
||||
guardrail, content_chunks, tool_argument_chunks=function_argument_chunks, legacy_function_call=True
|
||||
)
|
||||
|
||||
assert yielded_count == [0, 0, 2, 2, 2, 2, 2]
|
||||
assert _chat_text(collected) == "".join(content_chunks)
|
||||
assert _function_call_argument_text(collected) == "".join(function_argument_chunks)
|
||||
|
||||
|
||||
def test_tool_call_only_scan_key_is_not_skipped_as_empty():
|
||||
assert _is_redundant_scan(StreamingScanKey(texts=("",)), None) is True
|
||||
assert _is_redundant_scan(StreamingScanKey(texts=("",), tool_calls=("run_shell:{}",)), None) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_responses_output_item_done_round_keeps_text_window_withheld():
|
||||
guardrail = _MarkerBlockingGuardrail(
|
||||
guardrail_name="windowed-responses", event_hook="post_call", marker=ORIGINAL_MARKER
|
||||
)
|
||||
events = _responses_message_stream_events(["one ", f"{ORIGINAL_MARKER} "])
|
||||
|
||||
raw = await _run_windowed_responses(guardrail, events)
|
||||
|
||||
assert ORIGINAL_MARKER not in raw, f"unscanned window leaked: {raw!r}"
|
||||
assert BLOCK_MESSAGE in raw
|
||||
assert guardrail.scan_count >= 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_responses_incomplete_stream_scans_tool_call_before_release():
|
||||
guardrail = _MarkerBlockingGuardrail(
|
||||
guardrail_name="windowed-responses-tools",
|
||||
event_hook="post_call",
|
||||
marker=TOOL_ARGUMENTS_MARKER,
|
||||
field="tool_calls",
|
||||
)
|
||||
events = _responses_truncated_function_call_events("hi ", ['{"cmd": "', TOOL_ARGUMENTS_MARKER, '"}'])
|
||||
|
||||
raw = await _run_windowed_responses(guardrail, events)
|
||||
|
||||
assert '"hi "' in raw
|
||||
assert TOOL_ARGUMENTS_MARKER not in raw, f"unscanned tool call leaked: {raw!r}"
|
||||
assert BLOCK_MESSAGE in raw
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windowed_buffer_with_explicit_end_of_stream_only_stays_fully_buffered():
|
||||
guardrail = _CountingPassingGuardrail(guardrail_name="windowed-eos", event_hook="post_call")
|
||||
content_chunks = ["one ", "two ", "three ", "four ", "five ", "six "]
|
||||
|
||||
collected, yielded_count = await _run_windowed(guardrail, content_chunks, end_of_stream_only=True)
|
||||
|
||||
assert yielded_count == [0, 0, 0, 0, 0, 0, 0]
|
||||
assert _chat_text(collected) == "".join(content_chunks)
|
||||
assert guardrail.scan_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_buffered_mode_disabled_for_content_rewriting_guardrail():
|
||||
"""Buffered replay yields the withheld *original* chunks verbatim, which
|
||||
|
|
|
|||
|
|
@ -682,6 +682,22 @@ async def test_provider_specific_params_includes_embedding_toggle():
|
|||
assert field["default_value"] is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_specific_params_exposes_bedrock_streaming_flags():
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import get_provider_specific_params
|
||||
|
||||
provider_params = await get_provider_specific_params()
|
||||
|
||||
bedrock = provider_params["bedrock"]
|
||||
assert "guardrailIdentifier" in bedrock
|
||||
assert "guardrailVersion" in bedrock
|
||||
assert bedrock["streaming_buffer_release_on_scan"]["type"] == "boolean"
|
||||
assert bedrock["streaming_buffer_release_on_scan"]["default_value"] is False
|
||||
assert bedrock["streaming_buffer_until_moderated"]["default_value"] is True
|
||||
assert bedrock["streaming_end_of_stream_only"]["type"] == "boolean"
|
||||
assert bedrock["streaming_sampling_rate"]["type"] == "number"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_specific_params_includes_hide_secrets():
|
||||
"""hide-secrets lives in the enterprise package so it is not in
|
||||
|
|
|
|||
|
|
@ -416,6 +416,45 @@ class TestRotateVirtualKeyInSecretManager:
|
|||
assert call_kwargs["new_secret_name"] == "test-key-alias-new"
|
||||
assert call_kwargs["new_secret_value"] == "sk-new-key"
|
||||
|
||||
@pytest.mark.parametrize("key_alias", ["test-key-alias", None])
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotated_hook_without_request_body_syncs_secret_manager(
|
||||
self, monkeypatch: pytest.MonkeyPatch, key_alias: str | None
|
||||
):
|
||||
import litellm
|
||||
from litellm.proxy._types import GenerateKeyResponse, LiteLLM_VerificationToken
|
||||
from litellm.secret_managers.base_secret_manager import BaseSecretManager
|
||||
from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem
|
||||
|
||||
mock_secret_manager: Final = MagicMock(spec=BaseSecretManager)
|
||||
mock_secret_manager.async_rotate_secret = AsyncMock(return_value={"status": "success"})
|
||||
monkeypatch.setattr(litellm, "secret_manager_client", mock_secret_manager)
|
||||
monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.AWS_SECRET_MANAGER)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"_key_management_settings",
|
||||
KeyManagementSettings(store_virtual_keys=True, prefix_for_stored_virtual_keys="litellm/"),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||||
|
||||
existing_key_row: Final = LiteLLM_VerificationToken(token="hashed-old-token", key_alias=key_alias)
|
||||
response: Final = GenerateKeyResponse(token_id="hashed-new-token", key="sk-new-key", key_alias=key_alias)
|
||||
|
||||
await KeyManagementEventHooks.async_key_rotated_hook(
|
||||
data=None,
|
||||
existing_key_row=existing_key_row,
|
||||
response=response,
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
expected_secret_name: Final = f"litellm/{key_alias or 'virtual-key-hashed-old-token'}"
|
||||
mock_secret_manager.async_rotate_secret.assert_awaited_once_with(
|
||||
current_secret_name=expected_secret_name,
|
||||
new_secret_name=expected_secret_name,
|
||||
new_secret_value="sk-new-key",
|
||||
optional_params=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotate_virtual_key_when_store_virtual_keys_disabled(self):
|
||||
"""Test that rotation is skipped when store_virtual_keys is False."""
|
||||
|
|
@ -474,6 +513,112 @@ class TestRotateVirtualKeyInSecretManager:
|
|||
mock_secret_manager.async_rotate_secret.assert_not_called()
|
||||
|
||||
|
||||
class TestKeyUpdatedSecretManagerSync:
|
||||
|
||||
@staticmethod
|
||||
def _configure_secret_manager(
|
||||
monkeypatch: pytest.MonkeyPatch, stored_value: str | None, store_virtual_keys: bool = True
|
||||
) -> MagicMock:
|
||||
import litellm
|
||||
from litellm.secret_managers.base_secret_manager import BaseSecretManager
|
||||
from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem
|
||||
|
||||
mock_secret_manager: Final = MagicMock(spec=BaseSecretManager)
|
||||
mock_secret_manager.async_read_secret = AsyncMock(return_value=stored_value)
|
||||
mock_secret_manager.async_rotate_secret = AsyncMock(return_value={"status": "success"})
|
||||
monkeypatch.setattr(litellm, "secret_manager_client", mock_secret_manager)
|
||||
monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.AWS_SECRET_MANAGER)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"_key_management_settings",
|
||||
KeyManagementSettings(store_virtual_keys=store_virtual_keys, prefix_for_stored_virtual_keys="litellm/"),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||||
return mock_secret_manager
|
||||
|
||||
@pytest.mark.parametrize("existing_alias", ["old-alias", None])
|
||||
@pytest.mark.asyncio
|
||||
async def test_updated_hook_renames_secret_when_alias_changes(
|
||||
self, monkeypatch: pytest.MonkeyPatch, existing_alias: str | None
|
||||
):
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
|
||||
mock_secret_manager: Final = self._configure_secret_manager(monkeypatch, stored_value="sk-stored-key")
|
||||
existing_key_row: Final = LiteLLM_VerificationToken(token="hashed-token", key_alias=existing_alias)
|
||||
|
||||
await KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=UpdateKeyRequest(key="hashed-token", key_alias="new-alias"),
|
||||
existing_key_row=existing_key_row,
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
current_secret_name: Final = f"litellm/{existing_alias or 'virtual-key-hashed-token'}"
|
||||
mock_secret_manager.async_read_secret.assert_awaited_once_with(
|
||||
secret_name=current_secret_name, optional_params=None
|
||||
)
|
||||
mock_secret_manager.async_rotate_secret.assert_awaited_once_with(
|
||||
current_secret_name=current_secret_name,
|
||||
new_secret_name="litellm/new-alias",
|
||||
new_secret_value="sk-stored-key",
|
||||
optional_params=None,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("requested_alias", ["same-alias", None])
|
||||
@pytest.mark.asyncio
|
||||
async def test_updated_hook_leaves_secret_alone_when_alias_unchanged(
|
||||
self, monkeypatch: pytest.MonkeyPatch, requested_alias: str | None
|
||||
):
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
|
||||
mock_secret_manager: Final = self._configure_secret_manager(monkeypatch, stored_value="sk-stored-key")
|
||||
|
||||
await KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=UpdateKeyRequest(key="hashed-token", key_alias=requested_alias, max_budget=10.0),
|
||||
existing_key_row=LiteLLM_VerificationToken(token="hashed-token", key_alias="same-alias"),
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
mock_secret_manager.async_read_secret.assert_not_awaited()
|
||||
mock_secret_manager.async_rotate_secret.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_updated_hook_skips_rename_when_secret_missing(self, monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
|
||||
mock_secret_manager: Final = self._configure_secret_manager(monkeypatch, stored_value=None)
|
||||
|
||||
await KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=UpdateKeyRequest(key="hashed-token", key_alias="new-alias"),
|
||||
existing_key_row=LiteLLM_VerificationToken(token="hashed-token", key_alias="old-alias"),
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
mock_secret_manager.async_rotate_secret.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_updated_hook_ignores_alias_change_when_store_virtual_keys_disabled(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest
|
||||
|
||||
mock_secret_manager: Final = self._configure_secret_manager(
|
||||
monkeypatch, stored_value="sk-stored-key", store_virtual_keys=False
|
||||
)
|
||||
|
||||
await KeyManagementEventHooks.async_key_updated_hook(
|
||||
data=UpdateKeyRequest(key="hashed-token", key_alias="new-alias"),
|
||||
existing_key_row=LiteLLM_VerificationToken(token="hashed-token", key_alias="old-alias"),
|
||||
response=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
mock_secret_manager.async_read_secret.assert_not_awaited()
|
||||
mock_secret_manager.async_rotate_secret.assert_not_awaited()
|
||||
|
||||
|
||||
class TestKeyUpdatedAuditLogObjectId:
|
||||
"""Tests that /key/update audit logs never store the raw virtual key (issue #31620)."""
|
||||
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue