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

This commit is contained in:
mrinal 2026-09-10 18:14:25 +00:00
commit 10a5761bb9
146 changed files with 9380 additions and 1025 deletions

View file

@ -113,7 +113,7 @@ jobs:
- name: Check ruff format
if: steps.changes.outputs.decision != 'skip'
run: |
git diff --name-only --diff-filter=ACMR "$GATE_BASE_SHA" HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
git diff --name-only --diff-filter=ACMR "$GATE_BASE_SHA" HEAD -- ':(glob)litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
if [ ! -s "$RUNNER_TEMP/ruff_format_files.txt" ]; then
echo "No changed litellm Python files to check with ruff format."
exit 0
@ -172,7 +172,7 @@ jobs:
- name: Check tests/e2e basedpyright (zero errors)
if: steps.changes.outputs.decision != 'skip'
run: |
if git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- 'tests/e2e/**/*.py' | grep -q .; then
if git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- ':(glob)tests/e2e/**/*.py' | grep -q .; then
uv run --no-sync basedpyright tests/e2e
else
echo "No changed tests/e2e Python files; skipping."

View file

@ -150,8 +150,8 @@ lint-install:
# Diff-scoped format check, mirroring test-linting.yml's "Check ruff format" step:
# only the litellm Python files changed vs the base are checked, so a pre-existing
# format issue elsewhere doesn't block an unrelated commit. Git pathspecs match
# recursively, so 'litellm/*.py' covers nested modules and the top-level files that
# CI's 'litellm/**/*.py' skips, which makes this target a superset of the CI step.
# recursively, so 'litellm/*.py' covers top-level files and nested modules alike,
# the same set CI's ':(glob)litellm/**/*.py' selects.
lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
@base_ref=$$($(RESOLVE_BASE)) && \
changed=$$(git diff --name-only --diff-filter=ACMR "$$base_ref...HEAD" -- 'litellm/*.py') && \

View file

@ -45,9 +45,11 @@ from typing import (
TYPE_CHECKING,
Union,
)
from collections.abc import Mapping
from litellm.types.integrations.datadog import DatadogInitParams
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm.litellm_core_utils.core_helpers import drop_params_env_flag
from litellm.types.integrations.pointfive import PointFiveInitParams
from litellm._logging import (
set_verbose,
_turn_on_debug,
@ -154,6 +156,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"smtp_email",
"deepeval",
"s3_v2",
"pointfive",
"aws_sqs",
"vector_store_pre_call_hook",
"dotprompt",
@ -439,6 +442,7 @@ s3_audit_callback_params: Optional[Dict] = None
datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] = None
datadog_params: Optional[Union[DatadogInitParams, Dict]] = None
newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None
pointfive_params: Optional[Union[PointFiveInitParams, Mapping[str, object]]] = None
aws_sqs_callback_params: Optional[Dict] = None
generic_logger_headers: Optional[Dict] = None
default_key_generate_params: Optional[Dict] = None
@ -475,6 +479,7 @@ prometheus_metrics_config: Optional[List] = None
prometheus_exclude_metrics: Optional[List[str]] = None
prometheus_exclude_labels: Optional[List[str]] = None
prometheus_emit_stream_label: bool = False
prometheus_emit_input_sequence_length_label: bool = False
prometheus_deployment_and_latency_caller_identity: Literal[
"api_key_alias",
"user_email",

View file

@ -1683,6 +1683,7 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT
RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS", "3")))
RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2"))
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000))
PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25")))
DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute
PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597))
RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500")))

View file

@ -54,18 +54,22 @@ def missing_streamable_http_client_error() -> ImportError:
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
METHOD_NOT_FOUND,
ClientResult,
GetPromptRequestParams,
GetPromptResult,
ListPromptsResult,
ListResourcesResult,
ListResourceTemplatesResult,
Prompt,
ResourceTemplate,
ServerNotification,
ServerRequest,
TextContent,
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
@ -777,8 +781,19 @@ class MCPClient:
"""List available prompts from the server."""
verbose_logger.debug("MCP client listing tools from %s", self.server_url or "stdio")
async def _list_prompts_operation(session: ClientSession):
return await session.list_prompts()
async def _list_prompts_operation(session: ClientSession) -> ListPromptsResult:
capabilities: Final = session.get_server_capabilities()
if capabilities is not None and capabilities.prompts is None:
return ListPromptsResult(prompts=[])
try:
return await session.list_prompts()
except McpError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
verbose_logger.debug(
"MCP client list_prompts is unsupported by %s: %s", self.server_url or "stdio", error
)
return ListPromptsResult(prompts=[])
try:
result: Final = await self.run_with_session(_list_prompts_operation)
@ -854,8 +869,19 @@ class MCPClient:
"""List available resources from the server."""
verbose_logger.debug("MCP client listing resources from %s", self.server_url or "stdio")
async def _list_resources_operation(session: ClientSession):
return await session.list_resources()
async def _list_resources_operation(session: ClientSession) -> ListResourcesResult:
capabilities: Final = session.get_server_capabilities()
if capabilities is not None and capabilities.resources is None:
return ListResourcesResult(resources=[])
try:
return await session.list_resources()
except McpError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
verbose_logger.debug(
"MCP client list_resources is unsupported by %s: %s", self.server_url or "stdio", error
)
return ListResourcesResult(resources=[])
try:
result: Final = await self.run_with_session(_list_resources_operation)
@ -890,8 +916,19 @@ class MCPClient:
"""List available resource templates from the server."""
verbose_logger.debug("MCP client listing resource templates from %s", self.server_url or "stdio")
async def _list_resource_templates_operation(session: ClientSession):
return await session.list_resource_templates()
async def _list_resource_templates_operation(session: ClientSession) -> ListResourceTemplatesResult:
capabilities: Final = session.get_server_capabilities()
if capabilities is not None and capabilities.resources is None:
return ListResourceTemplatesResult(resourceTemplates=[])
try:
return await session.list_resource_templates()
except McpError as error:
if error.error.code != METHOD_NOT_FOUND:
raise
verbose_logger.debug(
"MCP client list_resource_templates is unsupported by %s: %s", self.server_url or "stdio", error
)
return ListResourceTemplatesResult(resourceTemplates=[])
try:
result: Final = await self.run_with_session(_list_resource_templates_operation)

View file

@ -378,6 +378,27 @@
},
"description": "OpenTelemetry Logging Integration"
},
{
"id": "pointfive",
"displayName": "PointFive",
"logo": "pointfive.png",
"supports_key_team_logging": false,
"dynamic_params": {
"POINTFIVE_API_KEY": {
"type": "password",
"ui_name": "API Key",
"description": "PointFive API key, used to request an upload url for each batch of logs",
"required": true
},
"POINTFIVE_API_URL": {
"type": "text",
"ui_name": "API URL",
"description": "PointFive API endpoint. Leave blank to use https://api.pointfive.co/api/v1/ingestion",
"required": false
}
},
"description": "PointFive Logging Integration"
},
{
"id": "s3",
"displayName": "S3",

View file

@ -0,0 +1,5 @@
"""PointFive logging integration for LiteLLM."""
from litellm.integrations.pointfive.logger import PointFiveLogger
__all__ = ("PointFiveLogger",)

View file

@ -0,0 +1,304 @@
"""
PointFive logging integration.
Buffers ``StandardLoggingPayload`` records and ships each flush as one gzipped
newline-delimited JSON object, rather than one object per request. Uploads go through a
presigned URL issued by the PointFive API, so the proxy needs no cloud credentials and
runs unchanged wherever it is hosted.
"""
import asyncio
from collections.abc import Mapping
from datetime import datetime
from typing import Final
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.pointfive.payload import chunk_lines, encode_lines, serialize_records
from litellm.integrations.pointfive.upload_client import PointFiveUploadClient, PointFiveUploadError
from litellm.litellm_core_utils.redact_messages import (
redacted_standard_logging_payload,
should_redact_message_logging,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, httpxSpecialProvider
from litellm.secret_managers.main import get_secret_str
from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus
from litellm.types.integrations.pointfive import DEFAULT_API_URL, PointFiveInitParams, PointFiveUploadFailure
_ENV_REFERENCE_PREFIX: Final = "os.environ/"
def _resolved_secret(value: str | None) -> str | None:
"""
Resolve a config value that may name a secret, in any shape the secret manager accepts.
A reference that resolves to nothing stays unresolved rather than falling back to its own
text, so an unset ``os.environ/NAME`` reports a missing key instead of being sent as one.
"""
if value is None:
return None
resolved: Final = get_secret_str(value)
if resolved:
return resolved
return None if value.startswith(_ENV_REFERENCE_PREFIX) else value
def _configured_params() -> PointFiveInitParams:
"""Read ``litellm.pointfive_params``, validating a raw config dict on the way through."""
configured: Final = litellm.pointfive_params
if isinstance(configured, PointFiveInitParams):
return configured
if isinstance(configured, Mapping):
return PointFiveInitParams.model_validate(configured)
return PointFiveInitParams()
def _resolved_api_key(params: PointFiveInitParams) -> str | None:
"""Prefer the configured key, falling back to the environment the proxy UI writes."""
return _resolved_secret(params.api_key) or get_secret_str("POINTFIVE_API_KEY")
def _resolved_api_url(params: PointFiveInitParams) -> str:
"""Prefer the configured url, then the environment, then the public endpoint."""
return _resolved_secret(params.api_url) or get_secret_str("POINTFIVE_API_URL") or DEFAULT_API_URL
def _upload_client_for(params: PointFiveInitParams) -> PointFiveUploadClient:
"""
Build an upload client for the key and url configured right now.
Resolved per call rather than kept: the proxy ui writes new values into the
environment of a running proxy, and reading them once would need a restart to take
effect. ``get_async_httpx_client`` is cached, so this reuses the same connections.
"""
api_key: Final = _resolved_api_key(params)
if not api_key:
raise ValueError(
"pointfive logging requires an api key. Set POINTFIVE_API_KEY, or "
"litellm_settings.pointfive_params.api_key in config.yaml"
)
return PointFiveUploadClient(
api_key=api_key,
api_url=_resolved_api_url(params),
http_client=get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback),
max_retries=params.max_upload_retries,
)
class PointFiveLogger(CustomBatchLogger):
"""Batching callback that ships LiteLLM request logs to PointFive."""
preserve_events_added_during_flush = True
def __init__(
self,
params: PointFiveInitParams | None = None,
upload_client: PointFiveUploadClient | None = None,
start_periodic_flush: bool = True,
) -> None:
resolved: Final = params if params is not None else _configured_params()
self.max_batch_bytes: Final = resolved.max_batch_bytes
self.params: Final = resolved
self.given_upload_client: Final = upload_client
if upload_client is None:
_upload_client_for(resolved) # refuse to start without a key, rather than at the first flush
super().__init__(
flush_lock=asyncio.Lock(),
batch_size=resolved.batch_size,
flush_interval=resolved.flush_interval,
turn_off_message_logging=bool(resolved.turn_off_message_logging),
)
self._flushing: bool = False
self._batch_flush_task: asyncio.Task[None] | None = None
self._periodic_flush_task: asyncio.Task[None] | None = (
self._start_periodic_flush_task() if start_periodic_flush else None
)
@property
def upload_client(self) -> PointFiveUploadClient:
"""The client for the currently configured key and url, so a ui edit needs no restart."""
if self.given_upload_client is not None:
return self.given_upload_client
return _upload_client_for(self.params)
def _start_periodic_flush_task(self) -> asyncio.Task[None] | None:
"""Start the periodic flush only once an event loop is actually running."""
try:
loop: Final = asyncio.get_running_loop()
except RuntimeError:
return None
return loop.create_task(self.periodic_flush())
def _start_batch_flush_task(self) -> None:
"""
Upload a full batch in the background, so no request waits on PointFive.
Awaiting it here put the upload, its retries and their backoff on the caller's
path, and a hung api held a response open for as long as the attempts took.
"""
if self._batch_flush_task is not None and not self._batch_flush_task.done():
return
try:
loop: Final = asyncio.get_running_loop()
except RuntimeError:
return
self._batch_flush_task = loop.create_task(self.flush_queue(skip_if_flushing=True))
def _flush_task_is_alive(self) -> bool:
"""A task whose loop has been closed never runs again, yet never reports itself done."""
task: Final = self._periodic_flush_task
return task is not None and not task.done() and not task.get_loop().is_closed()
async def periodic_flush(self) -> None:
"""
Report in straight away, then flush on the interval as usual.
The inherited loop sleeps first, so a proxy that has just loaded the callback says
nothing for a whole interval, five minutes by default. PointFive shows the integration
as still waiting for its first call for all that time, which reads as a broken setup
rather than an idle one. An empty queue makes this first cycle a ping, so a proxy with
no traffic yet announces itself without uploading an object that holds no records.
"""
await self.flush_queue(skip_if_flushing=True)
await super().periodic_flush()
async def async_log_success_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: datetime,
end_time: datetime,
) -> None:
await self._enqueue(kwargs)
async def async_log_failure_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: datetime,
end_time: datetime,
) -> None:
await self._enqueue(kwargs)
async def _enqueue(self, kwargs: Mapping[str, object]) -> None:
"""Buffer one record, flushing early once the batch threshold is reached."""
try:
if not self._flush_task_is_alive():
self._periodic_flush_task = self._start_periodic_flush_task()
record: Final = self._record_for(kwargs)
if record is None:
verbose_logger.debug("pointfive: event carried no standard_logging_object, skipping")
return
self.log_queue.append(record)
self._drop_overflow()
if len(self.log_queue) >= self.batch_size:
self._start_batch_flush_task()
except Exception: # noqa: BLE001 # logging must never break the request path
verbose_logger.exception("pointfive: failed to queue an event")
def _record_for(self, kwargs: Mapping[str, object]) -> Mapping[str, object] | None:
"""
The record to buffer, redacted the way the framework would have redacted it.
A success reaches a callback already redacted, an async failure does not, so both
the excluded-field list and this callback's own setting are applied here, then the
global, per-request and header settings that only the framework's predicate knows.
"""
details: Final = self.redact_standard_logging_payload_from_model_call_details(
dict(kwargs) # mutable-ok: both framework helpers take the call details as a dict
)
payload: Final = details.get("standard_logging_object")
if not isinstance(payload, dict):
return None
if should_redact_message_logging(details):
return redacted_standard_logging_payload(payload)
return payload
def _drop_overflow(self) -> None:
"""
Hold the queue to its cap as records arrive, not only after a flush has failed.
Never while a flush is running: it holds a snapshot taken by length, and trimming
the front underneath it would make the post-flush drain remove records that arrived
during the upload and were never sent. The next arrival after the flush trims.
"""
if self._flushing:
return
overflow: Final = len(self.log_queue) - self.max_queue_size
if overflow <= 0:
return
del self.log_queue[:overflow]
verbose_logger.warning("pointfive: queue over %s records, dropped %s oldest", self.max_queue_size, overflow)
async def flush_queue(self, skip_if_flushing: bool = False) -> None:
"""
Flush as usual, or report liveness when there is nothing to send.
``CustomBatchLogger`` skips an empty queue entirely, so without this an idle proxy
would look identical to a dead one.
``skip_if_flushing`` is what a full batch, and the loop's opening cycle, pass. Uploading one takes seconds, and
every event arriving meanwhile crosses the threshold too, so each would queue on the
flush lock and then ship the handful of records left behind it. That turns one burst
into a stream of tiny objects, which is what batching exists to avoid. The running
flush already carries what is queued, and the interval catches whatever it missed.
"""
if not self.log_queue:
await self._ping()
return
if skip_if_flushing and self._flushing:
return
self._flushing = True
try:
await super().flush_queue()
finally:
self._flushing = False
async def async_health_check(self) -> IntegrationHealthCheckStatus:
"""Answer the proxy ui test button by asking the api whether it accepts this key."""
try:
failure: Final = await self.upload_client.ping()
except ValueError as missing_key:
return IntegrationHealthCheckStatus(status="unhealthy", error_message=str(missing_key))
if failure is not None:
return IntegrationHealthCheckStatus(status="unhealthy", error_message=failure.detail)
return IntegrationHealthCheckStatus(status="healthy", error_message=None)
async def _ping(self) -> None:
"""Report liveness, never failing the flush over it."""
try:
failure: Final = await self.upload_client.ping()
except ValueError as missing_key:
verbose_logger.warning("pointfive: liveness ping skipped, %s", missing_key)
return
if failure is not None:
verbose_logger.warning("pointfive: liveness ping failed, %s", failure.detail)
async def async_send_batch(self) -> None:
"""
Upload everything queued, split into objects of at most ``max_batch_bytes``.
A retryable failure propagates so ``CustomBatchLogger`` keeps the rest of the batch
for the next flush; the records already shipped or already refused leave the queue
first, so a retry re-sends at most the object that failed. A rejection the server
will refuse again drops that object, since holding it would block every record
queued behind it.
"""
pending: Final = tuple(self.log_queue)
if not pending:
return
client: Final = self.upload_client
chunks: Final = chunk_lines(serialize_records(pending), self.max_batch_bytes)
for index, chunk in enumerate(chunks):
outcome = await client.upload(await encode_lines(chunk))
if not isinstance(outcome, PointFiveUploadFailure):
continue
if outcome.retryable:
del self.log_queue[: sum(len(shipped) for shipped in chunks[:index])]
raise PointFiveUploadError(outcome.detail)
verbose_logger.error("pointfive: dropping %s records, %s", len(chunk), outcome.detail)

View file

@ -0,0 +1,53 @@
"""Turns buffered log records into the gzipped NDJSON objects that get uploaded."""
import gzip
from collections.abc import Iterator, Mapping, Sequence
from itertools import accumulate, groupby, islice
from typing import Final
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
_NEWLINE_BYTES: Final = 1
def serialize_records(records: Sequence[Mapping[str, object]]) -> tuple[str, ...]:
"""Serialize each record to one JSON line."""
return tuple(safe_dumps(record) for record in records)
def _encoded_size(line: str) -> int:
return len(line.encode("utf-8")) + _NEWLINE_BYTES
def _object_indices(sizes: Sequence[int], max_bytes: int) -> Iterator[int]:
"""Number each line with the object it belongs to, opening a new one on overflow."""
def advance(state: tuple[int, int], size: int) -> tuple[int, int]:
index, used = state
return (index + 1, size) if used and used + size > max_bytes else (index, used + size)
return (index for index, _ in islice(accumulate(sizes, advance, initial=(0, 0)), 1, None))
def chunk_lines(lines: Sequence[str], max_bytes: int) -> tuple[tuple[str, ...], ...]:
"""
Group serialized lines into objects of at most ``max_bytes`` uncompressed.
A line above the bound on its own still becomes its own object. A record cannot be
split, and holding it back would stall every record queued behind it.
"""
sizes: Final = tuple(_encoded_size(line) for line in lines)
numbered: Final = zip(_object_indices(sizes, max_bytes), lines, strict=True)
return tuple(tuple(line for _, line in group) for _, group in groupby(numbered, lambda pair: pair[0]))
async def encode_lines(lines: Sequence[str]) -> bytes:
"""
Join lines as NDJSON and gzip them off the event loop.
An object can be several megabytes, and compressing that inline would block the
proxy for as long as it takes.
"""
compress: Final = asyncify(gzip.compress)
return await compress("\n".join(lines).encode("utf-8"))

View file

@ -0,0 +1,194 @@
"""
Uploads one batch to PointFive through a presigned URL.
The proxy holds no cloud credentials. For every batch it asks the PointFive API for a
single-use presigned URL and PUTs the bytes there, so the same plugin runs unchanged on
AWS, GCP, Azure or on-prem. The server picks the object key, so the proxy never chooses
where its data lands.
"""
import asyncio
from collections.abc import Awaitable, Callable
from types import MappingProxyType
from typing import Final
import httpx
from pydantic import BaseModel, Field, ValidationError
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.integrations.pointfive import (
RETRYABLE_UPLOAD_STATUS_CODES,
PointFiveUploadFailure,
PointFiveUploadTarget,
)
UPLOAD_KIND: Final = "LITELLM"
UPLOAD_URL_PATH: Final = "/upload-url"
PING_PATH: Final = "/ping"
PUT_HEADERS: Final = MappingProxyType({"Content-Type": "application/x-ndjson", "Content-Encoding": "gzip"})
class _PresignRequest(BaseModel):
kind: str = UPLOAD_KIND
byte_count: int = Field(serialization_alias="byteCount")
class _PingRequest(BaseModel):
kind: str = UPLOAD_KIND
class _TargetPayload(BaseModel):
upload_url: str = Field(alias="uploadUrl")
object_key: str = Field(alias="objectKey")
class _ErrorPayload(BaseModel):
error: str = ""
class PointFiveUploadError(Exception):
"""A batch could not be uploaded and the failure is worth retrying."""
def _failure_for(response: httpx.Response, what: str) -> PointFiveUploadFailure:
detail: Final = f"{what} returned {response.status_code}"
reason: Final = _refusal_reason(response.text)
return PointFiveUploadFailure(
f"{detail}, {reason}" if reason else detail,
retryable=response.status_code in RETRYABLE_UPLOAD_STATUS_CODES,
)
def _refusal_reason(body: str) -> str:
try:
return _ErrorPayload.model_validate_json(body).error
except ValidationError:
return ""
def _parse_target(body: str) -> PointFiveUploadTarget | PointFiveUploadFailure:
try:
target: Final = _TargetPayload.model_validate_json(body)
except ValidationError:
return PointFiveUploadFailure("pointfive api returned an unreadable body", retryable=False)
return PointFiveUploadTarget(upload_url=target.upload_url, object_key=target.object_key)
class PointFiveUploadClient:
"""Presigns and uploads one batch at a time."""
def __init__(
self,
api_key: str,
api_url: str,
http_client: AsyncHTTPHandler,
max_retries: int,
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
validate_upload_url: Callable[[str], tuple[str, str]] = validate_url,
) -> None:
self.api_key: Final = api_key
self.api_url: Final = api_url.rstrip("/")
self.http_client: Final = http_client
self.max_retries: Final = max_retries
self.sleep: Final = sleep
self.validate_upload_url: Final = validate_upload_url
async def upload(self, body: bytes) -> str | PointFiveUploadFailure:
"""
Upload one gzipped batch, returning the object key it landed at.
Every attempt presigns again, so a retry never reuses a URL that has expired or
has already been consumed.
"""
for attempt in range(self.max_retries):
match await self._upload_once(body):
case PointFiveUploadFailure(retryable=True) as failure:
if attempt + 1 >= self.max_retries:
return PointFiveUploadFailure(
f"{failure.detail}, gave up after {self.max_retries} attempts", retryable=True
)
await self.sleep(float(1 << attempt))
case outcome:
return outcome
return PointFiveUploadFailure("max_upload_retries must be at least 1", retryable=False)
async def _upload_once(self, body: bytes) -> str | PointFiveUploadFailure:
target: Final = await self._presign(len(body))
if isinstance(target, PointFiveUploadFailure):
return target
rejection: Final = await self._put(target, body)
if rejection is not None:
return rejection
verbose_logger.debug("pointfive: uploaded %s gzipped bytes to %s", len(body), target.object_key)
return target.object_key
async def ping(self) -> PointFiveUploadFailure | None:
"""Report that the proxy is alive when it has nothing to upload."""
body: Final = await self._post(PING_PATH, _PingRequest())
if isinstance(body, PointFiveUploadFailure):
return body
return None
async def _presign(self, byte_count: int) -> PointFiveUploadTarget | PointFiveUploadFailure:
"""Ask the PointFive API for a presigned URL sized to this batch."""
body: Final = await self._post(UPLOAD_URL_PATH, _PresignRequest(byte_count=byte_count))
if isinstance(body, PointFiveUploadFailure):
return body
return _parse_target(body)
async def _post(self, path: str, request: BaseModel) -> str | PointFiveUploadFailure:
"""POST one JSON request to the PointFive ingestion API and return its raw body."""
try:
response: Final = await self.http_client.post(
self.api_url + path,
json=request.model_dump(by_alias=True),
headers={ # mutable-ok: AsyncHTTPHandler.post types headers as dict
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
},
)
except httpx.HTTPStatusError as e:
return _failure_for(e.response, "pointfive api")
except Exception as e: # noqa: BLE001 # a transport fault is worth another attempt
return PointFiveUploadFailure(f"pointfive api unreachable: {type(e).__name__}", retryable=True)
return response.text
async def _put(self, target: PointFiveUploadTarget, body: bytes) -> PointFiveUploadFailure | None:
"""
PUT the batch to the presigned URL, which carries its own authorization.
The server chose that URL, so it is treated like any other externally supplied
destination: the host is checked against blocked networks before connecting, and
a redirect is refused rather than followed. A presigned URL never legitimately
redirects, and following one would let a compromised endpoint point the proxy at
an internal service.
"""
destination: Final = self._destination(target.upload_url)
if isinstance(destination, PointFiveUploadFailure):
return destination
url, host = destination
headers: Final = dict(PUT_HEADERS, Host=host) if host else dict(PUT_HEADERS) # mutable-ok: put wants dict
try:
await self.http_client.put(url, data=body, headers=headers, follow_redirects=False)
except httpx.HTTPStatusError as e:
if e.response.is_redirect:
return PointFiveUploadFailure(
f"presigned upload redirected with {e.response.status_code}, refusing to follow", retryable=False
)
return _failure_for(e.response, "presigned upload")
except Exception as e: # noqa: BLE001 # a transport fault is worth another attempt
return PointFiveUploadFailure(f"presigned upload unreachable: {type(e).__name__}", retryable=True)
return None
def _destination(self, upload_url: str) -> tuple[str, str | None] | PointFiveUploadFailure:
if not getattr(litellm, "user_url_validation", True):
return upload_url, None
try:
return self.validate_upload_url(upload_url)
except SSRFError as e:
return PointFiveUploadFailure(f"presigned upload url refused: {e}", retryable=False)

View file

@ -246,6 +246,7 @@ class PrometheusLogger(CustomLogger):
# logger so toggling these flags only takes effect after a
# restart, keeping init-time and runtime label sets in sync.
self._cached_metric_labels: dict[str, list[str]] = {}
self._emit_input_sequence_length_label = litellm.prometheus_emit_input_sequence_length_label is True
_custom_buckets: Final = litellm.prometheus_latency_buckets
self.latency_buckets = tuple(_custom_buckets) if _custom_buckets is not None else LATENCY_BUCKETS
@ -1522,6 +1523,11 @@ class PrometheusLogger(CustomLogger):
# 2. Pyright does not allow us to run isinstance(standard_logging_payload, StandardLoggingPayload) <- this would be ideal
enum_values=enum_values,
label_context=label_context,
input_sequence_length=(
self._get_input_sequence_length(standard_logging_payload, kwargs, response_obj)
if self._emit_input_sequence_length_label
else None
),
)
# set x-ratelimit headers
@ -2192,6 +2198,36 @@ class PrometheusLogger(CustomLogger):
)
self.litellm_remaining_api_key_tokens_for_model.labels(**tokens_labels).set(remaining_tokens)
@staticmethod
def _get_input_sequence_length(
standard_logging_payload: StandardLoggingPayload,
kwargs: Mapping[str, object],
response_obj: object,
) -> str:
prompt_tokens: Final = standard_logging_payload.get("prompt_tokens")
if prompt_tokens:
return get_input_sequence_length_bucket(prompt_tokens)
combined_usage: Final = kwargs.get("combined_usage_object")
if (
combined_usage is not None
and getattr(kwargs.get("_litellm_upstream_reported_usage"), "total_tokens", None) is not None
):
return get_input_sequence_length_bucket(None)
reported_usage: Final = (
response_obj.get("usage") if isinstance(response_obj, dict) else getattr(response_obj, "usage", None)
)
if reported_usage is None and combined_usage is None:
return get_input_sequence_length_bucket(None)
usage_metadata: Final = standard_logging_payload["metadata"].get("usage_object")
if isinstance(usage_metadata, Mapping):
return get_input_sequence_length_bucket(usage_metadata.get("prompt_tokens"))
if combined_usage is None and isinstance(response_obj, dict):
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
normalized_usage: Final[Mapping[str, object]] = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj)
return get_input_sequence_length_bucket(normalized_usage.get("prompt_tokens"))
return get_input_sequence_length_bucket(prompt_tokens)
def _set_latency_metrics(
self,
kwargs: dict,
@ -2202,7 +2238,16 @@ class PrometheusLogger(CustomLogger):
user_api_team_alias: str | None,
enum_values: UserAPIKeyLabelValues,
label_context: PrometheusLabelFactoryContext | None = None,
input_sequence_length: str | None = None,
):
latency_enum_values: Final = (
replace(enum_values, input_sequence_length=input_sequence_length)
if input_sequence_length is not None
else enum_values
)
latency_label_context: Final = (
PrometheusLabelFactoryContext(latency_enum_values) if input_sequence_length is not None else label_context
)
# latency metrics
end_time: Final[datetime] = kwargs.get("end_time") or datetime.now()
start_time: Final[datetime | None] = kwargs.get("start_time")
@ -2220,8 +2265,8 @@ class PrometheusLogger(CustomLogger):
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_llm_api_time_to_first_token_metric"
),
enum_values=enum_values,
label_context=label_context,
enum_values=latency_enum_values,
label_context=latency_label_context,
)
self.litellm_llm_api_time_to_first_token_metric.labels(**_ttft_labels).observe(time_to_first_token_seconds)
self._track_end_user_metric_series(
@ -2241,8 +2286,8 @@ class PrometheusLogger(CustomLogger):
if api_call_total_time_seconds is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_llm_api_latency_metric"),
enum_values=enum_values,
label_context=label_context,
enum_values=latency_enum_values,
label_context=latency_label_context,
)
self.litellm_llm_api_latency_metric.labels(**_labels).observe(api_call_total_time_seconds)
self._track_end_user_metric_series(
@ -2272,8 +2317,8 @@ class PrometheusLogger(CustomLogger):
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_request_total_latency_metric"),
enum_values=enum_values,
label_context=label_context,
enum_values=latency_enum_values,
label_context=latency_label_context,
)
self.litellm_request_total_latency_metric.labels(**_labels).observe(_observed_total_time_seconds)
self._track_end_user_metric_series(

View file

@ -10,9 +10,11 @@ import asyncio
import time
from collections.abc import Mapping
from datetime import datetime
from typing import Final, cast
from typing import TYPE_CHECKING, Final, cast
from urllib.parse import quote
import httpx
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
@ -35,6 +37,9 @@ from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload
from .custom_batch_logger import CustomBatchLogger
if TYPE_CHECKING:
from botocore.credentials import Credentials
class S3Logger(CustomBatchLogger, BaseAWSLLM):
def __init__(
@ -232,6 +237,26 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
f"{get_aws_dns_suffix(self.s3_region_name)}/{encoded_key}"
)
def _sign_put(
self, credentials: "Credentials", url: str, json_string: str, headers: Mapping[str, str]
) -> dict[str, str]: # mutable-ok: [LIT001] AsyncHTTPHandler.put/HTTPHandler.put only accept dict headers
"""
``RefreshableCredentials`` (IMDS roles) may refresh between the access key, secret and token
reads SigV4 performs, producing a mixed-generation signature that S3 rejects with 403.
Freezing first makes the three values one atomic snapshot.
"""
from botocore.auth import S3SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import RefreshableCredentials
frozen: Final = (
credentials.get_frozen_credentials() if isinstance(credentials, RefreshableCredentials) else credentials
)
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=dict(headers))
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(frozen, "s3", aws_region_name).add_auth(aws_request)
return dict(aws_request.headers.items())
def _sse_headers(self) -> Mapping[str, str]:
candidates: Final = {
"x-amz-server-side-encryption": self.s3_server_side_encryption,
@ -317,26 +342,12 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
try:
import base64
import hashlib
from botocore.auth import S3SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
try:
from litellm.litellm_core_utils.asyncify import asyncify
asyncified_get_credentials: Final = asyncify(self.get_credentials)
credentials: Final = await asyncified_get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
aws_session_token=self.s3_aws_session_token,
aws_region_name=self.s3_region_name,
aws_session_name=self.s3_aws_session_name,
aws_profile_name=self.s3_aws_profile_name,
aws_role_name=self.s3_aws_role_name,
aws_web_identity_token=self.s3_aws_web_identity_token,
aws_sts_endpoint=self.s3_aws_sts_endpoint,
)
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
verbose_logger.debug("s3_v2 logger - s3_verify setting: %s", self.s3_verify)
@ -363,19 +374,28 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
**self._sse_headers(),
}
# Sign the request
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
await run_aws_signing(S3SigV4Auth(credentials, "s3", aws_region_name).add_auth, aws_request)
async def signed_put() -> httpx.Response:
credentials: Final = await asyncified_get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
aws_session_token=self.s3_aws_session_token,
aws_region_name=self.s3_region_name,
aws_session_name=self.s3_aws_session_name,
aws_profile_name=self.s3_aws_profile_name,
aws_role_name=self.s3_aws_role_name,
aws_web_identity_token=self.s3_aws_web_identity_token,
aws_sts_endpoint=self.s3_aws_sts_endpoint,
)
signed_headers: Final = await run_aws_signing(self._sign_put, credentials, url, json_string, headers)
try:
return await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
except httpx.HTTPStatusError as error:
return error.response
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
# Make the request with retry for transient S3 errors (500/503)
max_retries: Final = 3
for attempt in range(max_retries):
response = await self.async_httpx_client.put(url, data=json_string, headers=signed_headers)
if response.status_code in (500, 503) and attempt < max_retries - 1:
response = await signed_put()
if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
wait_time = 2**attempt # 1s, 2s
verbose_logger.warning(
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",
@ -479,20 +499,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
try:
import base64
import hashlib
from botocore.auth import S3SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
try:
verbose_logger.debug("s3_v2 logger - uploading data to s3 - %s", batch_logging_element.s3_object_key)
credentials: Final[Credentials] = self.get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
aws_session_token=self.s3_aws_session_token,
aws_region_name=self.s3_region_name,
)
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
@ -516,22 +526,24 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
**self._sse_headers(),
}
# Sign the request
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
httpx_client: Final = _get_httpx_client(
params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None)
)
# Make the request with retry for transient S3 errors (500/503)
def signed_put() -> httpx.Response:
credentials: Final = self.get_credentials(
aws_access_key_id=self.s3_aws_access_key_id,
aws_secret_access_key=self.s3_aws_secret_access_key,
aws_session_token=self.s3_aws_session_token,
aws_region_name=self.s3_region_name,
)
signed_headers: Final = self._sign_put(credentials, url, json_string, headers)
return httpx_client.put(url, data=json_string, headers=signed_headers)
max_retries: Final = 3
for attempt in range(max_retries):
response = httpx_client.put(url, data=json_string, headers=signed_headers)
if response.status_code in (500, 503) and attempt < max_retries - 1:
response = signed_put()
if response.status_code in (403, 500, 503) and attempt < max_retries - 1:
wait_time = 2**attempt # 1s, 2s
verbose_logger.warning(
"S3 upload returned %s, retrying in %ss (attempt %s/%s) key=%s",

View file

@ -43,6 +43,7 @@ from litellm.integrations.newrelic import NewRelicLogger
from litellm.integrations.openmeter import OpenMeterLogger
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.integrations.opik.opik import OpikLogger
from litellm.integrations.pointfive import PointFiveLogger
from litellm.integrations.posthog import PostHogLogger
from litellm.integrations.prometheus import PrometheusLogger
from litellm.integrations.s3_v2 import S3Logger
@ -95,6 +96,7 @@ class CustomLoggerRegistry:
"agentops": AgentOps,
"deepeval": DeepEvalLogger,
"s3_v2": S3Logger,
"pointfive": PointFiveLogger,
"aws_sqs": SQSLogger,
"dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler,
"dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3,

View file

@ -17,6 +17,7 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
"aws_session_tags",
"aws_bedrock_runtime_endpoint",
"aws_bedrock_project_id",
}

View file

@ -120,7 +120,6 @@ from litellm.types.utils import (
CachingDetails,
CallTypes,
CostBreakdown,
CostResponseTypes,
CustomPricingLiteLLMParams,
DynamicPromptManagementParamLiteral,
EmbeddingResponse,
@ -183,6 +182,7 @@ from ..integrations.lunary import LunaryLogger
from ..integrations.newrelic import NewRelicLogger
from ..integrations.openmeter import OpenMeterLogger
from ..integrations.opik.opik import OpikLogger
from ..integrations.pointfive import PointFiveLogger
from ..integrations.posthog import PostHogLogger
from ..integrations.prompt_layer import PromptLayerLogger
from ..integrations.s3 import S3Logger
@ -204,7 +204,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, LoggedRelayResponse
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
@ -2381,7 +2381,7 @@ class Logging(LiteLLMLoggingBaseClass):
self,
raw_bytes: list[bytes],
provider_config: "BasePassthroughConfig",
) -> Optional["CostResponseTypes"]:
) -> Optional["LoggedRelayResponse"]:
all_chunks: Final = provider_config._convert_raw_bytes_to_str_lines(raw_bytes)
complete_streaming_response: Final = provider_config.handle_logging_collected_chunks(
all_chunks=all_chunks,
@ -4377,6 +4377,14 @@ def _init_custom_logger_compatible_class(
_s3_v2_logger: Final = S3V2Logger()
_in_memory_loggers.append(_s3_v2_logger)
return _s3_v2_logger
elif logging_integration == "pointfive":
for callback in _in_memory_loggers:
if isinstance(callback, PointFiveLogger):
return callback
_pointfive_logger: Final = PointFiveLogger()
_in_memory_loggers.append(_pointfive_logger)
return _pointfive_logger
elif logging_integration == "aws_sqs":
for callback in _in_memory_loggers:
if isinstance(callback, SQSLogger):
@ -5065,6 +5073,10 @@ def get_custom_logger_compatible_class(
for callback in _in_memory_loggers:
if isinstance(callback, S3V2Logger):
return callback
elif logging_integration == "pointfive":
for callback in _in_memory_loggers:
if isinstance(callback, PointFiveLogger):
return callback
elif logging_integration == "aws_sqs":
for callback in _in_memory_loggers:
if isinstance(callback, SQSLogger):

View file

@ -441,7 +441,15 @@ class LoggingCallbackManager:
return result
def get_callback_objects(self) -> tuple[tuple[str, CustomLogger | Callable], ...]:
return tuple(
(self._get_callback_string(callback), callback)
for callback in self._get_all_callbacks()
if not isinstance(callback, str)
)
def _get_callback_string(self, callback: CustomLogger | Callable | str) -> str:
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.litellm_core_utils.custom_logger_registry import (
CustomLoggerRegistry,
)
@ -449,6 +457,8 @@ class LoggingCallbackManager:
"""Convert a callback to its string representation"""
if isinstance(callback, str):
return callback
elif isinstance(callback, OpenTelemetry) and callback.callback_name is not None:
return callback.callback_name
elif isinstance(callback, CustomLogger):
# Try to get the string representation from the registry
callback_str: Final = CustomLoggerRegistry.get_callback_str_from_class_type(type(callback))

View file

@ -6,8 +6,8 @@ import io
import json
import mimetypes
import re
from collections.abc import Iterable, Iterator, Mapping, Sequence
from itertools import groupby
from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence
from itertools import groupby, islice
from os import PathLike
from pathlib import Path
from types import MappingProxyType
@ -1320,17 +1320,128 @@ def flatten_top_level_schema_combinators(schema: Mapping[str, object]) -> Mappin
return _flatten_schema_against_root(schema, schema, frozenset(), 0, {}) # mutable-ok: fresh per-call $ref memo
def tool_with_flattened_parameters(tool: Mapping[str, object]) -> Mapping[str, object]:
_SUBSCHEMA_KEYWORDS: Final = frozenset(
{
"additionalItems",
"additionalProperties",
"contains",
"else",
"if",
"items",
"not",
"propertyNames",
"then",
"unevaluatedItems",
"unevaluatedProperties",
}
)
_SUBSCHEMA_LIST_KEYWORDS: Final = frozenset({"allOf", "anyOf", "items", "oneOf", "prefixItems"})
_SUBSCHEMA_MAP_KEYWORDS: Final = frozenset(
{"$defs", "definitions", "dependentSchemas", "patternProperties", "properties"}
)
_MAX_SCHEMA_NESTING: Final = 1024
def drop_non_python_regex_patterns(schema: Mapping[str, object]) -> Mapping[str, object]:
"""Drop every regex in a schema position that Python's ``re`` cannot compile.
OpenAI validates tool ``parameters`` against the 2020-12 metaschema with
``jsonschema``'s format checker, which hands each ``pattern`` value and each
``patternProperties`` key to ``re.compile``, so a regex written for an
ECMA-262 engine (Unicode property escapes such as ``\\p{Cc}``, as in Claude
Code's ``Artifact`` tool) is refused with "'...' is not a 'regex'" by every
model family on both the chat and Responses wires. Only schema positions are
walked (properties, items, combinators, ``$defs`` and the other applicators),
so a ``pattern`` key inside ``default``, ``examples``, ``const`` or vendor
extensions is data and stays. Outside strict mode the keyword is only a
hint, so dropping it costs the model a constraint and the caller nothing.
Compilable regexes and everything else pass through, the input is never
mutated, and the same object comes back when nothing was dropped. The walk
is level-order rather than recursive, rebuilt deepest level first, and stops
at more schema levels than a JSON parser admits, so a cyclic schema built in
code cannot spin it.
"""
rebuilt: dict[int, Mapping[str, object]] = {} # mutable-ok: per-call memo of rewritten nodes, deepest level first
for level in reversed(tuple(islice(_schema_levels(schema), _MAX_SCHEMA_NESTING))):
rebuilt.update(
(id(node), rewritten)
for node in level
if (rewritten := _node_without_non_python_regex(node, rebuilt)) is not node
)
return rebuilt.get(id(schema), schema)
def _schema_levels(schema: Mapping[str, object]) -> Iterator[tuple[Mapping[str, object], ...]]:
frontier: tuple[Mapping[str, object], ...] = (schema,) # rebind-ok: level-order cursor, one level a round
while frontier:
yield frontier
frontier = tuple(child for node in frontier for child in _subschemas(node))
def _subschemas(node: Mapping[str, object]) -> Iterator[Mapping[str, object]]:
for key, value in node.items():
if key in _SUBSCHEMA_MAP_KEYWORDS and isinstance(value, dict):
yield from (sub for sub in value.values() if isinstance(sub, dict))
elif key in _SUBSCHEMA_LIST_KEYWORDS and isinstance(value, list):
yield from (sub for sub in value if isinstance(sub, dict))
elif key in _SUBSCHEMA_KEYWORDS and isinstance(value, dict):
yield value
def _node_without_non_python_regex(
node: Mapping[str, object], rebuilt: Mapping[int, Mapping[str, object]]
) -> Mapping[str, object]:
kept: Final = { # mutable-ok: tool parameters are JSON dicts
key: _keyword_value_rebuilt(key, value, rebuilt)
for key, value in node.items()
if key != "pattern" or not isinstance(value, str) or _is_python_regex(value)
}
return node if len(kept) == len(node) and all(kept[key] is node[key] for key in kept) else kept
def _keyword_value_rebuilt(key: str, value: object, rebuilt: Mapping[int, Mapping[str, object]]) -> object:
if key in _SUBSCHEMA_MAP_KEYWORDS and isinstance(value, dict):
kept: Final = { # mutable-ok: tool parameters are JSON dicts
name: rebuilt.get(id(sub), sub)
for name, sub in value.items()
if key != "patternProperties" or not isinstance(name, str) or _is_python_regex(name)
}
return value if len(kept) == len(value) and all(kept[name] is value[name] for name in kept) else kept
if key in _SUBSCHEMA_LIST_KEYWORDS and isinstance(value, list):
items: Final = [rebuilt.get(id(sub), sub) for sub in value] # mutable-ok: tool parameters are JSON lists
return value if all(new is old for new, old in zip(items, value, strict=True)) else items
if key in _SUBSCHEMA_KEYWORDS and isinstance(value, dict):
return rebuilt.get(id(value), value)
return value
def _is_python_regex(pattern: str) -> bool:
try:
re.compile(pattern)
except (re.error, RecursionError):
return False
return True
def flatten_combinators_and_drop_non_python_regex_patterns(schema: Mapping[str, object]) -> Mapping[str, object]:
return flatten_top_level_schema_combinators(drop_non_python_regex_patterns(schema))
def tool_with_sanitized_parameters(
tool: Mapping[str, object],
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
) -> Mapping[str, object]:
function: Final = tool.get("function")
if not isinstance(function, dict):
return tool
parameters: Final = function.get("parameters")
if not isinstance(parameters, dict):
return tool
flattened: Final = flatten_top_level_schema_combinators(parameters)
if flattened is parameters:
sanitized: Final = sanitize(parameters)
if sanitized is parameters:
return tool
return {**tool, "function": {**function, "parameters": flattened}} # mutable-ok: request tools are JSON dicts
return {**tool, "function": {**function, "parameters": sanitized}} # mutable-ok: request tools are JSON dicts
def _get_image_mime_type_from_url(url: str) -> str | None:

View file

@ -162,6 +162,19 @@ def _redact_responses_api_output_dict(output_items, redacted_str: str):
output_item["arguments"] = redacted_str
def redacted_standard_logging_payload(payload: Mapping[str, object]) -> Mapping[str, object]:
"""
Return a copy of a ``StandardLoggingPayload`` with its messages and response redacted.
The success path redacts through ``perform_redaction`` before a callback ever sees the
payload, but the failure path does not, so a callback that batches both has to redact
the ones it is handed.
"""
redacted: Final = copy.deepcopy(dict(payload)) # mutable-ok: redacted in place below
_redact_standard_logging_object({"standard_logging_object": redacted}) # mutable-ok: the callee's shape
return redacted
def _redact_standard_logging_object(model_call_details: dict):
"""Redact messages and response inside standard_logging_object if present."""
standard_logging_object: Final = model_call_details.get("standard_logging_object")

View file

@ -1,6 +1,6 @@
import base64
import time
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Callable, Iterator, Mapping, Sequence
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypeAlias, TypedDict, Union, cast
@ -210,7 +210,7 @@ def apply_grounding_request_counts(
class ChunkProcessor:
def __init__(self, chunks: list, messages: list | None = None):
def __init__(self, chunks: list, messages: Sequence | None = None):
self.chunks = self._sort_chunks(chunks)
self.messages = messages
self.first_chunk = chunks[0]
@ -1004,8 +1004,9 @@ class ChunkProcessor:
chunks: Sequence["_UsageBearingChunk | ModelResponse"],
model: str,
completion_output: str,
messages: list | None = None,
messages: Sequence | None = None,
reasoning_tokens: int | None = None,
count_prompt_tokens: Callable[[], int] | None = None,
) -> Usage:
"""
Calculate usage for the given chunks.
@ -1030,7 +1031,9 @@ class ChunkProcessor:
cost: Final[float | None] = calculated_usage_per_chunk["cost"]
try:
returned_usage.prompt_tokens = prompt_tokens or token_counter(model=model, messages=messages)
returned_usage.prompt_tokens = prompt_tokens or (
count_prompt_tokens() if count_prompt_tokens else token_counter(model=model, messages=messages)
)
except Exception: # don't allow this failing to block a complete streaming response from being returned
print_verbose("token_counter failed, assuming prompt tokens is 0")
returned_usage.prompt_tokens = 0

View file

@ -179,6 +179,13 @@ def calculate_tiles_needed(
return total_tiles
def high_detail_image_token_upper_bound(base_tokens: int = 85) -> int:
largest_tile_count: Final = calculate_tiles_needed(
MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES, MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES
)
return base_tokens + (base_tokens * 2) * largest_tile_count
def _unpack_ints(fmt: str, buffer: bytes) -> tuple[int, ...]:
return struct.unpack(fmt, buffer)

View file

@ -73,8 +73,13 @@ _CLAUDE_CODE_OBJECT_MAPPING_ADAPTER: Final = TypeAdapter(dict[object, object])
_CLAUDE_CODE_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object])
_CLAUDE_CODE_USER_AGENT_PREFIXES: Final = ("claude-cli/", "claude-code/")
def is_claude_code_user_agent(user_agent: str) -> bool:
return user_agent.startswith("claude-cli/")
"""Claude Code sends its API calls through the Anthropic SDK as `claude-cli/<version>` and its own
fetches, such as gateway model discovery, as `claude-code/<version>`"""
return user_agent.startswith(_CLAUDE_CODE_USER_AGENT_PREFIXES)
def _validated_claude_code_mapping(value: object) -> dict[object, object] | None:
@ -1656,11 +1661,16 @@ def process_anthropic_headers(headers: httpx.Headers | dict) -> dict:
def _anthropic_model_entry(
model: ModelInfoResponse, created_at: str, display_names: Mapping[str, str]
model: ModelInfoResponse, created_at: str, display_names: Mapping[str, str], listed_ids: Mapping[str, str]
) -> Mapping[str, object]:
listed_id: Final = listed_ids.get(model["id"])
source: Final[Mapping[str, object]] = (
MappingProxyType({"source_model": model["id"]}) if listed_id is not None else MappingProxyType({})
)
return { # mutable-ok: JSON response body, serialized by the route and never mutated
"type": "model",
"id": model["id"],
"id": listed_id or model["id"],
**source,
"display_name": display_names.get(model["id"], model["id"]),
"created_at": created_at,
"max_input_tokens": model.get("max_input_tokens"),
@ -1671,6 +1681,7 @@ def _anthropic_model_entry(
def create_anthropic_model_list_response(
models: Sequence[ModelInfoResponse],
display_names: Mapping[str, str] = MappingProxyType({}),
listed_ids: Mapping[str, str] = MappingProxyType({}),
) -> Mapping[str, object]:
"""Build the Anthropic-native /v1/models envelope.
@ -1680,17 +1691,19 @@ def create_anthropic_model_list_response(
over from the OpenAI-shaped listing, named as the Messages API names them, and
are always present because the vendor shape declares them nullable, not optional.
display_names maps a listed model id to a configured human-readable name; ids
without an entry fall back to the id itself, matching the vendor behavior
without an entry fall back to the id itself, matching the vendor behavior.
listed_ids maps a model id to the id the caller should see it under (the Claude
Code view); ids without an entry are listed as they are
"""
created_at: Final = (
datetime.fromtimestamp(DEFAULT_MODEL_CREATED_AT_TIME, tz=timezone.utc).isoformat().replace("+00:00", "Z")
)
data: Final = [ # mutable-ok: JSON response body, serialized by the route and never mutated
_anthropic_model_entry(model, created_at, display_names) for model in models
_anthropic_model_entry(model, created_at, display_names, listed_ids) for model in models
]
return { # mutable-ok: JSON response body, serialized by the route and never mutated
"data": data,
"has_more": False,
"first_id": models[0]["id"] if models else None,
"last_id": models[-1]["id"] if models else None,
"first_id": data[0]["id"] if data else None,
"last_id": data[-1]["id"] if data else None,
}

View file

@ -7,8 +7,9 @@ from httpx._models import Headers, Response
import litellm
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_tool_reference_parts_from_tool_messages,
flatten_combinators_and_drop_non_python_regex_patterns,
hoist_images_from_tool_messages,
tool_with_flattened_parameters,
tool_with_sanitized_parameters,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_azure_openai_messages,
@ -39,14 +40,17 @@ else:
_NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
def flattened_tools_update(optional_params: Mapping[str, object]) -> Mapping[str, object]:
def sanitized_tools_update(optional_params: Mapping[str, object]) -> Mapping[str, object]:
tools: Final = optional_params.get("tools")
if not isinstance(tools, list):
return _NO_TOOLS_UPDATE
flattened: Final = [ # mutable-ok: request tools are a JSON list
tool_with_flattened_parameters(tool) if isinstance(tool, dict) else tool for tool in tools
sanitized: Final = [ # mutable-ok: request tools are a JSON list
tool_with_sanitized_parameters(tool, flatten_combinators_and_drop_non_python_regex_patterns)
if isinstance(tool, dict)
else tool
for tool in tools
]
return MappingProxyType({"tools": flattened})
return MappingProxyType({"tools": sanitized})
class AzureOpenAIConfig(BaseConfig):
@ -278,7 +282,7 @@ class AzureOpenAIConfig(BaseConfig):
"model": model,
"messages": azure_messages,
**optional_params,
**flattened_tools_update(optional_params),
**sanitized_tools_update(optional_params),
}
def transform_response(

View file

@ -20,7 +20,7 @@ from litellm.types.llms.openai import AllMessageValues
from litellm.utils import get_model_info, supports_reasoning
from ...openai.chat.o_series_transformation import OpenAIOSeriesConfig
from .gpt_transformation import flattened_tools_update
from .gpt_transformation import sanitized_tools_update
class AzureOpenAIO1Config(OpenAIOSeriesConfig):
@ -111,6 +111,6 @@ class AzureOpenAIO1Config(OpenAIOSeriesConfig):
model = model.replace("o_series/", "") # handle o_series/my-random-deployment-name
flattened_params: Final = { # mutable-ok: transform_request's contract takes a plain JSON params dict
**optional_params,
**flattened_tools_update(optional_params),
**sanitized_tools_update(optional_params),
}
return super().transform_request(model, messages, flattened_params, litellm_params, headers)

View file

@ -1,24 +1,102 @@
import re
from collections.abc import Callable, Collection, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Optional
import httpx
from httpx import Response
from pydantic import BaseModel, ValidationError
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.llms.base_llm.passthrough.transformation import (
BasePassthroughConfig,
RelayShape,
logged_relay_shape,
replace_path_segment,
strip_leading_model_segment,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse, ResponsesTerminalEvent
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import CallTypes, EmbeddingResponse, ImageResponse
if TYPE_CHECKING:
from httpx import URL
from litellm.types.utils import CostResponseTypes
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
class RelayedChatRequest(BaseModel):
messages: Sequence[Mapping[str, object]] | None = None
class RelayedCallDetails(BaseModel):
request_data: RelayedChatRequest | None = None
def _relayed_messages(litellm_logging_obj: Logging) -> Sequence[Mapping[str, object]] | None:
try:
details: Final = RelayedCallDetails.model_validate(litellm_logging_obj.model_call_details)
except ValidationError:
return None
return details.request_data.messages if details.request_data else None
RESPONSES_RELAY_SHAPE: Final = RelayShape("/responses", CallTypes.aresponses, ResponsesAPIResponse.model_validate)
OPENAI_RELAY_SHAPES: Final = (
RelayShape("/embeddings", CallTypes.aembedding, EmbeddingResponse.model_validate),
RESPONSES_RELAY_SHAPE,
RelayShape("/images/generations", CallTypes.aimage_generation, ImageResponse.model_validate),
)
def logged_responses_stream(all_chunks: Sequence[str], logging_obj: Logging) -> ResponsesTerminalEvent | None:
"""A streaming logging object assembles the logged response from the terminal event, not from its body."""
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
terminal_event: Final = OpenAIResponsesAPIConfig.parse_terminal_event_from_stream_chunks(all_chunks=all_chunks)
if terminal_event is None:
return None
logging_obj.call_type = (
RESPONSES_RELAY_SHAPE.call_type.value
) # rebind-ok: routes cost calculation to the relayed shape's pricing path
return terminal_event
AZURE_DEPLOYMENT_SEGMENT: Final = re.compile(r"(?<![^/])openai/deployments/([^/]+)")
def azure_router_model_in_endpoint(endpoint: str, router_models: Collection[str]) -> str | None:
parts: Final = endpoint.split("/")
if len(parts) < 2:
return None
return next((part for part in parts if part in router_models), None)
def foreign_azure_deployment(
endpoint: str, model_group: str, served_models: Callable[[], Collection[str]]
) -> str | None:
match: Final = AZURE_DEPLOYMENT_SEGMENT.search(endpoint)
if match is None:
return None
deployment: Final = match.group(1)
if deployment == model_group:
return None
served: Final = frozenset(name.casefold() for name in served_models())
return None if deployment.casefold() in served else deployment
def without_api_version(api_base: str) -> str:
url: Final = httpx.URL(api_base)
kept_params: Final = tuple((key, value) for key, value in url.params.multi_items() if key != "api-version")
return str(url.copy_with(params=httpx.QueryParams(kept_params)))
class AzurePassthroughConfig(BasePassthroughConfig):
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
return "stream" in request_data
return bool(request_data.get("stream"))
def get_complete_url(
self,
@ -36,14 +114,17 @@ class AzurePassthroughConfig(BasePassthroughConfig):
litellm_metadata: Final = litellm_params.get("litellm_metadata") or {}
model_group: Final = litellm_metadata.get("model_group")
if model_group and model_group in endpoint:
endpoint = endpoint.replace(model_group, model)
routed_endpoint: Final = replace_path_segment(endpoint, model_group, model) if model_group else endpoint
native_endpoint: Final = strip_leading_model_segment(routed_endpoint, (model,))
caller_api_version: Final = request_query_params.get("api-version") if request_query_params else None
relay_base: Final = without_api_version(base_target_url) if caller_api_version else base_target_url
complete_url: Final = BaseAzureLLM._get_base_azure_url(
api_base=base_target_url,
litellm_params=litellm_params,
route=endpoint,
default_api_version=litellm_params.get("api_version"),
api_base=relay_base,
litellm_params=MappingProxyType(
{**litellm_params, "api_version": caller_api_version or litellm_params.get("api_version")}
),
route=native_endpoint,
)
return (
httpx.URL(complete_url),
@ -92,13 +173,13 @@ class AzurePassthroughConfig(BasePassthroughConfig):
request_data: dict,
logging_obj: Logging,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> Optional["LoggedRelayResponse"]:
from litellm import encoding
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.types.utils import ModelResponse
if "chat/completions" not in endpoint:
return None
return logged_relay_shape(OPENAI_RELAY_SHAPES, httpx_response, logging_obj, endpoint)
openai_chat_config: Final = OpenAIGPTConfig()
@ -116,3 +197,27 @@ class AzurePassthroughConfig(BasePassthroughConfig):
)
return litellm_model_response
def handle_logging_collected_chunks(
self,
all_chunks: Sequence[str],
litellm_logging_obj: Logging,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["LoggedRelayResponse"]:
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
OpenAIPassthroughLoggingHandler,
)
if f"/{endpoint.strip('/')}".endswith(RESPONSES_RELAY_SHAPE.path_suffix):
return logged_responses_stream(all_chunks, litellm_logging_obj)
if "chat/completions" not in endpoint:
return None
return OpenAIPassthroughLoggingHandler()._build_complete_streaming_response( # pyright: ignore[reportPrivateUsage] # the only OpenAI SSE-to-ModelResponse assembler; reimplementing it would fork the parser
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
messages=_relayed_messages(litellm_logging_obj),
)

View file

@ -2,7 +2,6 @@ import copy
import enum
import re
from typing import TYPE_CHECKING, Final, cast
from urllib.parse import urlparse
import httpx
from httpx import Response
@ -15,7 +14,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
)
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
from litellm.llms.azure_ai.common_utils import (
api_key_header_for_base,
is_foundry_model_inference_base,
)
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
@ -146,11 +148,7 @@ class AzureAIStudioConfig(OpenAIConfig):
"""
Returns True if the request should use `api-key` header for authentication.
"""
parsed_url: Final = urlparse(api_base)
host: Final = parsed_url.hostname
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
return True
return False
return api_key_header_for_base(api_base) == "api-key"
def get_complete_url(
self,

View file

@ -19,6 +19,13 @@ def is_foundry_model_inference_base(api_base: str) -> bool:
return "/openai/deployments" not in parsed.path
def api_key_header_for_base(api_base: str | None) -> AzureAIApiKeyHeader:
host: Final = urlparse(api_base).hostname if api_base else None
if host and (host.endswith(".services.ai.azure.com") or host.endswith(".openai.azure.com")):
return "api-key"
return "Authorization"
def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None) -> str | None:
"""
Resolve an Entra ID / OAuth access token for an Azure AI Foundry deployment.

View file

@ -0,0 +1,232 @@
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
import httpx
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.llms.azure_ai.common_utils import (
AzureFoundryModelInfo,
api_key_header_for_base,
get_azure_ai_auth_headers,
)
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
from litellm.llms.base_llm.passthrough.transformation import (
BasePassthroughConfig,
RelayShape,
logged_relay_shape,
strip_leading_model_segment,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.rerank import RerankResponse
from litellm.types.utils import CallTypes, ImageResponse, StandardPassThroughResponseObject
if TYPE_CHECKING:
from httpx import URL, Response
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
EMPTY_QUERY: Final[Mapping[str, object]] = MappingProxyType({})
class PassthroughMetadata(BaseModel):
model_config = ConfigDict(extra="ignore")
model_group: str = ""
def model_group_from(litellm_params: Mapping[str, object]) -> str:
try:
return PassthroughMetadata.model_validate(litellm_params.get("litellm_metadata")).model_group
except ValidationError:
return ""
def api_version_from(litellm_params: Mapping[str, object]) -> str | None:
try:
return TypeAdapter(str | None).validate_python(litellm_params.get("api_version"))
except ValidationError:
return None
def foundry_root(api_base: str) -> str:
url: Final = httpx.URL(api_base)
segments: Final = tuple(segment for segment in url.path.split("/") if segment)
root_segments: Final = segments[: segments.index("models")] if "models" in segments else segments
return str(url.copy_with(path="/" + "/".join(root_segments), query=None)).rstrip("/")
def is_repeated_native_prefix(native_segments: tuple[str, ...], overlap: int) -> bool:
return overlap == len(native_segments) or native_segments[0] == "openai"
def without_repeated_native_prefix(root: str, native_endpoint: str) -> str:
url: Final = httpx.URL(root)
root_segments: Final = tuple(segment for segment in url.path.split("/") if segment)
native_segments: Final = tuple(segment.casefold() for segment in native_endpoint.split("/") if segment)
overlap: Final = next(
(
length
for length in range(min(len(root_segments), len(native_segments)), 0, -1)
if tuple(segment.casefold() for segment in root_segments[-length:]) == native_segments[:length]
and is_repeated_native_prefix(native_segments, length)
),
0,
)
kept_segments: Final = root_segments[: len(root_segments) - overlap]
return str(url.copy_with(path="/" + "/".join(kept_segments), query=None)).rstrip("/")
def relay_query_params(
request_query_params: Mapping[str, object] | None,
deployment_api_version: str | None,
api_base: str,
) -> Mapping[str, object] | None:
if request_query_params and "api-version" in request_query_params:
return request_query_params
api_version: Final = deployment_api_version or httpx.URL(api_base).params.get("api-version")
if api_version is None:
return request_query_params
return MappingProxyType({**(request_query_params or EMPTY_QUERY), "api-version": api_version})
def relayed_body(httpx_response: Response) -> str | dict:
try:
body: Final[object] = httpx_response.json()
except ValueError:
return httpx_response.text
return body if isinstance(body, dict) else httpx_response.text
FOUNDRY_RELAY_SHAPES: Final = (
RelayShape("/rerank", CallTypes.arerank, RerankResponse.model_validate),
RelayShape("/providers/blackforestlabs/v1/flux-2-pro", CallTypes.aimage_generation, ImageResponse.model_validate),
)
class AzureAIPassthroughConfig(AzureFoundryModelInfo, BasePassthroughConfig):
def __init__(self, ocr_config_for: Callable[[str], BaseOCRConfig | None] = get_azure_ai_ocr_config) -> None:
super().__init__()
self.ocr_config_for: Final = ocr_config_for
def is_streaming_request(self, endpoint: str, request_data: Mapping[str, object]) -> bool:
return bool(request_data.get("stream"))
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
endpoint: str,
request_query_params: Mapping[str, object] | None,
litellm_params: Mapping[str, object],
) -> tuple[URL, str]:
base_target_url: Final = self.get_api_base(api_base)
if base_target_url is None:
raise ValueError("Azure AI api base not found: set `api_base` on the deployment or AZURE_AI_API_BASE")
native_endpoint: Final = strip_leading_model_segment(endpoint, (model, model_group_from(litellm_params)))
root: Final = without_repeated_native_prefix(foundry_root(base_target_url), native_endpoint)
query_params: Final = relay_query_params(
request_query_params, api_version_from(litellm_params), base_target_url
)
return (self.format_url(native_endpoint, root, query_params), root)
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict[str, str]: # mutable-ok: base class contract returns dict for httpx
auth_headers: Final = get_azure_ai_auth_headers(
api_key=api_key,
litellm_params=litellm_params,
api_key_header=api_key_header_for_base(api_base),
)
return {**headers, **auth_headers} # mutable-ok: base class contract returns dict for httpx
def logging_non_streaming_response(
self,
model: str,
custom_llm_provider: str,
httpx_response: Response,
request_data: Mapping[str, object],
logging_obj: Logging,
endpoint: str,
) -> LoggedRelayResponse | OCRResponse | StandardPassThroughResponseObject | None:
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
chat_result: Final = AzurePassthroughConfig().logging_non_streaming_response( # pyright: ignore[reportUnknownMemberType] # the Azure config still types request_data as a bare dict
model=model,
custom_llm_provider=custom_llm_provider,
httpx_response=httpx_response,
request_data=dict(request_data), # mutable-ok: AzurePassthroughConfig wants a dict
logging_obj=logging_obj,
endpoint=endpoint,
)
if chat_result is not None:
return chat_result
ocr_result: Final = self.logged_ocr_response(model, httpx_response, logging_obj, endpoint)
if ocr_result is not None:
return ocr_result
foundry_result: Final = logged_relay_shape(FOUNDRY_RELAY_SHAPES, httpx_response, logging_obj, endpoint)
if foundry_result is not None:
return foundry_result
return StandardPassThroughResponseObject(response=relayed_body(httpx_response))
def logged_ocr_response(
self, model: str, httpx_response: Response, logging_obj: Logging, endpoint: str
) -> OCRResponse | None:
ocr_config: Final = self.ocr_config_for(model)
if ocr_config is None or httpx_response.status_code != 200:
return None
relayed_url: Final = httpx_response.request.url
relayed_origin: Final = str(relayed_url.copy_with(path="/", query=None, fragment=None)).rstrip("/")
ocr_url: Final = httpx.URL(
ocr_config.get_complete_url(
api_base=relayed_origin,
model=model,
optional_params={}, # mutable-ok: BaseOCRConfig wants a dict
)
)
known_prefixes: Final = (model, model_group_from(logging_obj.litellm_params))
native_endpoint: Final = strip_leading_model_segment(endpoint, known_prefixes)
if f"/{native_endpoint.strip('/')}" != ocr_url.path:
return None
try:
ocr_response: Final = ocr_config.transform_ocr_response(
model=model, raw_response=httpx_response, logging_obj=logging_obj
)
except (ValueError, AttributeError) as error:
verbose_logger.warning("azure_ai passthrough: OCR body from %s is not costable: %s", ocr_url, error)
return None
logging_obj.call_type = CallTypes.aocr.value # rebind-ok: routes cost calculation to the per-page OCR path
return ocr_response
def handle_logging_collected_chunks(
self,
all_chunks: Sequence[str],
litellm_logging_obj: Logging,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> LoggedRelayResponse | None:
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
return AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=all_chunks,
litellm_logging_obj=litellm_logging_obj,
model=model,
custom_llm_provider=custom_llm_provider,
endpoint=endpoint,
)

View file

@ -1,5 +1,14 @@
from __future__ import annotations
import re
from abc import abstractmethod
from typing import TYPE_CHECKING, Final, Optional, Union
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, TypeAlias
from pydantic import TypeAdapter, ValidationError
from litellm.types.utils import CallTypes
from ..base_utils import BaseLLMModelInfo
@ -7,9 +16,68 @@ if TYPE_CHECKING:
from httpx import URL, Headers, Response
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import CostResponseTypes
from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesTerminalEvent
from litellm.types.rerank import RerankResponse
from litellm.types.utils import CostResponseTypes, StandardPassThroughResponseObject
from ..chat.transformation import BaseLLMException
from ..ocr.transformation import OCRResponse
LoggedRelayResponse: TypeAlias = CostResponseTypes | RerankResponse | ResponsesAPIResponse | ResponsesTerminalEvent
RELAYED_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
def strip_leading_model_segment(endpoint: str, model_names: tuple[str, ...]) -> str:
path: Final = endpoint.lstrip("/")
for model_name in model_names:
if not model_name:
continue
if path == model_name:
return ""
if path.startswith(f"{model_name}/"):
return path[len(model_name) + 1 :]
return path
def replace_path_segment(endpoint: str, segment: str, replacement: str) -> str:
bounded_segment: Final = re.compile(rf"(?<![^/]){re.escape(segment)}(?![^/:])")
return bounded_segment.sub(lambda _: replacement, endpoint)
def relayed_json_object(httpx_response: Response) -> Mapping[str, object] | None:
if httpx_response.status_code != 200:
return None
try:
return RELAYED_JSON_OBJECT.validate_python(httpx_response.json())
except (ValueError, ValidationError):
return None
@dataclass(frozen=True, slots=True)
class RelayShape:
path_suffix: str
call_type: CallTypes
parse: Callable[[Mapping[str, object]], LoggedRelayResponse]
def logged_relay_shape(
shapes: Sequence[RelayShape], httpx_response: Response, logging_obj: LiteLLMLoggingObj, endpoint: str
) -> LoggedRelayResponse | None:
relayed_path: Final = f"/{endpoint.strip('/')}"
shape: Final = next((candidate for candidate in shapes if relayed_path.endswith(candidate.path_suffix)), None)
body: Final = relayed_json_object(httpx_response) if shape else None
if shape is None or body is None:
return None
try:
parsed: Final = shape.parse(body)
except ValidationError:
return None
logging_obj.call_type = (
shape.call_type.value
) # rebind-ok: routes cost calculation to the relayed shape's pricing path
return parsed
class BasePassthroughConfig(BaseLLMModelInfo):
@ -23,8 +91,8 @@ class BasePassthroughConfig(BaseLLMModelInfo):
self,
endpoint: str,
base_target_url: str,
request_query_params: dict | None,
) -> "URL":
request_query_params: Mapping[str, object] | None,
) -> URL:
"""
Helper function to add query params to the url
Args:
@ -58,7 +126,7 @@ class BasePassthroughConfig(BaseLLMModelInfo):
endpoint: str,
request_query_params: dict | None,
litellm_params: dict,
) -> tuple["URL", str]:
) -> tuple[URL, str]:
"""
Get the complete url for the request
Returns:
@ -88,9 +156,7 @@ class BasePassthroughConfig(BaseLLMModelInfo):
"""
return headers, None
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, "Headers"]
) -> "BaseLLMException":
def get_error_class(self, error_message: str, status_code: int, headers: dict | Headers) -> BaseLLMException:
from litellm.llms.base_llm.chat.transformation import BaseLLMException
return BaseLLMException(status_code=status_code, message=error_message, headers=headers)
@ -99,21 +165,21 @@ class BasePassthroughConfig(BaseLLMModelInfo):
self,
model: str,
custom_llm_provider: str,
httpx_response: "Response",
httpx_response: Response,
request_data: dict,
logging_obj: "LiteLLMLoggingObj",
logging_obj: LiteLLMLoggingObj,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> LoggedRelayResponse | OCRResponse | StandardPassThroughResponseObject | None:
pass
def handle_logging_collected_chunks(
self,
all_chunks: list[str],
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
custom_llm_provider: str,
endpoint: str,
) -> Optional["CostResponseTypes"]:
) -> LoggedRelayResponse | None:
return None
def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]:

View file

@ -6,7 +6,7 @@ import json
import os
import re
import urllib.parse
from collections.abc import Callable, Mapping
from collections.abc import Callable, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
@ -14,7 +14,8 @@ from threading import Lock
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
from pydantic import BaseModel, ValidationError
from pydantic import BaseModel, TypeAdapter, ValidationError
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
@ -31,6 +32,7 @@ from litellm.constants import (
from litellm.litellm_core_utils.aws_partition import contains_bedrock_arn, get_aws_dns_suffix
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.secret_managers.main import get_secret, get_secret_str
from litellm.types.llms.bedrock import AwsSessionTag
if TYPE_CHECKING:
from botocore.awsrequest import AWSPreparedRequest
@ -52,6 +54,47 @@ _STS_REGION_FROM_ENDPOINT_PATTERN: Final = re.compile(
SIGV4_COMPUTED_HEADERS: Final = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"})
_AWS_SESSION_TAGS_ADAPTER: Final[TypeAdapter[tuple[AwsSessionTag, ...]]] = TypeAdapter(tuple[AwsSessionTag, ...])
def _canonical_aws_session_tags(raw_tags: object) -> tuple[AwsSessionTag, ...] | None:
if raw_tags is None:
return None
try:
validated: Final = _AWS_SESSION_TAGS_ADAPTER.validate_python(raw_tags)
except ValidationError as e:
raise ValueError(
"Invalid 'aws_session_tags' value. Expected a list of {'Key': <str>, 'Value': <str>} dicts, "
f"e.g. [{{'Key': 'team', 'Value': 'genai'}}]. Got: {raw_tags!r}"
) from e
return tuple(sorted(validated, key=lambda tag: tag["Key"]))
class _AssumeRoleParams(TypedDict):
RoleArn: ReadOnly[str]
RoleSessionName: ReadOnly[str]
ExternalId: ReadOnly[NotRequired[str]]
Tags: ReadOnly[NotRequired[tuple[AwsSessionTag, ...]]]
def _assume_role_params(
aws_role_name: str,
aws_session_name: str,
aws_external_id: str | None,
aws_session_tags: Sequence[AwsSessionTag] | None,
) -> _AssumeRoleParams:
match (aws_external_id, tuple(aws_session_tags or ())):
case (None, ()):
return _AssumeRoleParams(RoleArn=aws_role_name, RoleSessionName=aws_session_name)
case (None, tags):
return _AssumeRoleParams(RoleArn=aws_role_name, RoleSessionName=aws_session_name, Tags=tags)
case (external_id, ()):
return _AssumeRoleParams(RoleArn=aws_role_name, RoleSessionName=aws_session_name, ExternalId=external_id)
case (external_id, tags):
return _AssumeRoleParams(
RoleArn=aws_role_name, RoleSessionName=aws_session_name, ExternalId=external_id, Tags=tags
)
class BedrockRequestTarget(BaseModel):
aws_region_name: str
@ -129,6 +172,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
"aws_sts_endpoint",
"aws_bedrock_runtime_endpoint",
"aws_external_id",
"aws_session_tags",
]
def _get_ssl_verify(self, ssl_verify: bool | str | None = None):
@ -146,7 +190,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
return get_ssl_verify(ssl_verify=ssl_verify)
def get_cache_key(self, credential_args: Mapping[str, str | bool | None]) -> str:
def get_cache_key(self, credential_args: Mapping[str, str | bool | tuple[AwsSessionTag, ...] | None]) -> str:
"""
Generate a unique cache key based on the credential arguments.
"""
@ -156,7 +200,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
def _get_or_set_cached_credentials(
self,
credential_args: Mapping[str, str | bool | None],
credential_args: Mapping[str, str | bool | tuple[AwsSessionTag, ...] | None],
credential_fetcher: Callable[[], tuple[Credentials, int | None]],
) -> Any:
"""
@ -231,6 +275,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_web_identity_token: str | None = None,
aws_sts_endpoint: str | None = None,
aws_external_id: str | None = None,
aws_session_tags: Sequence[AwsSessionTag] | None = None,
ssl_verify: bool | str | None = None,
):
"""
@ -267,6 +312,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
(aws_external_id, "AWS_EXTERNAL_ID"),
)
)
session_tags: Final = _canonical_aws_session_tags(aws_session_tags)
verbose_logger.debug(
"in get credentials\n"
@ -279,7 +325,8 @@ class BaseAWSLLM(SignsRequestsWithAWS):
"aws_role_name=%s\n"
"aws_web_identity_token=[set=%s]\n"
"aws_sts_endpoint=%s\n"
"aws_external_id=%s",
"aws_external_id=%s\n"
"aws_session_tags=%s",
aws_access_key_id is not None,
aws_secret_access_key is not None,
aws_session_token is not None,
@ -290,6 +337,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_web_identity_token is not None,
aws_sts_endpoint,
aws_external_id,
session_tags,
)
args: Final = {
@ -303,6 +351,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
"aws_web_identity_token": aws_web_identity_token,
"aws_sts_endpoint": aws_sts_endpoint,
"aws_external_id": aws_external_id,
"aws_session_tags": session_tags,
"ssl_verify": ssl_verify,
}
@ -345,6 +394,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_region_name=aws_region_name,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=session_tags,
ssl_verify=ssl_verify,
),
)
@ -989,6 +1039,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_sts_endpoint: str | None = None,
ssl_verify: bool | str | None = None,
aws_region_name: str | None = None,
aws_session_tags: Sequence[AwsSessionTag] | None = None,
) -> dict:
"""Handle cross-account role assumption for IRSA."""
import boto3
@ -1041,16 +1092,9 @@ class BaseAWSLLM(SignsRequestsWithAWS):
# Now assume the target role
verbose_logger.debug("Attempting to assume target role: %s with session: %s", aws_role_name, aws_session_name)
assume_role_params: Final = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
}
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
return sts_client_with_creds.assume_role(**assume_role_params)
return sts_client_with_creds.assume_role(
**_assume_role_params(aws_role_name, aws_session_name, aws_external_id, aws_session_tags)
)
def _handle_irsa_same_account(
self,
@ -1060,6 +1104,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_sts_endpoint: str | None = None,
ssl_verify: bool | str | None = None,
aws_region_name: str | None = None,
aws_session_tags: Sequence[AwsSessionTag] | None = None,
) -> dict:
"""Handle same-account role assumption for IRSA."""
import boto3
@ -1083,16 +1128,9 @@ class BaseAWSLLM(SignsRequestsWithAWS):
# Assume the role
verbose_logger.debug("Attempting to assume role: %s with session: %s", aws_role_name, aws_session_name)
assume_role_params: Final = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
}
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
return sts_client.assume_role(**assume_role_params)
return sts_client.assume_role(
**_assume_role_params(aws_role_name, aws_session_name, aws_external_id, aws_session_tags)
)
def _extract_credentials_and_ttl(self, sts_response: dict) -> tuple[Credentials, int | None]:
"""Extract credentials and TTL from STS response.
@ -1127,6 +1165,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_region_name: str | None,
aws_sts_endpoint: str | None,
aws_external_id: str | None,
aws_session_tags: tuple[AwsSessionTag, ...] | None,
ssl_verify: bool | str | None,
) -> tuple[Credentials, int | None]:
"""
@ -1153,6 +1192,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_region_name=aws_region_name,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
ssl_verify=ssl_verify,
)
@ -1168,6 +1208,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_sts_endpoint: str | None = None,
aws_external_id: str | None = None,
ssl_verify: bool | str | None = None,
aws_session_tags: Sequence[AwsSessionTag] | None = None,
) -> tuple[Credentials, int | None]:
"""
Authenticate with AWS Role
@ -1198,6 +1239,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
aws_region_name=aws_region_name,
aws_session_tags=aws_session_tags,
)
else:
sts_response = self._handle_irsa_same_account(
@ -1207,6 +1249,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
aws_region_name=aws_region_name,
aws_session_tags=aws_session_tags,
)
return self._extract_credentials_and_ttl(sts_response)
@ -1243,14 +1286,9 @@ class BaseAWSLLM(SignsRequestsWithAWS):
**sts_client_kwargs,
)
assume_role_params: Final = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
}
# Add ExternalId parameter if provided
if aws_external_id is not None:
assume_role_params["ExternalId"] = aws_external_id
assume_role_params: Final = _assume_role_params(
aws_role_name, aws_session_name, aws_external_id, aws_session_tags
)
try:
sts_response = sts_client.assume_role(**assume_role_params)
@ -1469,6 +1507,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
"aws_bedrock_runtime_endpoint", None
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_external_id: Final = optional_params.pop("aws_external_id", None)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
if bearer_token is not None:
return BearerRequestTarget(
@ -1487,6 +1526,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
return Boto3CredentialsInfo(
credentials=credentials,
@ -1632,6 +1672,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_web_identity_token: Final = optional_params.get("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.get("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.get("aws_external_id", None)
aws_session_tags: Final = optional_params.get("aws_session_tags", None)
aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model=model)
credentials: Final[Credentials] = self.get_credentials(
@ -1645,6 +1686,7 @@ class BaseAWSLLM(SignsRequestsWithAWS):
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
sigv4: Final = SigV4Auth(credentials, service_name, aws_region_name)

View file

@ -1,4 +1,4 @@
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, cast
@ -6,6 +6,7 @@ from openai.types.batch import BatchRequestCounts
from openai.types.batch import Metadata as OpenAIBatchMetadata
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.types.llms.bedrock import AwsSessionTag
from litellm.types.utils import LiteLLMBatch
if TYPE_CHECKING:
@ -116,6 +117,7 @@ class BedrockBatchesHandler:
aws_web_identity_token: str | None = None,
aws_sts_endpoint: str | None = None,
aws_external_id: str | None = None,
aws_session_tags: Sequence[AwsSessionTag] | None = None,
**kwargs: object, # kwargs-ok: litellm.cancel_batch forwards arbitrary user kwargs verbatim
) -> "LiteLLMBatch":
try:
@ -139,6 +141,7 @@ class BedrockBatchesHandler:
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
client: Final = boto3.client(
@ -163,6 +166,7 @@ class BedrockBatchesHandler:
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
try:
@ -283,7 +287,7 @@ class BedrockBatchesHandler:
``aws_session_token``, ``aws_profile_name``,
``aws_role_name``, ``aws_session_name``,
``aws_web_identity_token``, ``aws_sts_endpoint``,
``aws_external_id``). Unknown keys are ignored.
``aws_external_id``, ``aws_session_tags``). Unknown keys are ignored.
Returns:
``LiteLLMBatch`` shaped like an OpenAI Batch resource.
@ -317,6 +321,7 @@ class BedrockBatchesHandler:
aws_web_identity_token=kwargs.get("aws_web_identity_token"),
aws_sts_endpoint=kwargs.get("aws_sts_endpoint"),
aws_external_id=kwargs.get("aws_external_id"),
aws_session_tags=kwargs.get("aws_session_tags"),
)
client: Final = boto3.client(

View file

@ -357,6 +357,7 @@ class BedrockConverseLLM(BaseAWSLLM):
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
optional_params.pop("aws_region_name", None)
litellm_params["aws_region_name"] = aws_region_name # [DO NOT DELETE] important for async calls
@ -375,6 +376,7 @@ class BedrockConverseLLM(BaseAWSLLM):
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
)

View file

@ -93,6 +93,7 @@ _BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
"aws_session_tags",
)
@ -1663,6 +1664,7 @@ class CommonBatchFilesUtils:
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
aws_external_id=optional_params.get("aws_external_id"),
aws_session_tags=optional_params.get("aws_session_tags"),
)
# Prepare the request data

View file

@ -87,6 +87,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -117,6 +118,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
)
return credentials, aws_region_name

View file

@ -751,7 +751,9 @@ class AsyncHTTPHandler:
timeout: float | httpx.Timeout | None = None,
stream: bool = False,
content: _RequestContent | None = None,
follow_redirects: bool | None = None,
):
_follow_redirects: Final = follow_redirects if follow_redirects is not None else USE_CLIENT_DEFAULT
try:
if timeout is None:
timeout = self.timeout
@ -769,22 +771,30 @@ class AsyncHTTPHandler:
timeout=timeout,
content=request_content,
)
response: Final = await self.client.send(req)
response: Final = await self.client.send(req, follow_redirects=_follow_redirects)
response.raise_for_status()
return response
except (httpx.RemoteProtocolError, httpx.ConnectError):
# Retry the request with a new session if there is a connection error
new_client: Final = self.create_client(timeout=timeout, event_hooks=self.event_hooks)
try:
return await self.single_connection_post_request(
url=url,
client=new_client,
data=data,
retry_data, retry_content = _prepare_request_data_and_content(data, content)
retry: Final = new_client.build_request(
"PUT",
url,
data=retry_data,
json=json,
params=params,
headers=headers,
stream=stream,
timeout=timeout,
content=retry_content,
)
retried: Final = await new_client.send(retry, stream=stream, follow_redirects=_follow_redirects)
try:
retried.raise_for_status()
except httpx.HTTPStatusError as retried_error:
await _raise_masked_async_error(retried_error, stream)
return retried
finally:
await new_client.aclose()
except httpx.TimeoutException as e:

View file

@ -2,7 +2,8 @@
Common utilities for the DashScope LLM provider.
"""
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Final
from urllib.parse import urlparse
import httpx
@ -16,6 +17,27 @@ if TYPE_CHECKING:
)
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
DASHSCOPE_CHAT_COMPATIBLE_PATH: Final = "/compatible-mode/v1"
DASHSCOPE_RERANK_PATH: Final = "/compatible-api/v1/reranks"
def _rerank_base_for_chat_shaped_base(api_base: str | None) -> str | None:
if api_base is None:
return None
parsed: Final = urlparse(api_base)
host: Final = parsed.hostname or ""
on_aliyun_host: Final = host == "aliyuncs.com" or host.endswith(".aliyuncs.com")
if not on_aliyun_host or parsed.path.rstrip("/") != DASHSCOPE_CHAT_COMPATIBLE_PATH:
return None
return f"{parsed.scheme}://{parsed.netloc}{DASHSCOPE_RERANK_PATH}"
def resolve_dashscope_family_rerank_api_base(api_base: str | None, env_var: str, default_rerank_base: str) -> str:
remapped: Final = _rerank_base_for_chat_shaped_base(api_base)
if api_base is not None and remapped is None:
return api_base
return get_secret_str(env_var) or remapped or default_rerank_base
def get_dashscope_family_embedding_config(custom_llm_provider: str) -> "BaseEmbeddingConfig":
if custom_llm_provider == "qwencloud":

View file

@ -3,6 +3,7 @@ from typing import Final
from litellm.secret_managers.main import get_secret_str
from .chat.transformation import DashScopeChatConfig
from .common_utils import resolve_dashscope_family_rerank_api_base
from .embed.transformation import DashScopeEmbeddingConfig
from .image_generation.transformation import DashScopeImageGenerationConfig
from .rerank.transformation import DashScopeRerankConfig
@ -51,7 +52,9 @@ class QwenAIPlatformRerankConfig(DashScopeRerankConfig):
return _require_qwen_ai_platform_api_key(api_key)
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
return api_base or get_secret_str("QWEN_AI_PLATFORM_API_BASE_RERANK") or QWEN_AI_PLATFORM_RERANK_API_BASE
return resolve_dashscope_family_rerank_api_base(
api_base, "QWEN_AI_PLATFORM_API_BASE_RERANK", QWEN_AI_PLATFORM_RERANK_API_BASE
)
class QwenAIPlatformImageGenerationConfig(DashScopeImageGenerationConfig):

View file

@ -3,6 +3,7 @@ from typing import Final
from litellm.secret_managers.main import get_secret_str
from .chat.transformation import DashScopeChatConfig
from .common_utils import resolve_dashscope_family_rerank_api_base
from .embed.transformation import DashScopeEmbeddingConfig
from .image_generation.transformation import DashScopeImageGenerationConfig
from .rerank.transformation import DashScopeRerankConfig
@ -51,7 +52,9 @@ class QwenCloudRerankConfig(DashScopeRerankConfig):
return _require_qwencloud_api_key(api_key)
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
return api_base or get_secret_str("QWENCLOUD_API_BASE_RERANK") or QWENCLOUD_RERANK_API_BASE
return resolve_dashscope_family_rerank_api_base(
api_base, "QWENCLOUD_API_BASE_RERANK", QWENCLOUD_RERANK_API_BASE
)
class QwenCloudImageGenerationConfig(DashScopeImageGenerationConfig):

View file

@ -12,8 +12,12 @@ Endpoint
- https://dashscope.aliyuncs.com/compatible-api/v1/reranks
Note: chat/embed live under `/compatible-mode/v1/`, but DashScope's rerank
route is exposed under `/compatible-api/v1/reranks` per the docs. Override
with `DASHSCOPE_API_BASE_RERANK` to point at a different host or path.
route is exposed under `/compatible-api/v1/reranks` per the docs. A chat-shaped
`.aliyuncs.com/compatible-mode/v1` base reaching this config (the chat default
from `get_llm_provider`, or a `DASHSCOPE_API_BASE` env var) is redirected to
the same host's rerank route, since `/compatible-mode/v1/reranks` is a dead
route on every DashScope host. Override with `DASHSCOPE_API_BASE_RERANK` to
point at a different host or path.
Empirically, qwen3-rerank accepts `return_documents=true` and echoes
`results[].document.text` back, even though the public docs list the flag
@ -40,7 +44,7 @@ from litellm.types.rerank import (
RerankTokens,
)
from ..common_utils import DashScopeError
from ..common_utils import DashScopeError, resolve_dashscope_family_rerank_api_base
DEFAULT_RERANK_URL: Final = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
@ -67,9 +71,7 @@ class DashScopeRerankConfig(BaseRerankConfig):
return resolved_api_key
def _resolve_rerank_api_base(self, api_base: str | None) -> str:
if api_base is not None:
return api_base
return get_secret_str("DASHSCOPE_API_BASE_RERANK") or DEFAULT_RERANK_URL
return resolve_dashscope_family_rerank_api_base(api_base, "DASHSCOPE_API_BASE_RERANK", DEFAULT_RERANK_URL)
def get_complete_url(
self,

View file

@ -157,9 +157,6 @@ class JinaAIRerankConfig(BaseRerankConfig):
billed_units: RerankBilledUnits | None = None,
model_info: ModelInfo | None = None,
) -> tuple[float, float]:
"""
Jina AI reranker is priced at $0.000000018 per token.
"""
if (
model_info is None
or "input_cost_per_token" not in model_info

View file

@ -19,10 +19,12 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
_should_convert_tool_call_to_json_mode,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
drop_tool_reference_parts_from_tool_messages,
flatten_combinators_and_drop_non_python_regex_patterns,
get_tool_call_names,
hoist_images_from_tool_messages,
tool_with_flattened_parameters,
tool_with_sanitized_parameters,
)
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_convert_url_to_base64,
@ -432,7 +434,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
custom_llm_provider, api_base
)
def _flattened_tools_update_for_openai(
def _sanitized_tools_update_for_openai(
self,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
@ -440,22 +442,26 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"""
OpenAI's chat completions validator rejects tool `parameters` carrying
'oneOf'/'anyOf'/'allOf'/'enum'/'const'/'not' at the top level for every
model family, unlike the Responses API, where GPT-5+ accepts them.
model family, unlike the Responses API, where GPT-5+ accepts them, and
a `pattern` Python's `re` cannot compile for every model family on both.
A custom api_base on the `openai` provider is usually a proxy in front of
the same validator, so regexes are dropped there too, while the lossier
combinator flattening stays limited to api.openai.com hosts.
"""
tools: Final = optional_params.get("tools")
if not isinstance(tools, list):
return _NO_TOOLS_UPDATE
provider: Final = litellm_params.get("custom_llm_provider")
raw_api_base: Final = litellm_params.get("api_base")
if not self._targets_openai_hosted_endpoint(
provider if isinstance(provider, str) else None,
raw_api_base if isinstance(raw_api_base, str) else None,
):
if not isinstance(tools, list) or provider != "openai":
return _NO_TOOLS_UPDATE
flattened: Final = [ # mutable-ok: request tools are a JSON list
tool_with_flattened_parameters(tool) if isinstance(tool, dict) else tool for tool in tools
raw_api_base: Final = litellm_params.get("api_base")
sanitize: Final = (
flatten_combinators_and_drop_non_python_regex_patterns
if self._targets_openai_hosted_endpoint(provider, raw_api_base if isinstance(raw_api_base, str) else None)
else drop_non_python_regex_patterns
)
sanitized: Final = [ # mutable-ok: request tools are a JSON list
tool_with_sanitized_parameters(tool, sanitize) if isinstance(tool, dict) else tool for tool in tools
]
return MappingProxyType({"tools": flattened})
return MappingProxyType({"tools": sanitized})
def transform_request(
self,
@ -489,7 +495,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"model": model,
"messages": messages,
**optional_params,
**self._flattened_tools_update_for_openai(optional_params, litellm_params),
**self._sanitized_tools_update_for_openai(optional_params, litellm_params),
}
async def async_transform_request(
@ -521,7 +527,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"model": model,
"messages": transformed_messages,
**optional_params,
**self._flattened_tools_update_for_openai(optional_params, litellm_params),
**self._sanitized_tools_update_for_openai(optional_params, litellm_params),
}
else:
## allow for any object specific behaviour to be handled

View file

@ -1,4 +1,4 @@
from collections.abc import Mapping, Sequence
from collections.abc import Callable, Mapping, Sequence
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_type_hints
@ -15,6 +15,10 @@ from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
_safe_convert_created_field,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
flatten_combinators_and_drop_non_python_regex_patterns,
)
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
@ -40,7 +44,7 @@ else:
_NO_TOOL_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
_MODEL_FAMILIES_REJECTING_TOP_LEVEL_SCHEMA_COMBINATORS: Final = ("gpt-4", "gpt-3.5", "chatgpt-4o", "o1", "o3", "o4")
_PROVIDERS_WITH_COMBINATOR_REJECTING_VALIDATOR: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
_PROVIDERS_WITH_OPENAI_SCHEMA_VALIDATOR: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
_PROVIDERS_VALIDATING_TOOL_CALL_ITEM_IDS: Final = frozenset({LlmProviders.AZURE, LlmProviders.OPENAI})
@ -293,7 +297,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
model=model, input=validated_input, tools=tools
)
object_schema_tools: Final = self._tools_with_object_parameters(model=model, tools=stripped_tools)
sanitized_tools: Final = self._flatten_tool_schema_combinators_for_openai(
sanitized_tools: Final = self._sanitized_tool_schemas_for_openai(
model=model, tools=object_schema_tools, litellm_params=litellm_params
)
return self._drop_foreign_tool_call_item_ids(stripped_input), sanitized_tools
@ -378,35 +382,35 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
return item
return {key: value for key, value in item.items() if key != "id"} # mutable-ok: outgoing JSON request item
def _flatten_tool_schema_combinators_for_openai(
def _sanitized_tool_schemas_for_openai(
self,
model: str,
tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None,
litellm_params: GenericLiteLLMParams,
) -> Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None:
"""Flatten top-level schema combinators only where OpenAI's validator rejects them.
"""Rewrite tool schemas only where OpenAI's validator rejects them.
OpenAI-compatible backends reusing this config (and the ChatGPT backend
Codex talks to natively) accept them, and so do GPT-5 and later models,
which also call tools better with the union intact. Codex wraps MCP tools
inside namespace entries, so nested ``tools`` arrays are walked too.
Azure OpenAI shares the validator but names deployments arbitrarily, so
the router's declared ``model_info.base_model`` wins over the deployment
name and an unrecognized name without one is left untouched.
Every model family refuses a ``pattern`` Python's ``re`` cannot compile,
while top-level schema combinators are flattened only for the families
whose validator rejects them: OpenAI-compatible backends reusing this
config (and the ChatGPT backend Codex talks to natively) accept them,
and so do GPT-5 and later models, which also call tools better with the
union intact. Codex wraps MCP tools inside namespace entries, so nested
``tools`` arrays are walked too. Azure OpenAI shares the validator but
names deployments arbitrarily, so the router's declared
``model_info.base_model`` wins over the deployment name and an
unrecognized name without one keeps its combinators.
"""
if tools is None or self.custom_llm_provider not in _PROVIDERS_WITH_COMBINATOR_REJECTING_VALIDATOR:
if tools is None or self.custom_llm_provider not in _PROVIDERS_WITH_OPENAI_SCHEMA_VALIDATOR:
return tools
gate_model: Final = self._combinator_gate_model(model=model, litellm_params=litellm_params)
if not self._rejects_top_level_schema_combinators(gate_model):
return tools
flattened: Final = [ # mutable-ok: request tools are a JSON list
self._flattened_tool_or_passthrough(tool) for tool in tools
]
return cast("Sequence[ALL_RESPONSES_API_TOOL_PARAMS]", flattened) # cast-ok: spread keeps each tool's shape
@staticmethod
def _flattened_tool_or_passthrough(tool: object) -> object:
return OpenAIResponsesAPIConfig._flattened_tool_entry(tool) if isinstance(tool, dict) else tool
sanitize: Final = (
flatten_combinators_and_drop_non_python_regex_patterns
if self._rejects_top_level_schema_combinators(gate_model)
else drop_non_python_regex_patterns
)
sanitized: Final = self._sanitized_tools(tools, sanitize)
return cast("Sequence[ALL_RESPONSES_API_TOOL_PARAMS]", sanitized) # cast-ok: spread keeps each tool's shape
@staticmethod
def _rejects_top_level_schema_combinators(model: str) -> bool:
@ -421,35 +425,42 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
return base_model if isinstance(base_model, str) and base_model else model
@staticmethod
def _flattened_tool_entry(
def _sanitized_tool_entry(
entry: Mapping[str, object],
) -> dict[str, object]: # mutable-ok: request tools are JSON dicts
from litellm.litellm_core_utils.prompt_templates.common_utils import (
flatten_top_level_schema_combinators,
)
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
) -> Mapping[str, object]:
parameters: Final = entry.get("parameters")
nested_tools: Final = entry.get("tools")
sanitized_parameters: Final = sanitize(parameters) if isinstance(parameters, dict) else parameters
sanitized_nested_tools: Final = (
OpenAIResponsesAPIConfig._sanitized_tools(nested_tools, sanitize)
if isinstance(nested_tools, list)
else nested_tools
)
parameters_update: Final = (
MappingProxyType({"parameters": flatten_top_level_schema_combinators(parameters)})
if isinstance(parameters, dict)
MappingProxyType({"parameters": sanitized_parameters})
if sanitized_parameters is not parameters
else _NO_TOOL_UPDATE
)
tools_update: Final = (
MappingProxyType({"tools": OpenAIResponsesAPIConfig._flattened_nested_tools(nested_tools)})
if isinstance(nested_tools, list)
MappingProxyType({"tools": sanitized_nested_tools})
if sanitized_nested_tools is not nested_tools
else _NO_TOOL_UPDATE
)
if not parameters_update and not tools_update:
return entry
return {**entry, **parameters_update, **tools_update} # mutable-ok: request tools are JSON dicts
@staticmethod
def _flattened_nested_tools(
nested_tools: Sequence[object],
) -> list[object]: # mutable-ok: namespace tools are a JSON list
return [ # mutable-ok: namespace tools are a JSON list
OpenAIResponsesAPIConfig._flattened_tool_entry(item) if isinstance(item, dict) else item
for item in nested_tools
def _sanitized_tools(
tools: Sequence[object],
sanitize: Callable[[Mapping[str, object]], Mapping[str, object]],
) -> Sequence[object]:
sanitized: Final = [ # mutable-ok: request tools are a JSON list
OpenAIResponsesAPIConfig._sanitized_tool_entry(item, sanitize) if isinstance(item, dict) else item
for item in tools
]
return tools if all(new is old for new, old in zip(sanitized, tools, strict=True)) else sanitized
def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam:
"""
@ -620,15 +631,20 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
return event_pydantic_model.model_construct(**parsed_chunk)
@staticmethod
def parse_terminal_response_from_stream_chunks(all_chunks: list[str]) -> ResponsesAPIResponse | None:
def parse_terminal_event_from_stream_chunks(all_chunks: Sequence[str]) -> ResponsesTerminalEvent | None:
for chunk_str in reversed(all_chunks):
for event_model in (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent):
try:
return event_model.model_validate_json(chunk_str.removeprefix("data: ")).response
return event_model.model_validate_json(chunk_str.removeprefix("data: "))
except ValueError:
continue
return None
@staticmethod
def parse_terminal_response_from_stream_chunks(all_chunks: list[str]) -> ResponsesAPIResponse | None:
terminal_event: Final = OpenAIResponsesAPIConfig.parse_terminal_event_from_stream_chunks(all_chunks)
return None if terminal_event is None else terminal_event.response
@staticmethod
def get_event_model_class(event_type: str) -> type[BaseLiteLLMOpenAIResponseObject]:
"""

View file

@ -36,6 +36,7 @@ class SagemakerChatHandler(BaseAWSLLM):
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -63,6 +64,7 @@ class SagemakerChatHandler(BaseAWSLLM):
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
return credentials, aws_region_name

View file

@ -59,6 +59,7 @@ class SagemakerLLM(BaseAWSLLM):
aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None)
aws_external_id: Final = optional_params.pop("aws_external_id", None)
aws_session_tags: Final = optional_params.pop("aws_session_tags", None)
### SET REGION NAME ###
if aws_region_name is None:
@ -86,6 +87,7 @@ class SagemakerLLM(BaseAWSLLM):
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
aws_external_id=aws_external_id,
aws_session_tags=aws_session_tags,
)
return credentials, aws_region_name

View file

@ -19,7 +19,7 @@ import random
import sys
import time
import traceback
from collections.abc import AsyncIterator, Coroutine, Iterable, Mapping, Sequence
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Mapping, Sequence
from concurrent import futures
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from copy import deepcopy
@ -8595,7 +8595,7 @@ def config_completion(**kwargs):
)
def stream_chunk_builder_text_completion(chunks: list, messages: list | None = None) -> TextCompletionResponse:
def stream_chunk_builder_text_completion(chunks: list, messages: Sequence | None = None) -> TextCompletionResponse:
id: Final = chunks[0]["id"]
object: Final = chunks[0]["object"]
created: Final = chunks[0]["created"]
@ -8712,10 +8712,11 @@ def _stamp_streaming_usage_cost(usage: Usage, response: ModelResponse, logging_o
def stream_chunk_builder(
chunks: list,
messages: list | None = None,
messages: Sequence | None = None,
start_time=None,
end_time=None,
logging_obj: Optional["Logging"] = None,
count_prompt_tokens: Callable[[], int] | None = None,
) -> ModelResponse | TextCompletionResponse | None:
try:
if chunks is None:
@ -8789,6 +8790,7 @@ def stream_chunk_builder(
completion_output=completion_output,
messages=messages,
reasoning_tokens=0,
count_prompt_tokens=count_prompt_tokens,
)
setattr(response, "usage", usage)
@ -8966,6 +8968,7 @@ def stream_chunk_builder(
completion_output=completion_output,
messages=messages,
reasoning_tokens=reasoning_tokens,
count_prompt_tokens=count_prompt_tokens,
)
setattr(response, "usage", usage)

View file

@ -30295,7 +30295,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@ -30340,7 +30340,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.1-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@ -30432,7 +30432,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.2-2025-12-11": {
"cache_read_input_token_cost": 1.75e-07,
@ -30478,7 +30478,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.2-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@ -31474,7 +31474,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 2.5e-06,
"output_cost_per_token_above_272k_tokens_flex": 1.125e-05,
"cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07
@ -31526,7 +31526,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 2.5e-06,
"output_cost_per_token_above_272k_tokens_flex": 1.125e-05,
"cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07
@ -31578,7 +31578,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 3e-05,
"output_cost_per_token_above_272k_tokens_flex": 0.000135
},
@ -31629,7 +31629,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 3e-05,
"output_cost_per_token_above_272k_tokens_flex": 0.000135
},
@ -33709,13 +33709,14 @@
"supports_tool_choice": true
},
"jina-reranker-v2-base-multilingual": {
"input_cost_per_token": 1.8e-08,
"input_cost_per_token": 5e-08,
"litellm_provider": "jina_ai",
"max_input_tokens": 1024,
"max_output_tokens": 1024,
"max_tokens": 1024,
"mode": "rerank",
"output_cost_per_token": 1.8e-08
"output_cost_per_token": 0.0,
"source": "https://api.jina.ai/v1/models"
},
"jp.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
@ -39969,6 +39970,45 @@
"supports_tool_choice": true,
"supports_vision": true
},
"openrouter/openai/gpt-5.6-sol": {
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
"default_reasoning_effort": "medium",
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
"litellm_provider": "openrouter",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"reasoning_effort_levels": [
"none",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://openrouter.ai/openai/gpt-5.6-sol",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"openrouter/openai/gpt-oss-120b": {
"input_cost_per_token": 3.7e-08,
"litellm_provider": "openrouter",

View file

@ -179,16 +179,7 @@ def _gateway_dcr_challenge_target(
mcp_servers: list[str] | None,
client_ip: str | None,
) -> str | None:
"""The single path-named server this request targets, iff it resolves to a
gateway-managed oauth2 server — the one per-server shape the gateway's own keyless
DCR flow serves end to end, so the 401 challenge may advertise the per-server
protected-resource metadata (whose ``authorization_servers`` names the gateway).
Multi-server CSV paths, header/path mismatches, unknown names, and every
client-forwarded or delegated mode return ``None``: those cells keep their existing
challenge (or absence of one), and a challenge is never emitted for a name the
public discovery routes would 404, so this reveals exactly the server set the
per-server protected-resource metadata already reveals."""
"""Resolve a single path target whose sign-in metadata advertises the gateway."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
@ -217,7 +208,7 @@ def _is_gateway_dcr_challenge_scope(
the caller is not a cold-start DCR client), on the scopes the gateway's keyless
flow serves: the aggregate ``/mcp`` endpoint, an ``x-mcp-servers``-scoped request
(the resource the client configured is still ``/mcp``), or a per-server path whose
single target is a gateway-managed oauth2 server. Every other named target keeps
single target advertises gateway-owned sign-in. Every other named target keeps
its existing behavior, failing closed to the original admission error."""
if not _is_litellm_auth_admission_error(exc):
return False
@ -236,7 +227,7 @@ def _gateway_dcr_challenge(
) -> HTTPException:
"""The RFC 9728 challenge pointing the client at the protected-resource metadata
matching the scope it requested: the per-server document (same URL spelling the
request arrived on) when the single target is a gateway-managed oauth2 server,
request arrived on) when the single target advertises gateway-owned sign-in,
else the gateway's aggregate document. Either way the client discovers the gateway
as its authorization server and starts the same sign-in flow.

View file

@ -2310,8 +2310,7 @@ async def _build_oauth_protected_resource_response(
it. Only the legacy ``is_oauth_passthrough`` opt-in rewrites ``resource`` to
the gateway's own URL so clients present the bearer token back to the gateway.
An explicitly named gateway-managed oauth2 server (interactive with
gateway-vaulted per-user tokens, or M2M) advertises the gateway's own
An explicitly named server with gateway-owned sign-in advertises the gateway's own
authorization server (``{base}/mcp``): a keyless DCR client that configured the
per-server URL completes the same sign-in flow the aggregate ``/mcp`` endpoint
supports and is admitted with a gateway session bearer. The per-server relay
@ -2401,11 +2400,6 @@ async def _build_oauth_protected_resource_response(
if obo_response is not None:
return obo_response
# An OBO server with no configured issuer falls through to the gateway default so discovery still
# returns metadata; every other non-oauth2 named server 404s to avoid enumeration.
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
if explicitly_named and mcp_server is not None and mcp_server.advertises_gateway_authorization_server:
return {
"authorization_servers": [f"{request_base_url}/mcp"],
@ -2413,6 +2407,9 @@ async def _build_oauth_protected_resource_response(
"scopes_supported": (mcp_server.scopes if mcp_server.scopes else []),
}
if mcp_server is None or mcp_server.auth_type != MCPAuth.oauth2_token_exchange:
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth-protected resource")
return {
"authorization_servers": [
(f"{request_base_url}/{mcp_server_name}" if mcp_server_name else f"{request_base_url}")

View file

@ -411,7 +411,7 @@ def relative_request_url(request: Request) -> str:
def resolve_scoped_resource_server(request: Request, resource: str | None) -> MCPServer | None:
"""Resolve an RFC 8707 ``resource`` value to the single gateway-managed oauth2 server it
"""Resolve an RFC 8707 ``resource`` value to the single gateway-owned server it
names, or ``None`` for every other shape: absent, the aggregate resource, a foreign
host, an unparseable value, a multi-server path, an unknown name, or any server mode the
keyless gateway flow does not serve (whose protected-resource metadata never directs a
@ -443,7 +443,7 @@ def resolve_scoped_resource_server(request: Request, resource: str | None) -> MC
if len(names) != 1:
return None
server: Final = global_mcp_server_manager.get_mcp_server_by_name(names[0])
if server is None or not server.is_gateway_managed_oauth2:
if server is None or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server):
return None
return server
@ -729,11 +729,15 @@ async def _flow_target(
server: Final = global_mcp_server_manager.get_mcp_server_by_id(flow.resource_server_id)
if (
server is None
or not server.is_gateway_managed_oauth2
or not (server.is_gateway_managed_oauth2 or server.advertises_gateway_authorization_server)
or not await lookup_server_reachability(flow.user_id, server.server_id)
):
return "stale", None
state: Final = "m2m" if MCPServerManager.effective_oauth2_flow(server) == "client_credentials" else "interactive"
state: Final = (
"interactive"
if server.is_gateway_managed_oauth2 and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
else "m2m"
)
return state, server

View file

@ -22,7 +22,16 @@ Response headers returned (all values are masked for safety):
x-mcp-debug-auth-resolution
Which auth priority was used for the outbound MCP call:
``per-request-header``, ``m2m-client-credentials``, ``static-token``,
``oauth2-passthrough``, or ``no-auth``.
``oauth2-passthrough``, ``stored-user-token``, ``token-exchange``,
``id-jag``, ``aws-sigv4``, ``extra-headers``, or ``no-auth``.
``unresolved`` means no outcome was available before the first response
frame; ``multiple`` means several servers resolved credentials;
``not-applicable`` covers stdio; ``resolution-failed`` is a resolver error.
x-mcp-debug-auth-resolutions
For multiple servers, a JSON map of server IDs to resolution labels.
At most 32 entries are included; x-mcp-debug-auth-resolutions-truncated
is true when additional servers were omitted. No credentials are included.
x-mcp-debug-outbound-url
The upstream MCP server URL that will receive the request.
@ -58,10 +67,16 @@ header is free for OAuth2 discovery::
Symptom: ``x-mcp-debug-oauth2-token`` shows ``(none)`` and
``x-mcp-debug-auth-resolution`` shows ``no-auth``.
This means the client didn't go through the OAuth2 flow. Check that:
1. The ``Authorization`` header is NOT set as a static header in the client config.
2. The ``.well-known/oauth-protected-resource`` endpoint returns valid metadata.
3. The MCP server in LiteLLM config has ``auth_type: oauth2``.
``no-auth`` means the resolved upstream client carries no authentication.
An absent inbound OAuth2 token does not imply the user skipped OAuth: the gateway
can retrieve a stored per-user token, reported as ``stored-user-token``.
``unresolved`` is used when a stream starts before credential resolution, or a
request (such as initialization or a cached tool listing) resolves no credential.
Debug reporting does not fetch credentials or delay a streaming frame to resolve them.
``extra-headers`` identifies supplied headers that won over the resolver or were
the only headers supplied; their values are never inspected to guess a scheme.
``per-request-header`` denotes a legacy credential override, including a BYOK
credential supplied by the gateway; it does not imply a caller-supplied token.
**Common issue: M2M token used instead of user token**
@ -69,8 +84,8 @@ Symptom: ``x-mcp-debug-auth-resolution`` shows ``m2m-client-credentials``.
This means the server has ``client_id``/``client_secret``/``token_url``
configured and LiteLLM is fetching a machine-to-machine token instead of
using the per-user OAuth2 token. If you want per-user tokens, remove the
client credentials from the server config.
using the per-user OAuth2 token. For gateway-stored per-user tokens,
configure ``oauth2_flow: authorization_code``.
Usage from Claude Code::
@ -85,14 +100,16 @@ Usage with curl::
http://localhost:4000/mcp/atlassian_mcp
"""
from typing import TYPE_CHECKING, Final
import json
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import Final
from starlette.requests import HTTPConnection
from starlette.types import Message, Send
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution
# Header the client sends to opt into debug mode
MCP_DEBUG_REQUEST_HEADER: Final = "x-litellm-mcp-debug"
@ -101,6 +118,83 @@ MCP_DEBUG_REQUEST_HEADER: Final = "x-litellm-mcp-debug"
_RESPONSE_HEADER_PREFIX: Final = "x-mcp-debug"
MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: Final = "litellm.mcp.auth_diagnostics"
def record_auth_resolution(server_id: str, source: AuthResolution) -> None:
from mcp.server.lowlevel.server import request_ctx
context: Final[object] = request_ctx.get(None)
request: Final[object] = getattr(context, "request", None)
if isinstance(request, HTTPConnection):
diagnostics: Final[object] = request.scope.get(MCP_AUTH_DIAGNOSTICS_SCOPE_KEY)
if isinstance(diagnostics, MCPAuthDiagnostics):
diagnostics.record(server_id, source)
class MCPAuthDiagnostics:
def __init__(self) -> None:
self._outcomes: tuple[tuple[str, AuthResolution], ...] = ()
def record(self, server_id: str, resolution: AuthResolution) -> None:
self._outcomes = tuple(item for item in self._outcomes if item[0] != server_id) + ((server_id, resolution),)
def resolution(self) -> str:
match self._outcomes:
case ():
return AuthResolution.unresolved.value
case ((_, source),):
return source.value
case _:
return AuthResolution.multiple.value
def headers(self) -> Mapping[str, str]:
if len(self._outcomes) <= 1:
return MappingProxyType({"x-mcp-debug-auth-resolution": self.resolution()})
return MappingProxyType(
{
"x-mcp-debug-auth-resolution": AuthResolution.multiple.value,
"x-mcp-debug-auth-resolutions": json.dumps(
{
server_id: source.value for server_id, source in self._outcomes[:32]
}, # mutable-ok: JSON encoder requires a concrete dict
separators=(",", ":"),
ensure_ascii=True,
),
**(
MappingProxyType({"x-mcp-debug-auth-resolutions-truncated": "true"})
if len(self._outcomes) > 32
else MappingProxyType({})
),
}
)
class _DiagnosticSend:
def __init__(self, send: Send, headers: Mapping[str, str], resolution: Callable[[], Mapping[str, str]]) -> None:
self._send = send
self._headers = headers
self._resolution = resolution
self._start: Message | None = None
async def __call__(self, message: Message) -> None:
if message["type"] == "http.response.start":
self._start = message
return
if self._start is not None:
start: Final = self._start
self._start = None
headers: Final = MappingProxyType({**self._headers, **self._resolution()})
await self._send(
{ # mutable-ok: ASGI send consumes a mutable message mapping
**start,
"headers": tuple(start.get("headers", ()))
+ tuple((key.encode(), value.encode()) for key, value in headers.items()),
}
)
await self._send(message)
class MCPDebug:
"""
Static helper class for MCP OAuth2 debug headers.
@ -144,37 +238,6 @@ class MCPDebug:
return val.strip().lower() in ("true", "1", "yes")
return False
@staticmethod
def resolve_auth_resolution(
server: "MCPServer",
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
oauth2_headers: dict[str, str] | None,
) -> str:
"""
Determine which auth priority will be used for the outbound MCP call.
Returns one of: ``per-request-header``, ``m2m-client-credentials``,
``static-token``, ``oauth2-passthrough``, or ``no-auth``.
"""
from litellm.types.mcp import MCPAuth
has_server_specific: Final = bool(
mcp_server_auth_headers
and (
mcp_server_auth_headers.get(server.alias or "") or mcp_server_auth_headers.get(server.server_name or "")
)
)
if has_server_specific or mcp_auth_header:
return "per-request-header"
if server.has_client_credentials:
return "m2m-client-credentials"
if server.authentication_token:
return "static-token"
if oauth2_headers and server.auth_type == MCPAuth.oauth2:
return "oauth2-passthrough"
return "no-auth"
@staticmethod
def build_debug_headers(
*,
@ -244,12 +307,21 @@ class MCPDebug:
return debug
@staticmethod
def wrap_send_with_debug_headers(send: Send, debug_headers: dict[str, str]) -> Send:
def wrap_send_with_debug_headers(
send: Send,
debug_headers: Mapping[str, str],
resolution: Callable[[], Mapping[str, str]] | None = None,
*,
request_method: str | None = None,
) -> Send:
"""
Return a new ASGI ``send`` callable that injects *debug_headers*
into the ``http.response.start`` message.
"""
if resolution is not None and request_method == "POST":
return _DiagnosticSend(send, debug_headers, resolution)
async def _send_with_debug(message: Message) -> None:
if message["type"] == "http.response.start":
headers: Final = list(message.get("headers", []))
@ -266,8 +338,6 @@ class MCPDebug:
raw_headers: dict[str, str] | None,
scope: dict,
mcp_servers: list[str] | None,
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
oauth2_headers: dict[str, str] | None,
client_ip: str | None,
) -> dict[str, str]:
@ -288,16 +358,13 @@ class MCPDebug:
server_url: str | None = None
server_auth_type: str | None = None
auth_resolution = "no-auth"
auth_resolution: Final = AuthResolution.unresolved.value
for server_name in mcp_servers or []:
server = global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
if server:
server_url = server.url
server_auth_type = server.auth_type
auth_resolution = MCPDebug.resolve_auth_resolution(
server, mcp_auth_header, mcp_server_auth_headers, oauth2_headers
)
break
scope_headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)

View file

@ -80,6 +80,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
raise_classified_list_failure,
upstream_auth_challenge,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import record_auth_resolution
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
@ -108,12 +109,14 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import resolve_credentials_with_source
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_provider import (
build_token_exchanger,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
AuthorizationCodeConfig,
AuthResolution,
ClientCredentialsConfig,
CredError,
IdJagConfig,
@ -3832,13 +3835,21 @@ class MCPServerManager:
(authorization_code's browser-OAuth 401, token_exchange's RFC 9728 challenge) or maps any
other ``CredError`` onto its public HTTP status; it never returns an error as a value.
"""
match await provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
case Ok(auth):
match await resolve_credentials_with_source(provider, to_subject(user_api_key_auth, subject_token), spec):
case Ok(credential):
auth: Final = credential.auth
# NoOpAuth has no header_name and so never conflicts.
header_name: Final[str | None] = getattr(auth, "header_name", None)
if header_name is None or not extra_headers:
source: Final = (
AuthResolution.extra_headers
if credential.source == AuthResolution.no_auth and extra_headers
else credential.source
)
record_auth_resolution(server.server_id, source)
return auth, extra_headers
if not has_header(extra_headers, header_name):
record_auth_resolution(server.server_id, credential.source)
return auth, extra_headers
if isinstance(
spec.config,
@ -3853,11 +3864,14 @@ class MCPServerManager:
# one-shot 401 refetch is lost with it). Drop only the header the resolved
# credential is about to occupy, so a static credential the operator aimed at a
# DIFFERENT header still reaches upstream.
record_auth_resolution(server.server_id, credential.source)
return auth, without_header(extra_headers, header_name)
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
# header or static_headers) is intentional and wins; v1 applies those last.
record_auth_resolution(server.server_id, AuthResolution.extra_headers)
return None, extra_headers
case Error(err):
record_auth_resolution(server.server_id, AuthResolution.failed)
if err.tag == "unauthorized" and isinstance(spec.config, AuthorizationCodeConfig):
# authorization_code's missing per-user token -> the per-server browser-OAuth
# challenge, built here where the full MCPServer is in hand.
@ -3960,6 +3974,7 @@ class MCPServerManager:
Returns:
Configured MCP client instance.
"""
record_auth_resolution(server.server_id, AuthResolution.unresolved)
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
transport: Final = resolved_server.transport or MCPTransport.sse
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
@ -4032,6 +4047,7 @@ class MCPServerManager:
env=resolved_env,
)
record_auth_resolution(server.server_id, AuthResolution.not_applicable)
return MCPClient(
server_url="", # Not used for stdio
transport_type=transport,
@ -4086,6 +4102,20 @@ class MCPServerManager:
aws_session_name=resolved_server.aws_session_name,
)
legacy_source: Final = (
AuthResolution.aws_sigv4
if aws_auth is not None
else AuthResolution.extra_headers
if extra_headers and has_header(extra_headers, auth_header_name or "Authorization")
else AuthResolution.per_request_header
if mcp_auth_header
else AuthResolution.static_token
if auth_value
else AuthResolution.extra_headers
if extra_headers
else AuthResolution.no_auth
)
record_auth_resolution(server.server_id, legacy_source)
return MCPClient(
server_url=server_url,
transport_type=transport,

View file

@ -65,6 +65,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
AuthorizationCodeConfig,
AuthResolution,
AuthSpecKind,
AwsSigV4Config,
Byok,
@ -76,6 +77,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
NoneConfig,
PassthroughConfig,
PrivateKeyJwtAuth,
ResolvedCredential,
ServerSpec,
SharedKey,
Subject,
@ -448,3 +450,32 @@ def _client_auth_fingerprint(client_auth: ClientAuth) -> str:
def _not_implemented(kind: AuthSpecKind) -> Result[httpx.Auth, CredError]:
return Error(CredError.of_not_implemented(f"{kind.value}: resolver arm not implemented yet"))
async def resolve_credentials_with_source(
provider: UpstreamCredentialProvider, subject: Subject, server: ServerSpec
) -> Result[ResolvedCredential, CredError]:
match await provider.resolve_credentials(subject, server):
case Error(err):
return Error(err)
case Ok(auth):
if isinstance(auth, NoOpAuth):
return Ok(ResolvedCredential(auth, AuthResolution.no_auth))
match server.config:
case NoneConfig():
return Ok(ResolvedCredential(auth, AuthResolution.no_auth))
case ApiKeyConfig():
return Ok(ResolvedCredential(auth, AuthResolution.static_token))
case PassthroughConfig():
return Ok(ResolvedCredential(auth, AuthResolution.oauth2_passthrough))
case ClientCredentialsConfig():
return Ok(ResolvedCredential(auth, AuthResolution.client_credentials))
case TokenExchangeConfig():
return Ok(ResolvedCredential(auth, AuthResolution.token_exchange))
case IdJagConfig():
return Ok(ResolvedCredential(auth, AuthResolution.id_jag))
case AuthorizationCodeConfig():
return Ok(ResolvedCredential(auth, AuthResolution.stored_user_token))
case AwsSigV4Config():
return Ok(ResolvedCredential(auth, AuthResolution.aws_sigv4))
assert_never(server.config)

View file

@ -26,10 +26,11 @@ union (see `result.py`), not `expression.Result`.
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from dataclasses import dataclass, field
from enum import Enum
from typing import Annotated, Final, Literal
import httpx
from expression import case, tag, tagged_union
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
from typing_extensions import assert_never
@ -46,6 +47,29 @@ from litellm.types.mcp import (
)
class AuthResolution(str, Enum):
no_auth = "no-auth"
stored_user_token = "stored-user-token"
static_token = "static-token"
per_request_header = "per-request-header"
oauth2_passthrough = "oauth2-passthrough"
client_credentials = "m2m-client-credentials"
token_exchange = "token-exchange"
id_jag = "id-jag"
aws_sigv4 = "aws-sigv4"
extra_headers = "extra-headers"
not_applicable = "not-applicable"
unresolved = "unresolved"
failed = "resolution-failed"
multiple = "multiple"
@dataclass(frozen=True, slots=True)
class ResolvedCredential:
auth: httpx.Auth = field(repr=False)
source: AuthResolution
class AuthSpecKind(str, Enum):
"""The server's statically-declared upstream-auth mode — derived from its `config`.

View file

@ -49,7 +49,11 @@ from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_server_name,
_mcp_proxy_mode, # pyright: ignore[reportPrivateUsage] # server-owned request mode
)
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
from litellm.proxy._experimental.mcp_server.mcp_debug import (
MCP_AUTH_DIAGNOSTICS_SCOPE_KEY,
MCPAuthDiagnostics,
MCPDebug,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
get_passthrough_www_authenticate,
@ -4472,13 +4476,15 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
scope=dict(scope),
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
client_ip=_client_ip,
)
if _debug_headers:
send = MCPDebug.wrap_send_with_debug_headers(send, _debug_headers)
diagnostics: Final = MCPAuthDiagnostics() if _debug_headers else None
if diagnostics is not None:
scope[MCP_AUTH_DIAGNOSTICS_SCOPE_KEY] = diagnostics
send = MCPDebug.wrap_send_with_debug_headers(
send, _debug_headers, diagnostics.headers, request_method=scope.get("method")
)
# Ensure session managers are initialized
if not _SESSION_MANAGERS_INITIALIZED:

View file

@ -3726,6 +3726,15 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
],
)
pointfive: CallbackOnUI = CallbackOnUI(
litellm_callback_name="pointfive",
ui_callback_name="PointFive",
litellm_callback_params=[ # mutable-ok: the registry field is typed list
"POINTFIVE_API_KEY",
"POINTFIVE_API_URL",
],
)
class SpendLogsRouterMetadata(TypedDict):
"""

View file

@ -94,6 +94,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
team_membership_auth_cache_key,
team_membership_reservation_cache_key,
)
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
TOOL_CAPABLE_CALL_TYPES,
@ -3450,36 +3451,37 @@ async def _fetch_key_object_from_db_with_reconnect(
"""
Fetch key object from DB and retry once if a DB connection error can be healed.
"""
try:
return await prisma_client.get_data(
token=hashed_token,
table_name="combined_view",
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e:
if PrismaDBExceptionHandler.is_database_transport_error(e):
did_reconnect = False
if hasattr(prisma_client, "attempt_db_reconnect"):
auth_reconnect_timeout = getattr(prisma_client, "_db_auth_reconnect_timeout_seconds", 2.0)
if not isinstance(auth_reconnect_timeout, (int, float)):
auth_reconnect_timeout = 2.0
auth_reconnect_lock_timeout = getattr(prisma_client, "_db_auth_reconnect_lock_timeout_seconds", 0.1)
if not isinstance(auth_reconnect_lock_timeout, (int, float)):
auth_reconnect_lock_timeout = 0.1
did_reconnect = await prisma_client.attempt_db_reconnect(
reason="auth_get_key_object_lookup_failure",
timeout_seconds=auth_reconnect_timeout,
lock_timeout_seconds=auth_reconnect_lock_timeout,
)
if did_reconnect:
return await prisma_client.get_data(
token=hashed_token,
table_name="combined_view",
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
raise
async with db_lookup_gate.current():
try:
return await prisma_client.get_data(
token=hashed_token,
table_name="combined_view",
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e:
if PrismaDBExceptionHandler.is_database_transport_error(e):
did_reconnect = False
if hasattr(prisma_client, "attempt_db_reconnect"):
auth_reconnect_timeout = getattr(prisma_client, "_db_auth_reconnect_timeout_seconds", 2.0)
if not isinstance(auth_reconnect_timeout, (int, float)):
auth_reconnect_timeout = 2.0
auth_reconnect_lock_timeout = getattr(prisma_client, "_db_auth_reconnect_lock_timeout_seconds", 0.1)
if not isinstance(auth_reconnect_lock_timeout, (int, float)):
auth_reconnect_lock_timeout = 0.1
did_reconnect = await prisma_client.attempt_db_reconnect(
reason="auth_get_key_object_lookup_failure",
timeout_seconds=auth_reconnect_timeout,
lock_timeout_seconds=auth_reconnect_lock_timeout,
)
if did_reconnect:
return await prisma_client.get_data(
token=hashed_token,
table_name="combined_view",
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
raise
def jwt_key_mapping_cache_key(jwt_claim_name: str, jwt_claim_value: str) -> str:

View file

@ -27,6 +27,7 @@ from litellm.litellm_core_utils.url_utils import (
provider_url_destination_candidates,
validate_url,
)
from litellm.llms.azure.passthrough.transformation import azure_router_model_in_endpoint
from litellm.proxy._types import *
from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_metadata
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
@ -316,6 +317,7 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
"aws_profile_name",
"aws_session_name",
"aws_external_id",
"aws_session_tags",
"vertex_credentials",
# Azure managed-identity / federated-auth token. The Azure provider
# transformer reads ``azure_ad_token`` (top-level or via
@ -2003,9 +2005,20 @@ def get_model_from_request(
bedrock_model: Final = _model_from_bedrock_route(route)
return model if bedrock_model is None else bedrock_model
if route.lower().startswith(("/azure/", "/azure_ai/")):
azure_model: Final = _router_model_from_azure_route(route, llm_router)
return model if azure_model is None else azure_model
return model
def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str | None:
if llm_router is None:
return None
endpoint: Final = re.sub(r"^/azure(?:_ai)?/", "", route, flags=re.IGNORECASE)
return azure_router_model_in_endpoint(endpoint, frozenset(llm_router.get_model_names()))
def _model_from_bedrock_route(route: str) -> str | None:
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_extract_model_from_bedrock_endpoint,

View file

@ -92,6 +92,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_set_request_parsed_body,
populate_request_with_path_params,
)
from litellm.proxy.common_utils.model_listing_utils import claude_code_requested_group
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
@ -183,6 +184,44 @@ def _get_model_from_request_context(
)
_CLAUDE_MODEL_ROUTES: Final = frozenset(
f"/{prefix}{endpoint}" for prefix in ("", "v1/") for endpoint in ("messages", "chat/completions", "responses")
)
_CLAUDE_MODEL_NORMALIZED: Final = "litellm.claude_model_normalized"
async def _normalize_claude_model(
request_data: dict, valid_token: UserAPIKeyAuth, request: Request | None, route: str
) -> None:
from litellm.proxy.proxy_server import llm_router, prisma_client, proxy_config, proxy_logging_obj
if route not in _CLAUDE_MODEL_ROUTES or llm_router is None:
return
if request is not None and request.scope.get(_CLAUDE_MODEL_NORMALIZED) is True:
return
requested: Final = _get_model_from_request_context(request_data, route, request, llm_router)
if not isinstance(requested, str) or requested != request_data.get("model"):
return
if not requested.startswith("claude-router-") and not requested.lower().endswith("[1m]"):
return
settings: Final = await proxy_config.get_hierarchical_router_settings(
user_api_key_dict=valid_token, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
)
aliases: Final = settings.get("model_group_alias") if isinstance(settings, Mapping) else None
source: Final = claude_code_requested_group(
requested, llm_router, valid_token.team_id, (valid_token.aliases, valid_token.team_model_aliases, aliases)
)
if request is not None:
request.scope[_CLAUDE_MODEL_NORMALIZED] = True
if source is None:
return
request_data["model"] = source
_safe_set_request_parsed_body(request=request, parsed_body=request_data)
if request is not None:
request._json = request_data
request._body = orjson.dumps(request_data)
def _get_model_names_for_budget_checks(
model: str | list[str] | None,
) -> list[str]:
@ -2768,6 +2807,7 @@ async def _authorize_authenticated_request(
"""
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj, request=request)
await _normalize_claude_model(request_data, user_api_key_auth_obj, request, route)
# Single authorization point. Builder paths MUST NOT call common_checks.
# Route through the same exception handler the builder uses so
@ -3131,6 +3171,7 @@ async def _enforce_key_and_fallback_model_access(
Key-level model allowlist and client fallbacks (same as standard auth).
Not included in common_checks — common_checks enforces team/user/project model access only.
"""
await _normalize_claude_model(request_data, valid_token, request, route)
config: Final = valid_token.config
if config != {}:

View file

@ -490,7 +490,7 @@ lite codex exec "summarize the repo"
Each command resolves your LiteLLM key (logging in via SSO when none is stored and you are at a terminal; otherwise it expects `LITELLM_PROXY_API_KEY` or `--api-key`), checks the key against the proxy so bad credentials fail immediately instead of deep inside the agent, exports the environment variables the agent reads, then replaces itself with the agent process.
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. It also gets `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` (again unless you already set it) so Claude Code v2.1.129+ fills its `/model` picker from the proxy's `/v1/models`; Claude Code only lists entries whose id contains `claude` or `anthropic`, and older versions ignore the variable. Export it as `0` to turn discovery off. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol). OpenCode additionally gets `OPENCODE_CONFIG_CONTENT` holding a generated `litellm` provider (`@ai-sdk/openai-compatible`, the proxy `/v1` URL, `{env:OPENAI_API_KEY}`) with one model entry per chat model your key can see on `/v1/models`, so its model picker mirrors the proxy without a hand-maintained `opencode.json`; OpenCode merges that over your own config files, and if you already export `OPENCODE_CONFIG_CONTENT` yours is left alone. When the list cannot be fetched, `lite opencode` says so on stderr and launches anyway.
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. It also gets `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` (again unless you already set it) so Claude Code v2.1.129+ fills its `/model` picker from the proxy's `/v1/models`; Claude Code only lists entries whose id contains `claude` or `anthropic`, so the proxy lists every other group to Claude Code as `claude-router-<UTF-8 hex of the group name>` and marks a group whose input window reaches 1M with `[1m]`, and a request on such an id is served by the group. Older Claude Code versions ignore the variable. Export it as `0` to turn discovery off. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol). OpenCode additionally gets `OPENCODE_CONFIG_CONTENT` holding a generated `litellm` provider (`@ai-sdk/openai-compatible`, the proxy `/v1` URL, `{env:OPENAI_API_KEY}`) with one model entry per chat model your key can see on `/v1/models`, so its model picker mirrors the proxy without a hand-maintained `opencode.json`; OpenCode merges that over your own config files, and if you already export `OPENCODE_CONFIG_CONTENT` yours is left alone. When the list cannot be fetched, `lite opencode` says so on stderr and launches anyway.
pi ignores base-URL environment variables entirely, so `lite pi` (kept out of the `lite --help` command listing for now, but fully functional) wires it up differently: before handoff it fetches the models your key can use from the proxy's `/v1/models` (plus each model's context window and output cap from `/model_group/info`, when available) and syncs them into a `litellm` provider entry in pi's `~/.pi/agent/models.json` (honoring `PI_CODING_AGENT_DIR`), then starts pi on that provider's first model via an injected `--model litellm/<id>`. Only that one provider entry is rewritten; the rest of the file, including any other custom providers, is left alone. The entry references the key as `$LITELLM_PROXY_API_KEY`, which the wrapper exports for the session, so the token itself never lands on disk and plain `pi` outside the wrapper simply shows the litellm models as unavailable. Your own flags come after the injected pin, so `lite pi --model litellm/<other-id>` wins, and inside the TUI the `/model` picker lists every synced litellm model.
@ -548,7 +548,7 @@ lite --base-url https://your-proxy.example.com configure claude --api-key sk-...
claude
```
With `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) the key is written into `env.ANTHROPIC_AUTH_TOKEN`. Without one, your `lite login` credential is used the way `--config-claude` uses it, through `apiKeyHelper`, so a later `lite login` (or a `--pkce` renewal) picks up on its own and nothing secret lands in the file; a missing or stale login is refreshed first. Either way the command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (the ones whose id contains `claude` or `anthropic`) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key, which has to be on `/v1/models` for the key. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute up` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window and sends no thinking parameters for it, so either name the group like a Claude model id or append `[1m]` to opt into the 1M window. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control
With `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) the key is written into `env.ANTHROPIC_AUTH_TOKEN`. Without one, your `lite login` credential is used the way `--config-claude` uses it, through `apiKeyHelper`, so a later `lite login` (or a `--pkce` renewal) picks up on its own and nothing secret lands in the file; a missing or stale login is refreshed first. Either way the command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-<UTF-8 hex of the group name>` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key, which has to be on `/v1/models` for the key. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute up` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control
Plain `lite configure`, with no agent named, asks the same things interactively: which agents to wire (Claude Code today) and which of the proxy's models to start on, picked from `/v1/models` with a type-to-filter prompt

View file

@ -1,16 +1,23 @@
"""`lite configure claude` and `lite unconfigure claude`: persistent Claude Code wiring, undoable."""
import os
import re
import sys
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
import click
from InquirerPy import inquirer
from InquirerPy.base.control import Choice
from litellm.proxy.common_utils.model_listing_utils import (
CLAUDE_CODE_CLIENT,
CLAUDE_CODE_PICKER_PATTERN,
GATEWAY_CLIENT_HEADER,
)
from .auth import CliContextObj, context_secret_vault, get_stored_api_key
from .claude_settings import (
STARTING_MODEL_ROLE,
@ -30,14 +37,16 @@ from .claude_settings import (
settings_file_owners,
unconfigure_claude_settings,
)
from .pi import ListingFailure, PiSyncError, fetch_model_ids
from .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing
from .up import ensure_fresh_login
_LISTED_MODELS_SHOWN: Final = 20
_CLAUDE_TARGET: Final = "claude"
_TARGETS: Final = ((_CLAUDE_TARGET, "Claude Code (CLI)"),)
_KEEP_DEFAULT_MODEL: Final = "Keep Claude Code's own default"
_CLAUDE_CODE_PICKER_FILTER: Final = re.compile(r"claude|anthropic", re.IGNORECASE)
_CLAUDE_CODE_VIEW: Final = MappingProxyType(
{"anthropic-version": "2023-06-01", GATEWAY_CLIENT_HEADER: CLAUDE_CODE_CLIENT}
)
_MODEL_OPTION_HELP: Final = (
f"Proxy model to set as {STARTING_MODEL_ROLE}. Must be listed on /v1/models for the key; without it, "
"Claude Code keeps its own default and a pin an earlier configure made is let go of. Nothing pins Claude "
@ -65,7 +74,16 @@ def resolve_credential(ctx: click.Context, api_key: str | None) -> tuple[ClaudeC
return ApiKeyHelper(resolve_api_key_helper(base_url)), stored
def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, tuple[str, ...]]:
@dataclass(frozen=True, slots=True)
class _Listing:
models: tuple[ListedModel, ...]
@property
def ids(self) -> tuple[str, ...]:
return tuple(model.id for model in self.models)
def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, _Listing]:
"""Every configure path begins the same way: the local ownership check first, so a `lite up`
session is refused before any login prompt or request, then the credential, then the listing."""
settings_path: Final = claude_settings_path(os.environ)
@ -88,21 +106,28 @@ def _listing_error(base_url: str, error: PiSyncError) -> str:
return f"{error.message} The proxy at {base_url} answered, so check that it is a LiteLLM proxy and is healthy."
def _listed_models(base_url: str, key: str) -> tuple[str, ...]:
listed: Final = fetch_model_ids(base_url, key)
def _listed_models(base_url: str, key: str) -> _Listing:
listed: Final = fetch_model_listing(base_url, key, headers=_CLAUDE_CODE_VIEW)
if isinstance(listed, PiSyncError):
raise click.ClickException(_listing_error(base_url, listed))
return listed
return _Listing(listed)
def _starting_model(model: str, listing: _Listing) -> str | None:
source: Final = next((listed.id for listed in listing.models if listed.source_model == model), None)
return source or next((listed.id for listed in listing.models if listed.id == model), None)
def _model_choice(model: str | None) -> ModelChoice:
return StartOn(model) if model is not None else UnpinModel()
def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listed: Sequence[str], model: str | None) -> None:
def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listing: _Listing, model: str | None) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
if model is not None and model not in listed:
listed: Final = listing.ids
starting: Final = _starting_model(model, listing) if model is not None else None
if model is not None and starting is None:
shown: Final = ", ".join(listed[:_LISTED_MODELS_SHOWN])
more: Final = f", and {len(listed) - _LISTED_MODELS_SHOWN} more" if len(listed) > _LISTED_MODELS_SHOWN else ""
raise click.ClickException(
@ -113,29 +138,32 @@ def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listed: Sequ
configure_claude_settings(
base_url,
credential,
_model_choice(model),
_model_choice(starting),
settings_path,
configure_state_path(settings_path),
settings_file_owners(settings_path),
)
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
in_picker: Final = sum(1 for listed_model in listed if _CLAUDE_CODE_PICKER_FILTER.search(listed_model))
in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.search(listed_model))
click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.")
click.echo(
"Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN."
if isinstance(credential, StaticToken)
else "Credential: your `lite login`, read through apiKeyHelper on every request, so a later login renews it."
)
click.echo(
f"Starting model: {model} ({STARTING_MODEL_ROLE}); switch any time with /model."
if model is not None
f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model."
if starting is not None
else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or "
"pass --model to start on a proxy model."
)
click.echo(
f"/model will list {in_picker} of the proxy's {len(listed)} models (Claude Code shows only ids containing "
"'claude' or 'anthropic')."
f"/model will list all {len(listed)} of the proxy's models."
if in_picker == len(listed)
else f"/model will list {in_picker} of the proxy's {len(listed)} models: Claude Code shows only ids containing "
"'claude' or 'anthropic', and this proxy does not list the rest under such names."
)
click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.")
if isinstance(credential, StaticToken) and settings_path.is_symlink():
@ -173,8 +201,10 @@ def interactive_configure(
targets: Final = pick_targets()
if _CLAUDE_TARGET not in targets:
return
credential, listed = _start(ctx, None)
_apply_claude(ctx, credential, listed, pick_model(listed))
credential, listing = _start(ctx, None)
_apply_claude(
ctx, credential, listing, pick_model(tuple(model.source_model or model.id for model in listing.models))
)
@click.group(name="configure", invoke_without_command=True)
@ -218,8 +248,8 @@ def configure_claude(ctx: click.Context, api_key: str | None, model: str | None)
setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back.
Assumes the proxy is already running.
"""
credential, listed = _start(ctx, api_key)
_apply_claude(ctx, credential, listed, model)
credential, listing = _start(ctx, api_key)
_apply_claude(ctx, credential, listing, model)
@unconfigure_group.command(name="claude")

View file

@ -13,10 +13,11 @@ from dataclasses import dataclass
from enum import StrEnum
from pathlib import Path
from types import MappingProxyType
from typing import Final
from typing import Annotated, Final
import requests
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, model_validator
from pydantic.types import StringConstraints
PI_CONFIG_DIR_ENV: Final = "PI_CODING_AGENT_DIR"
PI_PROVIDER_NAME: Final = "litellm"
@ -51,12 +52,25 @@ class ModelLimits:
max_tokens: int | None
class _Model(BaseModel):
id: str
_NonEmptyString = Annotated[str, StringConstraints(min_length=1)]
class ListedModel(BaseModel):
model_config = ConfigDict(frozen=True)
id: _NonEmptyString
source_model: _NonEmptyString | None = None
class _ModelList(BaseModel):
data: tuple[_Model, ...]
data: tuple[ListedModel, ...]
@model_validator(mode="after")
def unique_id_mappings(self) -> "_ModelList":
mappings: Final = frozenset((model.id, model.source_model or model.id) for model in self.data)
if len(frozenset(model.id for model in self.data)) != len(mappings):
raise ValueError("model ids must not map to multiple source models")
return self
class _ModelGroup(BaseModel):
@ -69,17 +83,18 @@ class _ModelGroupList(BaseModel):
data: tuple[_ModelGroup, ...]
def fetch_model_ids(
def fetch_model_listing(
base_url: str,
api_key: str,
*,
get: Callable[..., requests.Response] = requests.get,
) -> tuple[str, ...] | PiSyncError:
headers: Mapping[str, str] = MappingProxyType({}),
) -> tuple[ListedModel, ...] | PiSyncError:
url: Final = base_url.rstrip("/") + "/v1/models"
try:
resp: Final = get(
url,
headers={"Authorization": f"Bearer {api_key}"}, # mutable-ok: requests headers require a dict
headers={"Authorization": f"Bearer {api_key}", **headers}, # mutable-ok: requests headers require a dict
timeout=10,
)
except requests.RequestException as e:
@ -94,10 +109,21 @@ def fetch_model_ids(
listing: Final = _ModelList.model_validate(resp.json())
except (ValueError, ValidationError) as e:
return PiSyncError(f"Unexpected /v1/models response from the proxy: {e}", kind=ListingFailure.BAD_BODY)
ids: Final = tuple(dict.fromkeys(model.id for model in listing.data))
if not ids:
models: Final = tuple(dict.fromkeys(listing.data))
if not models:
return PiSyncError("The proxy returned no models for your key.", kind=ListingFailure.EMPTY)
return ids
return models
def fetch_model_ids(
base_url: str,
api_key: str,
*,
get: Callable[..., requests.Response] = requests.get,
headers: Mapping[str, str] = MappingProxyType({}),
) -> tuple[str, ...] | PiSyncError:
listed: Final = fetch_model_listing(base_url, api_key, get=get, headers=headers)
return listed if isinstance(listed, PiSyncError) else tuple(dict.fromkeys(model.id for model in listed))
_NO_LIMITS: Final[Mapping[str, ModelLimits]] = MappingProxyType({})
@ -222,11 +248,13 @@ __all__ = (
"LITELLM_PROXY_API_KEY_ENV",
"PI_CONFIG_DIR_ENV",
"PI_PROVIDER_NAME",
"ListedModel",
"ListingFailure",
"ModelLimits",
"PiSyncError",
"fetch_model_ids",
"fetch_model_limits",
"fetch_model_listing",
"models_json_path",
"provider_block",
"sync_models_json",

View file

@ -39,7 +39,7 @@ def _is_form_content_type(content_type: str) -> bool:
return _normalize_media_type(content_type) in _FORM_CONTENT_TYPES
def _is_json_content_type(content_type: str) -> bool:
def is_json_content_type(content_type: str) -> bool:
"""True iff the body should be parsed as JSON."""
return _normalize_media_type(content_type) == "application/json"
@ -406,7 +406,7 @@ async def get_request_body(request: Request) -> dict[str, Any]:
"""
if request.method == "POST":
content_type: Final = request.headers.get("content-type", "")
if _is_json_content_type(content_type):
if is_json_content_type(content_type):
return await _read_request_body(request)
elif _is_form_content_type(content_type):
return await get_form_data(request)

View file

@ -10,12 +10,24 @@ legacy internal names with `general_settings.use_team_public_model_name: false`.
from __future__ import annotations
from collections.abc import Mapping, Sequence
import re
from collections.abc import Container, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, cast
import litellm
if TYPE_CHECKING:
from litellm.router import Router
from litellm.types.proxy.model_listing import ModelInfoResponse
CLAUDE_CODE_PICKER_PATTERN: Final = re.compile(r"claude|anthropic", re.IGNORECASE)
GATEWAY_CLIENT_HEADER: Final = "x-gateway-client"
CLAUDE_CODE_CLIENT: Final = "claude-code"
_CLAUDE_CODE_ALIAS_PREFIX: Final = "claude-router-"
_ONE_MILLION_SUFFIX: Final = "[1m]"
_ONE_MILLION_TOKENS: Final = 1_000_000
def configured_display_names(
@ -40,6 +52,115 @@ def configured_display_names(
)
def _unmarked(name: str) -> str:
return name[: -len(_ONE_MILLION_SUFFIX)] if name.lower().endswith(_ONE_MILLION_SUFFIX) else name
def _compatibility_id(model_id: str) -> str:
return f"{_CLAUDE_CODE_ALIAS_PREFIX}{model_id.encode().hex()}"
def _decoded_compatibility_id(view_id: str) -> str | None:
encoded: Final = _unmarked(view_id).removeprefix(_CLAUDE_CODE_ALIAS_PREFIX)
if encoded == _unmarked(view_id):
return None
try:
model_id: Final = bytes.fromhex(encoded).decode()
except (ValueError, UnicodeDecodeError):
return None
return model_id if _compatibility_id(model_id) == _unmarked(view_id) else None
def claude_code_model_id(
model_id: str,
max_input_tokens: float | None,
routing_names: Container[str],
) -> str:
"""The collision-free id Claude Code's picker lists a model under."""
if "*" in model_id:
return model_id
shaped: Final = model_id if CLAUDE_CODE_PICKER_PATTERN.search(model_id) else _compatibility_id(model_id)
one_million: Final = max_input_tokens is not None and max_input_tokens >= _ONE_MILLION_TOKENS
marked: Final = (
f"{shaped}{_ONE_MILLION_SUFFIX}" if one_million and not shaped.lower().endswith(_ONE_MILLION_SUFFIX) else shaped
)
return next(
(
name
for name in (marked, shaped)
if name == model_id or claude_code_group_name(name, routing_names) == model_id
),
model_id,
)
def claude_code_group_name(view_id: str, routing_names: Container[str]) -> str | None:
"""Decode a canonical compatibility id only when no configured route claims it."""
if view_id in routing_names:
return None
unmarked: Final = _unmarked(view_id)
if unmarked != view_id and unmarked in routing_names:
return unmarked
model_id: Final = _decoded_compatibility_id(view_id)
return model_id if model_id and model_id in routing_names else None
def is_claude_code_client(headers: Mapping[str, str]) -> bool:
"""Claude Code itself, or a client asking for its view of the listing the way Ramp Router's does"""
from litellm.llms.anthropic.common_utils import is_claude_code_user_agent
return (
is_claude_code_user_agent(headers.get("user-agent", ""))
or headers.get(GATEWAY_CLIENT_HEADER, "").lower() == CLAUDE_CODE_CLIENT
)
def claude_code_view_ids(
rows: Sequence[ModelInfoResponse],
headers: Mapping[str, str],
routing_names: Container[str],
) -> Mapping[str, str]:
"""served id -> Claude Code id for the requested listing view"""
if not is_claude_code_client(headers):
return MappingProxyType({})
return MappingProxyType(
{row["id"]: claude_code_model_id(row["id"], row.get("max_input_tokens"), routing_names) for row in rows}
)
@dataclass(frozen=True, slots=True)
class ClaudeCodeRoutingNames:
"""Existing routes always own their names, including aliases and wildcard routes."""
llm_router: Router | None
team_id: str | None = None
alias_maps: tuple[object, ...] = ()
def __contains__(self, name: object) -> bool:
if not isinstance(name, str):
return False
if name in litellm.model_alias_map or any(
isinstance(aliases, Mapping) and name in aliases for aliases in self.alias_maps
):
return True
if self.llm_router is None:
return False
return (
name in self.llm_router.model_group_alias
or self.llm_router.has_model_id(name)
or bool(self.llm_router.get_candidate_model_ids_for_route(name, self.team_id))
)
def claude_code_requested_group(
requested: str,
llm_router: Router,
team_id: str | None,
alias_maps: tuple[object, ...] = (),
) -> str | None:
return claude_code_group_name(requested, ClaudeCodeRoutingNames(llm_router, team_id, alias_maps))
class TeamModelNameTranslator:
"""Translates internal team routing keys to their public names for the model
listing/retrieve responses. Stateless; the live router and general_settings

View file

@ -0,0 +1,23 @@
import asyncio
from typing import Final
from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY
class LoopBoundSemaphore:
__slots__ = ("_loop", "_semaphore", "_value")
def __init__(self, value: int) -> None:
self._value: Final = value
self._loop: asyncio.AbstractEventLoop | None = None
self._semaphore: asyncio.Semaphore | None = None
def current(self) -> asyncio.Semaphore:
loop: Final = asyncio.get_running_loop()
if self._semaphore is None or self._loop is not loop:
self._semaphore = asyncio.Semaphore(self._value)
self._loop = loop
return self._semaphore
db_lookup_gate: Final = LoopBoundSemaphore(PROXY_DB_LOOKUP_MAX_CONCURRENCY)

View file

@ -23,6 +23,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import Litellm_EntityType
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.table_repositories import (
BudgetWindowSpendRepository,
@ -121,30 +122,33 @@ class SpendCounterReseed:
if SpendCounterReseed._is_key_or_team_window_counter(counter_key):
return None
try:
if counter_key.startswith("spend:key:"):
token: Final = counter_key[len("spend:key:") :]
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token})
elif counter_key.startswith("spend:team_member:"):
suffix: Final = counter_key[len("spend:team_member:") :]
if ":" not in suffix:
async with db_lookup_gate.current():
if counter_key.startswith("spend:key:"):
token: Final = counter_key[len("spend:key:") :]
row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token})
elif counter_key.startswith("spend:team_member:"):
suffix: Final = counter_key[len("spend:team_member:") :]
if ":" not in suffix:
return None
user_id, team_id = suffix.rsplit(":", 1)
row = await TeamMembershipRepository(prisma_client).table.find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}
)
elif counter_key.startswith("spend:team:"):
team_id = counter_key[len("spend:team:") :]
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
elif counter_key.startswith("spend:user:"):
user_id = counter_key[len("spend:user:") :]
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
return None
elif counter_key.startswith("spend:org:"):
org_id: Final = counter_key[len("spend:org:") :]
row = await OrganizationRepository(prisma_client).table.find_unique(
where={"organization_id": org_id}
)
else:
return None
user_id, team_id = suffix.rsplit(":", 1)
row = await TeamMembershipRepository(prisma_client).table.find_unique(
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}
)
elif counter_key.startswith("spend:team:"):
team_id = counter_key[len("spend:team:") :]
row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
elif counter_key.startswith("spend:user:"):
user_id = counter_key[len("spend:user:") :]
row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"):
return None
elif counter_key.startswith("spend:org:"):
org_id: Final = counter_key[len("spend:org:") :]
row = await OrganizationRepository(prisma_client).table.find_unique(where={"organization_id": org_id})
else:
return None
except Exception:
verbose_proxy_logger.exception("SpendCounterReseed.from_db: failed for %s", counter_key)
return None

View file

@ -210,6 +210,7 @@ services = (
"arize",
"galileo",
"newrelic",
"pointfive",
"sqs",
]
| str
@ -297,6 +298,7 @@ async def health_services_endpoint(
"arize",
"galileo",
"newrelic",
"pointfive",
"sqs",
]:
raise HTTPException(
@ -321,7 +323,7 @@ async def health_services_endpoint(
service == "openmeter"
or service == "braintrust"
or service == "generic_api"
or (service_in_success_callbacks and service != "langfuse")
or (service_in_success_callbacks and service not in ("langfuse", "pointfive"))
):
_ = await litellm.acompletion(
model="openai/litellm-mock-response-model",
@ -413,6 +415,27 @@ async def health_services_endpoint(
),
}
elif service == "pointfive":
if not _is_proxy_admin(user_api_key_dict):
non_admin_detail: Final[_ServiceTestErrorDetail] = {
"error": "Only proxy admins can trigger the PointFive liveness ping."
}
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=non_admin_detail)
from litellm.integrations.pointfive import PointFiveLogger
try:
pointfive_logger: Final = PointFiveLogger(start_periodic_flush=False)
except ValueError as missing_key:
# No key configured is the answer the operator asked for, not a server error.
no_key: Final[_ServiceTestSuccessResponse] = {"status": "unhealthy", "message": str(missing_key)}
return no_key
response = await pointfive_logger.async_health_check()
pointfive_health: Final[_ServiceTestSuccessResponse] = {
"status": response["status"],
"message": (response["error_message"] if response["status"] == "unhealthy" else "PointFive is healthy")
or "PointFive is healthy",
}
return pointfive_health
if service == "webhook":
user_info: Final = CallInfo(
token=user_api_key_dict.token or "",

View file

@ -111,6 +111,7 @@ from litellm.router_utils.auto_router_model_naming import (
validate_strategy_router_model_write,
)
from litellm.router_utils.auto_router_tuning_baseline import is_mutable_tuned_candidate, tuning_quota_violation
from litellm.types.llms.bedrock import AwsSessionTag
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
AutoRouterClassifierDefaultPromptResponse,
UpdateUsefulLinksRequest,
@ -897,6 +898,12 @@ async def patch_model(
existing_litellm_params=db_model.litellm_params,
)
ModelManagementAuthChecks.can_user_set_aws_session_tags(
litellm_params=patch_data.litellm_params,
user_api_key_dict=user_api_key_dict,
existing_litellm_params=db_model.litellm_params,
)
_raise_on_strategy_router_write_violation(
incoming_params=patch_data.litellm_params,
existing_params=db_model.litellm_params,
@ -1650,6 +1657,10 @@ async def _update_existing_team_model_assignment(
# No team_model_add/delete calls required; public name is already registered
def _canonical_session_tags(tags: Sequence[AwsSessionTag]) -> tuple[tuple[str, str], ...]:
return tuple(sorted((tag["Key"], tag["Value"]) for tag in tags))
class ModelManagementAuthChecks:
"""
Common auth checks for model management endpoints
@ -1704,6 +1715,28 @@ class ModelManagementAuthChecks:
param="litellm_credential_name",
)
@staticmethod
def can_user_set_aws_session_tags(
litellm_params: GenericLiteLLMParams | None,
user_api_key_dict: UserAPIKeyAuth,
existing_litellm_params: GenericLiteLLMParams | None = None,
) -> Literal[True]:
if litellm_params is None or litellm_params.aws_session_tags is None:
return True
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return True
existing_tags: Final = existing_litellm_params.aws_session_tags if existing_litellm_params is not None else None
if existing_tags is not None and _canonical_session_tags(existing_tags) == _canonical_session_tags(
litellm_params.aws_session_tags
):
return True
raise ProxyException(
message=f"Only a proxy admin can set aws_session_tags on a model. Your role={user_api_key_dict.user_role}.",
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,
param="aws_session_tags",
)
@staticmethod
async def allow_team_model_action(
model_params: Deployment | updateDeployment,
@ -2037,6 +2070,11 @@ async def add_new_model(
user_api_key_dict=user_api_key_dict,
)
ModelManagementAuthChecks.can_user_set_aws_session_tags(
litellm_params=model_params.litellm_params,
user_api_key_dict=user_api_key_dict,
)
_raise_on_strategy_router_write_violation(
incoming_params=model_params.litellm_params,
existing_params=None,
@ -2221,6 +2259,12 @@ async def update_model(
existing_litellm_params=deployment.litellm_params,
)
ModelManagementAuthChecks.can_user_set_aws_session_tags(
litellm_params=model_params.litellm_params,
user_api_key_dict=user_api_key_dict,
existing_litellm_params=deployment.litellm_params,
)
_raise_on_strategy_router_write_violation(
incoming_params=model_params.litellm_params,
existing_params=deployment.litellm_params,

View file

@ -34,6 +34,7 @@ from litellm.constants import (
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
@ -53,6 +54,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_set_request_parsed_body,
get_form_data,
get_request_body,
is_json_content_type,
)
from litellm.proxy.common_utils.sse_keepalive import (
wrap_passthrough_sse_bytes_with_keepalive_pings,
@ -78,6 +80,7 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
from litellm.types.router import LiteLLMParamsTypedDict
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
@ -120,6 +123,24 @@ def is_passthrough_request_using_router_model(request_body: dict, llm_router: li
return False
class RelayRejection(TypedDict):
error: ReadOnly[str]
def _deployment_model_name(litellm_params: LiteLLMParamsTypedDict) -> str:
model: Final = litellm_params.get("model", "")
try:
return get_llm_provider(model=model, custom_llm_provider=litellm_params.get("custom_llm_provider"))[0]
except litellm.BadRequestError:
return model
def _models_served_by_group(llm_router: litellm.Router, model_group: str) -> frozenset[str]:
return frozenset(
_deployment_model_name(row["litellm_params"]) for row in llm_router.get_model_list(model_name=model_group) or ()
)
def is_passthrough_request_streaming(request_body: object) -> bool:
"""
Returns True if the request is streaming.
@ -412,7 +433,7 @@ async def vllm_proxy_route(
content=None,
data=None,
files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None),
json=(request_body if is_json_content_type(request.headers.get("content-type", "")) else None),
params=None,
headers=None,
cookies=None,
@ -1499,6 +1520,14 @@ async def _relay_upstream_bytes(upstream: AsyncGenerator[bytes, bytes]) -> Async
await upstream.aclose()
async def _relay_upstream_response(upstream: httpx.Response) -> Response:
return Response(
content=await upstream.aread(),
status_code=upstream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(headers=upstream.headers, custom_headers=None),
)
async def _relay_azure_router_model(
llm_router: litellm.Router,
model: str,
@ -1508,30 +1537,37 @@ async def _relay_azure_router_model(
is_streaming_request: bool,
user_api_key_dict: UserAPIKeyAuth,
) -> Response:
result: Final = await llm_router.allm_passthrough_route(
model=model,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=is_streaming_request,
content=None,
data=None,
files=None,
json=(request_body if request.headers.get("content-type") == "application/json" else None),
params=None,
headers=None,
cookies=None,
litellm_metadata=get_passthrough_router_request_metadata(user_api_key_dict),
foreign_deployment: Final = foreign_azure_deployment(
endpoint, model, lambda: _models_served_by_group(llm_router, model)
)
if foreign_deployment is not None:
rejection: Final[RelayRejection] = {
"error": f"deployment '{foreign_deployment}' in the path is not served by model group '{model}'; "
"put the model group name in the deployments segment"
}
raise HTTPException(status_code=400, detail=rejection)
try:
result: Final = await llm_router.allm_passthrough_route(
model=model,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=_safe_get_request_headers(request),
stream=is_streaming_request,
content=None,
data=None,
files=None,
json=(request_body if is_json_content_type(request.headers.get("content-type", "")) else None),
params=None,
headers=None,
cookies=None,
litellm_metadata=get_passthrough_router_request_metadata(user_api_key_dict),
)
except httpx.HTTPStatusError as upstream_error:
return await _relay_upstream_response(upstream_error.response)
if not is_streaming_request:
upstream: Final = cast(httpx.Response, result)
return Response(
content=await upstream.aread(),
status_code=upstream.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(headers=upstream.headers, custom_headers=None),
)
return await _relay_upstream_response(cast(httpx.Response, result))
if inspect.isasyncgen(result):
sse_headers: Final = {"content-type": "text/event-stream"}

View file

@ -4,6 +4,7 @@ OpenAI Passthrough Logging Handler
Handles cost tracking and logging for OpenAI passthrough endpoints, specifically /chat/completions.
"""
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Final
from urllib.parse import urlparse
@ -16,6 +17,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
)
from litellm.litellm_core_utils.token_counter import high_detail_image_token_upper_bound
from litellm.llms.openai.openai import OpenAIConfig
from litellm.llms.openai.openai import OpenAIConfig as OpenAIConfigType
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
@ -96,6 +98,47 @@ def _is_openai_compatible_url(url_route: str | None) -> bool:
return False
def _is_remote_high_detail_image(part: object) -> bool:
if not isinstance(part, Mapping) or part.get("type") != "image_url":
return False
image_url: Final = part.get("image_url")
if not isinstance(image_url, Mapping):
return False
url: Final = image_url.get("url")
return (
isinstance(url, str) and url.lower().startswith(("http://", "https://")) and image_url.get("detail") == "high"
)
def _content_parts(message: Mapping[str, object]) -> Sequence[object]:
content: Final = message.get("content")
return content if isinstance(content, list) else ()
def _without_remote_high_detail_images(message: Mapping[str, object]) -> Mapping[str, object]:
if not isinstance(message.get("content"), list):
return message
kept_parts: Final = [ # mutable-ok: token_counter reads message content only when it is a list
part for part in _content_parts(message) if not _is_remote_high_detail_image(part)
]
return {**message, "content": kept_parts} # mutable-ok: token_counter rejects any message that is not a dict
def count_relayed_prompt_tokens(model: str, messages: Sequence[Mapping[str, object]] | None) -> int:
if messages is None:
return 0
remote_high_detail_images: Final = sum(
1 for message in messages for part in _content_parts(message) if _is_remote_high_detail_image(part)
)
local_messages: Final = [ # mutable-ok: token_counter takes a list of messages
_without_remote_high_detail_images(message) for message in messages
]
return (
litellm.token_counter(model=model, messages=local_messages)
+ high_detail_image_token_upper_bound() * remote_high_detail_images
)
class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""
OpenAI-specific passthrough logging handler that provides cost tracking for /chat/completions endpoints.
@ -512,9 +555,10 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
def _build_complete_streaming_response(
self,
all_chunks: list[str],
all_chunks: Sequence[str],
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
messages: Sequence[Mapping[str, object]] | None = None,
) -> ModelResponse | TextCompletionResponse | None:
"""
Builds complete response from raw chunks for OpenAI streaming responses.
@ -558,7 +602,11 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
return None
# Build complete response from chunks
complete_streaming_response: Final = litellm.stream_chunk_builder(chunks=all_openai_chunks)
complete_streaming_response: Final = litellm.stream_chunk_builder(
chunks=all_openai_chunks,
messages=messages,
count_prompt_tokens=lambda: count_relayed_prompt_tokens(model, messages),
)
return complete_streaming_response

View file

@ -279,7 +279,7 @@ from litellm.constants import (
WEEKLY_SPEND_REPORT_JOB_ID,
)
from litellm.exceptions import RejectedRequestError
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.integrations.custom_guardrail import CustomGuardrail, ModifyResponseException
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.litellm_core_utils.agentic_loop_settings import (
@ -377,8 +377,11 @@ from litellm.proxy.common_utils.load_config_utils import (
)
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
from litellm.proxy.common_utils.model_listing_utils import (
ClaudeCodeRoutingNames,
TeamModelNameTranslator,
claude_code_view_ids,
configured_display_names,
is_claude_code_client,
)
from litellm.proxy.common_utils.openai_endpoint_utils import (
remove_sensitive_info_from_deployment,
@ -6690,6 +6693,14 @@ class ProxyConfig:
return parsed
return None
async def get_hierarchical_router_settings(
self,
user_api_key_dict: UserAPIKeyAuth | None,
prisma_client: PrismaClient | None,
proxy_logging_obj: ProxyLogging | None = None,
) -> dict | None:
return await self._get_hierarchical_router_settings(user_api_key_dict, prisma_client, proxy_logging_obj)
async def _get_hierarchical_router_settings(
self,
user_api_key_dict: Optional["UserAPIKeyAuth"],
@ -8683,6 +8694,7 @@ _STREAM_KEEPALIVE: Final = object()
_KEEPALIVE_MIN_SECONDS: Final = 1.0
_KEEPALIVE_MAX_SECONDS: Final = 300.0
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_EMPTY_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
async def _iter_with_keepalive(
@ -10520,6 +10532,24 @@ async def model_list(
wants_anthropic_format: Final = (
http_request is not None and http_request.headers.get("anthropic-version") is not None
)
client_headers: Final[Mapping[str, str]] = http_request.headers if http_request is not None else _EMPTY_HEADERS
view_router_settings: Final = (
await proxy_config.get_hierarchical_router_settings(user_api_key_dict, prisma_client, proxy_logging_obj)
if wants_anthropic_format and is_claude_code_client(client_headers)
else None
)
view_aliases: Final = (
view_router_settings.get("model_group_alias") if isinstance(view_router_settings, Mapping) else None
)
routing_names: Final = ClaudeCodeRoutingNames(
llm_router,
team_id or user_api_key_dict.team_id,
(
user_api_key_dict.aliases,
user_api_key_dict.team_model_aliases,
view_aliases,
),
)
# Validate scope parameter if provided
if scope is not None and scope != "expand":
@ -10606,6 +10636,11 @@ async def model_list(
return create_anthropic_model_list_response(
admin_listing,
display_names=configured_display_names(admin_entries, llm_router),
listed_ids=claude_code_view_ids(
admin_listing,
client_headers,
routing_names,
),
)
return dict(
@ -10654,6 +10689,11 @@ async def model_list(
return create_anthropic_model_list_response(
listing,
display_names=configured_display_names(entries, llm_router),
listed_ids=claude_code_view_ids(
listing,
client_headers,
routing_names,
),
)
return dict(
@ -17706,6 +17746,66 @@ async def delete_callback(
)
def _normalize_callback_alias(callback_name: str) -> str:
callback_aliases: Final = (
("opentelemetry", "otel"),
("s3_v2", "s3"),
("aws_sqs", "sqs"),
("custom_callback_api", "generic_api"),
)
return next(
(canonical_name for alias, canonical_name in callback_aliases if alias == callback_name),
callback_name,
)
def _callback_module_name(callback: CustomLogger | Callable[..., object]) -> str:
if inspect.ismethod(callback):
return callback.__func__.__module__
if inspect.isfunction(callback):
return callback.__module__
return type(callback).__module__
def _is_litellm_internal_callback(callback_name: str, callback: CustomLogger | Callable[..., object]) -> bool:
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
module_owner: Final = _callback_module_name(callback).partition(".")[0]
is_registered_integration: Final = callback_name in CustomLoggerRegistry.CALLBACK_CLASS_STR_TO_CLASS_TYPE
return not is_registered_integration and module_owner in ("litellm", "litellm_enterprise")
def _is_instance_of_configured_callback(
callback_name: str, callback: CustomLogger | Callable[..., object], configured_classes: tuple[type, ...]
) -> bool:
"""Self-naming OTel-family instances (`arize`, `weave_otel`) match by name, so a configured `logfire` (a bare
`OpenTelemetry`) does not hide YAML-configured siblings of the same class."""
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
class_derived_name: Final = CustomLoggerRegistry.get_callback_str_from_class_type(type(callback))
return isinstance(callback, configured_classes) and callback_name in (class_derived_name, type(callback).__name__)
def _hidden_runtime_callback_names(configured_callback_names: frozenset[str]) -> frozenset[str]:
from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry
configured_classes: Final = tuple(
CustomLoggerRegistry.CALLBACK_CLASS_STR_TO_CLASS_TYPE[name]
for name in configured_callback_names
if name in CustomLoggerRegistry.CALLBACK_CLASS_STR_TO_CLASS_TYPE
)
configured_modules: Final = frozenset(name.rsplit(".", 1)[0] for name in configured_callback_names if "." in name)
internal_callback_names: Final = frozenset({"cache", "vector_store_pre_call_hook"})
return internal_callback_names | frozenset(
callback_name
for callback_name, callback in litellm.logging_callback_manager.get_callback_objects()
if isinstance(callback, CustomGuardrail)
or _is_litellm_internal_callback(callback_name, callback)
or _is_instance_of_configured_callback(callback_name, callback, configured_classes)
or _callback_module_name(callback) in configured_modules
)
@router.get(
"/get/config/callbacks",
tags=["config.yaml"],
@ -17738,10 +17838,10 @@ async def get_config(
# Normalize string callbacks to lists
def normalize_callback(callback):
if isinstance(callback, str):
return [callback]
elif callback is None:
return []
return callback
return (callback,)
if callback is None:
return ()
return tuple(callback) if isinstance(callback, (list, dict)) else ()
_success_callbacks = normalize_callback(_success_callbacks)
_failure_callbacks = normalize_callback(_failure_callbacks)
@ -17772,6 +17872,30 @@ async def get_config(
for _callback in _success_and_failure_callbacks:
_data_to_return.append(process_callback(_callback, "success_and_failure", environment_variables))
configured_callback_names: Final = frozenset(
_normalize_callback_alias(callback)
for callback in (_success_callbacks + _failure_callbacks + _success_and_failure_callbacks)
)
runtime_callbacks_by_type: Final = litellm.logging_callback_manager.get_callbacks_by_type()
hidden_callback_names: Final = _hidden_runtime_callback_names(configured_callback_names)
runtime_callback_rows: Final = tuple(
(_normalize_callback_alias(callback_name), callback_type)
for callback_type, callback_names in (
("success", runtime_callbacks_by_type["success"]),
("failure", runtime_callbacks_by_type["failure"]),
("success_and_failure", runtime_callbacks_by_type["success_and_failure"]),
)
for callback_name in callback_names
if callback_name not in hidden_callback_names
)
runtime_only_rows: Final = sorted(
frozenset(row for row in runtime_callback_rows if row[0] not in configured_callback_names)
)
_data_to_return.extend(
dict(process_callback(callback_name, callback_type, environment_variables), read_only=True)
for callback_name, callback_type in runtime_only_rows
)
_data_to_return = _apply_callback_role_gate(_data_to_return, is_full_admin)
# Check if slack alerting is on

View file

@ -99,6 +99,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
mask_sensitive_structure,
)
from litellm.litellm_core_utils.token_counter import offload_token_count
from litellm.llms.base_llm.passthrough.transformation import replace_path_segment
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
vector_store_request_metadata,
@ -150,6 +151,7 @@ from litellm.router_utils.common_utils import (
filter_team_based_models,
filter_web_search_deployments,
get_request_team_id,
provider_for_generic_call,
resolve_model_group_alias,
truncate_fallback_error_detail,
warn_on_provider_credential_mismatch,
@ -5199,7 +5201,7 @@ class Router:
# If get_llm_provider fails, fall back to using model_name as-is
replacement_model_name = model_name
kwargs["endpoint"] = kwargs["endpoint"].replace(model, replacement_model_name)
kwargs["endpoint"] = replace_path_segment(kwargs["endpoint"], model, replacement_model_name)
return kwargs
async def _ageneric_api_call_with_fallbacks_helper(self, model: str, original_generic_function: Callable, **kwargs):
@ -5235,16 +5237,7 @@ class Router:
kwargs=kwargs, model=model, model_name=model_name
)
# Get custom_llm_provider from deployment params
try:
custom_llm_provider = data.get("custom_llm_provider")
_, inferred_custom_llm_provider, _, _ = get_llm_provider(
model=data["model"],
custom_llm_provider=custom_llm_provider,
)
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
except Exception:
custom_llm_provider = None
custom_llm_provider: Final = provider_for_generic_call(data)
response_kwargs: Final = {
**data,
@ -5755,15 +5748,7 @@ class Router:
# Perform pre-call checks for routing strategy
self.routing_strategy_pre_call_checks(deployment=deployment)
try:
custom_llm_provider = data.get("custom_llm_provider")
_, inferred_custom_llm_provider, _, _ = get_llm_provider(
model=data["model"],
custom_llm_provider=custom_llm_provider,
)
custom_llm_provider = custom_llm_provider or inferred_custom_llm_provider
except Exception:
custom_llm_provider = None
custom_llm_provider: Final = provider_for_generic_call(data)
response: Final = original_function(
**{

View file

@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Final
if TYPE_CHECKING:
from litellm.types.llms.openai import OpenAIFileObject
import litellm
from litellm._logging import verbose_logger, verbose_router_logger
from litellm.constants import ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS
from litellm.exceptions import BadRequestError
@ -256,6 +257,32 @@ PROVIDER_SCOPED_CREDENTIAL_PARAMS: Final[Mapping[str, frozenset[str]]] = Mapping
)
def provider_for_generic_call(litellm_params: Mapping[str, object]) -> str | None:
"""
The provider the router hands a deployment's generic SDK call, or None when it cannot be resolved.
A model that carries its own provider prefix keeps that prefix even where get_llm_provider
would resolve it to a sibling provider (azure_ai/<openai model> on an Azure OpenAI host
resolves to azure): the SDK call still receives the prefixed model, and an explicit provider
that contradicts the prefix makes get_llm_provider re-prefix it into a deployment name that
does not exist upstream.
"""
declared: Final = litellm_params.get("custom_llm_provider")
if isinstance(declared, str) and declared:
return declared
model: Final = litellm_params.get("model")
if not isinstance(model, str) or not model:
return None
prefix: Final = model.split("/", 1)[0]
if "/" in model and prefix in litellm.provider_list:
return prefix
try:
_, inferred, _, _ = get_llm_provider(model=model)
except BadRequestError:
return None
return inferred
def warn_on_provider_credential_mismatch(model_name: str, litellm_params: Mapping[str, object]) -> str | None:
"""
Warn when a deployment carries one provider's credentials but resolves to another.

View file

@ -0,0 +1,45 @@
from dataclasses import dataclass
from typing import Final
from pydantic import Field
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
RETRYABLE_UPLOAD_STATUS_CODES: Final = frozenset({429, 500, 502, 503, 504})
DEFAULT_API_URL: Final = "https://api.pointfive.co/api/v1/ingestion"
class PointFiveInitParams(StandardCustomLoggerInitParams):
"""
Params for initializing a PointFive logger on litellm.
Defaults trade freshness for fewer, larger uploads: every flush becomes one object, so
the interval is minutes rather than seconds. ``batch_size`` also bounds how much a busy
proxy holds in memory between flushes, so it stays modest. ``max_batch_bytes`` bounds
how much a single object may hold, which matters most when message logging is left on,
since an unredacted payload is orders of magnitude larger than a redacted one.
"""
api_key: str | None = None
api_url: str | None = None
batch_size: int = Field(default=1_000, gt=0)
flush_interval: int = Field(default=300, gt=0)
max_batch_bytes: int = Field(default=8 * 1024 * 1024, gt=0)
max_upload_retries: int = Field(default=3, ge=1)
@dataclass(frozen=True, slots=True)
class PointFiveUploadTarget:
"""A single-use presigned destination for one batch, issued by the PointFive API."""
upload_url: str
object_key: str
@dataclass(frozen=True, slots=True)
class PointFiveUploadFailure:
"""Why a batch could not be uploaded, and whether a later attempt could still succeed."""
detail: str
retryable: bool

View file

@ -154,6 +154,22 @@ LATENCY_BUCKETS: Final = (
float("inf"),
)
UNKNOWN_INPUT_SEQUENCE_LENGTH: Final = "unknown"
INPUT_SEQUENCE_LENGTH_BUCKETS: Final = (
(1_000, "0-1k"),
(4_000, "1k-4k"),
(16_000, "4k-16k"),
(64_000, "16k-64k"),
(float("inf"), "64k+"),
)
def get_input_sequence_length_bucket(prompt_tokens: object) -> str:
if not isinstance(prompt_tokens, int) or isinstance(prompt_tokens, bool) or prompt_tokens < 0:
return UNKNOWN_INPUT_SEQUENCE_LENGTH
return next(label for upper, label in INPUT_SEQUENCE_LENGTH_BUCKETS if prompt_tokens < upper)
# Batch jobs can run for minutes to hours; buckets span 1 min → 24 h.
BATCH_DURATION_BUCKETS: Final = (
60.0,
@ -205,6 +221,7 @@ class UserAPIKeyLabelNames(Enum):
MCP_TOOL_NAME = "mcp_tool_name"
MCP_SERVER_NAME = "mcp_server_name"
SERVICE_TIER = "service_tier"
INPUT_SEQUENCE_LENGTH = "input_sequence_length"
DEFINED_PROMETHEUS_METRICS = Literal[
@ -857,6 +874,13 @@ class PrometheusMetricLabels:
"litellm_images_generated_metric",
}
)
_input_sequence_length_metrics: ClassVar[frozenset[str]] = frozenset(
{
"litellm_llm_api_latency_metric",
"litellm_llm_api_time_to_first_token_metric",
"litellm_request_total_latency_metric",
}
)
# Managed batch metrics
_batch_user_labels = [
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
@ -955,14 +979,23 @@ class PrometheusMetricLabels:
custom_labels.append(label)
if label_name in PrometheusMetricLabels._org_label_metrics:
for label in [
for label in (
UserAPIKeyLabelNames.ORG_ID.value,
UserAPIKeyLabelNames.ORG_ALIAS.value,
]:
):
if label not in default_labels and label not in custom_labels:
custom_labels.append(label)
return default_labels + custom_labels
input_sequence_length_labels: Final = (
(UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value,)
if (
label_name in PrometheusMetricLabels._input_sequence_length_metrics
and litellm.prometheus_emit_input_sequence_length_label is True
and UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in custom_labels
)
else ()
)
return [*default_labels, *custom_labels, *input_sequence_length_labels]
_USER_API_KEY_LABEL_VALUE_INIT_ALIASES: Final[Mapping[str, str]] = MappingProxyType(
@ -1015,6 +1048,7 @@ class UserAPIKeyLabelValues:
mcp_tool_name: str | None = None
mcp_server_name: str | None = None
service_tier: str | None = None
input_sequence_length: str | None = None
# Added for test compatibility.
def __init__(self, **kwargs: Any) -> None:

View file

@ -1107,6 +1107,11 @@ class BedrockTag(TypedDict):
value: str
class AwsSessionTag(TypedDict):
Key: str # writable-ok: boto3's STS stubs type assume_role Tags as writable TagTypeDef, which rejects ReadOnly
Value: str # writable-ok: boto3's STS stubs type assume_role Tags as writable TagTypeDef, which rejects ReadOnly
class BedrockCreateBatchRequest(TypedDict, total=False):
"""
Request structure for creating a Bedrock batch inference job.

View file

@ -1564,6 +1564,9 @@ class ResponseIncompleteEvent(BaseLiteLLMOpenAIResponseObject):
response: ResponsesAPIResponse
ResponsesTerminalEvent: TypeAlias = ResponseCompletedEvent | ResponseIncompleteEvent | ResponseFailedEvent
class ResponsePartAddedEvent(BaseLiteLLMOpenAIResponseObject):
type: Literal[ResponsesAPIStreamEvents.RESPONSE_PART_ADDED]
item_id: str

View file

@ -250,7 +250,23 @@ class MCPServer(BaseModel):
@property
def advertises_gateway_authorization_server(self) -> bool:
"""Whether named discovery should advertise the aggregate gateway authorization server."""
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
if self.auth_type == MCPAuth.oauth2:
return self.is_gateway_managed_oauth2 and not self.uses_per_server_oauth_relay
if self.auth_type not in (
None,
MCPAuth.none,
MCPAuth.api_key,
MCPAuth.bearer_token,
MCPAuth.basic,
MCPAuth.authorization,
MCPAuth.token,
MCPAuth.aws_sigv4,
):
return False
return not any(
header.lower() in ("authorization", "x-api-key", "api-key", "apikey")
for header in (self.extra_headers or ())
)
@property
def is_true_passthrough(self) -> bool:

View file

@ -4,7 +4,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc
import datetime
import enum
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
@ -21,6 +21,7 @@ if TYPE_CHECKING:
from .completion import CompletionRequest
from .embedding import EmbeddingRequest
from .llms.bedrock import AwsSessionTag
from .llms.openai import OpenAIFileObject
from .search import SearchProvider
from .utils import (
@ -288,6 +289,7 @@ class CredentialLiteLLMParams(BaseModel):
aws_web_identity_token: str | None = None
aws_sts_endpoint: str | None = None
aws_external_id: str | None = None
aws_session_tags: Sequence[AwsSessionTag] | None = None
aws_bedrock_runtime_endpoint: str | None = None
aws_bedrock_project_id: str | None = None
s3_bucket_name: str | None = None

View file

@ -8850,6 +8850,12 @@ class ProviderConfigManager:
)
return AzurePassthroughConfig()
elif LlmProviders.AZURE_AI == provider:
from litellm.llms.azure_ai.passthrough.transformation import (
AzureAIPassthroughConfig,
)
return AzureAIPassthroughConfig()
elif LlmProviders.GIGACHAT == provider:
from litellm.llms.gigachat.passthrough.transformation import (
GigaChatPassthroughConfig,

View file

@ -30295,7 +30295,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@ -30340,7 +30340,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.1-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@ -30432,7 +30432,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.2-2025-12-11": {
"cache_read_input_token_cost": 1.75e-07,
@ -30478,7 +30478,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
"supports_minimal_reasoning_effort": false
},
"gpt-5.2-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@ -31474,7 +31474,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 2.5e-06,
"output_cost_per_token_above_272k_tokens_flex": 1.125e-05,
"cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07
@ -31526,7 +31526,7 @@
"supports_none_reasoning_effort": true,
"default_reasoning_effort": "none",
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 2.5e-06,
"output_cost_per_token_above_272k_tokens_flex": 1.125e-05,
"cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07
@ -31578,7 +31578,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 3e-05,
"output_cost_per_token_above_272k_tokens_flex": 0.000135
},
@ -31629,7 +31629,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"input_cost_per_token_above_272k_tokens_flex": 3e-05,
"output_cost_per_token_above_272k_tokens_flex": 0.000135
},
@ -33709,13 +33709,14 @@
"supports_tool_choice": true
},
"jina-reranker-v2-base-multilingual": {
"input_cost_per_token": 1.8e-08,
"input_cost_per_token": 5e-08,
"litellm_provider": "jina_ai",
"max_input_tokens": 1024,
"max_output_tokens": 1024,
"max_tokens": 1024,
"mode": "rerank",
"output_cost_per_token": 1.8e-08
"output_cost_per_token": 0.0,
"source": "https://api.jina.ai/v1/models"
},
"jp.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
@ -39969,6 +39970,45 @@
"supports_tool_choice": true,
"supports_vision": true
},
"openrouter/openai/gpt-5.6-sol": {
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
"default_reasoning_effort": "medium",
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
"litellm_provider": "openrouter",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"reasoning_effort_levels": [
"none",
"low",
"medium",
"high",
"xhigh",
"max"
],
"source": "https://openrouter.ai/openai/gpt-5.6-sol",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"openrouter/openai/gpt-oss-120b": {
"input_cost_per_token": 3.7e-08,
"litellm_provider": "openrouter",

View file

@ -327,6 +327,7 @@ func resourceKeyUpdate(ctx context.Context, d *schema.ResourceData, m interface{
key.Metadata = metadata
if _, err := c.UpdateKey(key); err != nil {
d.Partial(true)
return diag.FromErr(fmt.Errorf("error updating key: %s", err))
}

View file

@ -194,6 +194,7 @@ class DummyCredentials:
("aws_web_identity_token", "dummy_web_identity_token"),
("aws_sts_endpoint", "dummy_sts_endpoint"),
("aws_external_id", "dummy_external_id"),
("aws_session_tags", [{"Key": "team", "Value": "genai"}]),
],
)
def test_dynamic_aws_params_propagation(model, param_name, param_value):

View file

@ -2920,7 +2920,7 @@ async def test_get_config_callbacks_with_all_types(client_no_auth):
assert result["status"] == "success"
assert "callbacks" in result
callbacks = result["callbacks"]
callbacks = [cb for cb in result["callbacks"] if not cb.get("read_only", False)]
# Verify we have all 5 callbacks (2 success + 1 failure + 2 success_and_failure)
assert len(callbacks) == 5

View file

@ -12,6 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import anyio
import httpx
import pytest
import respx
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
from mcp import McpError
from mcp.client.streamable_http import streamable_http_client
@ -1686,3 +1687,211 @@ async def test_empty_http_event_stream_uses_the_existing_request_deadline() -> N
timeout=3,
)
assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError)
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("prompts/list", "resources/list", "resources/templates/list"))
@pytest.mark.parametrize(
"outcome",
(
"absent",
"other_capability",
"supported",
"method_not_found",
"internal_error",
"unauthorized",
"timeout",
"initialize_not_found",
),
)
async def test_optional_discovery_capabilities_and_errors(
method: str, outcome: str, caplog: pytest.LogCaptureFixture
) -> None:
import logging
from unittest.mock import Mock
from mcp.types import JSONRPCRequest
capability: Final = "prompts" if method == "prompts/list" else "resources"
field: Final = {
"prompts/list": "prompts",
"resources/list": "resources",
"resources/templates/list": "resourceTemplates",
}[method]
advertised: Final = "resources" if capability == "prompts" else "prompts"
entry: Final = {
"prompts/list": {"name": "example"},
"resources/list": {"name": "example", "uri": "test://example"},
"resources/templates/list": {"name": "example", "uriTemplate": "test://{name}"},
}[method]
def respond(request: httpx.Request) -> httpx.Response:
if request.method == "DELETE":
return httpx.Response(200)
payload: Final = JSONRPCMessage.model_validate_json(request.content).root
if not isinstance(payload, JSONRPCRequest):
return httpx.Response(202)
if outcome == "initialize_not_found":
return httpx.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"error": {"code": -32601, "message": "Initialization rejected"},
},
)
if payload.method == "initialize":
return httpx.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": LATEST_PROTOCOL_VERSION,
"capabilities": {}
if outcome == "absent"
else {advertised if outcome == "other_capability" else capability: {}},
"serverInfo": {"name": "discovery", "version": "1"},
},
},
)
if outcome == "timeout":
raise httpx.ReadTimeout("Optional list timed out", request=request)
if outcome == "unauthorized":
return httpx.Response(401)
if outcome in ("method_not_found", "internal_error", "absent", "other_capability"):
return httpx.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"error": {
"code": -32603 if outcome == "internal_error" else -32601,
"message": "Optional list rejected",
},
},
)
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": {field: [entry]}})
responder: Final = Mock(side_effect=respond)
caplog.set_level(logging.DEBUG, logger="LiteLLM")
with respx.mock(base_url="https://example.com") as router:
router.route().mock(side_effect=responder)
client: Final = MCPClient(server_url="https://example.com/mcp")
operation: Final = {
"prompts/list": client.list_prompts,
"resources/list": client.list_resources,
"resources/templates/list": client.list_resource_templates,
}[method]
result: Final = await operation()
requests: Final = tuple(
JSONRPCMessage.model_validate_json(call.args[0].content).root
for call in responder.call_args_list
if call.args[0].method == "POST"
)
assert sum(isinstance(request, JSONRPCRequest) and request.method == method for request in requests) == (
0 if outcome in ("absent", "other_capability", "initialize_not_found") else 1
)
assert [item.name for item in result] == (["example"] if outcome == "supported" else [])
failures: Final = tuple(
record for record in caplog.records if record.name == "LiteLLM" and record.levelno >= logging.WARNING
)
if outcome in ("internal_error", "unauthorized", "timeout", "initialize_not_found"):
assert any(record.levelno == logging.ERROR and "failed" in record.message for record in failures)
else:
assert failures == ()
if outcome == "method_not_found":
assert any(
record.levelno == logging.DEBUG and "Optional list rejected" in record.message for record in caplog.records
)
@pytest.mark.asyncio
@pytest.mark.parametrize("supports_first", (True, False))
async def test_optional_discovery_uses_each_sessions_capabilities(supports_first: bool) -> None:
from unittest.mock import Mock
from mcp.types import JSONRPCRequest
capabilities: Final = iter(({"resources": {}}, {}) if supports_first else ({}, {"resources": {}}))
def respond(request: httpx.Request) -> httpx.Response:
if request.method == "DELETE":
return httpx.Response(200)
payload: Final = JSONRPCMessage.model_validate_json(request.content).root
if not isinstance(payload, JSONRPCRequest):
return httpx.Response(202)
result: Final = (
{
"protocolVersion": LATEST_PROTOCOL_VERSION,
"capabilities": next(capabilities),
"serverInfo": {"name": "changing", "version": "1"},
}
if payload.method == "initialize"
else {"resources": [{"name": "example", "uri": "test://example"}]}
)
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result})
responder: Final = Mock(side_effect=respond)
with respx.mock(base_url="https://example.com") as router:
router.route().mock(side_effect=responder)
client: Final = MCPClient(server_url="https://example.com/mcp")
first: Final = await client.list_resources()
second: Final = await client.list_resources()
assert [item.name for item in first] == (["example"] if supports_first else [])
assert [item.name for item in second] == ([] if supports_first else ["example"])
requests: Final = tuple(
JSONRPCMessage.model_validate_json(call.args[0].content).root
for call in responder.call_args_list
if call.args[0].method == "POST"
)
assert sum(isinstance(request, JSONRPCRequest) and request.method == "resources/list" for request in requests) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("prompts/list", "resources/list", "resources/templates/list"))
async def test_optional_discovery_preserves_cancellation(method: str) -> None:
from mcp.types import JSONRPCRequest
ready: Final = asyncio.Event()
pending: Final = asyncio.Event()
async def respond(request: httpx.Request) -> httpx.Response:
if request.method == "DELETE":
return httpx.Response(200)
payload: Final = JSONRPCMessage.model_validate_json(request.content).root
if not isinstance(payload, JSONRPCRequest):
return httpx.Response(202)
if payload.method == "initialize":
return httpx.Response(
200,
json={
"jsonrpc": "2.0",
"id": payload.id,
"result": {
"protocolVersion": LATEST_PROTOCOL_VERSION,
"capabilities": {"resources": {}, "prompts": {}},
"serverInfo": {"name": "pending", "version": "1"},
},
},
)
ready.set()
await pending.wait()
return httpx.Response(202)
with respx.mock(base_url="https://example.com") as router:
router.route().mock(side_effect=respond)
client: Final = MCPClient(server_url="https://example.com/mcp")
operation: Final = {
"prompts/list": client.list_prompts,
"resources/list": client.list_resources,
"resources/templates/list": client.list_resource_templates,
}[method]
task: Final = asyncio.create_task(operation())
try:
await asyncio.wait_for(ready.wait(), timeout=3)
finally:
task.cancel()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(task, timeout=3)

View file

@ -0,0 +1,759 @@
import asyncio
import gzip
import json
import logging
from collections.abc import Callable
import pytest
from litellm.integrations.pointfive.logger import PointFiveLogger
from litellm.integrations.pointfive.upload_client import PointFiveUploadError
from litellm.types.integrations.pointfive import DEFAULT_API_URL, PointFiveInitParams, PointFiveUploadFailure
OBJECT_KEY = "some/object.ndjson.gz"
class FakeUploadClient:
"""Records the objects a flush produced, so tests can read what would have shipped."""
def __init__(
self,
outcomes: list[str | PointFiveUploadFailure] | None = None,
ping_failure: PointFiveUploadFailure | None = None,
) -> None:
self.outcomes = outcomes or [OBJECT_KEY]
self.bodies: list[bytes] = []
self.on_upload: Callable[[], None] | None = None
self.ping_failure = ping_failure
self.pings = 0
async def ping(self) -> PointFiveUploadFailure | None:
self.pings += 1
return self.ping_failure
async def upload(self, body: bytes) -> str | PointFiveUploadFailure:
if self.on_upload is not None:
self.on_upload()
self.bodies.append(body)
return self.outcomes.pop(0) if len(self.outcomes) > 1 else self.outcomes[0]
def records(self) -> list[dict]:
return [json.loads(line) for body in self.bodies for line in gzip.decompress(body).decode().splitlines()]
def _logger(upload_client: FakeUploadClient, **params) -> PointFiveLogger:
return PointFiveLogger(params=PointFiveInitParams(**params), upload_client=upload_client)
def _event(request_id: str, size: int = 0) -> dict:
return {"standard_logging_object": {"id": request_id, "model": "gpt-4o", "blob": "x" * size}}
@pytest.mark.asyncio
async def test_a_flush_ships_one_object_holding_every_buffered_record():
"""One object per flush is the whole point: s3_v2 sends one per request."""
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=3)
for request_id in ("a", "b", "c"):
await logger.async_log_success_event(_event(request_id), None, None, None)
await _settle(logger)
assert len(upload_client.bodies) == 1
assert [record["id"] for record in upload_client.records()] == ["a", "b", "c"]
@pytest.mark.asyncio
async def test_records_are_held_until_the_batch_is_full():
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=3)
await logger.async_log_success_event(_event("a"), None, None, None)
assert upload_client.bodies == []
assert len(logger.log_queue) == 1
@pytest.mark.asyncio
async def test_failed_requests_are_logged_too():
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1)
await logger.async_log_failure_event(_event("failed"), None, None, None)
await _settle(logger)
assert [record["id"] for record in upload_client.records()] == ["failed"]
@pytest.mark.asyncio
async def test_an_event_without_a_standard_payload_is_skipped():
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1)
await logger.async_log_success_event({"kwargs": "but no payload"}, None, None, None)
assert upload_client.bodies == []
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_a_batch_over_the_byte_cap_ships_as_several_objects():
"""Record count cannot bound an object: an unredacted payload dwarfs a redacted one."""
cap = 600
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=4, max_batch_bytes=cap)
for request_id in ("a", "b", "c", "d"):
await logger.async_log_success_event(_event(request_id, size=200), None, None, None)
await _settle(logger)
assert len(upload_client.bodies) > 1
assert [record["id"] for record in upload_client.records()] == ["a", "b", "c", "d"]
assert all(len(gzip.decompress(body)) <= cap for body in upload_client.bodies)
@pytest.mark.asyncio
async def test_a_retryable_failure_keeps_the_batch_for_the_next_flush():
upload_client = FakeUploadClient([PointFiveUploadFailure("upload target is down", retryable=True)])
logger = _logger(upload_client, batch_size=2)
for request_id in ("a", "b"):
await logger.async_log_success_event(_event(request_id), None, None, None)
assert [record["id"] for record in logger.log_queue] == ["a", "b"]
@pytest.mark.asyncio
async def test_a_retryable_failure_surfaces_so_the_base_logger_can_preserve_it():
upload_client = FakeUploadClient([PointFiveUploadFailure("upload target is down", retryable=True)])
logger = _logger(upload_client, batch_size=99)
logger.log_queue.append(_event("a")["standard_logging_object"])
with pytest.raises(PointFiveUploadError, match="upload target is down"):
await logger.async_send_batch()
@pytest.mark.asyncio
async def test_a_rejected_batch_is_dropped_rather_than_blocking_the_queue():
"""Retrying a rejection forever would stall every record queued behind it."""
upload_client = FakeUploadClient([PointFiveUploadFailure("object too large", retryable=False)])
logger = _logger(upload_client, batch_size=2)
for request_id in ("a", "b"):
await logger.async_log_success_event(_event(request_id), None, None, None)
await _settle(logger)
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_records_already_queued_ship_with_the_event_that_triggers_the_flush():
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1)
logger.log_queue.append(_event("mid-flight")["standard_logging_object"])
await logger.async_log_success_event(_event("a"), None, None, None)
await _settle(logger)
assert [record["id"] for record in upload_client.records()] == ["mid-flight", "a"]
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_a_record_that_arrives_mid_flush_is_kept_for_the_next_one():
"""The queue is drained by count, so a record appended mid-upload must survive."""
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1)
upload_client.on_upload = lambda: logger.log_queue.append(_event("late")["standard_logging_object"])
await logger.async_log_success_event(_event("first"), None, None, None)
await _settle(logger)
assert [record["id"] for record in upload_client.records()] == ["first"]
assert [record["id"] for record in logger.log_queue] == ["late"]
def test_defaults_favour_fewer_larger_uploads_over_freshness():
upload_client = FakeUploadClient()
logger = _logger(upload_client)
assert logger.batch_size == 1_000
assert logger.flush_interval == 300
assert logger.max_batch_bytes == 8 * 1024 * 1024
def test_the_default_api_url_is_the_pointfive_ingress(monkeypatch):
"""api.pointfive.co is the host the ingress serves; .com does not resolve to it."""
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_env")
logger = PointFiveLogger()
assert logger.upload_client.api_url == "https://api.pointfive.co/api/v1/ingestion"
def test_the_api_key_can_come_from_the_environment(monkeypatch):
"""The proxy ui configures a callback by writing environment variables."""
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_from_env")
logger = PointFiveLogger()
assert logger.upload_client.api_key == "p5tu_from_env"
def test_the_api_url_can_come_from_the_environment(monkeypatch):
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_env")
monkeypatch.setenv("POINTFIVE_API_URL", "https://api.staging.pointfive.co/api/v1/ingestion")
logger = PointFiveLogger()
assert logger.upload_client.api_url == "https://api.staging.pointfive.co/api/v1/ingestion"
def test_config_yaml_wins_over_the_environment(monkeypatch):
"""A value set in config.yaml is explicit, so it outranks whatever the ui left behind."""
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_from_env")
monkeypatch.setenv("POINTFIVE_API_URL", "https://from-env.example/api/v1/ingestion")
logger = PointFiveLogger(
params=PointFiveInitParams(api_key="p5tu_from_config", api_url="https://from-config.example/api/v1/ingestion")
)
assert logger.upload_client.api_key == "p5tu_from_config"
assert logger.upload_client.api_url == "https://from-config.example/api/v1/ingestion"
def test_a_missing_api_key_fails_at_startup_not_at_the_first_flush(monkeypatch):
monkeypatch.delenv("POINTFIVE_API_KEY", raising=False)
with pytest.raises(ValueError, match="api key"):
PointFiveLogger(params=PointFiveInitParams())
def test_an_api_key_can_be_an_environment_reference(monkeypatch):
"""config.yaml spells secrets as `os.environ/NAME`, so the plugin must resolve one."""
monkeypatch.setenv("POINTFIVE_TEST_KEY", "p5tu_from_env")
logger = PointFiveLogger(params=PointFiveInitParams(api_key="os.environ/POINTFIVE_TEST_KEY"))
assert logger.upload_client.api_key == "p5tu_from_env"
def test_params_are_read_from_litellm_settings(monkeypatch):
import litellm
monkeypatch.setattr(litellm, "pointfive_params", {"api_key": "p5tu_configured", "batch_size": 7})
logger = PointFiveLogger()
assert logger.upload_client.api_key == "p5tu_configured"
assert logger.batch_size == 7
def test_an_out_of_range_setting_is_rejected():
with pytest.raises(ValueError, match="batch_size"):
PointFiveInitParams(api_key="p5tu_k", batch_size=0)
@pytest.mark.asyncio
async def test_an_idle_flush_reports_liveness_instead_of_uploading():
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=99)
await logger.flush_queue()
assert upload_client.pings == 1
assert upload_client.bodies == []
@pytest.mark.asyncio
async def test_a_flush_with_records_uploads_and_does_not_ping():
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1)
await logger.async_log_success_event(_event("a"), None, None, None)
await _settle(logger)
assert upload_client.pings == 0
assert len(upload_client.bodies) == 1
@pytest.mark.asyncio
async def test_a_failed_ping_does_not_raise():
"""Liveness is bookkeeping; a proxy must not see errors from it."""
upload_client = FakeUploadClient(ping_failure=PointFiveUploadFailure("api down", retryable=True))
logger = _logger(upload_client, batch_size=99)
await logger.flush_queue()
assert upload_client.pings == 1
@pytest.mark.asyncio
async def test_health_check_is_healthy_when_the_api_accepts_the_key():
upload_client = FakeUploadClient()
assert await _logger(upload_client).async_health_check() == {"status": "healthy", "error_message": None}
assert upload_client.pings == 1
@pytest.mark.asyncio
async def test_health_check_reports_why_the_api_refused():
"""The ui test button shows this message, so a rejected key has to say so rather than pass."""
upload_client = FakeUploadClient(ping_failure=PointFiveUploadFailure("key was revoked", retryable=False))
outcome = await _logger(upload_client).async_health_check()
assert outcome == {"status": "unhealthy", "error_message": "key was revoked"}
def test_the_client_follows_a_key_and_url_changed_after_startup(monkeypatch):
"""
The proxy ui writes new values into a running proxy's environment.
Reading them once at construction would leave the logger talking to the old endpoint
until someone restarted the proxy.
"""
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_first")
monkeypatch.setenv("POINTFIVE_API_URL", "https://first.example.invalid/api/v1/ingestion")
logger = PointFiveLogger(params=PointFiveInitParams())
assert logger.upload_client.api_key == "p5tu_first"
assert logger.upload_client.api_url == "https://first.example.invalid/api/v1/ingestion"
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_second")
monkeypatch.setenv("POINTFIVE_API_URL", "https://second.example.invalid/api/v1/ingestion")
assert logger.upload_client.api_key == "p5tu_second"
assert logger.upload_client.api_url == "https://second.example.invalid/api/v1/ingestion"
def test_a_configured_key_still_wins_over_the_environment(monkeypatch):
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_from_env")
logger = PointFiveLogger(params=PointFiveInitParams(api_key="p5tu_from_config"))
assert logger.upload_client.api_key == "p5tu_from_config"
@pytest.mark.asyncio
async def test_health_check_says_so_when_the_key_was_removed(monkeypatch):
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_present")
logger = PointFiveLogger(params=PointFiveInitParams())
monkeypatch.delenv("POINTFIVE_API_KEY")
outcome = await logger.async_health_check()
assert outcome["status"] == "unhealthy"
assert "requires an api key" in (outcome["error_message"] or "")
def _pending_flush_tasks() -> tuple[asyncio.Task, ...]:
return tuple(task for task in asyncio.all_tasks() if "periodic_flush" in str(task.get_coro()))
@pytest.mark.asyncio
async def test_a_one_shot_logger_leaves_no_flush_task_behind():
"""
A health check builds a logger for a single answer and drops it.
Without this, every check would leave a flusher running that keeps pinging for the
lifetime of the proxy.
"""
before = _pending_flush_tasks()
logger = PointFiveLogger(params=PointFiveInitParams(), upload_client=FakeUploadClient(), start_periodic_flush=False)
assert logger._periodic_flush_task is None
assert _pending_flush_tasks() == before
@pytest.mark.asyncio
async def test_the_logger_flushes_periodically_by_default():
logger = PointFiveLogger(params=PointFiveInitParams(), upload_client=FakeUploadClient())
assert logger._periodic_flush_task is not None
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_params_already_built_are_used_as_they_are(monkeypatch):
"""config.yaml is validated once into a params object; a second validation would be wasted."""
import litellm
monkeypatch.setattr(litellm, "pointfive_params", PointFiveInitParams(max_batch_bytes=4096))
logger = PointFiveLogger(upload_client=FakeUploadClient())
assert logger.max_batch_bytes == 4096
@pytest.mark.asyncio
async def test_a_dead_flush_task_is_restarted_by_the_next_event():
"""A cancelled or crashed flusher would otherwise leave the queue growing forever."""
logger = _logger(FakeUploadClient())
logger._periodic_flush_task.cancel()
await asyncio.sleep(0) # let the cancellation land, so the task reports itself done
await logger.async_log_success_event(_event("after-cancel"), None, None, None)
assert logger._periodic_flush_task is not None
assert not logger._periodic_flush_task.done()
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_a_failure_while_queueing_never_breaks_the_request():
"""Logging sits on the request path, so a fault here must not surface to the caller."""
class ExplodingQueue(list):
def append(self, _item):
raise RuntimeError("queue is broken")
upload_client = FakeUploadClient()
logger = _logger(upload_client)
logger.log_queue = ExplodingQueue()
await logger.async_log_success_event(_event("boom"), None, None, None)
logger.log_queue = []
await logger.async_log_success_event(_event("after-the-fault"), None, None, None)
assert [record["id"] for record in logger.log_queue] == ["after-the-fault"]
@pytest.mark.asyncio
async def test_a_flush_with_nothing_queued_uploads_nothing():
upload_client = FakeUploadClient()
logger = _logger(upload_client)
await logger.async_send_batch()
assert upload_client.bodies == []
@pytest.mark.asyncio
async def test_the_idle_ping_is_skipped_when_the_key_was_removed(monkeypatch, caplog):
"""A key pulled mid-flight must not turn the periodic flush into an exception."""
monkeypatch.setenv("POINTFIVE_API_KEY", "p5tu_present")
logger = PointFiveLogger(params=PointFiveInitParams())
monkeypatch.delenv("POINTFIVE_API_KEY")
with caplog.at_level(logging.WARNING):
await logger.flush_queue()
assert "liveness ping skipped" in caplog.text
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_a_full_batch_stands_down_while_a_flush_is_already_running():
"""
Under load every event landing mid-upload also crosses the batch threshold.
Letting each one flush turns a single burst into a stream of tiny objects, which is
what batching exists to avoid, so a full batch defers to the flush already running.
"""
upload_client = FakeUploadClient()
# No periodic task: this test drives the flushes itself, and the loop's opening cycle
# would otherwise ship the queue it seeds below.
logger = PointFiveLogger(
params=PointFiveInitParams(batch_size=2),
upload_client=upload_client,
start_periodic_flush=False,
)
release = asyncio.Event()
finish_upload = upload_client.upload
async def held_upload(body: bytes):
await release.wait()
return await finish_upload(body)
upload_client.upload = held_upload
logger.log_queue.extend(_event(f"first-{index}")["standard_logging_object"] for index in range(2))
flushing = asyncio.create_task(logger.flush_queue())
await asyncio.sleep(0)
# Bounded: without the guard these block on the flush lock the held upload owns.
for index in range(6):
await asyncio.wait_for(logger.async_log_success_event(_event(f"mid-{index}"), None, None, None), timeout=2)
assert upload_client.bodies == []
release.set()
await flushing
assert len(upload_client.bodies) == 1
assert [record["id"] for record in upload_client.records()] == ["first-0", "first-1"]
assert [record["id"] for record in logger.log_queue] == [f"mid-{index}" for index in range(6)]
async def _settle(logger) -> None:
"""
Wait out the flush a full batch schedules, the way the proxy's loop would.
The upload runs off the request path now, and either the batch task or the periodic
loop can be the one carrying it, so this waits for whichever is in flight to finish.
"""
for _ in range(200):
await asyncio.sleep(0.001)
task = logger._batch_flush_task
if task is not None and not task.done():
await task
if not logger._flushing:
return
raise AssertionError("the flush never finished")
async def _until(done: Callable[[], bool], ticks: int = 400) -> None:
"""Wait for a condition the flush path reaches only after gzip finishes on a worker thread."""
for _ in range(ticks):
if done():
return
await asyncio.sleep(0.005)
raise AssertionError("condition never became true")
@pytest.mark.asyncio
async def test_a_new_logger_announces_itself_without_waiting_for_the_interval():
"""
Configuring the callback must make the integration connect, with no traffic and no test click.
The inherited loop sleeps a whole interval before its first flush, which left a freshly
configured proxy silent for five minutes and the integration looking unconfigured.
"""
upload_client = FakeUploadClient()
logger = _logger(upload_client)
await asyncio.sleep(0) # let the flush task reach its first cycle
assert upload_client.pings == 1
assert upload_client.bodies == []
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_the_first_cycle_ships_records_rather_than_announcing():
"""Announcing is only for an empty queue: records already waiting must go out as an upload."""
upload_client = FakeUploadClient()
logger = PointFiveLogger(
params=PointFiveInitParams(batch_size=100),
upload_client=upload_client,
start_periodic_flush=False,
)
logger.log_queue.append(_event("queued-before-start")["standard_logging_object"])
logger._periodic_flush_task = logger._start_periodic_flush_task()
await _until(lambda: bool(upload_client.bodies))
assert upload_client.pings == 0
assert [record["id"] for record in upload_client.records()] == ["queued-before-start"]
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_a_failed_request_ships_redacted_when_message_logging_is_off():
"""
Failure events skip the framework's redaction, so the callback has to redact what it buffers.
Without this the prompt of every failed request reaches PointFive in full, even though
the integration was configured not to send message content.
"""
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1, turn_off_message_logging=True)
event = _event("failed-request")
event["standard_logging_object"]["messages"] = [{"role": "user", "content": "my secret prompt"}]
event["standard_logging_object"]["response"] = "the secret answer"
await logger.async_log_failure_event(event, None, None, None)
await _settle(logger)
shipped = upload_client.records()[0]
assert "my secret prompt" not in json.dumps(shipped)
assert "the secret answer" not in json.dumps(shipped)
assert event["standard_logging_object"]["messages"][0]["content"] == "my secret prompt"
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_a_dead_loop_does_not_strand_the_flusher():
"""A task whose loop was closed never runs and never reports done, so it must be replaced."""
logger = PointFiveLogger(params=PointFiveInitParams(), upload_client=FakeUploadClient(), start_periodic_flush=False)
stranded_loop = asyncio.new_event_loop()
forever = asyncio.sleep(3600)
logger._periodic_flush_task = stranded_loop.create_task(forever)
stranded_loop.close()
forever.close()
await logger.async_log_success_event(_event("after-loop-close"), None, None, None)
assert logger._periodic_flush_task is not None
assert logger._periodic_flush_task.get_loop() is asyncio.get_running_loop()
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_a_retry_does_not_resend_objects_that_already_landed():
"""
A failure part way through a multi-object flush used to hand the whole batch back.
Every record already shipped, and every record already refused for good, went out
again on the next flush, so PointFive received duplicates of both.
"""
upload_client = FakeUploadClient(outcomes=[OBJECT_KEY, PointFiveUploadFailure("service busy", retryable=True)])
logger = PointFiveLogger(
params=PointFiveInitParams(max_batch_bytes=1), # one record per object
upload_client=upload_client,
start_periodic_flush=False,
)
logger.log_queue.extend(_event(request_id)["standard_logging_object"] for request_id in ("first", "second"))
with pytest.raises(PointFiveUploadError):
await logger.async_send_batch()
assert [record["id"] for record in logger.log_queue] == ["second"]
@pytest.mark.asyncio
async def test_the_queue_stops_growing_at_its_cap_without_waiting_for_a_failure():
"""The base class trims only after a failed send, so a proxy that keeps flushing never trims."""
logger = _logger(FakeUploadClient(), batch_size=10_000)
logger.max_queue_size = 3
for request_id in ("a", "b", "c", "d", "e"):
await logger.async_log_success_event(_event(request_id), None, None, None)
assert [record["id"] for record in logger.log_queue] == ["c", "d", "e"]
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_a_failed_request_honours_the_global_redaction_setting(monkeypatch):
"""
Redaction can be turned on globally or per request, not only on this callback.
The async failure path hands the payload over untouched, so a tenant could trigger a
provider failure and ship prompts that the operator had already asked to be redacted.
"""
import litellm
monkeypatch.setattr(litellm, "turn_off_message_logging", True)
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1)
event = _event("globally-redacted")
event["standard_logging_object"]["messages"] = [{"role": "user", "content": "my secret prompt"}]
await logger.async_log_failure_event(event, None, None, None)
await _settle(logger)
assert "my secret prompt" not in json.dumps(upload_client.records()[0])
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_excluded_fields_are_dropped_from_a_failed_request(monkeypatch):
"""standard_logging_payload_excluded_fields drops a field entirely; failures skipped it too."""
import litellm
monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["messages"])
upload_client = FakeUploadClient()
logger = _logger(upload_client, batch_size=1, turn_off_message_logging=True)
event = _event("field-excluded")
event["standard_logging_object"]["messages"] = [{"role": "user", "content": "my secret prompt"}]
await logger.async_log_failure_event(event, None, None, None)
await _settle(logger)
shipped = upload_client.records()[0]
assert "messages" not in shipped
assert shipped["id"] == "field-excluded"
logger._periodic_flush_task.cancel()
def _held_upload(upload_client: FakeUploadClient, release: asyncio.Event) -> None:
finish = upload_client.upload
async def held(body: bytes):
await release.wait()
return await finish(body)
upload_client.upload = held
@pytest.mark.asyncio
async def test_a_full_batch_does_not_hold_the_request():
"""
The upload belongs off the request path.
Awaiting it inline meant a hung PointFive api held the caller's response open for as
long as the attempts and their backoff took.
"""
upload_client = FakeUploadClient()
release = asyncio.Event()
_held_upload(upload_client, release)
logger = _logger(upload_client, batch_size=1)
await asyncio.wait_for(logger.async_log_success_event(_event("first"), None, None, None), timeout=2)
assert upload_client.bodies == []
release.set()
await _settle(logger)
assert [record["id"] for record in upload_client.records()] == ["first"]
logger._periodic_flush_task.cancel()
@pytest.mark.asyncio
async def test_records_arriving_during_a_flush_survive_the_queue_cap():
"""
The flush drains by count, so trimming the front underneath it loses records.
Records that arrived while the upload was in flight would be deleted by that drain
without ever being sent.
"""
upload_client = FakeUploadClient()
release = asyncio.Event()
_held_upload(upload_client, release)
logger = PointFiveLogger(
params=PointFiveInitParams(batch_size=2),
upload_client=upload_client,
start_periodic_flush=False,
)
logger.max_queue_size = 2
logger.log_queue.extend(_event(request_id)["standard_logging_object"] for request_id in ("a", "b"))
flushing = asyncio.create_task(logger.flush_queue())
await asyncio.sleep(0.01)
for request_id in ("c", "d", "e"):
await logger.async_log_success_event(_event(request_id), None, None, None)
release.set()
await flushing
assert [record["id"] for record in upload_client.records()] == ["a", "b"]
assert [record["id"] for record in logger.log_queue] == ["c", "d", "e"]
logger._periodic_flush_task.cancel()
def test_an_unset_env_reference_is_never_used_as_the_key(monkeypatch):
"""
A config that names a missing variable has no key, and must say so.
Falling back to the reference text sent the literal "os.environ/NAME" as the bearer
token, so the callback started and every upload was rejected for the wrong reason.
"""
monkeypatch.delenv("POINTFIVE_API_KEY", raising=False)
monkeypatch.delenv("POINTFIVE_MISSING_KEY", raising=False)
with pytest.raises(ValueError, match="requires an api key"):
PointFiveLogger(
params=PointFiveInitParams(api_key="os.environ/POINTFIVE_MISSING_KEY"),
start_periodic_flush=False,
)
def test_an_unset_url_reference_falls_back_to_the_public_endpoint(monkeypatch):
"""An unresolved url reference must not become the destination the proxy uploads to."""
from litellm.integrations.pointfive.logger import _resolved_api_url
monkeypatch.delenv("POINTFIVE_API_URL", raising=False)
monkeypatch.delenv("POINTFIVE_MISSING_URL", raising=False)
assert _resolved_api_url(PointFiveInitParams(api_url="os.environ/POINTFIVE_MISSING_URL")) == DEFAULT_API_URL

View file

@ -0,0 +1,80 @@
import gzip
import json
import pytest
from litellm.integrations.pointfive.payload import chunk_lines, encode_lines, serialize_records
UNBOUNDED = 10_000_000
def test_each_record_becomes_one_json_line():
lines = serialize_records([{"id": "a"}, {"id": "b"}, {"id": "c"}])
assert len(lines) == 3
assert [json.loads(line)["id"] for line in lines] == ["a", "b", "c"]
def test_non_serializable_values_do_not_raise():
"""An odd payload must not kill the flush."""
lines = serialize_records([{"id": "a", "when": object()}])
assert json.loads(lines[0])["id"] == "a"
def test_records_that_fit_stay_in_one_object():
lines = serialize_records([{"id": f"r{i}"} for i in range(50)])
assert chunk_lines(lines, UNBOUNDED) == (lines,)
def test_objects_are_capped_by_uncompressed_size():
lines = serialize_records([{"id": f"r{i}", "blob": "x" * 100} for i in range(10)])
line_bytes = len(lines[0].encode("utf-8")) + 1
chunks = chunk_lines(lines, line_bytes * 3)
assert [len(chunk) for chunk in chunks] == [3, 3, 3, 1]
def test_oversized_single_record_is_sent_alone_not_stalled():
"""A record too big for the cap must still go out, or it blocks everything behind it."""
lines = serialize_records([{"id": "small"}, {"id": "huge", "blob": "x" * 5000}, {"id": "small2"}])
chunks = chunk_lines(lines, 200)
assert sum(len(chunk) for chunk in chunks) == 3
huge = [chunk for chunk in chunks if any("huge" in line for line in chunk)]
assert len(huge) == 1
assert len(huge[0]) == 1
def test_no_records_produces_no_objects():
assert chunk_lines((), UNBOUNDED) == ()
def test_every_record_appears_exactly_once():
lines = serialize_records([{"id": f"r{i}"} for i in range(37)])
chunks = chunk_lines(lines, len(lines[0]) * 4)
assert [line for chunk in chunks for line in chunk] == list(lines)
@pytest.mark.asyncio
async def test_encode_lines_round_trips_through_gzip():
lines = serialize_records([{"id": f"r{i}"} for i in range(5)])
encoded = await encode_lines(lines)
assert encoded[:2] == b"\x1f\x8b"
assert gzip.decompress(encoded).decode("utf-8") == "\n".join(lines)
@pytest.mark.asyncio
async def test_encode_lines_compresses_repetitive_records():
lines = serialize_records([{"id": f"r{i}", "model": "gpt-4o", "cost": 0.01} for i in range(200)])
encoded = await encode_lines(lines)
assert len(encoded) < len(gzip.decompress(encoded)) / 2

View file

@ -0,0 +1,378 @@
import json
from collections.abc import Sequence
import httpx
import pytest
import litellm
from litellm.integrations.pointfive.upload_client import PointFiveUploadClient
from litellm.litellm_core_utils.url_utils import validate_url
from litellm.types.integrations.pointfive import PointFiveUploadFailure
API_URL = "https://api.pointfive.co/api/v1/ingestion"
UPLOAD_URL = "https://uploads.example.invalid/some/object.ndjson.gz?signature=sig"
OBJECT_KEY = "some/object.ndjson.gz"
BODY = b"gzipped-bytes"
def _presigned(status_code: int = 200) -> httpx.Response:
return _response(
status_code, {"uploadUrl": UPLOAD_URL, "objectKey": OBJECT_KEY, "expiresAt": "2026-08-25T14:35:00Z"}
)
def _response(status_code: int, payload: object) -> httpx.Response:
return httpx.Response(status_code, text=json.dumps(payload))
def _refused(status_code: int, error: str) -> httpx.Response:
"""The body PointFive sends with every refusal."""
return _response(status_code, {"success": False, "error": error})
def _accepted() -> httpx.Response:
return httpx.Response(200, text="")
def _no_content() -> httpx.Response:
return httpx.Response(204, text="")
class FakeHTTPClient:
"""
Stands in for AsyncHTTPHandler, including its habit of raising on error statuses.
Scripted results are consumed in order, and the last one repeats, so a test that
cares about a single behaviour passes a single result.
"""
def __init__(
self,
presign: Sequence[httpx.Response | Exception] | None = None,
put: Sequence[httpx.Response | Exception] | None = None,
) -> None:
self.presign = list(presign) if presign else [_presigned()] # mutable-ok: results are consumed by popping
self.put_results = list(put) if put else [_accepted()] # mutable-ok: results are consumed by popping
self.presign_calls: list[dict] = []
self.put_calls: list[dict] = []
async def post(self, url, json=None, headers=None, **_):
self.presign_calls.append({"url": url, "json": json, "headers": headers or {}})
return _next_result(self.presign, url)
async def put(self, url, data=None, headers=None, follow_redirects=None, **_):
self.put_calls.append(
{"url": url, "data": data, "headers": headers or {}, "follow_redirects": follow_redirects}
)
return _next_result(self.put_results, url)
def _next_result(results: list, url: str) -> httpx.Response:
result = results.pop(0) if len(results) > 1 else results[0]
if isinstance(result, Exception):
raise result
if result.status_code >= 300:
request = httpx.Request("POST", url)
raise httpx.HTTPStatusError(
"boom",
request=request,
response=httpx.Response(result.status_code, text=result.text, headers=result.headers),
)
return result
async def _no_backoff(_seconds: float) -> None:
return None
def _trusting_validator(url: str) -> tuple[str, str]:
"""Stands in for validate_url so the fixture hosts need no DNS; the SSRF tests use the real one."""
return url, httpx.URL(url).host
def _client(
http_client: FakeHTTPClient,
max_retries: int = 3,
api_url: str = API_URL,
validate_upload_url=_trusting_validator,
) -> PointFiveUploadClient:
return PointFiveUploadClient(
api_key="p5tu_testkey",
api_url=api_url,
http_client=http_client,
max_retries=max_retries,
sleep=_no_backoff,
validate_upload_url=validate_upload_url,
)
def _presigned_for(upload_url: str) -> httpx.Response:
return _response(200, {"uploadUrl": upload_url, "objectKey": OBJECT_KEY, "expiresAt": "2026-08-25T14:35:00Z"})
@pytest.mark.asyncio
async def test_uploads_the_body_to_the_url_the_api_returned():
http_client = FakeHTTPClient()
outcome = await _client(http_client).upload(BODY)
assert outcome == OBJECT_KEY
assert http_client.put_calls[0]["url"] == UPLOAD_URL
assert http_client.put_calls[0]["data"] == BODY
@pytest.mark.asyncio
async def test_presign_request_is_authenticated_and_sized():
http_client = FakeHTTPClient()
await _client(http_client).upload(BODY)
call = http_client.presign_calls[0]
assert call["url"] == "https://api.pointfive.co/api/v1/ingestion/upload-url"
assert call["headers"]["Authorization"] == "Bearer p5tu_testkey"
assert call["json"] == {"kind": "LITELLM", "byteCount": len(BODY)}
@pytest.mark.asyncio
async def test_a_trailing_slash_on_the_api_url_is_tolerated():
"""A pasted URL often ends in a slash; it must not produce a double slash in the path."""
http_client = FakeHTTPClient()
await _client(http_client, api_url=API_URL + "/").upload(BODY)
assert http_client.presign_calls[0]["url"] == "https://api.pointfive.co/api/v1/ingestion/upload-url"
@pytest.mark.asyncio
async def test_no_bearer_token_is_sent_to_the_presigned_url():
"""The URL carries its own authorization, so the api key must not travel with it."""
http_client = FakeHTTPClient()
await _client(http_client).upload(BODY)
assert "Authorization" not in http_client.put_calls[0]["headers"]
@pytest.mark.asyncio
async def test_the_upload_pins_the_host_and_never_follows_a_redirect():
http_client = FakeHTTPClient()
await _client(http_client).upload(BODY)
call = http_client.put_calls[0]
assert call["headers"]["Host"] == "uploads.example.invalid"
assert call["follow_redirects"] is False
@pytest.mark.asyncio
async def test_a_redirected_upload_is_refused_rather_than_followed():
"""A presigned URL never redirects legitimately; following one is how a bad endpoint reaches inside."""
http_client = FakeHTTPClient(put=[httpx.Response(301, headers={"location": "http://169.254.169.254/"})])
outcome = await _client(http_client).upload(BODY)
assert outcome == PointFiveUploadFailure(
"presigned upload redirected with 301, refusing to follow", retryable=False
)
assert len(http_client.put_calls) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize(
"upload_url",
[
"http://169.254.169.254/latest/meta-data",
"https://10.0.0.7/internal/bucket/object",
"http://127.0.0.1:9000/bucket/object",
],
)
async def test_an_upload_url_inside_the_network_is_refused_before_any_bytes_leave(upload_url):
http_client = FakeHTTPClient(presign=[_presigned_for(upload_url)])
outcome = await _client(http_client, validate_upload_url=validate_url).upload(BODY)
assert isinstance(outcome, PointFiveUploadFailure)
assert not outcome.retryable
assert outcome.detail.startswith("presigned upload url refused: ")
assert http_client.put_calls == []
@pytest.mark.asyncio
async def test_an_operator_can_switch_destination_validation_off(monkeypatch):
"""litellm.user_url_validation is the proxy-wide switch every SSRF guard honours."""
monkeypatch.setattr(litellm, "user_url_validation", False)
http_client = FakeHTTPClient(presign=[_presigned_for("http://10.0.0.7/bucket/object")])
outcome = await _client(http_client, validate_upload_url=validate_url).upload(BODY)
assert outcome == OBJECT_KEY
assert http_client.put_calls[0]["url"] == "http://10.0.0.7/bucket/object"
assert "Host" not in http_client.put_calls[0]["headers"]
@pytest.mark.asyncio
async def test_the_object_is_declared_as_gzipped_ndjson():
http_client = FakeHTTPClient()
await _client(http_client).upload(BODY)
assert http_client.put_calls[0]["headers"]["Content-Encoding"] == "gzip"
assert http_client.put_calls[0]["headers"]["Content-Type"] == "application/x-ndjson"
@pytest.mark.asyncio
async def test_each_retry_presigns_again():
"""A retry must never reuse a URL that was consumed or has expired."""
http_client = FakeHTTPClient(put=[httpx.Response(503), _accepted()])
outcome = await _client(http_client).upload(BODY)
assert outcome == OBJECT_KEY
assert len(http_client.presign_calls) == 2
assert len(http_client.put_calls) == 2
@pytest.mark.asyncio
async def test_retryable_upload_failure_gives_up_after_max_retries():
http_client = FakeHTTPClient(put=[httpx.Response(503)])
outcome = await _client(http_client, max_retries=2).upload(BODY)
assert outcome == PointFiveUploadFailure("presigned upload returned 503, gave up after 2 attempts", retryable=True)
assert len(http_client.put_calls) == 2
@pytest.mark.asyncio
async def test_rejected_upload_is_not_retried():
http_client = FakeHTTPClient(put=[httpx.Response(403)])
outcome = await _client(http_client).upload(BODY)
assert outcome == PointFiveUploadFailure("presigned upload returned 403", retryable=False)
assert len(http_client.put_calls) == 1
@pytest.mark.asyncio
async def test_bad_api_key_is_not_retried():
http_client = FakeHTTPClient(presign=[httpx.Response(401)])
outcome = await _client(http_client).upload(BODY)
assert outcome == PointFiveUploadFailure("pointfive api returned 401", retryable=False)
assert http_client.put_calls == []
@pytest.mark.asyncio
async def test_the_reason_for_a_refusal_is_surfaced():
"""A 403 means the key no longer maps to an integration; the operator needs to read why."""
http_client = FakeHTTPClient(presign=[_refused(403, "no integration accepts uploads from this api key")])
outcome = await _client(http_client).upload(BODY)
assert outcome == PointFiveUploadFailure(
"pointfive api returned 403, no integration accepts uploads from this api key", retryable=False
)
assert http_client.put_calls == []
@pytest.mark.asyncio
async def test_api_server_error_is_retried():
http_client = FakeHTTPClient(presign=[httpx.Response(503), _presigned()])
outcome = await _client(http_client).upload(BODY)
assert outcome == OBJECT_KEY
assert len(http_client.presign_calls) == 2
@pytest.mark.asyncio
async def test_too_many_requests_is_retried():
http_client = FakeHTTPClient(presign=[httpx.Response(429), _presigned()])
outcome = await _client(http_client).upload(BODY)
assert outcome == OBJECT_KEY
assert len(http_client.presign_calls) == 2
@pytest.mark.asyncio
async def test_unreachable_api_is_retried_then_reported_as_retryable():
http_client = FakeHTTPClient(presign=(ConnectionError("down"),))
outcome = await _client(http_client, max_retries=2).upload(BODY)
assert isinstance(outcome, PointFiveUploadFailure)
assert outcome.retryable
assert "unreachable" in outcome.detail
assert len(http_client.presign_calls) == 2
@pytest.mark.asyncio
async def test_malformed_api_body_is_not_retried():
http_client = FakeHTTPClient(presign=[_response(200, {"objectKey": "k"})])
outcome = await _client(http_client).upload(BODY)
assert outcome == PointFiveUploadFailure("pointfive api returned an unreadable body", retryable=False)
assert http_client.put_calls == []
@pytest.mark.asyncio
async def test_a_body_that_is_not_json_is_reported_as_unreadable():
http_client = FakeHTTPClient(presign=(httpx.Response(200, text="<html>gateway</html>"),))
outcome = await _client(http_client).upload(BODY)
assert outcome == PointFiveUploadFailure("pointfive api returned an unreadable body", retryable=False)
assert http_client.put_calls == []
@pytest.mark.asyncio
async def test_ping_reports_a_live_shipper():
http_client = FakeHTTPClient(presign=(_no_content(),))
failure = await _client(http_client).ping()
assert failure is None
assert http_client.presign_calls[0]["url"] == "https://api.pointfive.co/api/v1/ingestion/ping"
assert http_client.presign_calls[0]["json"] == {"kind": "LITELLM"}
@pytest.mark.asyncio
async def test_ping_surfaces_a_revoked_key():
http_client = FakeHTTPClient(presign=(_refused(403, "no integration accepts uploads from this api key"),))
failure = await _client(http_client).ping()
assert failure is not None
assert not failure.retryable
assert "no integration accepts uploads from this api key" in failure.detail
@pytest.mark.asyncio
async def test_ping_surfaces_an_unreachable_api():
http_client = FakeHTTPClient(presign=(ConnectionError("down"),))
failure = await _client(http_client).ping()
assert failure is not None
assert failure.retryable
@pytest.mark.asyncio
async def test_a_transport_fault_on_the_upload_itself_is_retryable():
http_client = FakeHTTPClient(put=(ConnectionError("reset"),))
outcome = await _client(http_client, max_retries=1).upload(BODY)
assert isinstance(outcome, PointFiveUploadFailure)
assert outcome.retryable
assert "presigned upload unreachable" in outcome.detail
@pytest.mark.asyncio
async def test_a_client_that_may_not_try_at_all_says_so():
"""max_upload_retries is validated as >= 1, so this guards the loop against a future zero."""
outcome = await _client(FakeHTTPClient(), max_retries=0).upload(BODY)
assert outcome == PointFiveUploadFailure("max_upload_retries must be at least 1", retryable=False)

View file

@ -84,6 +84,7 @@ async def test_async_post_call_success_hook_includes_client_ip_user_agent():
"litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None
):
logger = PrometheusLogger()
logger._emit_input_sequence_length_label = False
logger.litellm_proxy_total_requests_metric = MagicMock()
logger.get_labels_for_metric = MagicMock(
return_value=["client_ip", "user_agent"]

View file

@ -0,0 +1,428 @@
import asyncio
import datetime
from collections.abc import Mapping
from copy import deepcopy
from typing import Final, cast
import pytest
from prometheus_client import REGISTRY
from prometheus_client.samples import Sample
import litellm
from litellm.integrations.prometheus import PrometheusLogger
from litellm.types.integrations.prometheus import (
PrometheusMetricLabels,
UserAPIKeyLabelNames,
UserAPIKeyLabelValues,
get_input_sequence_length_bucket,
)
from litellm.types.utils import StandardLoggingPayload
LATENCY_METRICS: Final = (
"litellm_llm_api_latency_metric",
"litellm_llm_api_time_to_first_token_metric",
"litellm_request_total_latency_metric",
)
FLAG: Final = "prometheus_emit_input_sequence_length_label"
def _clear_prometheus_registry() -> None:
for collector in tuple(REGISTRY._collector_to_names): # pyright: ignore[reportPrivateUsage] # test registry reset
REGISTRY.unregister(collector)
@pytest.fixture(autouse=True)
def isolated_registry(monkeypatch: pytest.MonkeyPatch):
_clear_prometheus_registry()
monkeypatch.setattr(litellm, FLAG, False)
yield
_clear_prometheus_registry()
@pytest.mark.parametrize("metric", LATENCY_METRICS)
def test_input_sequence_length_label_is_opt_in(monkeypatch: pytest.MonkeyPatch, metric: str):
assert UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in PrometheusMetricLabels.get_labels(metric)
monkeypatch.setattr(litellm, FLAG, True)
assert UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value in PrometheusMetricLabels.get_labels(metric)
def test_input_sequence_length_label_stays_off_non_latency_metrics(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, FLAG, True)
assert UserAPIKeyLabelNames.INPUT_SEQUENCE_LENGTH.value not in PrometheusMetricLabels.get_labels(
"litellm_proxy_total_requests_metric"
)
@pytest.mark.parametrize(
"prompt_tokens, expected",
[
(None, "unknown"),
(0, "0-1k"),
(999, "0-1k"),
(1_000, "1k-4k"),
(3_999, "1k-4k"),
(4_000, "4k-16k"),
(15_999, "4k-16k"),
(16_000, "16k-64k"),
(63_999, "16k-64k"),
(64_000, "64k+"),
(10_000_000, "64k+"),
(-1, "unknown"),
],
)
def test_input_sequence_length_bucket_boundaries(prompt_tokens: int | None, expected: str):
assert get_input_sequence_length_bucket(prompt_tokens) == expected
def test_user_api_key_label_values_carries_input_sequence_length():
values: Final = UserAPIKeyLabelValues(input_sequence_length="4k-16k")
assert values.input_sequence_length == "4k-16k"
assert values.model_dump()["input_sequence_length"] == "4k-16k"
def _assert_latency_metrics(expected: str | None, stream: bool = True, queue_time: float = 0) -> None:
samples: Final = tuple(sample for metric in REGISTRY.collect() for sample in metric.samples)
for metric, duration in zip(LATENCY_METRICS, (2, 1, 3 + queue_time)):
counts: Final = tuple(sample for sample in samples if sample.name == f"{metric}_count")
sums: Final = tuple(sample for sample in samples if sample.name == f"{metric}_sum")
buckets: Final = tuple(sample for sample in samples if sample.name == f"{metric}_bucket")
if not stream and metric == "litellm_llm_api_time_to_first_token_metric":
assert not counts and not sums and not buckets
continue
assert len(counts) == len(sums) == 1
assert counts[0].value == 1
assert sums[0].value == pytest.approx(duration)
assert buckets and any(sample.labels["le"] == "+Inf" for sample in buckets)
assert all(sample.value == int(float(sample.labels["le"]) >= duration) for sample in buckets)
assert all(sample.labels.get("input_sequence_length") == expected for sample in (*counts, *sums, *buckets))
def _non_target_samples() -> tuple[Sample, ...]:
return tuple(
sample
for metric in REGISTRY.collect()
if metric.name not in LATENCY_METRICS
for sample in metric.samples
if "input_sequence_length" in sample.labels and not sample.name.endswith("_created")
)
def _standard_logging_payload(now: datetime.datetime, prompt_tokens: int) -> StandardLoggingPayload:
return cast(
StandardLoggingPayload,
{
"id": "t",
"call_type": "completion",
"response_cost": 0.001,
"status": "success",
"total_tokens": prompt_tokens + 20,
"prompt_tokens": prompt_tokens,
"completion_tokens": 20,
"startTime": now - datetime.timedelta(seconds=3),
"endTime": now,
"completionStartTime": now - datetime.timedelta(seconds=1),
"model": "gpt-4o-mini",
"model_id": "model-123",
"model_group": "gpt-4o-mini",
"api_base": "https://api.openai.com",
"custom_llm_provider": "openai",
"request_tags": [],
"stream": True,
"metadata": {
"user_api_key_hash": "h",
"user_api_key_alias": "a",
"user_api_key_team_id": "t",
"user_api_key_team_alias": "ta",
"user_api_key_user_id": "u",
"user_api_key_user_email": "e@x.com",
"user_api_key_org_id": None,
"user_api_key_org_alias": None,
"requester_metadata": None,
"user_api_key_end_user_id": None,
"usage_object": None,
},
"hidden_params": {"litellm_overhead_time_ms": None, "additional_headers": None},
},
)
def _success_kwargs(
now: datetime.datetime, prompt_tokens: int, requester_metadata: Mapping[str, object] | None = None
) -> Mapping[str, object]:
payload: Final = _standard_logging_payload(now, prompt_tokens)
return {
"model": "gpt-4o-mini",
"litellm_params": {"metadata": {}},
"standard_logging_object": {
**payload,
"metadata": {**payload["metadata"], "requester_metadata": requester_metadata},
},
"stream": True,
"start_time": now - datetime.timedelta(seconds=3),
"api_call_start_time": now - datetime.timedelta(seconds=2),
"completion_start_time": now - datetime.timedelta(seconds=1),
"end_time": now,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("flag_at_request_time", (True, False))
async def test_logger_emits_bucket_from_its_startup_label_set(
monkeypatch: pytest.MonkeyPatch, flag_at_request_time: bool
):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, True)
logger: Final = PrometheusLogger()
monkeypatch.setattr(litellm, FLAG, flag_at_request_time)
await logger.async_log_success_event(dict(_success_kwargs(now, prompt_tokens=4_000)), None, now, now)
_assert_latency_metrics("4k-16k")
@pytest.mark.asyncio
@pytest.mark.parametrize(
("response", "combined_usage", "expected"),
(
({"id": "moderation", "results": []}, None, "unknown"),
({"usage": None}, None, "unknown"),
({"usage": {}}, None, "unknown"),
({"usage": {"completion_tokens": 3}}, None, "unknown"),
({"usage": {"total_tokens": 5}}, None, "unknown"),
({"usage": {"prompt_tokens": 0}}, None, "0-1k"),
({"usage": {"prompt_tokens": 4_000}}, None, "4k-16k"),
({"usage": {"input_tokens": 0, "output_tokens": 3, "total_tokens": 3}}, None, "0-1k"),
({"usage": {"input_tokens": 4_000, "output_tokens": 3, "total_tokens": 4_003}}, None, "4k-16k"),
(litellm.ModelResponse(usage=litellm.Usage(prompt_tokens=0)), None, "0-1k"),
(None, litellm.Usage(prompt_tokens=0), "0-1k"),
),
)
@pytest.mark.parametrize("include_usage_metadata", (True, False))
async def test_logger_distinguishes_missing_usage_from_reported_zero(
monkeypatch: pytest.MonkeyPatch,
response: object,
combined_usage: object,
expected: str,
include_usage_metadata: bool,
):
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, True)
logger: Final = PrometheusLogger()
usage: Final = StandardLoggingPayloadSetup.get_usage_as_dict(
response_obj=response if isinstance(response, dict) else None
)
payload: Final = _standard_logging_payload(now, usage.get("prompt_tokens", 0))
await logger.async_log_success_event(
{
**_success_kwargs(now, prompt_tokens=usage.get("prompt_tokens", 0)),
"combined_usage_object": combined_usage,
"standard_logging_object": {
**payload,
"metadata": {**payload["metadata"], "usage_object": usage if include_usage_metadata else None},
},
},
response,
now,
now,
)
_assert_latency_metrics(expected)
@pytest.mark.asyncio
@pytest.mark.parametrize("total_tokens", (None, 0, 5_000))
@pytest.mark.parametrize("prompt_tokens", (0, 4_000))
async def test_upstream_total_only_usage_has_unknown_input_length(
monkeypatch: pytest.MonkeyPatch, total_tokens: int | None, prompt_tokens: int
):
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging, StandardLoggingPayloadSetup
from litellm.proxy.pass_through_endpoints.upstream_usage_headers import apply_upstream_reported_usage
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, True)
logger: Final = PrometheusLogger()
logging_obj: Final = Logging(
model="gpt-4o-mini",
messages=[],
stream=True,
call_type="pass_through_endpoint",
start_time=now,
litellm_call_id="test-call-id",
function_id="1",
)
headers: Final = httpx.Headers(
{
"x-litellm-response-cost": "0.001",
**({"x-litellm-total-tokens": str(total_tokens)} if total_tokens is not None else {}),
}
)
reported: Final = apply_upstream_reported_usage(logging_obj=logging_obj, headers=headers)
assert reported is not None
combined_usage: Final = logging_obj.model_call_details.get("combined_usage_object")
response: Final = {"usage": {"prompt_tokens": prompt_tokens}}
usage: Final = StandardLoggingPayloadSetup.get_usage_as_dict(response, combined_usage)
payload: Final = _standard_logging_payload(now, usage.get("prompt_tokens", 0))
await logger.async_log_success_event(
{
**logging_obj.model_call_details,
**_success_kwargs(now, usage.get("prompt_tokens", 0)),
"standard_logging_object": {**payload, "metadata": {**payload["metadata"], "usage_object": usage}},
},
response,
now,
now,
)
_assert_latency_metrics("unknown" if total_tokens is not None else get_input_sequence_length_bucket(prompt_tokens))
@pytest.mark.asyncio
async def test_logger_built_with_flag_off_emits_no_bucket_label(monkeypatch: pytest.MonkeyPatch):
now: Final = datetime.datetime.now()
logger: Final = PrometheusLogger()
monkeypatch.setattr(litellm, FLAG, True)
await logger.async_log_success_event(dict(_success_kwargs(now, prompt_tokens=4_000)), None, now, now)
_assert_latency_metrics(None)
@pytest.mark.asyncio
@pytest.mark.parametrize("flag_at_startup", (True, False))
@pytest.mark.parametrize("stream", (True, False))
@pytest.mark.parametrize(
"metadata",
(
None,
{},
{"input_sequence_length": None},
{"input_sequence_length": False},
{"input_sequence_length": True},
{"input_sequence_length": 0},
{"input_sequence_length": []},
{"input_sequence_length": {}},
{"input_sequence_length": ""},
{"input_sequence_length": "from-metadata"},
),
)
async def test_custom_input_length_label_is_scoped_to_target_histograms(
monkeypatch: pytest.MonkeyPatch, flag_at_startup: bool, stream: bool, metadata: Mapping[str, object] | None
):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, "custom_prometheus_metadata_labels", ["input_sequence_length"])
kwargs: Final = {
**_success_kwargs(now, prompt_tokens=4_000, requester_metadata=metadata),
"stream": stream,
"litellm_params": {"metadata": {"queue_time_seconds": 0.25}},
}
original_kwargs: Final = deepcopy(kwargs)
baseline_logger: Final = PrometheusLogger()
await baseline_logger.async_log_success_event(kwargs, None, now, now)
baseline_samples: Final = _non_target_samples()
_clear_prometheus_registry()
monkeypatch.setattr(litellm, FLAG, flag_at_startup)
logger: Final = PrometheusLogger()
monkeypatch.setattr(litellm, FLAG, not flag_at_startup)
await logger.async_log_success_event(kwargs, None, now, now)
assert kwargs == original_kwargs
custom_value: Final = (metadata or {}).get("input_sequence_length")
expected: Final = custom_value if isinstance(custom_value, str) else ("4k-16k" if flag_at_startup else "None")
_assert_latency_metrics(expected, stream=stream, queue_time=0.25)
assert all(logger.get_labels_for_metric(metric).count("input_sequence_length") == 1 for metric in LATENCY_METRICS)
non_target_samples: Final = _non_target_samples()
assert {
"litellm_requests_metric_total",
"litellm_spend_metric_total",
"litellm_total_tokens_metric_total",
"litellm_request_queue_time_seconds_count",
"litellm_deployment_success_responses_total",
}.issubset({sample.name for sample in non_target_samples})
assert non_target_samples == baseline_samples
queue_sum: Final = tuple(
sample for sample in non_target_samples if sample.name == "litellm_request_queue_time_seconds_sum"
)
assert len(queue_sum) == 1 and queue_sum[0].value == 0.25
@pytest.mark.asyncio
@pytest.mark.parametrize("enabled", (True, False))
async def test_concurrent_requests_keep_independent_buckets(monkeypatch: pytest.MonkeyPatch, enabled: bool):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, enabled)
logger: Final = PrometheusLogger()
cases: Final = (
(None, "unknown"),
(0, "0-1k"),
(1_000, "1k-4k"),
(4_000, "4k-16k"),
(16_000, "16k-64k"),
(64_000, "64k+"),
)
calls: Final = tuple(
(
{**_success_kwargs(now, prompt_tokens=tokens or 0), "stream": stream},
{"usage": {"prompt_tokens": tokens}} if tokens is not None else None,
)
for tokens, _ in cases
for stream in (True, False)
for _ in range(2)
)
original_calls: Final = deepcopy(calls)
await asyncio.gather(*(logger.async_log_success_event(kwargs, response, now, now) for kwargs, response in calls))
assert calls == original_calls
samples: Final = tuple(sample for metric in REGISTRY.collect() for sample in metric.samples)
for metric, duration in zip(LATENCY_METRICS, (2, 1, 3)):
expected_count: Final = 2 if metric == "litellm_llm_api_time_to_first_token_metric" else 4
counts: Final = tuple(sample for sample in samples if sample.name == f"{metric}_count")
sums: Final = tuple(sample for sample in samples if sample.name == f"{metric}_sum")
buckets: Final = tuple(sample for sample in samples if sample.name == f"{metric}_bucket")
expected: Final = (
{bucket: expected_count for _, bucket in cases} if enabled else {None: expected_count * len(cases)}
)
assert len(counts) == len(sums) == len(expected)
assert {sample.labels.get("input_sequence_length"): sample.value for sample in counts} == expected
assert {sample.labels.get("input_sequence_length"): sample.value for sample in sums} == {
bucket: count * duration for bucket, count in expected.items()
}
assert sum(sample.value for sample in buckets if sample.labels["le"] == "+Inf") == expected_count * len(cases)
assert all(
sample.value
== expected[sample.labels.get("input_sequence_length")] * int(float(sample.labels["le"]) >= duration)
for sample in buckets
)
@pytest.mark.asyncio
@pytest.mark.parametrize("enabled", (True, False))
async def test_failed_request_does_not_observe_latency(monkeypatch: pytest.MonkeyPatch, enabled: bool):
now: Final = datetime.datetime.now()
monkeypatch.setattr(litellm, FLAG, enabled)
monkeypatch.setattr(litellm, "custom_prometheus_metadata_labels", ["input_sequence_length"])
logger: Final = PrometheusLogger()
kwargs: Final = {
**_success_kwargs(now, prompt_tokens=4_000),
"standard_logging_object": {**_standard_logging_payload(now, 4_000), "status": "failure"},
"exception": RuntimeError("upstream request failed"),
}
await logger.async_log_failure_event(kwargs, None, now, now)
samples: Final = tuple(sample for metric in REGISTRY.collect() for sample in metric.samples)
assert not any(sample.name.startswith(LATENCY_METRICS) for sample in samples)
for metric in ("litellm_llm_api_failed_requests_metric_total", "litellm_deployment_failure_responses_total"):
counts: Final = tuple(sample for sample in samples if sample.name == metric)
assert len(counts) == 1
assert counts[0].value == 1
assert counts[0].labels["input_sequence_length"] == "None"

View file

@ -1,10 +1,19 @@
import asyncio
import re
import sys
import textwrap
import uuid
from contextlib import asynccontextmanager
from datetime import datetime
from unittest.mock import MagicMock, patch
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, call, patch
import httpx
import pytest
from litellm.integrations.s3_v2 import S3Logger
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
from litellm.types.utils import StandardLoggingPayload
@ -21,9 +30,7 @@ class TestS3V2UnitTests:
source_code = inspect.getsource(s3_v2)
# Verify that json.dumps is not used directly in the code
assert (
"json.dumps(" not in source_code
), "S3 v2 should not use json.dumps directly"
assert "json.dumps(" not in source_code, "S3 v2 should not use json.dumps directly"
@patch("asyncio.create_task")
@patch("litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush")
@ -86,12 +93,8 @@ class TestS3V2UnitTests:
call_args_minio = s3_logger_minio.async_httpx_client.put.call_args
assert call_args_minio is not None
url_minio = call_args_minio[0][0]
expected_minio_url = (
"https://minio.example.com:9000/litellm-logs/2025-09-14/test-key.json"
)
assert (
url_minio == expected_minio_url
), f"Expected MinIO URL {expected_minio_url}, got {url_minio}"
expected_minio_url = "https://minio.example.com:9000/litellm-logs/2025-09-14/test-key.json"
assert url_minio == expected_minio_url, f"Expected MinIO URL {expected_minio_url}, got {url_minio}"
# Test 3: Custom endpoint without bucket name (should fall back to default)
s3_logger_no_bucket = S3Logger(
@ -136,12 +139,8 @@ class TestS3V2UnitTests:
call_args_sync = mock_sync_client.put.call_args
assert call_args_sync is not None
url_sync = call_args_sync[0][0]
expected_sync_url = (
"https://custom.s3.endpoint.com/sync-bucket/2025-09-14/test-key.json"
)
assert (
url_sync == expected_sync_url
), f"Expected sync URL {expected_sync_url}, got {url_sync}"
expected_sync_url = "https://custom.s3.endpoint.com/sync-bucket/2025-09-14/test-key.json"
assert url_sync == expected_sync_url, f"Expected sync URL {expected_sync_url}, got {url_sync}"
# Test 5: Download method with custom endpoint
s3_logger_download = S3Logger(
@ -158,19 +157,15 @@ class TestS3V2UnitTests:
s3_logger_download.async_httpx_client = AsyncMock()
s3_logger_download.async_httpx_client.get.return_value = mock_download_response
result = asyncio.run(
s3_logger_download._download_object_from_s3(
"2025-09-14/download-test-key.json"
)
)
result = asyncio.run(s3_logger_download._download_object_from_s3("2025-09-14/download-test-key.json"))
call_args_download = s3_logger_download.async_httpx_client.get.call_args
assert call_args_download is not None
url_download = call_args_download[0][0]
expected_download_url = "https://download.s3.endpoint.com/download-bucket/2025-09-14/download-test-key.json"
assert (
url_download == expected_download_url
), f"Expected download URL {expected_download_url}, got {url_download}"
assert url_download == expected_download_url, (
f"Expected download URL {expected_download_url}, got {url_download}"
)
assert result == {"downloaded": "data"}
@ -216,12 +211,8 @@ class TestS3V2UnitTests:
call_args = s3_logger_virtual.async_httpx_client.put.call_args
assert call_args is not None
url = call_args[0][0]
expected_url = (
"https://test-bucket.s3.custom-endpoint.com/2025-09-14/test-key.json"
)
assert (
url == expected_url
), f"Expected virtual-hosted-style URL {expected_url}, got {url}"
expected_url = "https://test-bucket.s3.custom-endpoint.com/2025-09-14/test-key.json"
assert url == expected_url, f"Expected virtual-hosted-style URL {expected_url}, got {url}"
# Test 2: Path-style (default behavior with s3_use_virtual_hosted_style=False)
s3_logger_path = S3Logger(
@ -241,12 +232,8 @@ class TestS3V2UnitTests:
call_args_path = s3_logger_path.async_httpx_client.put.call_args
assert call_args_path is not None
url_path = call_args_path[0][0]
expected_path_url = (
"https://s3.custom-endpoint.com/test-bucket/2025-09-14/test-key.json"
)
assert (
url_path == expected_path_url
), f"Expected path-style URL {expected_path_url}, got {url_path}"
expected_path_url = "https://s3.custom-endpoint.com/test-bucket/2025-09-14/test-key.json"
assert url_path == expected_path_url, f"Expected path-style URL {expected_path_url}, got {url_path}"
# Test 3: Virtual-hosted-style with http protocol
s3_logger_http = S3Logger(
@ -266,12 +253,10 @@ class TestS3V2UnitTests:
call_args_http = s3_logger_http.async_httpx_client.put.call_args
assert call_args_http is not None
url_http = call_args_http[0][0]
expected_http_url = (
"http://http-bucket.minio.local:9000/2025-09-14/test-key.json"
expected_http_url = "http://http-bucket.minio.local:9000/2025-09-14/test-key.json"
assert url_http == expected_http_url, (
f"Expected virtual-hosted-style URL with http {expected_http_url}, got {url_http}"
)
assert (
url_http == expected_http_url
), f"Expected virtual-hosted-style URL with http {expected_http_url}, got {url_http}"
# Test 4: Sync upload method with virtual-hosted-style
s3_logger_sync_virtual = S3Logger(
@ -295,12 +280,10 @@ class TestS3V2UnitTests:
call_args_sync = mock_sync_client.put.call_args
assert call_args_sync is not None
url_sync = call_args_sync[0][0]
expected_sync_url = (
"https://sync-bucket.storage.example.com/2025-09-14/test-key.json"
expected_sync_url = "https://sync-bucket.storage.example.com/2025-09-14/test-key.json"
assert url_sync == expected_sync_url, (
f"Expected virtual-hosted-style sync URL {expected_sync_url}, got {url_sync}"
)
assert (
url_sync == expected_sync_url
), f"Expected virtual-hosted-style sync URL {expected_sync_url}, got {url_sync}"
# Test 5: Download method with virtual-hosted-style
s3_logger_download_virtual = S3Logger(
@ -316,34 +299,27 @@ class TestS3V2UnitTests:
mock_download_response.status_code = 200
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
s3_logger_download_virtual.async_httpx_client = AsyncMock()
s3_logger_download_virtual.async_httpx_client.get.return_value = (
mock_download_response
)
s3_logger_download_virtual.async_httpx_client.get.return_value = mock_download_response
result = asyncio.run(
s3_logger_download_virtual._download_object_from_s3(
"2025-09-14/download-test-key.json"
)
)
result = asyncio.run(s3_logger_download_virtual._download_object_from_s3("2025-09-14/download-test-key.json"))
call_args_download = s3_logger_download_virtual.async_httpx_client.get.call_args
assert call_args_download is not None
url_download = call_args_download[0][0]
expected_download_url = "https://download-bucket.download.endpoint.com/2025-09-14/download-test-key.json"
assert (
url_download == expected_download_url
), f"Expected virtual-hosted-style download URL {expected_download_url}, got {url_download}"
assert url_download == expected_download_url, (
f"Expected virtual-hosted-style download URL {expected_download_url}, got {url_download}"
)
assert result == {"downloaded": "data"}
@patch("asyncio.create_task")
@patch("litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush")
def test_s3_v2_put_url_encodes_spaces_in_object_key(
self, mock_periodic_flush, mock_create_task
):
import requests
def test_s3_v2_put_url_encodes_spaces_in_object_key(self, mock_periodic_flush, mock_create_task):
from unittest.mock import AsyncMock
import requests
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
mock_periodic_flush.return_value = None
@ -487,9 +463,7 @@ async def test_async_upload_exhausts_retries_on_persistent_503():
# All 3 attempts return 503
response_503 = MagicMock()
response_503.status_code = 503
response_503.raise_for_status = MagicMock(
side_effect=Exception("503 Service Unavailable")
)
response_503.raise_for_status = MagicMock(side_effect=Exception("503 Service Unavailable"))
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = AsyncMock(return_value=response_503)
@ -528,12 +502,12 @@ async def test_async_upload_no_retry_on_4xx():
s3_object_download_filename="test-no-retry.json",
)
response_403 = MagicMock()
response_403.status_code = 403
response_403.raise_for_status = MagicMock(side_effect=Exception("403 Forbidden"))
response_400 = MagicMock()
response_400.status_code = 400
response_400.raise_for_status = MagicMock(side_effect=Exception("400 Bad Request"))
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = AsyncMock(return_value=response_403)
logger.async_httpx_client.put = AsyncMock(return_value=response_400)
with patch.object(logger, "handle_callback_failure") as mock_failure:
await logger.async_upload_data_to_s3(test_element)
@ -543,6 +517,190 @@ async def test_async_upload_no_retry_on_4xx():
mock_failure.assert_called_once_with(callback_name="S3Logger")
_SIGV4_ACCESS_KEY = re.compile(r"Credential=(AKIA\d+)/")
@pytest.fixture
def rotating_profile(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> str:
"""
A real botocore profile whose credential_process hands out a new key generation on every call and
expires inside the advisory refresh window, so RefreshableCredentials re-runs it on every property read.
"""
counter = tmp_path / "generation"
script = tmp_path / "rotate_credentials.py"
script.write_text(
textwrap.dedent(
f"""
import json, sys
from datetime import datetime, timedelta, timezone
from pathlib import Path
counter = Path({str(counter)!r})
generation = int(counter.read_text()) if counter.exists() else 0
counter.write_text(str(generation + 1))
expiry = (datetime.now(timezone.utc) + timedelta(minutes=12)).strftime("%Y-%m-%dT%H:%M:%SZ")
json.dump(
{{
"Version": 1,
"AccessKeyId": f"AKIA{{generation}}",
"SecretAccessKey": f"secret-{{generation}}",
"SessionToken": f"token-{{generation}}",
"Expiration": expiry,
}},
sys.stdout,
)
"""
)
)
profile = f"rotating-{uuid.uuid4().hex}"
(tmp_path / "config").write_text(f"[profile {profile}]\ncredential_process = {sys.executable} {script}\n")
monkeypatch.setenv("AWS_CONFIG_FILE", str(tmp_path / "config"))
return profile
def _generation(request: httpx.Request) -> tuple[str, str]:
"""(access key generation, session token generation) SigV4 baked into one request."""
access_key = _SIGV4_ACCESS_KEY.search(request.headers["Authorization"])
assert access_key is not None
return access_key.group(1).removeprefix("AKIA"), request.headers["X-Amz-Security-Token"].removeprefix("token-")
@asynccontextmanager
async def _s3_logger_on_production_handler(profile: str, statuses: list[int]):
"""
S3Logger wired to the real AsyncHTTPHandler over an httpx MockTransport that answers with the given
statuses in order, so the handler's own raise_for_status behaviour is exercised end to end.
"""
requests: list[httpx.Request] = []
replies = iter(statuses)
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(next(replies), request=request, text="<Error><Code>SignatureDoesNotMatch</Code></Error>")
handler = AsyncHTTPHandler()
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_region_name="us-east-1",
s3_aws_profile_name=profile,
s3_flush_interval=3600,
)
logger.async_httpx_client = handler
with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
yield logger, requests, mock_sleep
await handler.client.aclose()
@pytest.mark.asyncio
async def test_async_upload_signs_with_one_frozen_credential_snapshot(rotating_profile: str, caplog):
"""
RefreshableCredentials refreshes on every property read once inside the advisory window, so signing
off the live object would mix the access key of one generation with the token of the next.
"""
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-frozen.json",
payload={"test": "frozen"},
s3_object_download_filename="test-frozen.json",
)
async with _s3_logger_on_production_handler(rotating_profile, [200]) as (logger, requests, _):
await logger.async_upload_data_to_s3(test_element)
assert len(requests) == 1
access_key_generation, token_generation = _generation(requests[0])
assert access_key_generation == token_generation
assert "Error uploading to s3" not in caplog.text
@pytest.mark.asyncio
async def test_async_upload_retries_403_with_fresh_credentials_and_signature(rotating_profile: str, caplog):
"""
A 403 (SignatureDoesNotMatch after an IMDS rotation) must be retried, and the retry must fetch
credentials again and carry a signature computed from that newer generation.
"""
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-403.json",
payload={"test": "403"},
s3_object_download_filename="test-403.json",
)
async with _s3_logger_on_production_handler(rotating_profile, [403, 200]) as (logger, requests, mock_sleep):
await logger.async_upload_data_to_s3(test_element)
assert len(requests) == 2
first_key, first_token = _generation(requests[0])
second_key, second_token = _generation(requests[1])
assert first_key == first_token
assert second_key == second_token
assert int(second_key) > int(first_key)
assert requests[1].headers["Authorization"] != requests[0].headers["Authorization"]
mock_sleep.assert_awaited_once_with(1)
assert "Error uploading to s3" not in caplog.text
@pytest.mark.asyncio
async def test_async_upload_exhausts_403_retries_through_production_http_handler(rotating_profile: str, caplog):
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-403-exhausted.json",
payload={"test": "403-exhausted"},
s3_object_download_filename="test-403-exhausted.json",
)
async with _s3_logger_on_production_handler(rotating_profile, [403, 403, 403]) as (logger, requests, mock_sleep):
await logger.async_upload_data_to_s3(test_element)
assert len(requests) == 3
assert mock_sleep.await_args_list == [call(1), call(2)]
assert "Error uploading to s3" in caplog.text
@pytest.mark.asyncio
async def test_async_upload_does_not_retry_404_through_production_http_handler(rotating_profile: str, caplog):
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-404.json",
payload={"test": "404"},
s3_object_download_filename="test-404.json",
)
async with _s3_logger_on_production_handler(rotating_profile, [404]) as (logger, requests, mock_sleep):
await logger.async_upload_data_to_s3(test_element)
assert len(requests) == 1
mock_sleep.assert_not_awaited()
assert "Error uploading to s3" in caplog.text
def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("AWS_PROFILE", rotating_profile)
logger = S3Logger(s3_bucket_name="test-bucket", s3_region_name="us-east-1", s3_flush_interval=3600)
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-sync-403.json",
payload={"test": "sync-403"},
s3_object_download_filename="test-sync-403.json",
)
requests: list[httpx.Request] = []
replies = iter([403, 200])
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(next(replies), request=request)
handler = HTTPHandler()
handler.client = httpx.Client(transport=httpx.MockTransport(respond))
with (
patch( # test-quality-ok: sync upload builds its HTTPHandler per call, there is no injection seam for it
"litellm.integrations.s3_v2._get_httpx_client", return_value=handler
),
patch("time.sleep") as mock_sleep,
):
logger.upload_data_to_s3(test_element)
assert len(requests) == 2
first_key, first_token = _generation(requests[0])
second_key, second_token = _generation(requests[1])
assert first_key == first_token
assert second_key == second_token
assert int(second_key) > int(first_key)
mock_sleep.assert_called_once_with(1)
def test_sync_upload_retries_on_s3_503():
"""
Test that the sync upload_data_to_s3 retries on transient S3 503.
@ -626,9 +784,7 @@ async def test_async_log_event_skips_when_standard_logging_object_missing():
# Nothing should have been queued (catches the case where code falls
# through without returning and appends None to the queue)
assert (
len(logger.log_queue) == 0
), "log_queue should be empty when standard_logging_object is missing"
assert len(logger.log_queue) == 0, "log_queue should be empty when standard_logging_object is missing"
@pytest.mark.asyncio
@ -767,20 +923,18 @@ async def test_s3_verify_false_handling(monkeypatch: pytest.MonkeyPatch):
litellm,
"s3_callback_params",
{
"s3_bucket_name": "test-bucket",
"s3_endpoint_url": "https://localhost:443",
"s3_aws_access_key_id": "minioadmin",
"s3_aws_secret_access_key": "minioadmin",
"s3_region_name": "us-east-1",
"s3_verify": False, # This should NOT be ignored
"s3_use_ssl": False, # This should also NOT be ignored
},
"s3_bucket_name": "test-bucket",
"s3_endpoint_url": "https://localhost:443",
"s3_aws_access_key_id": "minioadmin",
"s3_aws_secret_access_key": "minioadmin",
"s3_region_name": "us-east-1",
"s3_verify": False, # This should NOT be ignored
"s3_use_ssl": False, # This should also NOT be ignored
},
)
with patch("asyncio.create_task"):
with patch(
"litellm.integrations.s3_v2.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.integrations.s3_v2.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_get_client.return_value = mock_client
@ -788,22 +942,16 @@ async def test_s3_verify_false_handling(monkeypatch: pytest.MonkeyPatch):
logger = S3Logger()
# Verify s3_verify is False, not None
assert (
logger.s3_verify is False
), f"Expected s3_verify=False, got {logger.s3_verify}"
assert (
logger.s3_use_ssl is False
), f"Expected s3_use_ssl=False, got {logger.s3_use_ssl}"
assert logger.s3_verify is False, f"Expected s3_verify=False, got {logger.s3_verify}"
assert logger.s3_use_ssl is False, f"Expected s3_use_ssl=False, got {logger.s3_use_ssl}"
# Verify that get_async_httpx_client was called with ssl_verify=False
mock_get_client.assert_called_once()
call_kwargs = mock_get_client.call_args.kwargs
assert (
"params" in call_kwargs
), "params should be passed to get_async_httpx_client"
assert call_kwargs["params"] == {
"ssl_verify": False
}, f"Expected ssl_verify=False in params, got {call_kwargs.get('params')}"
assert "params" in call_kwargs, "params should be passed to get_async_httpx_client"
assert call_kwargs["params"] == {"ssl_verify": False}, (
f"Expected ssl_verify=False in params, got {call_kwargs.get('params')}"
)
@pytest.mark.asyncio
@ -820,17 +968,15 @@ async def test_s3_verify_none_handling(monkeypatch: pytest.MonkeyPatch):
litellm,
"s3_callback_params",
{
"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_bucket_name": "test-bucket",
"s3_aws_access_key_id": "test-key",
"s3_aws_secret_access_key": "test-secret",
"s3_region_name": "us-east-1",
},
)
with patch("asyncio.create_task"):
with patch(
"litellm.integrations.s3_v2.get_async_httpx_client"
) as mock_get_client:
with patch("litellm.integrations.s3_v2.get_async_httpx_client") as mock_get_client:
mock_client = AsyncMock()
mock_get_client.return_value = mock_client
@ -838,9 +984,7 @@ async def test_s3_verify_none_handling(monkeypatch: pytest.MonkeyPatch):
logger = S3Logger()
# Verify s3_verify is None (default)
assert (
logger.s3_verify is None
), f"Expected s3_verify=None, got {logger.s3_verify}"
assert logger.s3_verify is None, f"Expected s3_verify=None, got {logger.s3_verify}"
# Verify that get_async_httpx_client was called
mock_get_client.assert_called_once()
@ -868,13 +1012,13 @@ async def test_s3_verify_false_creates_httpx_client_with_verify_false(monkeypatc
litellm,
"s3_callback_params",
{
"s3_bucket_name": "test-bucket",
"s3_endpoint_url": "https://localhost:443",
"s3_aws_access_key_id": "minioadmin",
"s3_aws_secret_access_key": "minioadmin",
"s3_region_name": "us-east-1",
"s3_verify": False,
},
"s3_bucket_name": "test-bucket",
"s3_endpoint_url": "https://localhost:443",
"s3_aws_access_key_id": "minioadmin",
"s3_aws_secret_access_key": "minioadmin",
"s3_region_name": "us-east-1",
"s3_verify": False,
},
)
with patch("asyncio.create_task"):
@ -890,9 +1034,7 @@ async def test_s3_verify_false_creates_httpx_client_with_verify_false(monkeypatc
httpx_client = logger.async_httpx_client.client
# Check the _verify attribute (httpx internal)
if hasattr(httpx_client, "_verify"):
assert (
httpx_client._verify is False
), f"Expected httpx client _verify=False, got {httpx_client._verify}"
assert httpx_client._verify is False, f"Expected httpx client _verify=False, got {httpx_client._verify}"
@pytest.mark.asyncio
@ -910,13 +1052,13 @@ async def test_s3_verify_false_async_client(monkeypatch: pytest.MonkeyPatch):
litellm,
"s3_callback_params",
{
"s3_bucket_name": "test-bucket",
"s3_endpoint_url": "https://localhost:443",
"s3_aws_access_key_id": "minioadmin",
"s3_aws_secret_access_key": "minioadmin",
"s3_region_name": "us-east-1",
"s3_verify": False,
},
"s3_bucket_name": "test-bucket",
"s3_endpoint_url": "https://localhost:443",
"s3_aws_access_key_id": "minioadmin",
"s3_aws_secret_access_key": "minioadmin",
"s3_region_name": "us-east-1",
"s3_verify": False,
},
)
with patch("asyncio.create_task"):
@ -948,9 +1090,9 @@ async def test_s3_verify_false_async_client(monkeypatch: pytest.MonkeyPatch):
if hasattr(logger.async_httpx_client, "client"):
httpx_client = logger.async_httpx_client.client
if hasattr(httpx_client, "_verify"):
assert (
httpx_client._verify is False
), f"Expected async httpx client _verify=False, got {httpx_client._verify}"
assert httpx_client._verify is False, (
f"Expected async httpx client _verify=False, got {httpx_client._verify}"
)
@pytest.mark.asyncio
@ -1017,9 +1159,7 @@ def patch_asyncio_create_task():
(True, True, None, None, ""),
],
)
def test_s3_object_key_prefix_combinations(
use_team_prefix, use_key_prefix, team_alias, key_alias, expected_prefix
):
def test_s3_object_key_prefix_combinations(use_team_prefix, use_key_prefix, team_alias, key_alias, expected_prefix):
"""
Validate correct S3 prefix composition for team alias + key alias combinations.
"""
@ -1490,9 +1630,7 @@ def test_s3_callback_params_override_does_not_mutate_inputs(monkeypatch):
logger = S3Logger(s3_callback_params_override=override)
assert logger.s3_bucket_name == "resolved-bucket"
assert override["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
assert (
litellm.s3_callback_params["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
)
assert litellm.s3_callback_params["s3_bucket_name"] == "os.environ/MY_AUDIT_BUCKET"
def test_s3_callback_params_override_none_falls_back_to_global(monkeypatch):
@ -1520,9 +1658,7 @@ def _expected_content_md5(payload: dict) -> str:
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
json_string = safe_dumps(payload)
return base64.b64encode(
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
).decode()
return base64.b64encode(hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()).decode()
def _require_non_security_md5(monkeypatch):
@ -1658,9 +1794,9 @@ def test_s3_server_side_encryption_read_from_callback_params(monkeypatch):
litellm,
"s3_callback_params",
{
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
},
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
},
)
logger = S3Logger()
assert logger.s3_server_side_encryption == "aws:kms"
@ -1789,10 +1925,10 @@ def test_s3_sse_kms_key_id_read_from_callback_params(monkeypatch):
litellm,
"s3_callback_params",
{
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
},
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
},
)
logger = S3Logger()
assert logger.s3_sse_kms_key_id == ("arn:aws:kms:us-east-1:111122223333:key/test-key-id")
@ -1863,10 +1999,10 @@ def test_kms_key_id_dropped_when_algorithm_is_not_kms(monkeypatch):
litellm,
"s3_callback_params",
{
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "AES256",
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
},
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "AES256",
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
},
)
logger = S3Logger()
assert logger.s3_server_side_encryption == "AES256"
@ -1884,10 +2020,10 @@ def test_non_string_algorithm_is_dropped_and_valid_key_id_is_rescued(monkeypatch
litellm,
"s3_callback_params",
{
"s3_bucket_name": "from-global",
"s3_server_side_encryption": True,
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
},
"s3_bucket_name": "from-global",
"s3_server_side_encryption": True,
"s3_sse_kms_key_id": "arn:aws:kms:us-east-1:111122223333:key/test-key-id",
},
)
logger = S3Logger()
assert logger.s3_server_side_encryption == "aws:kms"
@ -1902,10 +2038,10 @@ def test_non_string_key_id_is_dropped_and_valid_algorithm_is_kept(monkeypatch):
litellm,
"s3_callback_params",
{
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
"s3_sse_kms_key_id": 12345,
},
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
"s3_sse_kms_key_id": 12345,
},
)
logger = S3Logger()
assert logger.s3_server_side_encryption == "aws:kms"
@ -2045,6 +2181,7 @@ async def test_download_signs_object_key_with_space_the_way_s3_does():
headers=call.kwargs["headers"],
)
_RESERVED_CHAR_KEYS = (
"2026-08-21/time-05-29-36_resp_bGl0ZWxsbTpjdXN0b20=.json",
"session=logs/2026-08-21/time-05-29-36_abc.json",

View file

@ -2,11 +2,12 @@ import copy
import functools
import json
import os
import sys
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
from litellm.litellm_core_utils.prompt_templates.common_utils import (
ENCRYPTED_REASONING_SIGNATURE_PREFIX,
TOOL_RESULT_IMAGE_BOUNDARY,
@ -25,6 +26,8 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
update_messages_with_model_file_ids,
)
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
def test_get_format_from_file_id():
unified_file_id = "litellm_proxy:application/pdf;unified_id,cbbe3534-8bf8-4386-af00-f5f6b7e370bf"
@ -1442,7 +1445,7 @@ class TestFlattenTopLevelSchemaCombinators:
assert schema == snapshot
class TestToolWithFlattenedParameters:
class TestToolWithSanitizedParameters:
def _anyof_tool(self):
return {
"type": "function",
@ -1469,11 +1472,12 @@ class TestToolWithFlattenedParameters:
def test_flattens_anyof_parameters_into_new_tool(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
tool_with_flattened_parameters,
flatten_combinators_and_drop_non_python_regex_patterns,
tool_with_sanitized_parameters,
)
tool = self._anyof_tool()
result = tool_with_flattened_parameters(tool)
result = tool_with_sanitized_parameters(tool, flatten_combinators_and_drop_non_python_regex_patterns)
assert result is not tool
parameters = result["function"]["parameters"]
@ -1486,7 +1490,8 @@ class TestToolWithFlattenedParameters:
def test_clean_parameters_return_the_same_tool_object(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
tool_with_flattened_parameters,
flatten_combinators_and_drop_non_python_regex_patterns,
tool_with_sanitized_parameters,
)
tool = {
@ -1497,7 +1502,23 @@ class TestToolWithFlattenedParameters:
},
}
assert tool_with_flattened_parameters(tool) is tool
assert tool_with_sanitized_parameters(tool, flatten_combinators_and_drop_non_python_regex_patterns) is tool
def test_pattern_only_sanitizer_drops_the_regex_and_keeps_the_union(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
tool_with_sanitized_parameters,
)
tool = self._anyof_tool()
tool["function"]["parameters"]["properties"]["id"]["pattern"] = _ARTIFACT_FIELD_PATTERN
result = tool_with_sanitized_parameters(tool, drop_non_python_regex_patterns)
parameters = result["function"]["parameters"]
assert parameters["properties"]["id"] == {"type": "string"}
assert parameters["anyOf"] == self._anyof_tool()["function"]["parameters"]["anyOf"]
assert tool["function"]["parameters"]["properties"]["id"]["pattern"] == _ARTIFACT_FIELD_PATTERN
@pytest.mark.parametrize(
"tool",
@ -1510,10 +1531,127 @@ class TestToolWithFlattenedParameters:
)
def test_non_dict_function_or_parameters_return_the_same_tool_object(self, tool):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
tool_with_flattened_parameters,
flatten_combinators_and_drop_non_python_regex_patterns,
tool_with_sanitized_parameters,
)
assert tool_with_flattened_parameters(tool) is tool
assert tool_with_sanitized_parameters(tool, flatten_combinators_and_drop_non_python_regex_patterns) is tool
class TestDropNonPythonRegexPatterns:
"""Claude Code's Artifact tool declares ECMA-262 ``\\p{..}`` escapes that OpenAI's
validator, which compiles ``pattern`` values and ``patternProperties`` keys with
Python ``re``, refuses as "not a 'regex'"."""
def _schema(self, pattern):
return {
"type": "object",
"properties": {
"field": {"type": "string", "pattern": pattern},
"writes": {
"type": "array",
"items": {"properties": {"doc_id": {"type": "string", "pattern": pattern}}},
},
"query": {"anyOf": [{"type": "string", "pattern": pattern}, {"type": "null"}]},
"pair": {"type": "array", "prefixItems": [{"type": "string", "pattern": pattern}]},
"extra": {"type": "object", "additionalProperties": {"type": "string", "pattern": pattern}},
"tagged": {
"type": "object",
"patternProperties": {pattern: {"type": "string"}, "^x_": {"type": "integer"}},
},
},
"$defs": {"segment": {"type": "string", "pattern": pattern}},
"required": ["field"],
}
def test_drops_every_regex_python_re_rejects_from_every_schema_position(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
)
schema = self._schema(_ARTIFACT_FIELD_PATTERN)
result = drop_non_python_regex_patterns(schema)
assert '"pattern"' not in json.dumps(result)
properties = result["properties"]
assert properties["field"] == {"type": "string"}
assert properties["writes"]["items"]["properties"]["doc_id"] == {"type": "string"}
assert properties["query"]["anyOf"] == [{"type": "string"}, {"type": "null"}]
assert properties["pair"]["prefixItems"] == [{"type": "string"}]
assert properties["extra"]["additionalProperties"] == {"type": "string"}
assert properties["tagged"]["patternProperties"] == {"^x_": {"type": "integer"}}
assert result["$defs"]["segment"] == {"type": "string"}
assert result["required"] == ["field"]
assert schema == self._schema(_ARTIFACT_FIELD_PATTERN)
def test_keeps_regexes_python_re_compiles_and_returns_the_same_object(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
)
schema = self._schema(r'^(?!__.*__$)[^"\\./[\]]{1,200}$')
assert drop_non_python_regex_patterns(schema) is schema
def test_pattern_keys_inside_data_positions_are_not_regexes(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
)
schema = {
"type": "object",
"properties": {
"pattern": {"type": "string"},
"template": {"type": "object", "default": {"pattern": _ARTIFACT_FIELD_PATTERN}},
"samples": {"type": "array", "examples": [{"pattern": _ARTIFACT_FIELD_PATTERN}]},
"fixed": {"const": {"pattern": _ARTIFACT_FIELD_PATTERN}},
"vendor": {"type": "string", "x-litellm": {"pattern": _ARTIFACT_FIELD_PATTERN}},
},
"required": ["pattern"],
}
assert drop_non_python_regex_patterns(schema) is schema
def test_regex_nested_past_what_python_re_can_parse_is_dropped_not_raised(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
)
schema = {
"type": "object",
"properties": {"deep": {"type": "string", "pattern": "(" * 2000 + "a" + ")" * 2000}},
}
assert drop_non_python_regex_patterns(schema)["properties"]["deep"] == {"type": "string"}
def test_walks_schemas_deeper_than_the_interpreter_recursion_limit(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
)
depth = sys.getrecursionlimit()
leaf = {"type": "string", "pattern": _ARTIFACT_FIELD_PATTERN}
schema = functools.reduce(
lambda inner, _: {"type": "object", "properties": {"child": inner}}, range(depth), leaf
)
result = drop_non_python_regex_patterns(schema)
assert functools.reduce(lambda node, _: node["properties"]["child"], range(depth), result) == {"type": "string"}
assert functools.reduce(lambda node, _: node["properties"]["child"], range(depth), schema) is leaf
def test_leaves_levels_past_the_json_nesting_limit_alone(self):
from litellm.litellm_core_utils.prompt_templates.common_utils import (
drop_non_python_regex_patterns,
)
leaf = {"type": "string", "pattern": _ARTIFACT_FIELD_PATTERN}
schema = functools.reduce(
lambda inner, _: {"type": "object", "properties": {"child": inner}}, range(1100), leaf
)
assert drop_non_python_regex_patterns(schema) is schema
class TestRequestContainsImageContent:

View file

@ -1,6 +1,7 @@
#### What this tests ####
# This tests litellm.token_counter.token_counter() function
import asyncio
import base64
import importlib
import threading
import time
@ -25,6 +26,8 @@ from litellm.litellm_core_utils.token_counter import (
_get_exact_count_function,
_get_extrapolating_count_function,
_get_tiktoken_count_function,
calculate_img_tokens,
high_detail_image_token_upper_bound,
offload_token_count,
)
from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new
@ -1558,3 +1561,18 @@ def test_openai_file_block_without_inline_bytes_counts_what_it_carries():
assert _count_user_content([prompt, named]) == _count_user_content(
[prompt, {"type": "text", "text": "report.pdf"}]
)
def _png_data_url(width: int, height: int) -> str:
ihdr = b"\x89PNG\r\n\x1a\n" + (13).to_bytes(4, "big") + b"IHDR" + width.to_bytes(4, "big") + height.to_bytes(4, "big")
return "data:image/png;base64," + base64.b64encode(ihdr + b"\x08\x06\x00\x00\x00").decode()
@pytest.mark.parametrize(("width", "height"), [(1, 1), (768, 768), (2000, 768), (768, 2000), (4096, 4096), (8000, 3072)])
def test_high_detail_image_token_upper_bound_covers_every_image_size(width: int, height: int) -> None:
assert calculate_img_tokens(_png_data_url(width, height), mode="high") <= high_detail_image_token_upper_bound()
def test_high_detail_image_token_upper_bound_is_reached_by_the_largest_high_res_image() -> None:
assert calculate_img_tokens(_png_data_url(2000, 768), mode="high") == high_detail_image_token_upper_bound()
assert calculate_img_tokens(_png_data_url(1, 1), mode="high") < high_detail_image_token_upper_bound()

View file

@ -2267,3 +2267,25 @@ def test_create_anthropic_model_list_response_empty():
assert response["has_more"] is False
assert response["first_id"] is None
assert response["last_id"] is None
def test_create_anthropic_model_list_response_lists_ids_as_told():
"""listed_ids renames an entry for the caller while display_name and every other field stay keyed to the served
id, and the envelope's first/last ids follow the renamed entries."""
from litellm.llms.anthropic.common_utils import (
create_anthropic_model_list_response,
)
response = create_anthropic_model_list_response(
[
{"id": "gpt-4o", "object": "model", "created": 0, "owned_by": "openai", "max_input_tokens": 1000000},
{"id": "claude-haiku-4-5", "object": "model", "created": 0, "owned_by": "openai"},
],
display_names={"gpt-4o": "GPT 4o"},
listed_ids={"gpt-4o": "claude-router-gpt-4o[1m]"},
)
gpt, haiku = response["data"]
assert (gpt["id"], gpt["display_name"], gpt["max_input_tokens"]) == ("claude-router-gpt-4o[1m]", "GPT 4o", 1000000)
assert (haiku["id"], haiku["display_name"]) == ("claude-haiku-4-5", "claude-haiku-4-5")
assert (response["first_id"], response["last_id"]) == ("claude-router-gpt-4o[1m]", "claude-haiku-4-5")

View file

@ -198,6 +198,9 @@ def test_azure_gpt_5_takes_the_reasoning_path() -> None:
assert "reasoning_effort" in supported
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
class TestAzureToolSchemaCombinatorFlattening:
"""
Regression tests for LIT-6510: Azure's chat completions validator rejects
@ -259,6 +262,26 @@ class TestAzureToolSchemaCombinatorFlattening:
self._transform(AzureOpenAIConfig(), "gpt-4o", [tool])
assert tool == self._anyof_tool()
def test_transform_request_drops_non_python_regex_pattern(self):
tool = {
"type": "function",
"function": {
"name": "Artifact",
"parameters": {
"type": "object",
"properties": {"field": {"type": "string", "pattern": _ARTIFACT_FIELD_PATTERN}},
},
},
}
request = self._transform(AzureOpenAIConfig(), "gpt-4o", [tool])
assert request["tools"][0]["function"]["parameters"] == {
"type": "object",
"properties": {"field": {"type": "string"}},
}
assert tool["function"]["parameters"]["properties"]["field"]["pattern"] == _ARTIFACT_FIELD_PATTERN
def test_clean_object_schema_passes_through_as_same_object(self):
tool = {
"type": "function",

View file

@ -1,11 +1,20 @@
import json
from datetime import datetime
from unittest.mock import MagicMock
import httpx
import pytest
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
from litellm.types.utils import ModelResponse
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.token_counter import high_detail_image_token_upper_bound
from litellm.llms.azure.passthrough.transformation import (
AzurePassthroughConfig,
azure_router_model_in_endpoint,
foreign_azure_deployment,
)
from litellm.types.llms.openai import ResponseCompletedEvent, ResponsesAPIResponse
from litellm.types.utils import EmbeddingResponse, ModelResponse
def _azure_chat_completion_body():
@ -73,22 +82,408 @@ def test_azure_passthrough_logging_non_streaming_response_chat_completions():
assert result.usage.total_tokens == 18
def test_azure_passthrough_logging_non_streaming_response_unknown_endpoint_returns_none():
"""
Endpoints other than chat/completions (responses, messages, images) fall
through to None — matches base-class behavior and Bedrock's "unknown
endpoint" handling. Not a regression; just scoping.
"""
config = AzurePassthroughConfig()
logging_obj = MagicMock()
result = config.logging_non_streaming_response(
model="gpt-4.1-mini",
def _relay_logging_obj(model: str) -> Logging:
logging_obj = Logging(
model=model,
messages=[],
stream=False,
call_type="allm_passthrough_route",
start_time=datetime.now(),
litellm_call_id="call-1",
function_id="fn-1",
)
logging_obj.update_environment_variables(
model=model,
litellm_params={"api_base": "https://my-resource.openai.azure.com", "custom_llm_provider": "azure"},
optional_params={},
custom_llm_provider="azure",
httpx_response=_make_httpx_response(_azure_chat_completion_body()),
)
return logging_obj
def _relay_logging_result(model: str, endpoint: str, body, status_code: int = 200):
logging_obj = _relay_logging_obj(model)
response = httpx.Response(
status_code=status_code,
headers={"content-type": "application/json"},
content=json.dumps(body).encode("utf-8"),
request=httpx.Request(
"POST", f"https://my-resource.openai.azure.com/{endpoint}?api-version=2025-04-01-preview"
),
)
result = AzurePassthroughConfig().logging_non_streaming_response(
model=model,
custom_llm_provider="azure",
httpx_response=response,
request_data={},
logging_obj=logging_obj,
endpoint="openai/responses",
endpoint=endpoint,
)
return result, logging_obj
EMBEDDINGS_BODY = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 1000, "total_tokens": 1000},
}
RESPONSES_BODY = {
"id": "resp_1",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4.1-mini",
"output": [
{
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"usage": {"input_tokens": 1000, "output_tokens": 100, "total_tokens": 1100},
}
def test_azure_passthrough_embeddings_relay_is_costed_per_input_token():
result, logging_obj = _relay_logging_result(
"text-embedding-3-small", "openai/deployments/text-embedding-3-small/embeddings", EMBEDDINGS_BODY
)
per_token = litellm.get_model_info("azure/text-embedding-3-small")["input_cost_per_token"]
assert isinstance(result, EmbeddingResponse)
assert logging_obj.call_type == "aembedding"
assert per_token > 0
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(1000 * per_token)
def test_azure_passthrough_responses_relay_is_costed_per_token():
result, logging_obj = _relay_logging_result("gpt-4.1-mini", "openai/responses", RESPONSES_BODY)
info = litellm.get_model_info("azure/gpt-4.1-mini")
assert isinstance(result, ResponsesAPIResponse)
assert logging_obj.call_type == "aresponses"
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(
1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"]
)
def test_azure_passthrough_failed_embeddings_relay_is_not_costed():
result, logging_obj = _relay_logging_result(
"text-embedding-3-small",
"openai/deployments/text-embedding-3-small/embeddings",
{"error": {"code": "429", "message": "rate limited"}},
status_code=429,
)
assert result is None
assert logging_obj.call_type == "allm_passthrough_route"
def test_azure_passthrough_logging_non_streaming_response_unknown_endpoint_returns_none():
result, logging_obj = _relay_logging_result(
"gpt-4o-mini-tts", "openai/deployments/gpt-4o-mini-tts/audio/speech", {"audio": "..."}
)
assert result is None
assert logging_obj.call_type == "allm_passthrough_route"
def _sse_line(payload: dict) -> str:
return "data: " + json.dumps(payload)
def _azure_chat_completion_chunks() -> list[str]:
head = {"id": "chatcmpl-abc123", "object": "chat.completion.chunk", "created": 1700000000, "model": "gpt-4.1-mini"}
return [
_sse_line(
{
**head,
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "Hello!"}, "finish_reason": None}],
}
),
_sse_line(
{**head, "choices": [{"index": 0, "delta": {"content": " How can I assist?"}, "finish_reason": None}]}
),
_sse_line({**head, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}),
_sse_line({**head, "choices": [], "usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18}}),
"data: [DONE]",
]
def test_azure_passthrough_streaming_chat_chunks_build_the_complete_response():
response = AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=_azure_chat_completion_chunks(),
litellm_logging_obj=MagicMock(),
model="gpt-4.1-mini",
custom_llm_provider="azure",
endpoint="openai/deployments/gpt-4.1-mini/chat/completions",
)
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "Hello! How can I assist?"
assert response.usage.prompt_tokens == 10
assert response.usage.completion_tokens == 8
def test_azure_passthrough_streaming_chunks_without_usage_count_prompt_tokens_from_the_relayed_request():
messages = [{"role": "user", "content": "Say hi in three words"}]
logging_obj = MagicMock()
logging_obj.model_call_details = {"request_data": {"messages": messages, "stream": True}}
response = AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=[chunk for chunk in _azure_chat_completion_chunks() if '"usage"' not in chunk],
litellm_logging_obj=logging_obj,
model="gpt-4.1-mini",
custom_llm_provider="azure",
endpoint="openai/deployments/gpt-4.1-mini/chat/completions",
)
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "Hello! How can I assist?"
assert response.usage.prompt_tokens > 0
assert response.usage.prompt_tokens == litellm.token_counter(model="gpt-4.1-mini", messages=messages)
assert response.usage.completion_tokens > 0
def test_azure_passthrough_streaming_chunks_count_remote_image_prompt_tokens_without_fetching_the_image():
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this"},
{"type": "image_url", "image_url": {"url": "http://127.0.0.1:9/doc.png", "detail": "high"}},
],
}
]
logging_obj = MagicMock()
logging_obj.model_call_details = {"request_data": {"messages": messages, "stream": True}}
response = AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=[chunk for chunk in _azure_chat_completion_chunks() if '"usage"' not in chunk],
litellm_logging_obj=logging_obj,
model="gpt-4.1-mini",
custom_llm_provider="azure",
endpoint="openai/deployments/gpt-4.1-mini/chat/completions",
)
text_only_messages = [{"role": "user", "content": [{"type": "text", "text": "Describe this"}]}]
assert isinstance(response, ModelResponse)
assert response.usage.prompt_tokens == (
litellm.token_counter(model="gpt-4.1-mini", messages=text_only_messages) + high_detail_image_token_upper_bound()
)
def test_azure_passthrough_streaming_chunks_for_unknown_endpoint_return_none():
response = AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=_azure_chat_completion_chunks(),
litellm_logging_obj=MagicMock(),
model="gpt-4.1-mini",
custom_llm_provider="azure",
endpoint="openai/deployments/gpt-4.1-mini/embeddings",
)
assert response is None
def _azure_responses_stream_chunks(terminal_event: str | None = "response.completed") -> list[str]:
in_progress = {**RESPONSES_BODY, "status": "in_progress", "output": [], "usage": None}
events = [
("response.created", {"type": "response.created", "sequence_number": 0, "response": in_progress}),
(
"response.output_text.delta",
{"type": "response.output_text.delta", "sequence_number": 1, "item_id": "msg_1", "delta": "hi"},
),
] + (
[(terminal_event, {"type": terminal_event, "sequence_number": 2, "response": RESPONSES_BODY})]
if terminal_event
else []
)
return [line for name, payload in events for line in (f"event: {name}", _sse_line(payload))]
def test_azure_passthrough_streaming_responses_chunks_are_costed_per_token():
logging_obj = _relay_logging_obj("gpt-4.1-mini")
response = AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=_azure_responses_stream_chunks(),
litellm_logging_obj=logging_obj,
model="gpt-4.1-mini",
custom_llm_provider="azure",
endpoint="openai/responses",
)
info = litellm.get_model_info("azure/gpt-4.1-mini")
assert isinstance(response, ResponseCompletedEvent)
assert response.response.usage.input_tokens == 1000
assert logging_obj.call_type == "aresponses"
assert logging_obj._response_cost_calculator(result=response.response) == pytest.approx(
1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"]
)
def test_azure_passthrough_streaming_responses_without_a_terminal_event_are_not_costed():
logging_obj = _relay_logging_obj("gpt-4.1-mini")
response = AzurePassthroughConfig().handle_logging_collected_chunks(
all_chunks=_azure_responses_stream_chunks(terminal_event=None),
litellm_logging_obj=logging_obj,
model="gpt-4.1-mini",
custom_llm_provider="azure",
endpoint="openai/responses",
)
assert response is None
assert logging_obj.call_type == "allm_passthrough_route"
def _complete_url(request_query_params: dict, litellm_params: dict) -> httpx.URL:
url, _ = AzurePassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com",
api_key="key",
model="gpt-4.1-mini",
endpoint="openai/deployments/gpt-4.1-mini/chat/completions",
request_query_params=request_query_params,
litellm_params=litellm_params,
)
return url
def test_azure_passthrough_url_forwards_the_callers_api_version():
url = _complete_url(request_query_params={"api-version": "2025-04-01-preview"}, litellm_params={})
assert url.path == "/openai/deployments/gpt-4.1-mini/chat/completions"
assert url.params["api-version"] == "2025-04-01-preview"
def test_azure_passthrough_url_prefers_the_callers_api_version_over_the_deployments():
url = _complete_url(
request_query_params={"api-version": "2025-04-01-preview"}, litellm_params={"api_version": "2024-10-21"}
)
assert url.params["api-version"] == "2025-04-01-preview"
def test_azure_passthrough_url_fills_in_the_deployments_api_version_when_the_caller_sends_none():
url = _complete_url(request_query_params={}, litellm_params={"api_version": "2024-10-21"})
assert url.params["api-version"] == "2024-10-21"
FULL_URL_API_BASE = (
"https://my-resource.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions?api-version=2024-10-21"
)
def _full_url_complete_url(request_query_params: dict) -> httpx.URL:
url, _ = AzurePassthroughConfig().get_complete_url(
api_base=FULL_URL_API_BASE,
api_key="key",
model="gpt-4.1-mini",
endpoint="chat/completions",
request_query_params=request_query_params,
litellm_params={},
)
return url
def test_azure_passthrough_url_prefers_the_callers_api_version_over_a_full_url_api_bases():
url = _full_url_complete_url(request_query_params={"api-version": "2025-04-01-preview"})
assert str(url) == (
"https://my-resource.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions"
"?api-version=2025-04-01-preview"
)
def test_azure_passthrough_url_keeps_a_full_url_api_bases_api_version_when_the_caller_sends_none():
url = _full_url_complete_url(request_query_params={})
assert url.params["api-version"] == "2024-10-21"
def test_azure_passthrough_url_strips_the_leading_router_model_segment():
url, _ = AzurePassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com",
api_key="key",
model="gpt-4.1-mini",
endpoint="gpt-4.1-mini/openai/deployments/gpt-4.1-mini/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={},
)
assert (
str(url)
== "https://my-resource.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions?api-version=2024-10-21"
)
def test_azure_passthrough_url_rewrites_the_model_group_only_as_a_whole_segment():
url, _ = AzurePassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com",
api_key="key",
model="gpt-4.1-mini",
endpoint="gpt/openai/deployments/gpt-4.1-mini/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={"litellm_metadata": {"model_group": "gpt"}},
)
assert (
str(url)
== "https://my-resource.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions?api-version=2024-10-21"
)
@pytest.mark.parametrize(
"request_data, expected",
[({"stream": True}, True), ({"stream": 1}, True), ({"stream": False}, False), ({}, False)],
)
def test_azure_passthrough_is_streaming_request_reads_the_stream_flag(request_data, expected):
assert (
AzurePassthroughConfig().is_streaming_request(
endpoint="openai/deployments/x/chat/completions", request_data=request_data
)
is expected
)
@pytest.mark.parametrize(
"endpoint, expected",
[
("gpt/openai/deployments/gpt/chat/completions", None),
("openai/deployments/gpt/chat/completions", None),
("gpt/openai/deployments/gpt-5.4-mini/chat/completions", None),
("gpt/openai/deployments/GPT-5.4-MINI/chat/completions", None),
("gpt/openai/deployments/Gpt/chat/completions", "Gpt"),
("gpt/models/chat/completions", None),
("gpt/openai/deployments/gpt-5.4/chat/completions", "gpt-5.4"),
("gpt/openai/deployments/other-group/chat/completions", "other-group"),
("openai/deployments/victim/gpt/chat/completions", "victim"),
("gpt/openai/deployments/GPT-5.4/chat/completions", "GPT-5.4"),
],
)
def test_foreign_azure_deployment_names_a_segment_outside_the_group(endpoint, expected):
assert foreign_azure_deployment(endpoint, "gpt", lambda: frozenset({"gpt-5.4-mini"})) == expected
def test_foreign_azure_deployment_skips_the_router_when_the_segment_is_the_group_itself():
def served_models():
raise AssertionError("the router must not be consulted for the group's own name")
assert foreign_azure_deployment("gpt/openai/deployments/gpt/chat/completions", "gpt", served_models) is None
@pytest.mark.parametrize(
"endpoint, expected",
[
("other-group/openai/deployments/other-group/chat/completions", "other-group"),
("openai/deployments/gpt/chat/completions", "gpt"),
("openai/deployments/my-azure-deployment/chat/completions", None),
("gpt", None),
],
)
def test_azure_router_model_in_endpoint_picks_the_first_router_model_segment(endpoint, expected):
assert azure_router_model_in_endpoint(endpoint, frozenset({"gpt", "other-group"})) == expected

View file

@ -1,4 +1,5 @@
from copy import deepcopy
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
@ -243,6 +244,9 @@ def test_provider_config_manager_o_series_selection():
assert not isinstance(default_config, AzureOpenAIOSeriesResponsesAPIConfig)
_ARTIFACT_FIELD_PATTERN: Final = r'^(?!__.*__$)[^\p{Cc}\p{Cf}\p{Zl}\p{Zp}"\\./[\]]{1,200}$'
class TestAzureResponsesAPIConfig:
def setup_method(self):
self.config = AzureOpenAIResponsesAPIConfig()
@ -599,6 +603,31 @@ class TestAzureResponsesAPIConfig:
assert result["tools"][0] is tool
assert "anyOf" in result["tools"][0]["parameters"]
def test_azure_drops_non_python_regex_pattern_while_keeping_gpt5_combinators(self):
tool = {
"type": "function",
"name": "Artifact",
"parameters": {
"type": "object",
"anyOf": [{"properties": {"field": {"type": "string", "pattern": _ARTIFACT_FIELD_PATTERN}}}],
"properties": {"field": {"type": "string", "pattern": _ARTIFACT_FIELD_PATTERN}},
},
}
result = self.config.transform_responses_api_request(
model="my-eastus-deployment",
input="hi",
response_api_optional_request_params={"tools": [tool]},
litellm_params=GenericLiteLLMParams(model_info={"base_model": "azure/gpt-5.4-mini"}),
headers={},
)
assert result["tools"][0]["parameters"] == {
"type": "object",
"anyOf": [{"properties": {"field": {"type": "string"}}}],
"properties": {"field": {"type": "string"}},
}
def test_azure_keeps_combinators_for_unrecognized_deployment_without_base_model(self):
tool = self._anyof_tool()

View file

@ -0,0 +1,617 @@
import json
from datetime import datetime
from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.azure_ai.passthrough.transformation import AzureAIPassthroughConfig
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.types.rerank import RerankResponse
from litellm.types.utils import EmbeddingResponse, ImageResponse, LlmProviders, ModelResponse
from litellm.utils import ProviderConfigManager
FOUNDRY_BASE = "https://my-resource.services.ai.azure.com"
RESPONSES_COMPLETED_EVENT = {
"type": "response.completed",
"sequence_number": 2,
"response": {
"id": "resp_1",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-5.4-mini",
"output": [
{
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"usage": {"input_tokens": 1000, "output_tokens": 100, "total_tokens": 1100},
},
}
class _SpendProbe(CustomLogger):
logged_call_type: str | None = None
logged_cost: float | None = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.logged_call_type = kwargs["call_type"]
self.logged_cost = kwargs["response_cost"]
@pytest.fixture(autouse=True)
def clear_azure_ai_env(monkeypatch):
for env_var in ("AZURE_AI_API_BASE", "AZURE_AI_API_KEY", "AZURE_AD_TOKEN", "AZURE_API_KEY"):
monkeypatch.delenv(env_var, raising=False)
monkeypatch.setattr(litellm, "api_base", None)
monkeypatch.setattr(litellm, "api_key", None)
def test_provider_config_manager_resolves_azure_ai_passthrough_config():
config = ProviderConfigManager.get_provider_passthrough_config(
model="Cohere-parse-v5", provider=LlmProviders.AZURE_AI
)
assert isinstance(config, AzureAIPassthroughConfig)
def test_router_model_prefix_is_stripped_and_native_path_kept_verbatim():
url, base = AzureAIPassthroughConfig().get_complete_url(
api_base=FOUNDRY_BASE,
api_key=None,
model="Cohere-parse-v5",
endpoint="Cohere-parse-v5/providers/cohere/v2/parse",
request_query_params=None,
litellm_params={},
)
assert str(url) == f"{FOUNDRY_BASE}/providers/cohere/v2/parse"
assert base == FOUNDRY_BASE
def test_model_group_prefix_is_stripped_when_router_metadata_names_it():
url, _ = AzureAIPassthroughConfig().get_complete_url(
api_base=FOUNDRY_BASE,
api_key=None,
model="Cohere-parse-v5",
endpoint="/parse-alias/providers/cohere/v2/parse",
request_query_params=None,
litellm_params={"litellm_metadata": {"model_group": "parse-alias"}},
)
assert str(url) == f"{FOUNDRY_BASE}/providers/cohere/v2/parse"
def test_model_inside_the_path_stays_and_query_params_are_forwarded():
url, _ = AzureAIPassthroughConfig().get_complete_url(
api_base=f"{FOUNDRY_BASE}/",
api_key=None,
model="gpt-5.4-mini",
endpoint="openai/deployments/gpt-5.4-mini/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={},
)
assert str(url) == f"{FOUNDRY_BASE}/openai/deployments/gpt-5.4-mini/chat/completions?api-version=2024-10-21"
def test_api_base_that_already_ends_in_models_is_cut_back_to_the_foundry_root():
url, base = AzureAIPassthroughConfig().get_complete_url(
api_base=f"{FOUNDRY_BASE}/models",
api_key="key",
model="gpt-5.4-mini",
endpoint="gpt-5.4-mini/models/chat/completions",
request_query_params={"api-version": "2024-05-01-preview"},
litellm_params={},
)
assert str(url) == f"{FOUNDRY_BASE}/models/chat/completions?api-version=2024-05-01-preview"
assert base == FOUNDRY_BASE
def test_full_url_api_base_that_already_ends_with_the_native_path_is_not_doubled():
model_router_url = (
"https://my-resource.cognitiveservices.azure.com/openai/deployments/model-router/chat/completions"
)
url, base = AzureAIPassthroughConfig().get_complete_url(
api_base=f"{model_router_url}?api-version=2025-01-01-preview",
api_key="key",
model="model_router/model-router",
endpoint="model-router/chat/completions",
request_query_params=None,
litellm_params={"litellm_metadata": {"model_group": "model-router"}},
)
assert str(url) == f"{model_router_url}?api-version=2025-01-01-preview"
assert base == "https://my-resource.cognitiveservices.azure.com/openai/deployments/model-router"
@pytest.mark.parametrize("relayed_deployment", ["gpt-4o", "GPT-4o"])
def test_deployment_root_api_base_is_not_repeated_when_the_relay_carries_the_deployment_path(relayed_deployment):
url, base = AzureAIPassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com/openai/deployments/gpt-4o",
api_key="key",
model="gpt-4o",
endpoint=f"aoai-gpt-4o/openai/deployments/{relayed_deployment}/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={"litellm_metadata": {"model_group": "aoai-gpt-4o"}},
)
assert str(url) == (
f"https://my-resource.openai.azure.com/openai/deployments/{relayed_deployment}/chat/completions"
"?api-version=2024-10-21"
)
assert base == "https://my-resource.openai.azure.com"
def test_deployment_named_like_the_first_native_segment_keeps_its_deployment_root():
url, base = AzureAIPassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com/openai/deployments/chat",
api_key="key",
model="chat",
endpoint="aoai-chat/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={"litellm_metadata": {"model_group": "aoai-chat"}},
)
assert str(url) == "https://my-resource.openai.azure.com/openai/deployments/chat/chat/completions?api-version=2024-10-21"
assert base == "https://my-resource.openai.azure.com/openai/deployments/chat"
def test_parse_relay_under_a_models_api_base_targets_the_foundry_root():
url, _ = AzureAIPassthroughConfig().get_complete_url(
api_base=f"{FOUNDRY_BASE}/models",
api_key="key",
model="Cohere-parse-v5",
endpoint="Cohere-parse-v5/providers/cohere/v2/parse",
request_query_params=None,
litellm_params={},
)
assert str(url) == f"{FOUNDRY_BASE}/providers/cohere/v2/parse"
def test_deployment_api_version_fills_in_when_the_caller_sends_none():
url, _ = AzureAIPassthroughConfig().get_complete_url(
api_base=FOUNDRY_BASE,
api_key="key",
model="gpt-5.4-mini",
endpoint="gpt-5.4-mini/models/chat/completions",
request_query_params=None,
litellm_params={"api_version": "2024-05-01-preview"},
)
assert str(url) == f"{FOUNDRY_BASE}/models/chat/completions?api-version=2024-05-01-preview"
def test_callers_api_version_beats_the_deployments():
url, _ = AzureAIPassthroughConfig().get_complete_url(
api_base=FOUNDRY_BASE,
api_key="key",
model="gpt-5.4-mini",
endpoint="gpt-5.4-mini/models/chat/completions",
request_query_params={"api-version": "2025-04-01-preview"},
litellm_params={"api_version": "2024-05-01-preview"},
)
assert str(url) == f"{FOUNDRY_BASE}/models/chat/completions?api-version=2025-04-01-preview"
def test_api_version_on_the_configured_api_base_is_the_last_fallback():
url, _ = AzureAIPassthroughConfig().get_complete_url(
api_base=f"{FOUNDRY_BASE}/models/chat/completions?api-version=2024-05-01-preview",
api_key="key",
model="gpt-5.4-mini",
endpoint="gpt-5.4-mini/models/chat/completions",
request_query_params=None,
litellm_params={},
)
assert str(url) == f"{FOUNDRY_BASE}/models/chat/completions?api-version=2024-05-01-preview"
def test_missing_api_base_raises_instead_of_building_a_relative_url():
with pytest.raises(ValueError, match="AZURE_AI_API_BASE"):
AzureAIPassthroughConfig().get_complete_url(
api_base=None,
api_key=None,
model="Cohere-parse-v5",
endpoint="Cohere-parse-v5/providers/cohere/v2/parse",
request_query_params=None,
litellm_params={},
)
def _auth_headers(api_key: str | None, api_base: str, litellm_params: dict | None = None) -> dict:
return AzureAIPassthroughConfig().validate_environment(
headers={"content-type": "application/json"},
model="Cohere-parse-v5",
messages=[],
optional_params={},
litellm_params=litellm_params or {},
api_key=api_key,
api_base=api_base,
)
def test_foundry_host_gets_the_api_key_header():
headers = _auth_headers(api_key="deployment-key", api_base=FOUNDRY_BASE)
assert headers == {"content-type": "application/json", "api-key": "deployment-key"}
def test_serverless_host_gets_a_bearer_token():
headers = _auth_headers(api_key="deployment-key", api_base="https://cohere-parse.eastus.models.ai.azure.com")
assert headers["Authorization"] == "Bearer deployment-key"
assert "api-key" not in headers
def test_entra_token_is_used_when_the_deployment_has_no_api_key():
headers = _auth_headers(api_key=None, api_base=FOUNDRY_BASE, litellm_params={"azure_ad_token": "entra-token"})
assert headers["Authorization"] == "Bearer entra-token"
def test_no_credentials_at_all_raises():
with pytest.raises(ValueError, match="Missing Azure AI credentials"):
_auth_headers(api_key=None, api_base=FOUNDRY_BASE)
@pytest.mark.parametrize(
"request_data, expected",
[({"stream": True}, True), ({"stream": 1}, True), ({"stream": False}, False), ({}, False)],
)
def test_is_streaming_request_reads_the_stream_flag(request_data, expected):
assert (
AzureAIPassthroughConfig().is_streaming_request(endpoint="models/chat/completions", request_data=request_data)
is expected
)
def _chat_completion_response() -> httpx.Response:
body = {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-5.4-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 8, "total_tokens": 18},
}
return httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
content=json.dumps(body).encode("utf-8"),
request=httpx.Request("POST", f"{FOUNDRY_BASE}/models/chat/completions"),
)
def test_chat_completions_relay_yields_a_model_response_for_cost_tracking():
result = AzureAIPassthroughConfig().logging_non_streaming_response(
model="gpt-5.4-mini",
custom_llm_provider="azure_ai",
httpx_response=_chat_completion_response(),
request_data={"model": "gpt-5.4-mini", "messages": [{"role": "user", "content": "hi"}]},
logging_obj=MagicMock(),
endpoint="models/chat/completions",
)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "hi"
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 8
def _non_chat_logging_result(content: bytes, content_type: str):
parse_response = httpx.Response(
status_code=200,
headers={"content-type": content_type},
content=content,
request=httpx.Request("POST", f"{FOUNDRY_BASE}/providers/cohere/v2/parse"),
)
return AzureAIPassthroughConfig().logging_non_streaming_response(
model="Cohere-parse-v5",
custom_llm_provider="azure_ai",
httpx_response=parse_response,
request_data={"model": "Cohere-parse-v5"},
logging_obj=MagicMock(),
endpoint="providers/cohere/v2/parse",
)
def test_non_chat_relay_with_a_non_json_body_logs_the_raw_text():
assert _non_chat_logging_result(b"page one", "text/plain") == {"response": "page one"}
def _relay_logging_obj(
model: str,
api_base: str,
stream: bool = False,
callbacks: list[CustomLogger] | None = None,
endpoint: str = "",
) -> Logging:
logging_obj = Logging(
model=model,
messages=[],
stream=stream,
call_type="allm_passthrough_route",
start_time=datetime.now(),
litellm_call_id="call-1",
function_id="fn-1",
dynamic_async_success_callbacks=callbacks,
)
logging_obj.update_environment_variables(
model=model,
litellm_params={"api_base": api_base, "custom_llm_provider": "azure_ai"},
optional_params={},
custom_llm_provider="azure_ai",
endpoint=endpoint,
)
return logging_obj
def _relay_logging_result(
config: AzureAIPassthroughConfig,
model: str,
native_path: str,
body,
api_base: str = FOUNDRY_BASE,
status_code: int = 200,
):
relayed_url = f"{FOUNDRY_BASE}/{native_path}?api-version=2024-05-01-preview"
logging_obj = _relay_logging_obj(model, api_base)
response = httpx.Response(
status_code=status_code,
headers={"content-type": "application/json"},
content=json.dumps(body).encode("utf-8"),
request=httpx.Request("POST", relayed_url),
)
result = config.logging_non_streaming_response(
model=model,
custom_llm_provider="azure_ai",
httpx_response=response,
request_data={"model": model},
logging_obj=logging_obj,
endpoint=f"{model}/{native_path}",
)
return result, logging_obj
MISTRAL_OCR_BODY = {
"pages": [{"index": 0, "markdown": "page one"}, {"index": 1, "markdown": "page two"}],
"model": "mistral-document-ai-2512",
"usage_info": {"pages_processed": 2, "doc_size_bytes": 4321},
}
def test_mistral_document_ai_relay_is_costed_per_page():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "mistral-document-ai-2512", "providers/mistral/azure/ocr", MISTRAL_OCR_BODY
)
per_page = litellm.get_model_info("azure_ai/mistral-document-ai-2512")["ocr_cost_per_page"]
assert isinstance(result, OCRResponse)
assert result.usage_info.pages_processed == 2
assert per_page > 0
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(2 * per_page)
def test_ocr_route_under_a_models_api_base_is_still_recognised():
result, _ = _relay_logging_result(
AzureAIPassthroughConfig(),
"mistral-document-ai-2512",
"providers/mistral/azure/ocr",
MISTRAL_OCR_BODY,
api_base=f"{FOUNDRY_BASE}/models",
)
assert isinstance(result, OCRResponse)
def test_relay_to_a_non_ocr_route_keeps_the_passthrough_object_and_call_type():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "mistral-document-ai-2512", "models/info", {"name": "mistral-document-ai-2512"}
)
assert result == {"response": {"name": "mistral-document-ai-2512"}}
assert logging_obj.call_type == "allm_passthrough_route"
COHERE_PARSE_BODY = {"id": "parse-1", "pages": [], "meta": {"billed_units": {"pages": 3}}}
def test_cohere_parse_relay_is_costed_per_billed_page():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "Cohere-parse-v5", "providers/cohere/v2/parse", COHERE_PARSE_BODY
)
per_page = litellm.get_model_info("azure_ai/Cohere-parse-v5")["ocr_cost_per_page"]
assert isinstance(result, OCRResponse)
assert result.usage_info.pages_processed == 3
assert logging_obj.call_type == "aocr"
assert per_page > 0
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(3 * per_page)
def test_deployment_without_an_ocr_config_is_never_costed_as_ocr():
config = AzureAIPassthroughConfig(ocr_config_for=lambda model: None)
result, logging_obj = _relay_logging_result(
config, "mistral-document-ai-2512", "providers/mistral/azure/ocr", MISTRAL_OCR_BODY
)
assert result == {"response": MISTRAL_OCR_BODY}
assert logging_obj.call_type == "allm_passthrough_route"
def test_accepted_ocr_job_without_a_result_body_is_not_costed():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(),
"mistral-document-ai-2512",
"providers/mistral/azure/ocr",
{"status": "running"},
status_code=202,
)
assert result == {"response": {"status": "running"}}
assert logging_obj.call_type == "allm_passthrough_route"
def test_unparseable_ocr_body_falls_back_to_the_passthrough_object():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(),
"mistral-document-ai-2512",
"providers/mistral/azure/ocr",
["not", "an", "ocr", "body"],
)
assert result == {"response": '["not", "an", "ocr", "body"]'}
assert logging_obj.call_type == "allm_passthrough_route"
EMBEDDINGS_BODY = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
"model": "embed-v-4-0",
"usage": {"prompt_tokens": 1200, "total_tokens": 1200},
}
RERANK_BODY = {
"id": "rerank-1",
"results": [{"index": 1, "relevance_score": 0.9}, {"index": 0, "relevance_score": 0.2}],
"meta": {"api_version": {"version": "2"}, "billed_units": {"search_units": 2}},
}
IMAGE_BODY = {"created": 1, "data": [{"b64_json": "AAAA"}]}
def test_foundry_embeddings_relay_is_costed_per_input_token():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "embed-v-4-0", "models/embeddings", EMBEDDINGS_BODY
)
per_token = litellm.get_model_info("azure_ai/embed-v-4-0")["input_cost_per_token"]
assert isinstance(result, EmbeddingResponse)
assert logging_obj.call_type == "aembedding"
assert per_token > 0
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(1200 * per_token)
def test_cohere_rerank_relay_is_costed_per_search_unit():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "cohere-rerank-v4.0-fast", "providers/cohere/v2/rerank", RERANK_BODY
)
per_query = litellm.get_model_info("azure_ai/cohere-rerank-v4.0-fast")["input_cost_per_query"]
assert isinstance(result, RerankResponse)
assert logging_obj.call_type == "arerank"
assert per_query > 0
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(2 * per_query)
def test_image_generation_relay_is_costed_per_image():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "FLUX.2-pro", "openai/deployments/FLUX.2-pro/images/generations", IMAGE_BODY
)
per_image = litellm.get_model_info("azure_ai/FLUX.2-pro")["output_cost_per_image"]
assert isinstance(result, ImageResponse)
assert logging_obj.call_type == "aimage_generation"
assert per_image > 0
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(per_image)
def test_flux_2_relay_through_the_provider_route_is_costed_per_image():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(), "FLUX.2-pro", "providers/blackforestlabs/v1/flux-2-pro", IMAGE_BODY
)
per_image = litellm.get_model_info("azure_ai/FLUX.2-pro")["output_cost_per_image"]
assert isinstance(result, ImageResponse)
assert logging_obj.call_type == "aimage_generation"
assert logging_obj._response_cost_calculator(result=result) == pytest.approx(per_image)
def test_rejected_rerank_relay_keeps_the_passthrough_object_and_call_type():
result, logging_obj = _relay_logging_result(
AzureAIPassthroughConfig(),
"cohere-rerank-v4.0-fast",
"providers/cohere/v2/rerank",
{"message": "invalid request"},
status_code=400,
)
assert result == {"response": {"message": "invalid request"}}
assert logging_obj.call_type == "allm_passthrough_route"
def test_streaming_chat_completion_chunks_are_costed_like_azure():
head = {"id": "chatcmpl-1", "object": "chat.completion.chunk", "created": 1, "model": "gpt-5.4-mini"}
chunks = [
"data: "
+ json.dumps(
{
**head,
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
}
),
"data: "
+ json.dumps({**head, "choices": [], "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}}),
"data: [DONE]",
]
response = AzureAIPassthroughConfig().handle_logging_collected_chunks(
all_chunks=chunks,
litellm_logging_obj=MagicMock(),
model="gpt-5.4-mini",
custom_llm_provider="azure_ai",
endpoint="chat/completions",
)
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content == "hi"
assert response.usage.total_tokens == 4
def test_streaming_responses_chunks_through_a_router_relay_are_costed_like_azure():
logging_obj = _relay_logging_obj("gpt-5.4-mini", FOUNDRY_BASE)
response = AzureAIPassthroughConfig().handle_logging_collected_chunks(
all_chunks=["event: response.completed", "data: " + json.dumps(RESPONSES_COMPLETED_EVENT)],
litellm_logging_obj=logging_obj,
model="gpt-5.4-mini",
custom_llm_provider="azure_ai",
endpoint="gpt/openai/responses",
)
info = litellm.get_model_info("azure_ai/gpt-5.4-mini")
assert response is not None
assert response.response.usage.output_tokens == 100
assert logging_obj.call_type == "aresponses"
assert logging_obj._response_cost_calculator(result=response.response) == pytest.approx(
1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"]
)
async def test_streaming_responses_relay_flush_reaches_the_success_callbacks_with_a_price():
probe = _SpendProbe()
logging_obj = _relay_logging_obj(
"gpt-5.4-mini", FOUNDRY_BASE, stream=True, callbacks=[probe], endpoint="gpt/openai/responses"
)
stream = "event: response.completed\ndata: " + json.dumps(RESPONSES_COMPLETED_EVENT) + "\n\n"
await logging_obj.async_flush_passthrough_collected_chunks(
raw_bytes=[stream.encode()], provider_config=AzureAIPassthroughConfig()
)
info = litellm.get_model_info("azure_ai/gpt-5.4-mini")
assert probe.logged_call_type == "allm_passthrough_route"
assert probe.logged_cost == pytest.approx(1000 * info["input_cost_per_token"] + 100 * info["output_cost_per_token"])

View file

@ -8,7 +8,7 @@ the tests don't hit AWS.
from __future__ import annotations
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock, patch
import pytest
@ -480,3 +480,93 @@ def test_litellm_cancel_batch_dispatches_to_bedrock(patched_boto3):
fake_client.stop_model_invocation_job.assert_called_once_with(jobIdentifier=JOB_ARN)
assert batch.status == "cancelled"
class _TagGatedSTSClient:
"""Stands in for STS behind a trust policy that only admits sessions carrying ``tags``."""
def __init__(self, tags: list[dict[str, str]], access_key_id: str) -> None:
self._tags = tags
self._access_key_id = access_key_id
def get_caller_identity(self):
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
def assume_role(self, **params):
from botocore.exceptions import ClientError
if list(params.get("Tags", ())) != self._tags:
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:TagSession"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": self._access_key_id,
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.now(timezone.utc) + timedelta(minutes=30),
}
}
def test_handle_model_invocation_job_status_builds_the_client_from_the_tagged_session(monkeypatch):
"""Status polling must assume the role with the deployment's session tags, like every other call."""
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False)
monkeypatch.delenv("AWS_ROLE_ARN", raising=False)
tags = [{"Key": "team", "Value": "genai"}]
bedrock_client_kwargs: list[dict] = []
fake_bedrock = MagicMock()
fake_bedrock.get_model_invocation_job.return_value = _fake_boto3_response()
def boto3_client(service_name, **kwargs):
if service_name == "sts":
return _TagGatedSTSClient(tags, "ASIABATCHSTATUSTAGGED")
bedrock_client_kwargs.append(kwargs)
return fake_bedrock
with patch("boto3.client", side_effect=boto3_client):
batch = BedrockBatchesHandler._handle_model_invocation_job_status(
batch_id=JOB_ARN,
aws_access_key_id="AKIABATCHSTATUSCALLER",
aws_secret_access_key="pod-caller-secret",
aws_role_name="arn:aws:iam::999999999999:role/litellm-batch-role",
aws_session_name="litellm-batch-session",
aws_session_tags=tags,
)
assert batch.status == "completed"
assert [kwargs["aws_access_key_id"] for kwargs in bedrock_client_kwargs] == ["ASIABATCHSTATUSTAGGED"]
def test_cancel_batch_stops_and_polls_the_job_with_the_tagged_session(monkeypatch):
"""Cancelling on a tag-gated role must forward the deployment's session tags to both the stop and status calls."""
import litellm
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False)
monkeypatch.delenv("AWS_ROLE_ARN", raising=False)
tags = [{"Key": "team", "Value": "genai"}]
bedrock_client_kwargs: list[dict] = []
fake_bedrock = MagicMock()
fake_bedrock.get_model_invocation_job.return_value = _fake_boto3_response(status="Stopped")
def boto3_client(service_name, **kwargs):
if service_name == "sts":
return _TagGatedSTSClient(tags, "ASIABATCHCANCELTAGGED")
bedrock_client_kwargs.append(kwargs)
return fake_bedrock
with patch("boto3.client", side_effect=boto3_client):
batch = litellm.cancel_batch(
batch_id=JOB_ARN,
custom_llm_provider="bedrock",
aws_access_key_id="AKIABATCHCANCELCALLER",
aws_secret_access_key="pod-caller-secret",
aws_role_name="arn:aws:iam::999999999999:role/litellm-batch-role",
aws_session_name="litellm-batch-session",
aws_session_tags=tags,
)
fake_bedrock.stop_model_invocation_job.assert_called_once_with(jobIdentifier=JOB_ARN)
assert batch.status == "cancelled"
assert [kwargs["aws_access_key_id"] for kwargs in bedrock_client_kwargs] == ["ASIABATCHCANCELTAGGED"] * 2

View file

@ -15,6 +15,7 @@ AWS_AUTH_PARAMS = {
"aws_sts_endpoint": "https://sts.us-west-2.amazonaws.com",
"aws_bedrock_runtime_endpoint": "https://bedrock-runtime.us-west-2.amazonaws.com",
"aws_external_id": "external",
"aws_session_tags": [{"Key": "team", "Value": "genai"}],
}

View file

@ -57,6 +57,7 @@ def test_aws_params_filtered_from_request_body():
"aws_sts_endpoint": "https://sts.amazonaws.com",
"aws_bedrock_runtime_endpoint": "https://bedrock-runtime.us-west-2.amazonaws.com",
"aws_external_id": "external-id-123",
"aws_session_tags": [{"Key": "team", "Value": "genai"}],
}
# Transform the request
@ -105,6 +106,9 @@ def test_aws_params_filtered_from_request_body():
assert (
"aws_external_id" not in result_json
), "AWS external ID should not be in request body"
assert (
"aws_session_tags" not in result_json
), "AWS session tags should not be in request body"
# Also check that the sensitive values themselves are not in the response
assert (

View file

@ -7,12 +7,15 @@ extension, and AWS credential resolution is stubbed so nothing reaches STS.
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta, timezone
from unittest.mock import MagicMock, patch
import boto3
import httpx
import pytest
from botocore.credentials import Credentials
from botocore.exceptions import ClientError
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.rust_bridge import chat_completions as bridge
@ -562,3 +565,54 @@ def test_bearer_token_auth_never_runs_the_sigv4_credential_chain(monkeypatch, co
assert response.choices[0].message.content == "hi"
assert client.post.call_args.kwargs["headers"]["Authorization"] == "Bearer bedrock-bearer-token"
def test_session_tags_sign_the_request_and_stay_out_of_the_body(monkeypatch):
"""The tagged STS session signs the Converse call and the tags never reach the request body (#34069)."""
monkeypatch.setenv("LITELLM_RUST", "0")
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False)
monkeypatch.delenv("AWS_ROLE_ARN", raising=False)
tags = [{"Key": "team", "Value": "genai"}]
class FakeSTSClient:
def get_caller_identity(self):
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
def assume_role(self, **params):
if list(params.get("Tags", ())) != tags:
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:TagSession"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": "ASIACONVERSETAGGED",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session-token",
"Expiration": datetime.now(timezone.utc) + timedelta(minutes=30),
}
}
client = _sync_client_returning_converse_response()
with patch.object(boto3, "client", return_value=FakeSTSClient()):
response = BedrockConverseLLM().completion(
**_completion_kwargs(
optional_params={
"maxTokens": 16,
"aws_region_name": "us-east-1",
"aws_access_key_id": "AKIACONVERSECALLER",
"aws_secret_access_key": "pod-caller-secret",
"aws_role_name": "arn:aws:iam::999999999999:role/litellm-converse-role",
"aws_session_name": "litellm-converse-session",
"aws_session_tags": tags,
},
litellm_params={},
client=client,
)
)
assert response.choices[0].message.content == "hi"
sent = client.post.call_args.kwargs
assert "Credential=ASIACONVERSETAGGED/" in sent["headers"]["Authorization"]
assert "aws_session_tags" not in sent["data"]

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