mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_health_check_db_storm
This commit is contained in:
commit
10a5761bb9
146 changed files with 9380 additions and 1025 deletions
4
.github/workflows/test-linting.yml
vendored
4
.github/workflows/test-linting.yml
vendored
|
|
@ -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."
|
||||
|
|
|
|||
4
Makefile
4
Makefile
|
|
@ -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') && \
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
5
litellm/integrations/pointfive/__init__.py
Normal file
5
litellm/integrations/pointfive/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""PointFive logging integration for LiteLLM."""
|
||||
|
||||
from litellm.integrations.pointfive.logger import PointFiveLogger
|
||||
|
||||
__all__ = ("PointFiveLogger",)
|
||||
304
litellm/integrations/pointfive/logger.py
Normal file
304
litellm/integrations/pointfive/logger.py
Normal 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)
|
||||
53
litellm/integrations/pointfive/payload.py
Normal file
53
litellm/integrations/pointfive/payload.py
Normal 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"))
|
||||
194
litellm/integrations/pointfive/upload_client.py
Normal file
194
litellm/integrations/pointfive/upload_client.py
Normal 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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
232
litellm/llms/azure_ai/passthrough/transformation.py
Normal file
232
litellm/llms/azure_ai/passthrough/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 != {}:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
23
litellm/proxy/db/db_lookup_gate.py
Normal file
23
litellm/proxy/db/db_lookup_gate.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 "",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
**{
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
45
litellm/types/integrations/pointfive.py
Normal file
45
litellm/types/integrations/pointfive.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
759
tests/test_litellm/integrations/pointfive/test_logger.py
Normal file
759
tests/test_litellm/integrations/pointfive/test_logger.py
Normal 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
|
||||
80
tests/test_litellm/integrations/pointfive/test_payload.py
Normal file
80
tests/test_litellm/integrations/pointfive/test_payload.py
Normal 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
|
||||
378
tests/test_litellm/integrations/pointfive/test_upload_client.py
Normal file
378
tests/test_litellm/integrations/pointfive/test_upload_client.py
Normal 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)
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue