mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_pricing_auto_update_action
This commit is contained in:
commit
9b1c0f8daa
219 changed files with 14919 additions and 3974 deletions
2
.github/CODEOWNERS
vendored
Normal file
2
.github/CODEOWNERS
vendored
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
/ui/ @yuneng-jiang @ryan-crabbe-berri
|
||||
/litellm/proxy/_experimental/out/ @yuneng-jiang @ryan-crabbe-berri
|
||||
47
.github/actions/setup-uv-with-retries/action.yml
vendored
Normal file
47
.github/actions/setup-uv-with-retries/action.yml
vendored
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
name: "Set up uv with retries"
|
||||
description: >-
|
||||
Install uv via astral-sh/setup-uv, retrying on transient failures. Even with
|
||||
an exact pinned version, the action resolves the artifact URL by fetching
|
||||
https://raw.githubusercontent.com/astral-sh/versions/main/v1/uv.ndjson in a
|
||||
single request with no retry, timeout, or fallback, so one connection-level
|
||||
network error ("fetch failed") fails the whole job before any test runs.
|
||||
Retrying the full step covers the manifest fetch and the binary download.
|
||||
|
||||
inputs:
|
||||
version:
|
||||
description: "uv version to install"
|
||||
required: true
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
steps:
|
||||
- name: Set up uv (attempt 1)
|
||||
id: attempt-1
|
||||
continue-on-error: true
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
version: ${{ inputs.version }}
|
||||
|
||||
- name: Wait before attempt 2
|
||||
if: steps.attempt-1.outcome == 'failure'
|
||||
shell: bash
|
||||
run: sleep 15
|
||||
|
||||
- name: Set up uv (attempt 2)
|
||||
id: attempt-2
|
||||
if: steps.attempt-1.outcome == 'failure'
|
||||
continue-on-error: true
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
version: ${{ inputs.version }}
|
||||
|
||||
- name: Wait before attempt 3
|
||||
if: steps.attempt-2.outcome == 'failure'
|
||||
shell: bash
|
||||
run: sleep 30
|
||||
|
||||
- name: Set up uv (attempt 3)
|
||||
if: steps.attempt-2.outcome == 'failure'
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0
|
||||
with:
|
||||
version: ${{ inputs.version }}
|
||||
2
.github/workflows/_test-unit-base.yml
vendored
2
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -63,7 +63,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ jobs:
|
|||
with:
|
||||
persist-credentials: false
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
- name: Create Pull Request
|
||||
|
|
|
|||
2
.github/workflows/check-ui-api-types.yml
vendored
2
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -31,7 +31,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/codspeed.yml
vendored
2
.github/workflows/codspeed.yml
vendored
|
|
@ -37,7 +37,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/mutation-test.yml
vendored
2
.github/workflows/mutation-test.yml
vendored
|
|
@ -39,7 +39,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/oss_daily_guardrails.yml
vendored
2
.github/workflows/oss_daily_guardrails.yml
vendored
|
|
@ -35,7 +35,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-code-quality.yml
vendored
2
.github/workflows/test-code-quality.yml
vendored
|
|
@ -38,7 +38,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
4
.github/workflows/test-linting.yml
vendored
4
.github/workflows/test-linting.yml
vendored
|
|
@ -33,7 +33,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
@ -172,7 +172,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-mcp.yml
vendored
2
.github/workflows/test-mcp.yml
vendored
|
|
@ -32,7 +32,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-semgrep.yml
vendored
2
.github/workflows/test-semgrep.yml
vendored
|
|
@ -31,7 +31,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
|
|
@ -74,7 +74,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-unit-proxy-legacy.yml
vendored
2
.github/workflows/test-unit-proxy-legacy.yml
vendored
|
|
@ -59,7 +59,7 @@ jobs:
|
|||
python-version: "3.12"
|
||||
|
||||
- name: Set up uv
|
||||
uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.49"
|
||||
version = "0.1.50"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.49"
|
||||
version = "0.1.50"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,6 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_DeletedVerificationToken" ADD COLUMN "key_type" TEXT;
|
||||
|
||||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "key_type" TEXT;
|
||||
|
||||
|
|
@ -422,6 +422,7 @@ model LiteLLM_VerificationToken {
|
|||
budget_reset_at DateTime?
|
||||
allowed_cache_controls String[] @default([])
|
||||
allowed_routes String[] @default([])
|
||||
key_type String?
|
||||
policies String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
model_spend Json @default("{}")
|
||||
|
|
@ -516,6 +517,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
budget_reset_at DateTime?
|
||||
allowed_cache_controls String[] @default([])
|
||||
allowed_routes String[] @default([])
|
||||
key_type String?
|
||||
policies String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
model_spend Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.76"
|
||||
version = "0.4.77"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.76"
|
||||
version = "0.4.77"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -515,6 +515,22 @@ class CustomGuardrail(CustomLogger):
|
|||
return True
|
||||
return False
|
||||
|
||||
def uses_apply_guardrail_interface(self) -> bool:
|
||||
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
|
||||
|
||||
def _deployment_pre_call_target(self) -> "CustomLogger":
|
||||
if not self.uses_apply_guardrail_interface():
|
||||
return self
|
||||
try:
|
||||
from litellm.proxy.utils import unified_guardrail
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
f"Guardrail {self.guardrail_name or type(self).__name__} implements apply_guardrail, which needs "
|
||||
"the litellm proxy dependencies to run at the deployment level. "
|
||||
"Install them with: pip install 'litellm[proxy]'"
|
||||
) from e
|
||||
return unified_guardrail
|
||||
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
|
||||
) -> Optional[dict]:
|
||||
|
|
@ -533,7 +549,10 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
# CHECK IF GUARDRAIL REJECTS THE REQUEST
|
||||
if call_type == CallTypes.completion or call_type == CallTypes.acompletion:
|
||||
result = await self.async_pre_call_hook(
|
||||
target = self._deployment_pre_call_target()
|
||||
if target is not self:
|
||||
kwargs["guardrail_to_apply"] = self
|
||||
result = await target.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id=kwargs.get("user_api_key_user_id"),
|
||||
team_id=kwargs.get("user_api_key_team_id"),
|
||||
|
|
@ -543,7 +562,7 @@ class CustomGuardrail(CustomLogger):
|
|||
),
|
||||
cache=dc,
|
||||
data=kwargs,
|
||||
call_type=call_type.value or "acompletion", # type: ignore
|
||||
call_type="completion" if call_type == CallTypes.completion else "acompletion",
|
||||
)
|
||||
|
||||
if result is not None and isinstance(result, dict):
|
||||
|
|
|
|||
|
|
@ -239,6 +239,18 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=self.get_labels_for_metric("litellm_output_audio_tokens_metric"),
|
||||
)
|
||||
|
||||
self.litellm_video_duration_seconds_metric = self._counter_factory(
|
||||
"litellm_video_duration_seconds_metric",
|
||||
"Seconds of video generated, from usage.duration_seconds on video generation calls",
|
||||
labelnames=self.get_labels_for_metric("litellm_video_duration_seconds_metric"),
|
||||
)
|
||||
|
||||
self.litellm_images_generated_metric = self._counter_factory(
|
||||
"litellm_images_generated_metric",
|
||||
"Number of images generated, from the image generation response",
|
||||
labelnames=self.get_labels_for_metric("litellm_images_generated_metric"),
|
||||
)
|
||||
|
||||
# Remaining Budget for Team
|
||||
self.litellm_remaining_team_budget_metric = self._gauge_factory(
|
||||
"litellm_remaining_team_budget_metric",
|
||||
|
|
@ -1336,6 +1348,12 @@ class PrometheusLogger(CustomLogger):
|
|||
label_context=label_context,
|
||||
)
|
||||
|
||||
self._increment_media_generation_metrics(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
||||
# MCP tool call metrics
|
||||
self._increment_mcp_tool_call_metrics(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
|
|
@ -1459,8 +1477,65 @@ class PrometheusLogger(CustomLogger):
|
|||
),
|
||||
]
|
||||
|
||||
for counter, metric_name, value in detail_metrics:
|
||||
if not isinstance(value, (int, float)) or value <= 0:
|
||||
PrometheusLogger._inc_sparse_usage_counters(
|
||||
self,
|
||||
detail_metrics,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
||||
def _increment_media_generation_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
label_context: PrometheusLabelFactoryContext | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Increment video-seconds and images-generated counters from
|
||||
``standard_logging_payload["metadata"]["usage_object"]``. Video
|
||||
providers report ``duration_seconds`` there; image generation calls
|
||||
report ``output_image_count``. Both are sparse: only emitted when the
|
||||
value is present and > 0, so token-only call types are unaffected.
|
||||
"""
|
||||
metadata = standard_logging_payload.get("metadata") or {}
|
||||
usage_object = metadata.get("usage_object") if isinstance(metadata, dict) else None
|
||||
if not isinstance(usage_object, dict):
|
||||
return
|
||||
|
||||
media_metrics: list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
|
||||
(
|
||||
self.litellm_video_duration_seconds_metric,
|
||||
"litellm_video_duration_seconds_metric",
|
||||
usage_object.get("duration_seconds"),
|
||||
),
|
||||
(
|
||||
self.litellm_images_generated_metric,
|
||||
"litellm_images_generated_metric",
|
||||
usage_object.get("output_image_count"),
|
||||
),
|
||||
]
|
||||
|
||||
PrometheusLogger._inc_sparse_usage_counters(
|
||||
self,
|
||||
media_metrics,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
||||
def _inc_sparse_usage_counters(
|
||||
self,
|
||||
counters_with_values: list[tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]],
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
label_context: PrometheusLabelFactoryContext | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Increment each ``(counter, metric_name, value)`` entry whose value is
|
||||
a positive number. Non-numeric values (including booleans from
|
||||
malformed provider usage dicts) and values <= 0 are skipped, keeping
|
||||
scrape output sparse.
|
||||
"""
|
||||
for counter, metric_name, value in counters_with_values:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)) or value <= 0:
|
||||
continue
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
self,
|
||||
|
|
@ -1716,6 +1791,35 @@ class PrometheusLogger(CustomLogger):
|
|||
amount=float(response_cost),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_remaining_from_v3_rate_limit_headers(
|
||||
standard_logging_payload: StandardLoggingPayload | None,
|
||||
rate_limit_type: Literal["requests", "tokens"],
|
||||
) -> int | None:
|
||||
"""
|
||||
Read the per-(key, model) remaining value emitted by the v3 rate
|
||||
limiter (``parallel_request_limiter_v3.py``), which writes
|
||||
``x-ratelimit-model_per_key-remaining-{requests,tokens}`` into
|
||||
``standard_logging_object.hidden_params.additional_headers`` instead
|
||||
of the ``litellm-key-remaining-*`` metadata keys the legacy limiter
|
||||
sets. The header carries no model group; it always refers to this
|
||||
request's model group, which is what the gauges are labeled with.
|
||||
Values are written in-process as plain ints (never HTTP-serialized
|
||||
strings), so anything else is rejected rather than coerced.
|
||||
"""
|
||||
if standard_logging_payload is None:
|
||||
return None
|
||||
hidden_params = standard_logging_payload.get("hidden_params")
|
||||
if hidden_params is None:
|
||||
return None
|
||||
additional_headers = hidden_params.get("additional_headers")
|
||||
if additional_headers is None:
|
||||
return None
|
||||
value = dict(additional_headers).get(f"x-ratelimit-model_per_key-remaining-{rate_limit_type}")
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
return None
|
||||
return value
|
||||
|
||||
def _set_virtual_key_rate_limit_metrics(
|
||||
self,
|
||||
user_api_key: Optional[str],
|
||||
|
|
@ -1733,11 +1837,20 @@ class PrometheusLogger(CustomLogger):
|
|||
model_group = get_model_group_from_litellm_kwargs(kwargs)
|
||||
remaining_requests_variable_name = f"litellm-key-remaining-requests-{model_group}"
|
||||
remaining_tokens_variable_name = f"litellm-key-remaining-tokens-{model_group}"
|
||||
standard_logging_payload: StandardLoggingPayload | None = kwargs.get("standard_logging_object")
|
||||
|
||||
remaining_requests = metadata.get(remaining_requests_variable_name)
|
||||
if remaining_requests is None:
|
||||
remaining_requests = self._get_remaining_from_v3_rate_limit_headers(
|
||||
standard_logging_payload=standard_logging_payload, rate_limit_type="requests"
|
||||
)
|
||||
if remaining_requests is None:
|
||||
remaining_requests = sys.maxsize
|
||||
remaining_tokens = metadata.get(remaining_tokens_variable_name)
|
||||
if remaining_tokens is None:
|
||||
remaining_tokens = self._get_remaining_from_v3_rate_limit_headers(
|
||||
standard_logging_payload=standard_logging_payload, rate_limit_type="tokens"
|
||||
)
|
||||
if remaining_tokens is None:
|
||||
remaining_tokens = sys.maxsize
|
||||
|
||||
|
|
|
|||
|
|
@ -2,26 +2,8 @@ from typing import Optional
|
|||
|
||||
from litellm.llms.openai.data_residency import infer_openai_data_residency
|
||||
|
||||
# Pre-define optional kwargs keys as frozenset for O(1) lookups
|
||||
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
|
||||
OPTIONAL_KWARGS_KEYS = frozenset(
|
||||
AWS_CREDENTIAL_KWARGS_KEYS = frozenset(
|
||||
{
|
||||
"azure_ad_token",
|
||||
"tenant_id",
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"azure_username",
|
||||
"azure_password",
|
||||
"azure_scope",
|
||||
"timeout",
|
||||
"gcs_bucket_name",
|
||||
"bucket_name",
|
||||
"vertex_credentials",
|
||||
"vertex_project",
|
||||
"vertex_location",
|
||||
"vertex_ai_project",
|
||||
"vertex_ai_location",
|
||||
"vertex_ai_credentials",
|
||||
"aws_region_name",
|
||||
"aws_access_key_id",
|
||||
"aws_secret_access_key",
|
||||
|
|
@ -34,14 +16,40 @@ OPTIONAL_KWARGS_KEYS = frozenset(
|
|||
"aws_external_id",
|
||||
"aws_bedrock_runtime_endpoint",
|
||||
"aws_bedrock_project_id",
|
||||
"tpm",
|
||||
"rpm",
|
||||
"itpm",
|
||||
"otpm",
|
||||
"use_xai_oauth",
|
||||
}
|
||||
)
|
||||
|
||||
# Pre-define optional kwargs keys as frozenset for O(1) lookups
|
||||
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
|
||||
OPTIONAL_KWARGS_KEYS = (
|
||||
frozenset(
|
||||
{
|
||||
"azure_ad_token",
|
||||
"tenant_id",
|
||||
"client_id",
|
||||
"client_secret",
|
||||
"azure_username",
|
||||
"azure_password",
|
||||
"azure_scope",
|
||||
"timeout",
|
||||
"gcs_bucket_name",
|
||||
"bucket_name",
|
||||
"vertex_credentials",
|
||||
"vertex_project",
|
||||
"vertex_location",
|
||||
"vertex_ai_project",
|
||||
"vertex_ai_location",
|
||||
"vertex_ai_credentials",
|
||||
"tpm",
|
||||
"rpm",
|
||||
"itpm",
|
||||
"otpm",
|
||||
"use_xai_oauth",
|
||||
}
|
||||
)
|
||||
| AWS_CREDENTIAL_KWARGS_KEYS
|
||||
)
|
||||
|
||||
# Backward-compatible alias for existing imports/tests.
|
||||
_OPTIONAL_KWARGS_KEYS = OPTIONAL_KWARGS_KEYS
|
||||
|
||||
|
|
|
|||
|
|
@ -73,6 +73,7 @@ from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
|
|||
from litellm.litellm_core_utils.redact_messages import (
|
||||
redact_message_input_output_from_custom_logger,
|
||||
redact_message_input_output_from_logging,
|
||||
redact_streaming_responses_for_custom_logger,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
|
|
@ -2576,6 +2577,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
model_call_details = callback.redact_standard_logging_payload_from_model_call_details(
|
||||
model_call_details=model_call_details
|
||||
)
|
||||
model_call_details = redact_streaming_responses_for_custom_logger(
|
||||
model_call_details=model_call_details, custom_logger=callback
|
||||
)
|
||||
##################################
|
||||
if self.stream is True:
|
||||
if "async_complete_streaming_response" in model_call_details:
|
||||
|
|
@ -5208,10 +5212,15 @@ def get_standard_logging_object_payload(
|
|||
call_type = kwargs.get("call_type")
|
||||
cache_hit = kwargs.get("cache_hit", False)
|
||||
# Extract usage as a plain dict, avoiding Pydantic round-trip
|
||||
usage_dict = StandardLoggingPayloadSetup.get_usage_as_dict(
|
||||
raw_usage_dict = StandardLoggingPayloadSetup.get_usage_as_dict(
|
||||
response_obj=response_obj,
|
||||
combined_usage_object=cast(Optional[Usage], kwargs.get("combined_usage_object")),
|
||||
)
|
||||
usage_dict = (
|
||||
{**raw_usage_dict, "output_image_count": len(init_response_obj.data)}
|
||||
if isinstance(init_response_obj, ImageResponse) and init_response_obj.data
|
||||
else raw_usage_dict
|
||||
)
|
||||
|
||||
id = response_obj.get("id", kwargs.get("litellm_call_id"))
|
||||
|
||||
|
|
|
|||
|
|
@ -38,10 +38,45 @@ def redact_message_input_output_from_custom_logger(
|
|||
litellm_logging_obj: LiteLLMLoggingObject, result, custom_logger: CustomLogger
|
||||
):
|
||||
if hasattr(custom_logger, "message_logging") and custom_logger.message_logging is not True:
|
||||
return perform_redaction(litellm_logging_obj.model_call_details, result)
|
||||
return perform_redaction(litellm_logging_obj.model_call_details, result, redact_streaming_responses=False)
|
||||
return result
|
||||
|
||||
|
||||
def redact_streaming_responses_for_custom_logger(model_call_details: dict, custom_logger: CustomLogger) -> dict:
|
||||
"""
|
||||
Returns a copy of model_call_details whose streaming response entries are redacted deepcopies
|
||||
when the custom logger has opted out of message logging. The shared model_call_details is left
|
||||
untouched so other callbacks still receive the unredacted response.
|
||||
"""
|
||||
if not (hasattr(custom_logger, "message_logging") and custom_logger.message_logging is not True):
|
||||
return model_call_details
|
||||
redacted_entries = {
|
||||
streaming_key: _redacted_streaming_response_copy(model_call_details[streaming_key])
|
||||
for streaming_key in ("complete_streaming_response", "async_complete_streaming_response")
|
||||
if model_call_details.get(streaming_key) is not None
|
||||
}
|
||||
if not redacted_entries:
|
||||
return model_call_details
|
||||
return {**model_call_details, **redacted_entries}
|
||||
|
||||
|
||||
def _redacted_streaming_response_copy(streaming_response):
|
||||
redacted_response = copy.deepcopy(streaming_response)
|
||||
_redact_streaming_response(redacted_response)
|
||||
return redacted_response
|
||||
|
||||
|
||||
def _redact_streaming_response(streaming_response):
|
||||
if hasattr(streaming_response, "choices"):
|
||||
for choice in streaming_response.choices:
|
||||
_redact_choice_content(choice)
|
||||
redact_vertex_ai_metadata_from_logged_object(streaming_response)
|
||||
elif hasattr(streaming_response, "output"):
|
||||
_redact_responses_api_output(streaming_response.output)
|
||||
if hasattr(streaming_response, "reasoning") and streaming_response.reasoning is not None:
|
||||
streaming_response.reasoning = None
|
||||
|
||||
|
||||
def _redact_choice_content(choice):
|
||||
"""Helper to redact content in a choice (message or delta)."""
|
||||
if isinstance(choice, litellm.Choices):
|
||||
|
|
@ -150,9 +185,13 @@ def _redact_model_response_dict_choices(choices, redacted_str: str):
|
|||
_redact_choice_content(choice)
|
||||
|
||||
|
||||
def perform_redaction(model_call_details: dict, result):
|
||||
def perform_redaction(model_call_details: dict, result, redact_streaming_responses: bool = True):
|
||||
"""
|
||||
Performs the actual redaction on the logging object and result.
|
||||
|
||||
redact_streaming_responses=False skips the in-place redaction of the shared streaming
|
||||
response entries; per-callback redaction hands each opted-out callback its own redacted
|
||||
copy via redact_streaming_responses_for_custom_logger instead.
|
||||
"""
|
||||
# Redact model_call_details
|
||||
model_call_details["messages"] = [{"role": "user", "content": "redacted-by-litellm"}]
|
||||
|
|
@ -162,17 +201,9 @@ def perform_redaction(model_call_details: dict, result):
|
|||
redact_vertex_ai_metadata_from_litellm_params(model_call_details)
|
||||
|
||||
# Redact streaming response
|
||||
if model_call_details.get("stream", False) is True and "complete_streaming_response" in model_call_details:
|
||||
_streaming_response = model_call_details["complete_streaming_response"]
|
||||
if hasattr(_streaming_response, "choices"):
|
||||
for choice in _streaming_response.choices:
|
||||
_redact_choice_content(choice)
|
||||
redact_vertex_ai_metadata_from_logged_object(_streaming_response)
|
||||
elif hasattr(_streaming_response, "output"):
|
||||
_redact_responses_api_output(_streaming_response.output)
|
||||
# Redact reasoning field in ResponsesAPIResponse
|
||||
if hasattr(_streaming_response, "reasoning") and _streaming_response.reasoning is not None:
|
||||
_streaming_response.reasoning = None
|
||||
if redact_streaming_responses and model_call_details.get("stream", False) is True:
|
||||
for _streaming_key in ("complete_streaming_response", "async_complete_streaming_response"):
|
||||
_redact_streaming_response(model_call_details.get(_streaming_key))
|
||||
|
||||
# Redact result
|
||||
if result is not None:
|
||||
|
|
|
|||
|
|
@ -1441,24 +1441,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
output_key=param,
|
||||
)
|
||||
elif param == "response_format" and isinstance(value, dict):
|
||||
if any(
|
||||
substring in model
|
||||
for substring in {
|
||||
"sonnet-4.5",
|
||||
"sonnet-4-5",
|
||||
"opus-4.1",
|
||||
"opus-4-1",
|
||||
"opus-4.5",
|
||||
"opus-4-5",
|
||||
"opus-4.6",
|
||||
"opus-4-6",
|
||||
"opus-4.7",
|
||||
"opus-4-7",
|
||||
"sonnet-4.6",
|
||||
"sonnet-4-6",
|
||||
"sonnet_4.6",
|
||||
"sonnet_4_6",
|
||||
}
|
||||
if AnthropicConfig._supports_model_capability(
|
||||
model,
|
||||
"supports_native_structured_output",
|
||||
self._resolved_provider,
|
||||
):
|
||||
_output_format = self.map_response_format_to_anthropic_output_format(value)
|
||||
if _output_format is not None:
|
||||
|
|
|
|||
|
|
@ -340,11 +340,15 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
def _get_model_capability(model: str, key: str) -> Optional[bool]:
|
||||
"""Read boolean capability ``key`` from the model map, or None when
|
||||
no entry declares it."""
|
||||
from litellm.utils import _get_bundled_model_cost_map
|
||||
|
||||
try:
|
||||
for cand in AnthropicModelInfo._model_map_lookup_candidates(model):
|
||||
value = litellm.model_cost.get(cand, {}).get(key)
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
candidates = AnthropicModelInfo._model_map_lookup_candidates(model)
|
||||
for model_cost in (litellm.model_cost, _get_bundled_model_cost_map()):
|
||||
for cand in candidates:
|
||||
value = model_cost.get(cand, {}).get(key)
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -375,6 +375,36 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
else:
|
||||
optional_params.pop("output_config", None)
|
||||
|
||||
@staticmethod
|
||||
def _drop_incompatible_temperature_for_thinking(
|
||||
model: str, optional_params: dict, custom_llm_provider: str
|
||||
) -> None:
|
||||
"""Anthropic rejects any ``temperature`` other than 1 while extended thinking
|
||||
is enabled ("temperature may only be set to 1 when thinking is enabled").
|
||||
|
||||
Clients like Claude Code send ``thinking``/``output_config.effort`` together
|
||||
with a pinned ``temperature`` (e.g. the safety classifier uses ``temperature=0``
|
||||
for determinism). When the request lands on a non-adaptive model, the effort
|
||||
interface is reshaped above into legacy ``thinking={type: enabled}`` (or kept
|
||||
as ``output_config.effort`` on Opus 4.5), and the leftover ``temperature`` would
|
||||
400. Preserving the thinking the caller asked for wins over an unhonorable
|
||||
sampling value (Anthropic forces ``temperature=1`` under thinking regardless),
|
||||
so drop it and let the API default apply.
|
||||
|
||||
Adaptive models (4.6+) own this natively and are left untouched.
|
||||
"""
|
||||
if AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider):
|
||||
return
|
||||
temperature = optional_params.get("temperature")
|
||||
if temperature is None or temperature == 1:
|
||||
return
|
||||
thinking = optional_params.get("thinking")
|
||||
output_config = optional_params.get("output_config")
|
||||
thinking_enabled = isinstance(thinking, dict) and thinking.get("type") == "enabled"
|
||||
effort_enabled = isinstance(output_config, dict) and output_config.get("effort") is not None
|
||||
if thinking_enabled or effort_enabled:
|
||||
optional_params.pop("temperature", None)
|
||||
|
||||
def transform_anthropic_messages_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -415,6 +445,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
custom_llm_provider=self._resolved_provider,
|
||||
)
|
||||
|
||||
self._drop_incompatible_temperature_for_thinking(
|
||||
model=model,
|
||||
optional_params=anthropic_messages_optional_request_params,
|
||||
custom_llm_provider=self._resolved_provider,
|
||||
)
|
||||
|
||||
system_param = anthropic_messages_optional_request_params.get("system")
|
||||
if self.should_strip_billing_metadata() and system_param is not None:
|
||||
filtered_system = self._filter_billing_headers_from_system(system_param)
|
||||
|
|
|
|||
|
|
@ -198,14 +198,16 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
|
|||
@staticmethod
|
||||
def translate_tool_choice_to_responses_api(
|
||||
tool_choice: AnthropicMessagesToolChoice,
|
||||
) -> Dict[str, Any]:
|
||||
) -> Union[str, dict[str, Any]]:
|
||||
"""Convert Anthropic tool_choice to Responses API tool_choice."""
|
||||
tc_type = tool_choice.get("type")
|
||||
if tc_type == "any":
|
||||
return {"type": "required"}
|
||||
return "required"
|
||||
elif tc_type == "tool":
|
||||
return {"type": "function", "name": tool_choice.get("name", "")}
|
||||
return {"type": "auto"}
|
||||
elif tc_type == "none":
|
||||
return "none"
|
||||
return "auto"
|
||||
|
||||
@staticmethod
|
||||
def translate_context_management_to_responses_api(
|
||||
|
|
|
|||
|
|
@ -877,6 +877,15 @@ class BaseAWSLLM:
|
|||
"Resource": "*",
|
||||
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
|
||||
},
|
||||
{
|
||||
"Sid": "BedrockMantleLiteLLM",
|
||||
"Effect": "Allow",
|
||||
"Action": [
|
||||
"bedrock-mantle:CreateInference",
|
||||
],
|
||||
"Resource": "*",
|
||||
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
|
||||
},
|
||||
],
|
||||
}
|
||||
assume_role_params = {
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from litellm.types.utils import LlmProviders
|
|||
|
||||
from ..common_utils import OpenAIError
|
||||
|
||||
OPENAI_RESPONSES_API_MIN_MAX_OUTPUT_TOKENS = 16
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
|
|
@ -59,6 +61,19 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
key="supports_none_reasoning_effort",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _enforce_min_max_output_tokens(max_output_tokens: "int | None") -> "int | None":
|
||||
"""Raise sub-minimum max_output_tokens up to the OpenAI Responses API minimum.
|
||||
|
||||
OpenAI's Responses API rejects max_output_tokens below 16 for every model
|
||||
(not gpt-5 specific), so a client like Claude Code that sends a max_tokens=1
|
||||
warmup probe on model switch would otherwise 400. Values that are None or
|
||||
already at/above the minimum are returned unchanged.
|
||||
"""
|
||||
if isinstance(max_output_tokens, int) and max_output_tokens < OPENAI_RESPONSES_API_MIN_MAX_OUTPUT_TOKENS:
|
||||
return OPENAI_RESPONSES_API_MIN_MAX_OUTPUT_TOKENS
|
||||
return max_output_tokens
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""
|
||||
All OpenAI Responses API params are supported
|
||||
|
|
@ -92,6 +107,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
"""
|
||||
params = dict(response_api_optional_params)
|
||||
|
||||
if "max_output_tokens" in params:
|
||||
params["max_output_tokens"] = self._enforce_min_max_output_tokens(params.get("max_output_tokens"))
|
||||
|
||||
if self._is_gpt_5_model(model=model):
|
||||
temperature = params.get("temperature")
|
||||
if temperature is not None and temperature != 1:
|
||||
|
|
|
|||
|
|
@ -92,7 +92,10 @@ from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
|
|||
from litellm.litellm_core_utils.request_timeout_resolver import (
|
||||
get_configured_request_timeout,
|
||||
)
|
||||
from litellm.litellm_core_utils.get_litellm_params import OPTIONAL_KWARGS_KEYS
|
||||
from litellm.litellm_core_utils.get_litellm_params import (
|
||||
AWS_CREDENTIAL_KWARGS_KEYS,
|
||||
OPTIONAL_KWARGS_KEYS,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.litellm_core_utils.get_provider_specific_headers import (
|
||||
ProviderSpecificHeaderUtils,
|
||||
|
|
@ -5322,7 +5325,7 @@ def completion( # type: ignore
|
|||
tpm=kwargs.get("tpm"),
|
||||
rpm=kwargs.get("rpm"),
|
||||
use_xai_oauth=kwargs.get("use_xai_oauth", False),
|
||||
aws_bedrock_project_id=kwargs.get("aws_bedrock_project_id"),
|
||||
**{key: kwargs[key] for key in AWS_CREDENTIAL_KWARGS_KEYS if key in kwargs},
|
||||
)
|
||||
cast(LiteLLMLoggingObj, logging).update_environment_variables(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -11331,6 +11331,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
|
|
@ -11362,6 +11363,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
|
|
@ -11424,6 +11426,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -11479,6 +11482,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
|
|
@ -11506,6 +11510,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
|
|
@ -11559,6 +11564,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_output_config": true
|
||||
|
|
@ -11586,6 +11592,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_output_config": true
|
||||
|
|
@ -11614,6 +11621,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"provider_specific_entry": {
|
||||
|
|
@ -11648,6 +11656,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"provider_specific_entry": {
|
||||
|
|
@ -11682,6 +11691,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -11718,6 +11728,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -11788,6 +11799,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
|
|||
budget_reset_at: Optional[datetime] = None
|
||||
allowed_cache_controls: Optional[list] = []
|
||||
allowed_routes: Optional[list] = []
|
||||
key_type: str | None = None
|
||||
permissions: Dict = {}
|
||||
model_spend: Dict = {}
|
||||
model_max_budget: Dict = {}
|
||||
|
|
|
|||
|
|
@ -18,6 +18,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti
|
|||
is_bridge_envelope_shaped,
|
||||
resolve_bridge_envelope,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
EnvelopeIdentity,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
UI_TEAM_ID,
|
||||
LiteLLM_TeamTable,
|
||||
|
|
@ -543,7 +546,7 @@ class MCPRequestHandler:
|
|||
header_key = server.alias or server.server_name
|
||||
if header_key is None:
|
||||
raise HTTPException(status_code=500, detail="Server misconfigured: MCP server has no routable name")
|
||||
admitted = await MCPRequestHandler._reload_admitted_key(result.identity.key_hash)
|
||||
admitted = await MCPRequestHandler._reload_admitted_principal(result.identity)
|
||||
await MCPRequestHandler._enforce_admitted_live_policy(admitted=admitted, request=request, route=route)
|
||||
injected = {header_key: {"Authorization": result.upstream_authorization.get_secret_value()}}
|
||||
new_headers = {**(mcp_server_auth_headers or {}), **injected}
|
||||
|
|
@ -572,6 +575,89 @@ class MCPRequestHandler:
|
|||
route=route,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _reload_admitted_principal(identity: EnvelopeIdentity) -> UserAPIKeyAuth:
|
||||
"""Reload the live litellm record the envelope's subject references.
|
||||
|
||||
Dispatches on the sealed subject type: a ``key_hash`` reloads the virtual key that
|
||||
minted the envelope (the scripted two-header client that presents a litellm key at the
|
||||
token endpoint), a ``user_id`` reloads the user that authenticated interactively (the
|
||||
DCR client, whose SSO login at the bridged authorize yields a user, not a key). Both
|
||||
return a ``UserAPIKeyAuth`` the caller runs through the centralized policy gate, so
|
||||
team/project/org/budget/SCIM enforcement is identical to the principal presenting
|
||||
itself directly."""
|
||||
match identity.subject_type:
|
||||
case "key_hash":
|
||||
return await MCPRequestHandler._reload_admitted_key(identity.subject)
|
||||
case "user_id":
|
||||
return await MCPRequestHandler._reload_admitted_user(identity.subject)
|
||||
case _:
|
||||
assert_never(identity.subject_type)
|
||||
|
||||
@staticmethod
|
||||
async def _reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
|
||||
"""Reload the live user an interactively-minted envelope references and admit them as
|
||||
themselves.
|
||||
|
||||
The DCR client authenticates via SSO at the bridged authorize, which yields a user
|
||||
subject rather than a virtual key, so the envelope admits under the user's own
|
||||
identity: the reloaded ``user_id`` and the user's own MCP object permission ride on the
|
||||
returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` the key path uses then
|
||||
computes which servers the user may reach, so the user's litellm MCP grants and access groups
|
||||
gate the request exactly as a key's do. Only the user's OWN object permission is bound: a
|
||||
``UserAPIKeyAuth`` carries a single ``team_id`` while a user may belong to many teams, so
|
||||
team-inherited MCP grants for a user are a follow-up (they need a many-teams union
|
||||
``get_allowed_mcp_servers`` does not do off one auth object). The caller's centralized policy
|
||||
gate enforces the user's live budget and org state, and a SCIM-deactivated owner fails closed.
|
||||
|
||||
Error handling mirrors the key path's retryable-503 contract, but ``get_user_object`` defeats a
|
||||
type-based check: where ``get_key_object`` raises a typed ``ProxyException`` for a missing key
|
||||
and lets a DB outage propagate raw, ``get_user_object`` catches every DB failure and re-raises a
|
||||
bare ``ValueError``, so a missing user and a real outage look identical and the original error
|
||||
survives only as ``__context__``. ``_raise_503_if_db_unavailable`` therefore walks the cause
|
||||
chain: a transient DB outage still surfaces as a retryable 503, while a missing user, or any
|
||||
other non-outage resolution failure, fails closed as a 401 rather than an opaque 500. The
|
||||
object-permission load shares this one boundary, so an outage there is classified the same
|
||||
way (``get_object_permission`` itself swallows a failed load to ``None``, matching how
|
||||
``get_key_object`` best-effort-loads a key's object permission)."""
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission, get_user_object
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Server misconfigured: no database connection")
|
||||
try:
|
||||
user_object = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
|
||||
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
|
||||
# get_object_permission resolver the key and team paths use; no permission logic is duplicated.
|
||||
object_permission = user_object.object_permission if user_object is not None else None
|
||||
if user_object is not None and object_permission is None and user_object.object_permission_id:
|
||||
object_permission = await get_object_permission(
|
||||
object_permission_id=user_object.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
|
||||
except Exception as e: # noqa: BLE001 # a DB outage anywhere in the resolution is a retryable 503, not an opaque 500; anything else fails closed as 401
|
||||
MCPRequestHandler._raise_503_if_db_unavailable(e)
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
|
||||
if user_object is None:
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential")
|
||||
if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False:
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential")
|
||||
return UserAPIKeyAuth(
|
||||
user_id=user_object.user_id,
|
||||
user_role=user_object.user_role,
|
||||
object_permission=object_permission,
|
||||
object_permission_id=user_object.object_permission_id,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth:
|
||||
"""Reload the live key record an admitted envelope references and re-check live policy.
|
||||
|
|
@ -615,10 +701,14 @@ class MCPRequestHandler:
|
|||
"""Raise a retryable 503 when ``e`` means the auth database is unreachable, else return so the
|
||||
caller applies its own fail-closed mapping. A DB outage must not masquerade as an auth failure
|
||||
(401) or surface as an opaque 500; the caller retries. Mirrors ``UserAPIKeyAuthExceptionHandler``,
|
||||
which renders a service-unavailable database error as 503 on the standard pipeline."""
|
||||
which renders a service-unavailable database error as 503 on the standard pipeline.
|
||||
|
||||
Classifies across the ``__cause__``/``__context__`` chain, not just ``e`` itself: ``get_user_object``
|
||||
re-raises every DB failure as a bare ``ValueError``, so a type-based check on the top exception
|
||||
would miss a real outage wrapped inside it."""
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(e):
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="Service Unavailable, the authentication database is temporarily unreachable. Please retry shortly.",
|
||||
|
|
|
|||
694
litellm/proxy/_experimental/mcp_server/bridge_token_flow.py
Normal file
694
litellm/proxy/_experimental/mcp_server/bridge_token_flow.py
Normal file
|
|
@ -0,0 +1,694 @@
|
|||
"""Bridge token flow: litellm identity resolution and the DCR-bridge oauth_delegate mint/refresh pipeline."""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Literal, Optional
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import SecretStr
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _BridgeAuthorizationCode
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
EnvelopeIdentity,
|
||||
EnvelopeKeys,
|
||||
RefreshCredential,
|
||||
UpstreamTokenGrant,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
def _litellm_key_from_request(request: Request) -> Optional[str]:
|
||||
"""Return the LiteLLM API key presented on the request, or ``None``.
|
||||
|
||||
Accepts the key from ``x-litellm-api-key`` (what MCP clients such as Claude Desktop/Code
|
||||
send) as well as ``Authorization``; either may carry a bare token or ``Bearer <token>``.
|
||||
``x-litellm-api-key`` wins when both are present, since ``Authorization`` may instead carry
|
||||
an OAuth/upstream bearer.
|
||||
"""
|
||||
for header_value in (
|
||||
request.headers.get("x-litellm-api-key"),
|
||||
request.headers.get("Authorization") or request.headers.get("authorization"),
|
||||
):
|
||||
if not header_value:
|
||||
continue
|
||||
value = header_value.strip()
|
||||
if value.lower().startswith("bearer "):
|
||||
value = value[7:].strip()
|
||||
if value:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool:
|
||||
"""``True`` when the presented key is neither blocked nor past its expiry.
|
||||
|
||||
The OAuth token endpoint is unauthenticated, so the presented key is validated here before it is
|
||||
trusted; a revoked or expired key must not mint a bridge envelope or write a stored credential.
|
||||
``get_key_object`` resolves a row without these checks (the main ``user_api_key_auth`` pipeline
|
||||
enforces them downstream, which this endpoint bypasses), so they are applied here. Deleted keys
|
||||
are already rejected upstream, where ``get_key_object`` raises on a row that no longer exists.
|
||||
|
||||
This is an active-state gate only; it deliberately does not require a ``user_id``. A valid
|
||||
team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating
|
||||
on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token
|
||||
store) derive it separately via :func:`_active_key_user_id`.
|
||||
|
||||
Total by design: ``expires`` is typed ``str | datetime``, and an unparseable string would make
|
||||
``datetime.fromisoformat`` raise. Since the callers run this outside their key-resolution
|
||||
``try``, an uncaught parse error would surface as a 500 instead of the endpoint's fail-closed
|
||||
behavior, so a malformed expiry is treated as inactive (return ``False``) rather than raising.
|
||||
"""
|
||||
if key_obj.blocked is True:
|
||||
return False
|
||||
expires = key_obj.expires
|
||||
if expires is not None:
|
||||
if isinstance(expires, datetime):
|
||||
expiry = expires
|
||||
else:
|
||||
try:
|
||||
expiry = datetime.fromisoformat(expires)
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
|
||||
expiry = expiry.replace(tzinfo=timezone.utc)
|
||||
if expiry < datetime.now(timezone.utc):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None:
|
||||
"""The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no
|
||||
``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which
|
||||
needs a user to key the stored credential; the bridge mint uses the key hash and does not."""
|
||||
return key_obj.user_id if _key_is_active(key_obj) else None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ResolvedKey:
|
||||
"""An active litellm key resolved from the token request: its hash (the value ``get_key_object``
|
||||
and the cache/DB layer key the record by) and the live record."""
|
||||
|
||||
key_hash: str
|
||||
key: "UserAPIKeyAuth"
|
||||
|
||||
|
||||
_KeyResolutionFailure = Literal["no_active_key", "unavailable", "unresolvable"]
|
||||
"""Why a token request yielded no active litellm key, kept distinct so a caller statuses each truthfully
|
||||
instead of blaming the client for a gateway problem:
|
||||
- ``no_active_key``: none was presented, or the presented key is unknown / blocked / expired (the
|
||||
caller's request is at fault)
|
||||
- ``unavailable``: the auth database was transiently unreachable while resolving (retryable)
|
||||
- ``unresolvable``: the gateway cannot resolve identity right now (no DB connection, or an unexpected
|
||||
error) -- a gateway fault, not the caller's
|
||||
The classification mirrors admission's ``_reload_admitted_key`` so the mint (ingress) and admission
|
||||
(egress) never disagree on the status of the same outage."""
|
||||
|
||||
|
||||
async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyResolutionFailure":
|
||||
"""Resolve the presented litellm key to an active key record, or say precisely why not.
|
||||
|
||||
Single resolution path the OAuth token endpoint reuses, resolving authoritatively via
|
||||
``get_key_object`` (cache first, then DB). The failure is a value, not a bare ``None``, so a caller
|
||||
can tell "the client sent no usable credential" (a request error) apart from "the gateway could not
|
||||
check" (an infrastructure error) and status each truthfully; collapsing both to ``None`` is what let
|
||||
a DB outage read as a 400. A resolved key is still gated by ``_key_is_active``, so a blocked or
|
||||
expired key is ``no_active_key`` while a valid team-scoped or service-account key (no ``user_id``)
|
||||
resolves. Classification mirrors admission's ``_reload_admitted_key``: no DB connection is a gateway
|
||||
fault, a ``ProxyException`` / ``HTTPException`` from ``get_key_object`` is an unknown or invalid key,
|
||||
a database-service-unavailable error is a retryable outage, and anything else is an unexpected
|
||||
gateway fault."""
|
||||
token = _litellm_key_from_request(request)
|
||||
if not token:
|
||||
return "no_active_key"
|
||||
from litellm.proxy._types import hash_token # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
|
||||
return await _reload_active_key_by_hash(hash_token(token))
|
||||
|
||||
|
||||
async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResolutionFailure":
|
||||
"""Reload the live key record for ``key_hash`` (cache first, then DB) and gate it on active state,
|
||||
returning the resolved key or a precise failure. Shared by the token request's presented-key
|
||||
resolution (:func:`_resolve_active_litellm_key`, which hashes the presented key) and the refresh
|
||||
path (which already holds the hash sealed in the refresh envelope), so both re-validate identity
|
||||
through one active-key gate and one failure classification. Classification mirrors admission's
|
||||
``_reload_admitted_key``: no DB connection is a gateway fault, a ``ProxyException`` / ``HTTPException``
|
||||
from ``get_key_object`` is an unknown or invalid key, a database-service-unavailable error is a
|
||||
retryable outage, and anything else is an unexpected gateway fault. A blocked or expired key is
|
||||
``no_active_key``, so a revoked key can neither mint nor refresh a bridge envelope."""
|
||||
from litellm.proxy._types import (
|
||||
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
get_key_object,
|
||||
)
|
||||
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
PrismaDBExceptionHandler,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return "unresolvable"
|
||||
try:
|
||||
key_obj = await get_key_object(
|
||||
hashed_token=key_hash,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
return "no_active_key"
|
||||
except Exception as exc: # noqa: BLE001 # classify: a DB outage is retryable, anything else is an opaque gateway fault
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(exc):
|
||||
return "unavailable"
|
||||
verbose_logger.debug(
|
||||
"_reload_active_key_by_hash: unexpected key-resolution error (%s)",
|
||||
type(exc).__name__,
|
||||
)
|
||||
return "unresolvable"
|
||||
if not _key_is_active(key_obj):
|
||||
return "no_active_key"
|
||||
return _ResolvedKey(key_hash=key_hash, key=key_obj)
|
||||
|
||||
|
||||
async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None":
|
||||
"""Re-validate a live litellm user by id, returning ``None`` when the user is active or a precise
|
||||
failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a
|
||||
user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a
|
||||
deactivated user cannot keep refreshing, mirroring how admission re-validates the same user subject on
|
||||
the egress side. No DB connection is a gateway fault (``unresolvable``) and a
|
||||
database-service-unavailable error is a retryable outage (``unavailable``). Everything else fails
|
||||
closed as ``no_active_key`` (the caller maps it to invalid_grant): a ``ProxyException`` /
|
||||
``HTTPException``, a SCIM-deactivated user, and, unlike the key path, a missing user. ``get_user_object``
|
||||
catches every DB failure and re-raises a bare ``ValueError`` (a deleted user and a real outage look
|
||||
identical, the original error surviving only as ``__context__``), so the outage check walks the cause
|
||||
chain, and a missing user falls through to ``no_active_key`` rather than an opaque gateway fault."""
|
||||
from litellm.proxy._types import (
|
||||
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
PrismaDBExceptionHandler,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return "unresolvable"
|
||||
try:
|
||||
user_object = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
return "no_active_key"
|
||||
except Exception as exc: # noqa: BLE001 # a DB outage is retryable; a missing user (get_user_object's wrapped ValueError) or any other resolution failure fails closed as no_active_key, never a 500
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(exc):
|
||||
return "unavailable"
|
||||
verbose_logger.debug("_reload_active_user_by_id: user-resolution error (%s)", type(exc).__name__)
|
||||
return "no_active_key"
|
||||
if user_object is None:
|
||||
return "no_active_key"
|
||||
if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False:
|
||||
return "no_active_key"
|
||||
return None
|
||||
|
||||
|
||||
async def _key_owner_scim_deactivated(key: "UserAPIKeyAuth") -> bool:
|
||||
"""True only when the key's owning user was explicitly SCIM-deactivated, so a refresh revokes an
|
||||
offboarded owner's key exactly as admission does via ``_reject_if_admitted_owner_scim_deactivated``.
|
||||
A key with no owner, a missing owner record, or a failed lookup fails OPEN (returns ``False``),
|
||||
matching admission and the standard builder: a key may outlive its owner record, and a transient DB
|
||||
blip must not revoke a live key. Only an explicit ``scim_active`` of ``False`` gates renewal."""
|
||||
if key.user_id is None:
|
||||
return False
|
||||
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
return False
|
||||
try:
|
||||
owner = await get_user_object(
|
||||
user_id=key.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 # fail open: a missing owner (get_user_object's wrapped ValueError) or a DB blip must not revoke a live key
|
||||
verbose_logger.debug("refresh: key-owner SCIM lookup failed, not revoking (%s)", type(exc).__name__)
|
||||
return False
|
||||
return owner is not None and isinstance(owner.metadata, dict) and owner.metadata.get("scim_active") is False
|
||||
|
||||
|
||||
async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResolutionFailure | None":
|
||||
"""Re-validate that the subject sealed in a refresh envelope is still live, dispatching on its type:
|
||||
a key_hash reloads the virtual key, a user_id reloads the user. Returns ``None`` when the subject is
|
||||
active or a precise failure otherwise, so revocation gates renewal for either identity source the same
|
||||
way admission gates the egress: a blocked or expired key, a SCIM-deactivated key owner (mirroring
|
||||
admission's owner check, so an offboarded user cannot keep renewing a still-active key), and a
|
||||
deactivated or deleted user all fail closed to ``no_active_key``."""
|
||||
match identity.subject_type:
|
||||
case "key_hash":
|
||||
reloaded = await _reload_active_key_by_hash(identity.subject)
|
||||
if not isinstance(reloaded, _ResolvedKey):
|
||||
return reloaded
|
||||
if await _key_owner_scim_deactivated(reloaded.key):
|
||||
return "no_active_key"
|
||||
return None
|
||||
case "user_id":
|
||||
return await _reload_active_user_by_id(identity.subject)
|
||||
case _:
|
||||
assert_never(identity.subject_type)
|
||||
|
||||
|
||||
async def _extract_user_id_from_request(request: Request) -> str | None:
|
||||
"""The litellm ``user_id`` for the token request, so a per-user token is stored under the same
|
||||
identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome
|
||||
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
|
||||
the bridge mint, which must status those outcomes differently, consumes
|
||||
:func:`_resolve_active_litellm_key` directly."""
|
||||
resolved = await _resolve_active_litellm_key(request)
|
||||
if not isinstance(resolved, _ResolvedKey):
|
||||
return None
|
||||
return _active_key_user_id(resolved.key)
|
||||
|
||||
|
||||
_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]
|
||||
"""Why an upstream token response cannot back a bridge envelope:
|
||||
- ``no_access_token``: the response carries no usable ``access_token``
|
||||
- ``expired_lifetime``: the response reports a parseable, non-positive ``expires_in``, i.e. an upstream
|
||||
token that is already dead, so sealing it would forward a bearer the edge cannot use
|
||||
An absent or unparseable ``expires_in`` is NOT a rejection; the lifetime is merely unknown and the
|
||||
envelope caps it, the by-design behaviour for an upstream that omits the field."""
|
||||
|
||||
|
||||
def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']":
|
||||
"""Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent
|
||||
or unparseable, so the envelope caps it), or ``"expired"`` (a non-positive value the upstream reports
|
||||
as already elapsed). Telling "we do not know the lifetime" apart from "the upstream says it is
|
||||
already dead" is what stops an explicitly-expired token from silently receiving the envelope's 1h
|
||||
cap. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a
|
||||
positive sub-second lifetime in ``(0, 1)`` is not truncated to ``0`` and misread as elapsed; the
|
||||
envelope works in whole seconds, so such a lifetime clamps up to its 1s floor. ``bool`` is excluded
|
||||
(an ``int`` subclass but never a real lifetime), and the conversions can raise on ``NaN`` /
|
||||
``Infinity`` / oversized input, which reads as unparseable rather than surfacing as a 500."""
|
||||
if raw_expires_in is None or isinstance(raw_expires_in, bool) or not isinstance(raw_expires_in, (int, float, str)):
|
||||
return "unspecified"
|
||||
try:
|
||||
numeric = float(raw_expires_in)
|
||||
seconds = int(numeric)
|
||||
except (ValueError, TypeError, OverflowError):
|
||||
return "unspecified"
|
||||
if numeric <= 0:
|
||||
return "expired"
|
||||
return max(1, seconds)
|
||||
|
||||
|
||||
def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenGrant | _UpstreamGrantRejection":
|
||||
"""Validate an upstream OAuth token response into a typed grant, or say why it cannot back an
|
||||
envelope. Each field is isinstance-checked so nothing untyped from ``response.json()`` reaches the
|
||||
grant. ``expires_in`` is read three ways (see :func:`_classify_upstream_lifetime`): an unknown
|
||||
lifetime leaves the grant ``expires_in`` ``None`` for the envelope to cap, a positive value is
|
||||
honoured, and an explicit already-elapsed value is a rejection rather than a silent fall-through to
|
||||
the cap."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
UpstreamTokenGrant,
|
||||
)
|
||||
|
||||
if not isinstance(token_response, dict):
|
||||
return "no_access_token"
|
||||
access = token_response.get("access_token")
|
||||
if not isinstance(access, str) or not access:
|
||||
return "no_access_token"
|
||||
lifetime = _classify_upstream_lifetime(token_response.get("expires_in"))
|
||||
if lifetime == "expired":
|
||||
return "expired_lifetime"
|
||||
token_type = token_response.get("token_type")
|
||||
scope = token_response.get("scope")
|
||||
return UpstreamTokenGrant(
|
||||
access_token=SecretStr(access),
|
||||
token_type=token_type if isinstance(token_type, str) and token_type else "Bearer",
|
||||
# The upstream refresh_token is deliberately NOT sealed: the edge never consumes it (it forwards
|
||||
# only token_type + access_token), so it would be dead weight embedding a long-lived upstream
|
||||
# credential in the client-held bearer, and it enlarges the envelope. Refresh support is a
|
||||
# follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap.
|
||||
refresh_token=None,
|
||||
scope=scope if isinstance(scope, str) and scope else None,
|
||||
expires_in=lifetime if isinstance(lifetime, int) else None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DCR-bridge oauth_delegate mint: a three-phase pipeline whose failures are values.
|
||||
#
|
||||
# prepare (before the upstream exchange) -> validate every precondition and resolve identity+keys
|
||||
# exchange (the single-use upstream code is consumed here, in exchange_token_with_server)
|
||||
# finish (after the exchange) -> seal the upstream grant into the client-held envelope
|
||||
#
|
||||
# Every precondition lives in ``prepare``, which runs BEFORE the exchange, so no failure can burn the
|
||||
# single-use code or rotate a refresh token, for either grant type -- that whole class of bug is gone
|
||||
# by construction rather than guarded case by case. Failures are values mapped to an OAuth-shaped
|
||||
# response in one place (``_bridge_mint_error_response``), so status codes and the RFC 6749 §5.2 body
|
||||
# shape are uniform. Adding a failure mode is a new literal plus a match arm the type checker forces.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_BridgeMintError = Literal[
|
||||
"no_identity",
|
||||
"invalid_refresh",
|
||||
"identity_unavailable",
|
||||
"identity_unresolvable",
|
||||
"not_configured",
|
||||
"no_upstream_token",
|
||||
"upstream_token_expired",
|
||||
"too_large",
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BridgeMintReady:
|
||||
"""Everything the seal needs, resolved once before the exchange: the identity to bind the envelope
|
||||
to and the master-key-derived envelope keys. The identity is a key_hash subject for the scripted
|
||||
two-header client (resolved from the litellm key it presents) or a user_id subject for the
|
||||
interactive SSO client (the user recovered from the gateway authorization code), so one phase-3 seal
|
||||
serves both. Resolving identity here means ``_finish_bridge_mint`` has no preconditions left to
|
||||
fail."""
|
||||
|
||||
identity: "EnvelopeIdentity"
|
||||
keys: "EnvelopeKeys"
|
||||
|
||||
|
||||
def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
|
||||
"""Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape
|
||||
(top-level ``error``, no-store headers) for every case, with a status truthful about where the
|
||||
failure is. The caller's request is 400, a transient gateway outage is 503, a gateway
|
||||
misconfiguration is 500, and an upstream problem is 502. The identity-resolution statuses match how
|
||||
admission statuses the same conditions on the egress side, so mint and admit never disagree under
|
||||
one outage."""
|
||||
match error:
|
||||
case "no_identity":
|
||||
status, code, desc = (
|
||||
400,
|
||||
"invalid_request",
|
||||
"this server issues a gateway-bound credential; complete the interactive sign-in, or "
|
||||
"send a litellm credential (x-litellm-api-key or Authorization) on the token request",
|
||||
)
|
||||
case "invalid_refresh":
|
||||
status, code, desc = (
|
||||
400,
|
||||
"invalid_grant",
|
||||
"the refresh credential is not a valid, live refresh envelope for this server; "
|
||||
"re-run authorization_code to obtain a new one",
|
||||
)
|
||||
case "identity_unavailable":
|
||||
status, code, desc = (
|
||||
503,
|
||||
"temporarily_unavailable",
|
||||
"the authentication database is temporarily unreachable; retry shortly",
|
||||
)
|
||||
case "identity_unresolvable":
|
||||
status, code, desc = (
|
||||
500,
|
||||
"server_error",
|
||||
"the gateway could not resolve the litellm identity for this request",
|
||||
)
|
||||
case "not_configured":
|
||||
status, code, desc = (
|
||||
500,
|
||||
"server_error",
|
||||
"the gateway is not configured to mint a gateway-bound credential (master_key is not set)",
|
||||
)
|
||||
case "no_upstream_token":
|
||||
status, code, desc = (
|
||||
502,
|
||||
"server_error",
|
||||
"the upstream token response has no usable access_token",
|
||||
)
|
||||
case "upstream_token_expired":
|
||||
status, code, desc = (
|
||||
502,
|
||||
"server_error",
|
||||
"the upstream token response reports an already-expired lifetime",
|
||||
)
|
||||
case "too_large":
|
||||
status, code, desc = (
|
||||
502,
|
||||
"server_error",
|
||||
"the upstream token is too large to seal into a gateway-bound credential",
|
||||
)
|
||||
case _:
|
||||
assert_never(error)
|
||||
return JSONResponse(
|
||||
status_code=status, content={"error": code, "error_description": desc}, headers=TOKEN_NO_CACHE_HEADERS
|
||||
)
|
||||
|
||||
|
||||
def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
|
||||
"""Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays
|
||||
truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that
|
||||
cannot resolve identity is 500."""
|
||||
match failure:
|
||||
case "no_active_key":
|
||||
return "no_identity"
|
||||
case "unavailable":
|
||||
return "identity_unavailable"
|
||||
case "unresolvable":
|
||||
return "identity_unresolvable"
|
||||
case _:
|
||||
assert_never(failure)
|
||||
|
||||
|
||||
def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _BridgeMintError:
|
||||
"""Lift an upstream-response rejection into the mint taxonomy; both are upstream faults (502)."""
|
||||
match rejection:
|
||||
case "no_access_token":
|
||||
return "no_upstream_token"
|
||||
case "expired_lifetime":
|
||||
return "upstream_token_expired"
|
||||
case _:
|
||||
assert_never(rejection)
|
||||
|
||||
|
||||
async def _prepare_bridge_mint(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
bridge_identity: "_BridgeAuthorizationCode | None" = None,
|
||||
) -> "_BridgeMintReady | _BridgeMintError":
|
||||
"""Phase 1 for the authorization_code grant, BEFORE the upstream exchange: confirm the gateway can
|
||||
mint (master_key set), resolve the litellm identity, and derive the envelope keys. Returns a ready
|
||||
context or a precise failure value. Running before the exchange is what makes every failure here fail
|
||||
closed without consuming the single-use code.
|
||||
|
||||
Two identity sources, one envelope. The interactive DCR client authenticates via SSO at the bridged
|
||||
authorize, so its identity arrives as ``bridge_identity`` (the user recovered from the gateway
|
||||
authorization code) and mints a user subject. The scripted two-header client presents a litellm key
|
||||
on the token request instead, so its identity is the active key's hash and mints a key_hash subject.
|
||||
A missing or invalid presented key keeps its resolution origin so the mapper statuses it truthfully;
|
||||
neither source present is ``no_identity``. The refresh_token grant has its own phase-1
|
||||
(:func:`_prepare_bridge_refresh`), which recovers identity from the presented refresh envelope."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
envelope_keys_from_master_key,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
key_hash_identity,
|
||||
user_identity,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
master_key,
|
||||
)
|
||||
|
||||
if not master_key:
|
||||
return "not_configured"
|
||||
keys = envelope_keys_from_master_key(master_key)
|
||||
if bridge_identity is not None:
|
||||
identity = user_identity(server_id=mcp_server.server_id, user_id=bridge_identity.litellm_user_id)
|
||||
return _BridgeMintReady(identity=identity, keys=keys)
|
||||
resolved = await _resolve_active_litellm_key(request)
|
||||
if not isinstance(resolved, _ResolvedKey):
|
||||
return _key_resolution_failure_to_mint_error(resolved)
|
||||
identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=resolved.key_hash)
|
||||
return _BridgeMintReady(identity=identity, keys=keys)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BridgeRefreshReady:
|
||||
"""A validated refresh request: the identity+keys to mint the renewed pair under, the upstream refresh
|
||||
token (unwrapped from the client's refresh envelope) to exchange with the upstream IdP, and the scope
|
||||
sealed alongside it at mint. The upstream refresh token is a ``SecretStr`` like every other credential
|
||||
in this layer, so a repr or a traceback that captures this value never exposes the raw upstream refresh
|
||||
token in plaintext. ``upstream_scope`` carries the originally-granted scope so the renewal re-requests
|
||||
it when the client (a DCR/MCP client that typically omits scope on refresh) sends none, keeping the
|
||||
renewed token's scope stable against an upstream that would otherwise narrow or drop it."""
|
||||
|
||||
ready: "_BridgeMintReady"
|
||||
upstream_refresh_token: SecretStr
|
||||
upstream_scope: str | None = None
|
||||
|
||||
|
||||
def _refresh_key_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
|
||||
"""Lift an identity-resolution failure on the refresh path into the mint taxonomy. Unlike the mint
|
||||
path, a resolved-but-inactive (or unknown) key is ``invalid_grant`` rather than ``invalid_request``:
|
||||
the client did present an identity (sealed in the refresh envelope), but it is no longer live, so the
|
||||
refresh is invalid and the client must re-authenticate. A transient outage is still 503 and a gateway
|
||||
fault still 500, matching the mint path and admission."""
|
||||
match failure:
|
||||
case "no_active_key":
|
||||
return "invalid_refresh"
|
||||
case "unavailable":
|
||||
return "identity_unavailable"
|
||||
case "unresolvable":
|
||||
return "identity_unresolvable"
|
||||
case _:
|
||||
assert_never(failure)
|
||||
|
||||
|
||||
async def _prepare_bridge_refresh(
|
||||
mcp_server: MCPServer, refresh_value: str | None
|
||||
) -> "_BridgeRefreshReady | _BridgeMintError":
|
||||
"""Phase 1 for the refresh_token grant, BEFORE the upstream exchange: open the client's refresh
|
||||
envelope, re-validate the sealed litellm identity so a revoked key cannot keep refreshing, and
|
||||
recover the upstream refresh token to exchange. Identity comes entirely from the sealed envelope, not
|
||||
the HTTP request, so the request object is not needed here. The client presents a refresh envelope,
|
||||
never a raw upstream refresh token, so a missing value, a non-envelope, an unopenable envelope, or one
|
||||
minted for another server is ``invalid_grant``. Running before the exchange means a rejected refresh
|
||||
never consumes or rotates the upstream refresh token."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
BridgeRefreshOpened,
|
||||
envelope_keys_from_master_key,
|
||||
open_bridge_refresh_envelope,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
master_key,
|
||||
)
|
||||
|
||||
if not master_key:
|
||||
return "not_configured"
|
||||
if not refresh_value:
|
||||
return "invalid_refresh"
|
||||
keys = envelope_keys_from_master_key(master_key)
|
||||
opened = open_bridge_refresh_envelope(refresh_value, keys, datetime.now(timezone.utc), mcp_server.server_id)
|
||||
if not isinstance(opened, BridgeRefreshOpened):
|
||||
return "invalid_refresh"
|
||||
failure = await _revalidate_active_subject(opened.identity)
|
||||
if failure is not None:
|
||||
return _refresh_key_failure_to_mint_error(failure)
|
||||
return _BridgeRefreshReady(
|
||||
ready=_BridgeMintReady(identity=opened.identity, keys=keys),
|
||||
upstream_refresh_token=opened.refresh.refresh_token,
|
||||
upstream_scope=opened.refresh.scope,
|
||||
)
|
||||
|
||||
|
||||
def _finish_bridge_mint(
|
||||
ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime
|
||||
) -> "JSONResponse | _BridgeMintError":
|
||||
"""Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held access envelope
|
||||
using the pre-resolved identity and keys, and, when the upstream returned a refresh token, seal a
|
||||
long-lived refresh envelope alongside it so the client can renew without re-authenticating. Shared by
|
||||
the authorization_code and refresh_token paths, so a renewal that the upstream rotates re-issues a
|
||||
fresh refresh envelope. The only hard failures here are properties of the upstream access token (no
|
||||
usable token, an already-expired lifetime, or a token too large to seal); a refresh token that cannot
|
||||
be sealed degrades to an access-only response rather than failing the whole exchange."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
build_bridge_token_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
)
|
||||
|
||||
grant = _bridge_grant_from_token_response(token_response)
|
||||
if not isinstance(grant, UpstreamTokenGrant):
|
||||
return _upstream_rejection_to_mint_error(grant)
|
||||
sealed = build_bridge_token_response(ready.identity, grant, ready.keys, now)
|
||||
if not isinstance(sealed, SealedEnvelope):
|
||||
return "too_large"
|
||||
# Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the
|
||||
# client is never told the bearer lives past the point admission (which uses that exp) rejects it.
|
||||
expires_in = max(0, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp()))
|
||||
refresh_envelope = _mint_refresh_envelope_value(ready.identity, token_response, ready.keys, now, mcp_server)
|
||||
body = {
|
||||
"access_token": sealed.token.get_secret_value(),
|
||||
"token_type": "Bearer",
|
||||
"expires_in": expires_in,
|
||||
# A refresh envelope rides along only when the upstream returned a refresh token to seal; when it
|
||||
# rotates on renewal, the client receives the new one and the old envelope's upstream token dies.
|
||||
**({"refresh_token": refresh_envelope} if refresh_envelope is not None else {}),
|
||||
}
|
||||
return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
|
||||
|
||||
|
||||
def _upstream_refresh_credential(token_response: object) -> "RefreshCredential | None":
|
||||
"""Extract the upstream refresh grant from a token response, or ``None`` when there is none to seal.
|
||||
Each field is isinstance-checked so nothing untyped reaches the refresh envelope; ``refresh_expires_in``
|
||||
(the refresh token's own lifetime, when the upstream reports it) is classified like ``expires_in`` and
|
||||
bounds the refresh envelope's TTL. An upstream that reports the refresh token itself as already elapsed
|
||||
(``refresh_expires_in`` non-positive) yields ``None`` rather than a refresh envelope: sealing a dead
|
||||
token would hand the client a full-TTL-capped envelope the IdP will reject, so the exchange degrades to
|
||||
an access-only response (the client re-authenticates at access expiry), mirroring how
|
||||
:func:`_bridge_grant_from_token_response` refuses an already-elapsed access token instead of capping it."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
RefreshCredential,
|
||||
)
|
||||
|
||||
if not isinstance(token_response, dict):
|
||||
return None
|
||||
refresh = token_response.get("refresh_token")
|
||||
if not isinstance(refresh, str) or not refresh:
|
||||
return None
|
||||
lifetime = _classify_upstream_lifetime(token_response.get("refresh_expires_in"))
|
||||
if lifetime == "expired":
|
||||
return None
|
||||
scope = token_response.get("scope")
|
||||
return RefreshCredential(
|
||||
refresh_token=SecretStr(refresh),
|
||||
scope=scope if isinstance(scope, str) and scope else None,
|
||||
expires_in=lifetime if isinstance(lifetime, int) else None,
|
||||
)
|
||||
|
||||
|
||||
def _mint_refresh_envelope_value(
|
||||
identity: "EnvelopeIdentity", token_response: object, keys: "EnvelopeKeys", now: datetime, mcp_server: MCPServer
|
||||
) -> str | None:
|
||||
"""Seal the upstream refresh grant (if any) into a refresh envelope and return its bearer string, or
|
||||
``None`` when the upstream returned no refresh token or the refresh token is too large to seal. A
|
||||
too-large refresh token degrades to an access-only response (logged) rather than failing an exchange
|
||||
that already succeeded upstream: the client simply re-authenticates when the access envelope expires."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
build_bridge_refresh_token_response,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
SealedEnvelope,
|
||||
)
|
||||
|
||||
refresh_credential = _upstream_refresh_credential(token_response)
|
||||
if refresh_credential is None:
|
||||
return None
|
||||
sealed = build_bridge_refresh_token_response(identity, refresh_credential, keys, now)
|
||||
if isinstance(sealed, SealedEnvelope):
|
||||
return sealed.token.get_secret_value()
|
||||
verbose_logger.warning(
|
||||
"bridge mint: the upstream refresh token is too large to seal into a refresh envelope for "
|
||||
"server=%s; issuing an access-only response, so the client re-authenticates at access expiry",
|
||||
mcp_server.server_id,
|
||||
)
|
||||
return None
|
||||
|
|
@ -10,7 +10,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
|||
import httpx
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -21,6 +21,24 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
|||
TokenEndpointAuthConfigError,
|
||||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
|
||||
_bridge_mint_error_response,
|
||||
_BridgeMintReady,
|
||||
_BridgeRefreshReady,
|
||||
_extract_user_id_from_request,
|
||||
_finish_bridge_mint,
|
||||
_prepare_bridge_mint,
|
||||
_prepare_bridge_refresh,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults import (
|
||||
CallerRejected,
|
||||
CredentialSource,
|
||||
UpstreamProtocolFault,
|
||||
classify_upstream_dcr_rejection,
|
||||
classify_upstream_token_rejection,
|
||||
dcr_fault_detail,
|
||||
render_token_fault,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
TOKEN_NO_CACHE_HEADERS,
|
||||
get_request_base_url,
|
||||
|
|
@ -37,7 +55,7 @@ from litellm.types.mcp import MCPAuth, MCPCredentials
|
|||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable
|
||||
|
||||
# TTL cache for upstream OAuth metadata fetched from pass-through MCP servers.
|
||||
# Keeps us from hammering the upstream IdP on each discovery request.
|
||||
|
|
@ -91,6 +109,8 @@ def encode_state_with_base_url(
|
|||
code_challenge: Optional[str] = None,
|
||||
code_challenge_method: Optional[str] = None,
|
||||
client_redirect_uri: Optional[str] = None,
|
||||
litellm_user_id: str | None = None,
|
||||
mcp_server_id: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Encode the base_url, original state, and PKCE parameters using encryption.
|
||||
|
|
@ -101,6 +121,11 @@ def encode_state_with_base_url(
|
|||
code_challenge: PKCE code challenge from client
|
||||
code_challenge_method: PKCE code challenge method from client
|
||||
client_redirect_uri: Original redirect_uri from client
|
||||
litellm_user_id: The SSO-authenticated litellm user captured at the bridge authorize
|
||||
(interactive dcr_bridge oauth_delegate only); the callback seals it into the gateway
|
||||
authorization code so the token mint can bind the envelope to this user
|
||||
mcp_server_id: The bridge server the interactive flow targets, sealed alongside
|
||||
litellm_user_id so the gateway code cannot be replayed against another server
|
||||
|
||||
Returns:
|
||||
An encrypted string that encodes all values
|
||||
|
|
@ -111,6 +136,8 @@ def encode_state_with_base_url(
|
|||
"code_challenge": code_challenge,
|
||||
"code_challenge_method": code_challenge_method,
|
||||
"client_redirect_uri": client_redirect_uri,
|
||||
"litellm_user_id": litellm_user_id,
|
||||
"mcp_server_id": mcp_server_id,
|
||||
}
|
||||
state_json = json.dumps(state_data, sort_keys=True)
|
||||
encrypted_state = encrypt_value_helper(state_json)
|
||||
|
|
@ -138,6 +165,68 @@ def decode_state_hash(encrypted_state: str) -> dict:
|
|||
return state_data
|
||||
|
||||
|
||||
_BRIDGE_AUTH_CODE_PREFIX = "llm_bcode_"
|
||||
|
||||
|
||||
class _BridgeAuthorizationCode(BaseModel):
|
||||
"""The identity and upstream code the gateway seals into the authorization code it hands a DCR
|
||||
client for an interactive dcr_bridge oauth_delegate sign-in, recovered at the token endpoint."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
upstream_code: str = Field(min_length=1)
|
||||
litellm_user_id: str = Field(min_length=1)
|
||||
mcp_server_id: str = Field(min_length=1)
|
||||
|
||||
|
||||
def is_bridge_authorization_code(code: str) -> bool:
|
||||
"""Cheap prefix check that ``code`` is a gateway-sealed bridge authorization code rather than a
|
||||
raw upstream code, so the token endpoint can route without decrypting."""
|
||||
return code.startswith(_BRIDGE_AUTH_CODE_PREFIX)
|
||||
|
||||
|
||||
def seal_bridge_authorization_code(upstream_code: str, litellm_user_id: str, mcp_server_id: str) -> str:
|
||||
"""Seal the upstream authorization code and the SSO-captured litellm user into a gateway
|
||||
authorization code. The DCR client only echoes this opaque value back at the token endpoint; the
|
||||
gateway decrypts it there to recover the user (to bind the envelope) and the upstream code (to
|
||||
exchange with the upstream), so a litellm identity captured in the browser at authorize survives
|
||||
to the back-channel token call with nothing stored server-side. Encrypted with the repo's
|
||||
authenticated symmetric helper (the same family the OAuth state uses), so the client can neither
|
||||
read nor forge it."""
|
||||
payload = json.dumps(
|
||||
{"upstream_code": upstream_code, "litellm_user_id": litellm_user_id, "mcp_server_id": mcp_server_id},
|
||||
sort_keys=True,
|
||||
)
|
||||
return _BRIDGE_AUTH_CODE_PREFIX + encrypt_value_helper(payload)
|
||||
|
||||
|
||||
def open_bridge_authorization_code(code: str) -> _BridgeAuthorizationCode | None:
|
||||
"""Recover the sealed identity and upstream code, or ``None`` when ``code`` is not a gateway
|
||||
bridge code or does not decrypt / validate. Total over hostile input: a raw upstream code (the
|
||||
scripted two-header path) returns ``None`` and the caller falls through to the existing
|
||||
behavior."""
|
||||
if not is_bridge_authorization_code(code):
|
||||
return None
|
||||
decrypted = decrypt_value_helper(
|
||||
code[len(_BRIDGE_AUTH_CODE_PREFIX) :], "bridge_authorization_code", return_original_value=False
|
||||
)
|
||||
if not isinstance(decrypted, str):
|
||||
return None
|
||||
try:
|
||||
return _BridgeAuthorizationCode.model_validate_json(decrypted)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _redirect_to_litellm_login(request: Request) -> RedirectResponse:
|
||||
"""Send an unauthenticated browser through litellm login before the interactive bridge authorize
|
||||
can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code,
|
||||
so a session is required; without one there is nothing to bind. After login the user re-initiates
|
||||
the connection, which then finds the session cookie (the seamless return-to round-trip, which is
|
||||
origin-validated against the control-plane URL, is a follow-up)."""
|
||||
base_url = get_request_base_url(request)
|
||||
return RedirectResponse(f"{base_url}/sso/key/generate")
|
||||
|
||||
|
||||
# LIT-4197: some upstream authorization servers reject an over-long ``state``
|
||||
# (the encrypted OAuth session blob routinely exceeds their limit). The upstream
|
||||
# only needs an opaque value it echoes back on ``/callback``, so we forward a
|
||||
|
|
@ -304,90 +393,6 @@ def _validate_token_response(
|
|||
)
|
||||
|
||||
|
||||
def _litellm_key_from_request(request: Request) -> Optional[str]:
|
||||
"""Return the LiteLLM API key presented on the request, or ``None``.
|
||||
|
||||
Accepts the key from ``x-litellm-api-key`` (what MCP clients such as Claude Desktop/Code
|
||||
send) as well as ``Authorization``; either may carry a bare token or ``Bearer <token>``.
|
||||
``x-litellm-api-key`` wins when both are present, since ``Authorization`` may instead carry
|
||||
an OAuth/upstream bearer.
|
||||
"""
|
||||
for header_value in (
|
||||
request.headers.get("x-litellm-api-key"),
|
||||
request.headers.get("Authorization") or request.headers.get("authorization"),
|
||||
):
|
||||
if not header_value:
|
||||
continue
|
||||
value = header_value.strip()
|
||||
if value.lower().startswith("bearer "):
|
||||
value = value[7:].strip()
|
||||
if value:
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> Optional[str]:
|
||||
"""The key's ``user_id``, or ``None`` if the key is blocked or expired.
|
||||
|
||||
The OAuth token endpoint is unauthenticated, so the presented key is validated here before its
|
||||
identity is trusted to key a stored credential; a revoked or expired key must not be able to
|
||||
write or overwrite the per-user OAuth token. ``get_key_object`` resolves a row without these
|
||||
checks (the main ``user_api_key_auth`` pipeline enforces them downstream, which this endpoint
|
||||
bypasses), so they are applied here. Deleted keys are already rejected upstream, where
|
||||
``get_key_object`` raises on a row that no longer exists.
|
||||
"""
|
||||
if key_obj.blocked is True:
|
||||
return None
|
||||
expires = key_obj.expires
|
||||
if expires is not None:
|
||||
expiry = expires if isinstance(expires, datetime) else datetime.fromisoformat(expires)
|
||||
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
|
||||
expiry = expiry.replace(tzinfo=timezone.utc)
|
||||
if expiry < datetime.now(timezone.utc):
|
||||
return None
|
||||
return key_obj.user_id
|
||||
|
||||
|
||||
async def _extract_user_id_from_request(request: Request) -> Optional[str]:
|
||||
"""Resolve the LiteLLM ``user_id`` at the OAuth token endpoint so a per-user token is stored
|
||||
under the same identity the egress later reads it by (``user_api_key_auth.user_id``).
|
||||
|
||||
Resolves authoritatively via ``get_key_object`` (cache first, then DB) instead of a raw cache
|
||||
peek. On a multi-replica gateway the token-exchange request can land on a worker whose in-memory
|
||||
cache never saw the key, and a cross-replica Redis hit deserializes to a plain ``dict`` rather
|
||||
than a ``UserAPIKeyAuth``; the previous code read only ``Authorization`` and did
|
||||
``getattr(cached, "user_id")`` with no ``model_type`` rehydration and no DB fallback, so it
|
||||
silently returned ``None`` and the token was never persisted, which makes the egress 401 on every
|
||||
reconnect. The resolved key is validated (``_active_key_user_id``) before its identity is trusted,
|
||||
so a blocked or expired key cannot write. Returns ``None`` when no key is present, the key cannot
|
||||
be resolved, or it is blocked/expired.
|
||||
"""
|
||||
token = _litellm_key_from_request(request)
|
||||
if not token:
|
||||
return None
|
||||
try:
|
||||
from litellm.proxy._types import hash_token # noqa: PLC0415
|
||||
from litellm.proxy.auth.auth_checks import get_key_object # noqa: PLC0415
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
key_obj = await get_key_object(
|
||||
hashed_token=hash_token(token),
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
return _active_key_user_id(key_obj)
|
||||
except Exception as exc:
|
||||
verbose_logger.debug(
|
||||
"_extract_user_id_from_request: could not resolve a LiteLLM user_id for the presented "
|
||||
"key (%s); per-user token will not be stored server-side.",
|
||||
type(exc).__name__,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def _store_per_user_token_server_side(
|
||||
server: MCPServer,
|
||||
user_id: str,
|
||||
|
|
@ -620,12 +625,31 @@ async def authorize_with_server(
|
|||
parsed = urlparse(redirect_uri)
|
||||
base_url = urlunparse(parsed._replace(query=""))
|
||||
request_base_url = get_request_base_url(request)
|
||||
|
||||
# Interactive dcr_bridge oauth_delegate sign-in: this arm runs the gateway /callback and /token in
|
||||
# the loop, so the gateway can capture the litellm user here (from the browser's UI session) and
|
||||
# carry it to the back-channel token mint. Seal the SSO user and the target server into the state;
|
||||
# the callback reads them back to mint the gateway authorization code. A DCR client cannot present a
|
||||
# litellm key, so the browser session is the only identity source; without one there is nothing to
|
||||
# bind, so send the user through login first. Every other oauth2 server keeps the identity-less state.
|
||||
litellm_user_id: str | None = None
|
||||
if mcp_server.is_dcr_bridge and mcp_server.is_oauth_delegate:
|
||||
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # inline import avoids a module-load circular import
|
||||
_user_id_from_session_cookie,
|
||||
)
|
||||
|
||||
litellm_user_id = _user_id_from_session_cookie(request)
|
||||
if litellm_user_id is None:
|
||||
return _redirect_to_litellm_login(request)
|
||||
|
||||
encoded_state = encode_state_with_base_url(
|
||||
base_url=base_url,
|
||||
original_state=state,
|
||||
code_challenge=code_challenge,
|
||||
code_challenge_method=code_challenge_method,
|
||||
client_redirect_uri=redirect_uri,
|
||||
litellm_user_id=litellm_user_id,
|
||||
mcp_server_id=mcp_server.server_id if litellm_user_id else None,
|
||||
)
|
||||
relay_state = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES)
|
||||
|
||||
|
|
@ -654,6 +678,13 @@ async def authorize_with_server(
|
|||
return response
|
||||
|
||||
|
||||
def _token_credential_source(mcp_server: MCPServer) -> CredentialSource:
|
||||
"""Mirrors the resolved-client rule in :func:`exchange_token_with_server`: when the server has a
|
||||
stored client_id the gateway presents its own credentials upstream, so a credential rejection is
|
||||
the operator's fault, not the caller's."""
|
||||
return "gateway_stored" if mcp_server.client_id else "caller_supplied"
|
||||
|
||||
|
||||
async def exchange_token_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -688,25 +719,61 @@ async def exchange_token_with_server(
|
|||
except TokenEndpointAuthConfigError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
bridge_identity: _BridgeAuthorizationCode | None = None
|
||||
bridge_mint_ready: _BridgeMintReady | None = None
|
||||
bridge_upstream_refresh: SecretStr | None = None
|
||||
bridge_upstream_scope: str | None = None
|
||||
refresh_request_scope: str | None = None
|
||||
is_bridge = mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge
|
||||
|
||||
if grant_type == "refresh_token":
|
||||
if not refresh_token:
|
||||
# Phase 1 for a bridge refresh: open the client's refresh envelope, re-validate the sealed
|
||||
# identity, and unwrap the real upstream refresh token BEFORE building token_data, so the exchange
|
||||
# sends the upstream token and never the envelope. A failure returns without touching the upstream.
|
||||
if is_bridge:
|
||||
prepared_refresh = await _prepare_bridge_refresh(mcp_server, refresh_token)
|
||||
if not isinstance(prepared_refresh, _BridgeRefreshReady):
|
||||
return _bridge_mint_error_response(prepared_refresh)
|
||||
bridge_mint_ready = prepared_refresh.ready
|
||||
bridge_upstream_refresh = prepared_refresh.upstream_refresh_token
|
||||
bridge_upstream_scope = prepared_refresh.upstream_scope
|
||||
# A bridge server sends the unwrapped upstream refresh token recovered from the client's refresh
|
||||
# envelope above; every other server sends the client's own refresh token verbatim.
|
||||
upstream_refresh_token = (
|
||||
bridge_upstream_refresh.get_secret_value() if bridge_upstream_refresh is not None else refresh_token
|
||||
)
|
||||
if not upstream_refresh_token:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="refresh_token is required for refresh_token grant",
|
||||
)
|
||||
token_data: dict = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
"refresh_token": upstream_refresh_token,
|
||||
**client_auth.body,
|
||||
}
|
||||
if scope:
|
||||
token_data["scope"] = scope
|
||||
refresh_request_scope = scope or bridge_upstream_scope
|
||||
if refresh_request_scope:
|
||||
token_data["scope"] = refresh_request_scope
|
||||
else:
|
||||
if not code:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="code is required for authorization_code grant",
|
||||
)
|
||||
# Interactive dcr_bridge oauth_delegate: the client presents the gateway authorization code the
|
||||
# callback sealed. Recover the SSO user and the real upstream code from it; the upstream exchange
|
||||
# below uses the upstream code, and the mint binds the envelope to the recovered user. Bind the
|
||||
# sealed server to this request so a code minted for one bridge server cannot be spent at another.
|
||||
# A raw upstream code (scripted path) opens to None and the code is used as-is.
|
||||
bridge_identity = open_bridge_authorization_code(code)
|
||||
if bridge_identity is not None:
|
||||
if bridge_identity.mcp_server_id != mcp_server.server_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Authorization code was issued for a different MCP server",
|
||||
)
|
||||
code = bridge_identity.upstream_code
|
||||
bridge_token_relay = _dcr_bridge_relays_client_registration(mcp_server)
|
||||
if bridge_token_relay and not redirect_uri:
|
||||
raise HTTPException(
|
||||
|
|
@ -726,32 +793,49 @@ async def exchange_token_with_server(
|
|||
}
|
||||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
|
||||
# Phase 1 for a bridge authorization_code mint: resolve identity (the SSO user recovered above, or
|
||||
# the presented litellm key) and the envelope keys BEFORE the exchange consumes the single-use code.
|
||||
if is_bridge:
|
||||
prepared = await _prepare_bridge_mint(request, mcp_server, bridge_identity)
|
||||
if not isinstance(prepared, _BridgeMintReady):
|
||||
return _bridge_mint_error_response(prepared)
|
||||
bridge_mint_ready = prepared
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json", **client_auth.headers},
|
||||
data=token_data,
|
||||
)
|
||||
try:
|
||||
response = await async_client.post(
|
||||
mcp_server.token_url,
|
||||
headers={"Accept": "application/json", **client_auth.headers},
|
||||
data=token_data,
|
||||
)
|
||||
if response is not None:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
fault = classify_upstream_token_rejection(
|
||||
exc.response,
|
||||
credential_source=_token_credential_source(mcp_server),
|
||||
log_context=mcp_server.server_id,
|
||||
)
|
||||
upstream_rejected_bridge_refresh = (
|
||||
is_bridge
|
||||
and grant_type == "refresh_token"
|
||||
and isinstance(fault, CallerRejected)
|
||||
and fault.code == "invalid_grant"
|
||||
)
|
||||
if upstream_rejected_bridge_refresh:
|
||||
verbose_logger.info(
|
||||
"bridge refresh: the upstream rejected the sealed refresh token for server=%s with "
|
||||
"invalid_grant (revoked or expired at the IdP); returning invalid_grant so the client "
|
||||
"re-runs authorization_code rather than an opaque upstream error",
|
||||
mcp_server.server_id,
|
||||
)
|
||||
return _bridge_mint_error_response("invalid_refresh")
|
||||
return render_token_fault(fault)
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream token endpoint returned no response",
|
||||
)
|
||||
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
if "invalid_target" in exc.response.text:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: the upstream authorization server rejected the token request with "
|
||||
"invalid_target; it may require RFC 8707 resource indicators, which the gateway "
|
||||
"does not send yet (tracked as LIT-4339)",
|
||||
mcp_server.server_id,
|
||||
)
|
||||
raise
|
||||
token_response = response.json()
|
||||
access_token = token_response["access_token"]
|
||||
|
||||
# Validate token response against server-configured rules before any storage.
|
||||
# This rejects tokens from wrong Slack workspaces, Atlassian orgs, etc.
|
||||
|
|
@ -791,8 +875,23 @@ async def exchange_token_with_server(
|
|||
mcp_server.server_id,
|
||||
)
|
||||
|
||||
# A DCR-bridge oauth_delegate server hands the client a gateway-bound envelope (identity plus the
|
||||
# upstream token) instead of the raw upstream token, so the one bearer both admits the caller and
|
||||
# forwards the upstream credential. Only this mode mints; every other server returns the raw token.
|
||||
if bridge_mint_ready is not None:
|
||||
if refresh_request_scope and isinstance(token_response, dict) and not token_response.get("scope"):
|
||||
token_response = {**token_response, "scope": refresh_request_scope}
|
||||
# Phase 3: seal the upstream grant into the client-held envelope; failures map through the same
|
||||
# OAuth-shaped response as the phase-1 preconditions.
|
||||
minted = _finish_bridge_mint(bridge_mint_ready, mcp_server, token_response, datetime.now(timezone.utc))
|
||||
return minted if isinstance(minted, JSONResponse) else _bridge_mint_error_response(minted)
|
||||
|
||||
raw_access_token = token_response.get("access_token") if isinstance(token_response, dict) else None
|
||||
if not isinstance(raw_access_token, str) or not raw_access_token:
|
||||
return render_token_fault(UpstreamProtocolFault(note="the upstream token response has no usable access_token"))
|
||||
|
||||
result = {
|
||||
"access_token": access_token,
|
||||
"access_token": raw_access_token,
|
||||
"token_type": token_response.get("token_type", "Bearer"),
|
||||
}
|
||||
|
||||
|
|
@ -1048,21 +1147,6 @@ async def _persist_dcr_client_registration(
|
|||
return "failed"
|
||||
|
||||
|
||||
_MAX_UPSTREAM_ERROR_CHARS = 500
|
||||
|
||||
|
||||
def _safe_upstream_error_detail(response: httpx.Response) -> str:
|
||||
"""Bounded plaintext summary of an upstream registration failure for the client.
|
||||
|
||||
RFC 7591 error bodies are small JSON objects (``error`` / ``error_description``); relaying the
|
||||
text lets the client read the real reason instead of a bare 500, and the length bound keeps a
|
||||
hostile or oversized upstream body from bloating the gateway response."""
|
||||
body = response.text
|
||||
if not body:
|
||||
return response.reason_phrase or "upstream registration failed"
|
||||
return body[:_MAX_UPSTREAM_ERROR_CHARS]
|
||||
|
||||
|
||||
async def register_client_with_server(
|
||||
request: Request,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -1122,19 +1206,24 @@ async def register_client_with_server(
|
|||
}
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register)
|
||||
response = await async_client.post(
|
||||
mcp_server.registration_url,
|
||||
headers=headers,
|
||||
json=register_data,
|
||||
)
|
||||
try:
|
||||
response = await async_client.post(
|
||||
mcp_server.registration_url,
|
||||
headers=headers,
|
||||
json=register_data,
|
||||
)
|
||||
if response is not None:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
status_code, detail = dcr_fault_detail(
|
||||
classify_upstream_dcr_rejection(exc.response, log_context=mcp_server.server_id)
|
||||
)
|
||||
raise HTTPException(status_code=status_code, detail=detail) from exc
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no response",
|
||||
)
|
||||
if bridge_relay and response.status_code >= 400:
|
||||
raise HTTPException(status_code=response.status_code, detail=_safe_upstream_error_detail(response))
|
||||
response.raise_for_status()
|
||||
|
||||
token_response = response.json()
|
||||
|
||||
|
|
@ -1362,7 +1451,20 @@ async def callback(
|
|||
# states while permitting same-origin / allowlisted clients.
|
||||
redirect_uri = _get_validated_client_redirect_uri(request, state_data)
|
||||
|
||||
params = {"code": code, "state": original_state}
|
||||
# Interactive dcr_bridge oauth_delegate: the state carries the litellm user the authorize step
|
||||
# captured. Instead of forwarding the raw upstream code (which the client would present at the
|
||||
# token endpoint with no way to prove who signed in), seal the user and the upstream code into a
|
||||
# gateway authorization code and forward THAT. The token endpoint decrypts it to bind the
|
||||
# envelope to this user. Every other flow forwards the raw code unchanged.
|
||||
litellm_user_id = state_data.get("litellm_user_id")
|
||||
mcp_server_id = state_data.get("mcp_server_id")
|
||||
forwarded_code = code
|
||||
if isinstance(litellm_user_id, str) and litellm_user_id and isinstance(mcp_server_id, str) and mcp_server_id:
|
||||
forwarded_code = seal_bridge_authorization_code(
|
||||
upstream_code=code, litellm_user_id=litellm_user_id, mcp_server_id=mcp_server_id
|
||||
)
|
||||
|
||||
params = {"code": forwarded_code, "state": original_state}
|
||||
complete_returned_url = _append_query_params(redirect_uri, params)
|
||||
response = RedirectResponse(url=complete_returned_url, status_code=302)
|
||||
_clear_oauth_state_cookie(response, request, state)
|
||||
|
|
|
|||
38
litellm/proxy/_experimental/mcp_server/faults/__init__.py
Normal file
38
litellm/proxy/_experimental/mcp_server/faults/__init__.py
Normal file
|
|
@ -0,0 +1,38 @@
|
|||
"""Typed fault values for upstream OAuth/DCR failures (phase 1 of the MCP error-handling framework).
|
||||
|
||||
The invariant this package exists to enforce: an upstream failure is classified ONCE into a single
|
||||
fault value, and the response status, wire error code, and prose are all derived from that value.
|
||||
Deriving all three from one classification makes contradictory pairings (a caller-fault error code on
|
||||
a server-fault status) unrepresentable, and gives the trust-boundary rule one enforcement point:
|
||||
spec-defined machine fields may cross to callers, upstream prose and raw bodies go to server logs.
|
||||
"""
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.faults.classify import (
|
||||
classify_upstream_dcr_rejection,
|
||||
classify_upstream_token_rejection,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.render_oauth import (
|
||||
dcr_fault_detail,
|
||||
render_token_fault,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.types import (
|
||||
CallerRejected,
|
||||
CredentialSource,
|
||||
GatewayRejected,
|
||||
UpstreamOAuthFault,
|
||||
UpstreamProtocolFault,
|
||||
UpstreamReportedFault,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CallerRejected",
|
||||
"CredentialSource",
|
||||
"GatewayRejected",
|
||||
"UpstreamOAuthFault",
|
||||
"UpstreamProtocolFault",
|
||||
"UpstreamReportedFault",
|
||||
"classify_upstream_dcr_rejection",
|
||||
"classify_upstream_token_rejection",
|
||||
"dcr_fault_detail",
|
||||
"render_token_fault",
|
||||
]
|
||||
133
litellm/proxy/_experimental/mcp_server/faults/classify.py
Normal file
133
litellm/proxy/_experimental/mcp_server/faults/classify.py
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
"""The single place that reads upstream OAuth/DCR failure responses.
|
||||
|
||||
Every accessor here is total: an upstream that lies about its content encoding, sends an undecodable
|
||||
body, or omits the spec fields yields a classified fault, never an exception. Nothing outside this
|
||||
module should touch a failed upstream response's body.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.faults.types import (
|
||||
GATEWAY_CAPABILITY_CODES,
|
||||
GATEWAY_CREDENTIAL_CODES,
|
||||
MAX_WIRE_FIELD_CHARS,
|
||||
CallerRejected,
|
||||
CredentialSource,
|
||||
GatewayRejected,
|
||||
UpstreamOAuthFault,
|
||||
UpstreamProtocolFault,
|
||||
UpstreamReportedFault,
|
||||
)
|
||||
|
||||
|
||||
def _safe_text(response: httpx.Response) -> str:
|
||||
try:
|
||||
return response.text
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
|
||||
def _safe_json(response: httpx.Response) -> object:
|
||||
try:
|
||||
return response.json()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _bounded_field(value: object) -> str | None:
|
||||
if not isinstance(value, str) or not value:
|
||||
return None
|
||||
return value[:MAX_WIRE_FIELD_CHARS]
|
||||
|
||||
|
||||
def _log_out_of_contract(endpoint_kind: str, response: httpx.Response, log_context: str) -> None:
|
||||
verbose_logger.warning(
|
||||
"MCP upstream %s endpoint (%s) returned HTTP %s outside the OAuth error contract (first %s chars): %s",
|
||||
endpoint_kind,
|
||||
log_context,
|
||||
response.status_code,
|
||||
MAX_WIRE_FIELD_CHARS,
|
||||
_safe_text(response)[:MAX_WIRE_FIELD_CHARS],
|
||||
)
|
||||
|
||||
|
||||
def _classify_oauth_error_code(
|
||||
code: str,
|
||||
description: str | None,
|
||||
error_uri: str | None,
|
||||
credential_source: CredentialSource,
|
||||
log_context: str,
|
||||
) -> UpstreamOAuthFault:
|
||||
"""Blame assignment for a contract-conformant OAuth error code, shared by the token and DCR
|
||||
classifiers. Codes by which the upstream blames itself keep that blame; ``invalid_target`` is a
|
||||
gateway capability gap (RFC 8707 resource indicators, LIT-4339) no matter whose credentials were
|
||||
presented; credential-indicting codes follow the credential source; everything else, including
|
||||
codes we do not recognize, is the caller's to act on. The upstream's HTTP status is deliberately
|
||||
never consulted: status derives from this classification at render time, which is what keeps
|
||||
status and code from contradicting each other."""
|
||||
if code == "server_error" or code == "temporarily_unavailable":
|
||||
return UpstreamReportedFault(code=code)
|
||||
if code in GATEWAY_CAPABILITY_CODES:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: the upstream authorization server rejected the request with "
|
||||
"invalid_target; it may require RFC 8707 resource indicators, which the gateway "
|
||||
"does not send yet (tracked as LIT-4339)",
|
||||
log_context,
|
||||
)
|
||||
return GatewayRejected(code=code)
|
||||
if credential_source == "gateway_stored" and code in GATEWAY_CREDENTIAL_CODES:
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: upstream authorization server rejected the gateway's configured client "
|
||||
"credentials (%s): %s",
|
||||
log_context,
|
||||
code,
|
||||
description or "<no description>",
|
||||
)
|
||||
return GatewayRejected(code=code)
|
||||
return CallerRejected(code=code, description=description, error_uri=error_uri)
|
||||
|
||||
|
||||
def classify_upstream_token_rejection(
|
||||
response: httpx.Response,
|
||||
credential_source: CredentialSource,
|
||||
log_context: str,
|
||||
) -> UpstreamOAuthFault:
|
||||
"""Classify a token-endpoint rejection into exactly one fault: a body with an RFC 6749 §5.2
|
||||
``error`` field goes through blame assignment (:func:`_classify_oauth_error_code`); anything
|
||||
without a usable ``error`` field is an upstream protocol fault."""
|
||||
parsed = _safe_json(response)
|
||||
fields = parsed if isinstance(parsed, dict) else {}
|
||||
code = _bounded_field(fields.get("error"))
|
||||
if code is None:
|
||||
_log_out_of_contract("token", response, log_context)
|
||||
return UpstreamProtocolFault(note=f"upstream token endpoint returned HTTP {response.status_code}")
|
||||
return _classify_oauth_error_code(
|
||||
code,
|
||||
description=_bounded_field(fields.get("error_description")),
|
||||
error_uri=_bounded_field(fields.get("error_uri")),
|
||||
credential_source=credential_source,
|
||||
log_context=log_context,
|
||||
)
|
||||
|
||||
|
||||
def classify_upstream_dcr_rejection(response: httpx.Response, log_context: str) -> UpstreamOAuthFault:
|
||||
"""Classify a dynamic-client-registration rejection. RFC 7591 §3.2.2 errors carry
|
||||
``error`` / ``error_description`` and go through the same blame assignment as token errors
|
||||
(registration sends no client credentials, so credential codes stay caller-actionable); anything
|
||||
without a usable ``error`` field is an upstream protocol fault."""
|
||||
parsed = _safe_json(response)
|
||||
fields = parsed if isinstance(parsed, dict) else {}
|
||||
code = _bounded_field(fields.get("error"))
|
||||
if code is None:
|
||||
_log_out_of_contract("registration", response, log_context)
|
||||
return UpstreamProtocolFault(note=f"upstream registration failed with HTTP {response.status_code}")
|
||||
return _classify_oauth_error_code(
|
||||
code,
|
||||
description=_bounded_field(fields.get("error_description")),
|
||||
error_uri=None,
|
||||
credential_source="caller_supplied",
|
||||
log_context=log_context,
|
||||
)
|
||||
|
|
@ -0,0 +1,89 @@
|
|||
"""Render upstream OAuth/DCR faults onto the wire. The only place that chooses statuses and bodies
|
||||
for these faults, so every consumer emits the same contract: RFC 6749 §5.2-shaped JSON with the §5.1
|
||||
no-store headers on token endpoints, HTTPException details on registration. Status, code, and prose
|
||||
all derive from the fault tag; exhaustive matches keep a new fault arm from shipping unrendered.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi.responses import JSONResponse
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.faults.types import UpstreamOAuthFault
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS
|
||||
|
||||
|
||||
def _gateway_rejected_description(code: str) -> str:
|
||||
if code == "invalid_target":
|
||||
return (
|
||||
"the upstream authorization server rejected the request (invalid_target); "
|
||||
"it may require RFC 8707 resource indicators, which the gateway does not send yet"
|
||||
)
|
||||
return (
|
||||
f"the upstream authorization server rejected the gateway's configured client credentials "
|
||||
f"({code}); verify the MCP server's client_id and client_secret"
|
||||
)
|
||||
|
||||
|
||||
def _upstream_reported_status_and_description(code: str) -> tuple[int, str]:
|
||||
if code == "temporarily_unavailable":
|
||||
return 503, "the upstream authorization server is temporarily unavailable; retry shortly"
|
||||
return 502, "the upstream authorization server reported an internal error"
|
||||
|
||||
|
||||
def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse:
|
||||
"""RFC 6749 §5.2 response for a token-endpoint fault. Caller-actionable rejections relay the
|
||||
upstream's code on the status that code implies (401 for invalid_client per §5.2, else 400);
|
||||
gateway-side faults are 502 ``server_error`` with gateway-authored prose so a caller is never
|
||||
blamed for, or shown the internals of, a failure only the operator can fix."""
|
||||
match fault.tag:
|
||||
case "caller_rejected":
|
||||
content = {
|
||||
"error": fault.code,
|
||||
**({"error_description": fault.description} if fault.description else {}),
|
||||
**({"error_uri": fault.error_uri} if fault.error_uri else {}),
|
||||
}
|
||||
status_code = 401 if fault.code == "invalid_client" else 400
|
||||
return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS)
|
||||
case "gateway_rejected":
|
||||
return JSONResponse(
|
||||
status_code=502,
|
||||
content={
|
||||
"error": "server_error",
|
||||
"error_description": _gateway_rejected_description(fault.code),
|
||||
},
|
||||
headers=TOKEN_NO_CACHE_HEADERS,
|
||||
)
|
||||
case "upstream_reported_fault":
|
||||
status_code, description = _upstream_reported_status_and_description(fault.code)
|
||||
return JSONResponse(
|
||||
status_code=status_code,
|
||||
content={"error": fault.code, "error_description": description},
|
||||
headers=TOKEN_NO_CACHE_HEADERS,
|
||||
)
|
||||
case "upstream_protocol_fault":
|
||||
return JSONResponse(
|
||||
status_code=502,
|
||||
content={"error": "server_error", "error_description": fault.note},
|
||||
headers=TOKEN_NO_CACHE_HEADERS,
|
||||
)
|
||||
case _:
|
||||
assert_never(fault.tag)
|
||||
|
||||
|
||||
def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]:
|
||||
"""Status and detail string for a registration fault, raised as HTTPException by the caller.
|
||||
RFC 7591 §3.2.2 defines registration errors as 400, so a contract-conformant rejection is 400
|
||||
regardless of the status the upstream chose; everything else is a 502 upstream fault."""
|
||||
match fault.tag:
|
||||
case "caller_rejected":
|
||||
detail = f"{fault.code}: {fault.description}" if fault.description else fault.code
|
||||
return 400, detail
|
||||
case "gateway_rejected":
|
||||
return 502, _gateway_rejected_description(fault.code)
|
||||
case "upstream_reported_fault":
|
||||
return _upstream_reported_status_and_description(fault.code)
|
||||
case "upstream_protocol_fault":
|
||||
return 502, fault.note
|
||||
case _:
|
||||
assert_never(fault.tag)
|
||||
79
litellm/proxy/_experimental/mcp_server/faults/types.py
Normal file
79
litellm/proxy/_experimental/mcp_server/faults/types.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
"""Fault taxonomy for upstream OAuth token and DCR registration failures.
|
||||
|
||||
Each fault is a frozen model on a ``tag`` literal. The tag alone decides the HTTP status, the wire
|
||||
error code, and whose prose the caller sees, so those three facts can never disagree the way they can
|
||||
when an upstream's status and error code are relayed independently.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
MAX_WIRE_FIELD_CHARS = 500
|
||||
"""Bound on every upstream-derived string that crosses to a caller or into a log line."""
|
||||
|
||||
CredentialSource: TypeAlias = Literal["gateway_stored", "caller_supplied"]
|
||||
"""Whose client credentials the gateway presented upstream: the MCP server's stored configuration or
|
||||
credentials the caller supplied on the request. Decides whether a credential rejection is the
|
||||
caller's problem to fix or the gateway operator's."""
|
||||
|
||||
GATEWAY_CREDENTIAL_CODES: frozenset[str] = frozenset({"invalid_client", "unauthorized_client"})
|
||||
"""RFC 6749 error codes that indict the OAuth client's credentials or grant authorization. When the
|
||||
gateway presented its own stored credentials, these are gateway-side faults the caller cannot act on;
|
||||
when the caller supplied the credentials, they are the caller's to fix."""
|
||||
|
||||
GATEWAY_CAPABILITY_CODES: frozenset[str] = frozenset({"invalid_target"})
|
||||
"""Codes that indict a gateway capability regardless of whose credentials were presented:
|
||||
``invalid_target`` means the upstream wants RFC 8707 resource indicators, which the gateway does not
|
||||
send yet (LIT-4339). Never the caller's fault."""
|
||||
|
||||
UPSTREAM_FAULT_CODES: frozenset[str] = frozenset({"server_error", "temporarily_unavailable"})
|
||||
"""Codes by which the upstream blames itself. Relaying them as caller faults would invert blame, so
|
||||
they classify as upstream-reported faults and render on the 5xx their meaning implies."""
|
||||
|
||||
|
||||
class CallerRejected(BaseModel):
|
||||
"""The upstream spoke the OAuth error contract and the failure is actionable by our caller
|
||||
(e.g. ``invalid_grant``: re-run authorization). The code and its bounded prose relay on the
|
||||
4xx status the code itself implies."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["caller_rejected"] = "caller_rejected"
|
||||
code: str
|
||||
description: str | None = None
|
||||
error_uri: str | None = None
|
||||
|
||||
|
||||
class GatewayRejected(BaseModel):
|
||||
"""The upstream rejected the request for a cause only the gateway operator can address: the
|
||||
server's stored client credentials or a gateway capability gap. Not actionable by the caller:
|
||||
rendered as 502 with gateway-authored prose naming the code; the upstream's prose goes to
|
||||
server logs only."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["gateway_rejected"] = "gateway_rejected"
|
||||
code: str
|
||||
|
||||
|
||||
class UpstreamReportedFault(BaseModel):
|
||||
"""The upstream blamed itself in the OAuth vocabulary. Rendered on the 5xx the code implies
|
||||
(``server_error`` 502, ``temporarily_unavailable`` 503) so blame and status agree."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["upstream_reported_fault"] = "upstream_reported_fault"
|
||||
code: Literal["server_error", "temporarily_unavailable"]
|
||||
|
||||
|
||||
class UpstreamProtocolFault(BaseModel):
|
||||
"""The upstream broke the error contract: no JSON ``error`` field, an undecodable body, or a
|
||||
success response without a usable token. Rendered as 502 with a gateway-authored note; the
|
||||
upstream body never crosses to the caller."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["upstream_protocol_fault"] = "upstream_protocol_fault"
|
||||
note: str
|
||||
|
||||
|
||||
UpstreamOAuthFault: TypeAlias = CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault
|
||||
|
|
@ -21,11 +21,16 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import
|
|||
EnvelopeKeys,
|
||||
EnvelopeMintError,
|
||||
OpenedEnvelope,
|
||||
OpenedRefreshEnvelope,
|
||||
RefreshCredential,
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
is_envelope,
|
||||
is_refresh_envelope,
|
||||
mint_envelope,
|
||||
mint_refresh_envelope,
|
||||
open_envelope,
|
||||
open_refresh_envelope,
|
||||
)
|
||||
|
||||
_SIGNING_KEY_DOMAIN = b"litellm-mcp-bridge:envelope-signing:"
|
||||
|
|
@ -92,6 +97,67 @@ def build_bridge_token_response(
|
|||
return mint_envelope(identity, grant, keys, now)
|
||||
|
||||
|
||||
def build_bridge_refresh_token_response(
|
||||
identity: EnvelopeIdentity,
|
||||
refresh: RefreshCredential,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> SealedEnvelope | EnvelopeMintError:
|
||||
"""Seal ``refresh`` for ``identity`` into the long-lived refresh envelope the token endpoint returns
|
||||
alongside the access envelope, so the client can renew without re-authenticating. A thin, pure
|
||||
wrapper over :func:`mint_refresh_envelope`; returns the mint error as a value for the caller to map.
|
||||
"""
|
||||
return mint_refresh_envelope(identity, refresh, keys, now)
|
||||
|
||||
|
||||
class BridgeRefreshOpened(BaseModel):
|
||||
"""A valid refresh envelope presented to the token endpoint: the identity to re-validate and renew
|
||||
under, and the upstream refresh grant to exchange."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["opened"] = "opened"
|
||||
identity: EnvelopeIdentity
|
||||
refresh: RefreshCredential
|
||||
|
||||
|
||||
class BridgeRefreshInvalid(BaseModel):
|
||||
"""The presented refresh grant is not a valid refresh envelope for this server (not refresh-shaped,
|
||||
will not open, or minted for a different server); the token endpoint fails the refresh closed."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
tag: Literal["invalid"] = "invalid"
|
||||
|
||||
|
||||
BridgeRefreshResult: TypeAlias = BridgeRefreshOpened | BridgeRefreshInvalid
|
||||
|
||||
|
||||
def open_bridge_refresh_envelope(
|
||||
refresh_value: str,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
expected_server_id: str,
|
||||
) -> BridgeRefreshResult:
|
||||
"""Open a refresh envelope a bridge ``oauth_delegate`` client presented on a refresh_token grant.
|
||||
|
||||
The token-endpoint mirror of :func:`resolve_bridge_envelope`: strips an optional ``Bearer`` scheme,
|
||||
then returns ``BridgeRefreshOpened`` with the recovered identity and upstream refresh grant, or
|
||||
``BridgeRefreshInvalid`` for anything that is not a valid refresh envelope for this server. Never
|
||||
raises; total over hostile input via :func:`open_refresh_envelope`. ``expected_server_id`` binds the
|
||||
envelope to the server the request targets, so a refresh envelope minted for one server cannot renew
|
||||
against another. A raw upstream refresh token (not envelope-shaped) is ``BridgeRefreshInvalid``: this
|
||||
mode never hands the client a bare upstream refresh token, so it must never accept one.
|
||||
"""
|
||||
candidate = _strip_bearer(refresh_value)
|
||||
if not is_refresh_envelope(candidate):
|
||||
return BridgeRefreshInvalid()
|
||||
opened = open_refresh_envelope(candidate, keys, now)
|
||||
if not isinstance(opened, OpenedRefreshEnvelope):
|
||||
return BridgeRefreshInvalid()
|
||||
if opened.identity.server_id != expected_server_id:
|
||||
return BridgeRefreshInvalid()
|
||||
return BridgeRefreshOpened(identity=opened.identity, refresh=opened.refresh)
|
||||
|
||||
|
||||
class NotBridgeEnvelope(BaseModel):
|
||||
"""The bearer is not an envelope; admission continues on its normal path."""
|
||||
|
||||
|
|
@ -128,10 +194,12 @@ def _strip_bearer(value: str) -> str:
|
|||
|
||||
|
||||
def is_bridge_envelope_shaped(authorization_value: str) -> bool:
|
||||
"""Cheap, keyless test that an ``Authorization`` value carries an envelope (optional
|
||||
``Bearer`` scheme stripped). The admission edge engages the bridge arm only for an
|
||||
envelope, so a plain upstream bearer falls through to normal oauth2 admission."""
|
||||
return is_envelope(_strip_bearer(authorization_value))
|
||||
"""Cheap, keyless test that an ``Authorization`` value carries an envelope of either kind (optional
|
||||
``Bearer`` scheme stripped). The admission edge engages the bridge arm for an access envelope (to
|
||||
admit) and for a refresh envelope (to reject it explicitly, since a refresh credential is never
|
||||
usable at the tool-call edge); a plain upstream bearer falls through to normal oauth2 admission."""
|
||||
candidate = _strip_bearer(authorization_value)
|
||||
return is_envelope(candidate) or is_refresh_envelope(candidate)
|
||||
|
||||
|
||||
def resolve_bridge_envelope(
|
||||
|
|
@ -148,6 +216,10 @@ def resolve_bridge_envelope(
|
|||
envelope, and ``BridgeEnvelopeInvalid`` for an envelope-shaped bearer that will not
|
||||
open. Never raises: it is total over hostile input via :func:`open_envelope`.
|
||||
|
||||
A refresh envelope is ``BridgeEnvelopeInvalid`` here: it is a valid gateway credential but only ever
|
||||
presented back to the token endpoint, never usable to authenticate a tool call, so admission must
|
||||
fail it closed rather than let it fall through to another arm.
|
||||
|
||||
``expected_server_id`` is the ``server_id`` of the MCP server the request targets; an
|
||||
opened envelope whose sealed ``server_id`` does not match is rejected as
|
||||
``BridgeEnvelopeInvalid``. Binding here (rather than leaving it to the caller) prevents
|
||||
|
|
@ -157,6 +229,8 @@ def resolve_bridge_envelope(
|
|||
unlike ``hmac.compare_digest`` on ``str``, does not raise on a non-ASCII server_id.
|
||||
"""
|
||||
candidate = _strip_bearer(authorization_value)
|
||||
if is_refresh_envelope(candidate):
|
||||
return BridgeEnvelopeInvalid()
|
||||
if not is_envelope(candidate):
|
||||
return NotBridgeEnvelope()
|
||||
opened = open_envelope(candidate, keys, now)
|
||||
|
|
|
|||
|
|
@ -44,18 +44,33 @@ from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
|
|||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value, encrypt_value
|
||||
|
||||
ENVELOPE_PREFIX = "llm_env_"
|
||||
"""Marker prefix on every serialized envelope so the edge can cheaply tell an envelope
|
||||
"""Marker prefix on every serialized ACCESS envelope so the edge can cheaply tell an envelope
|
||||
from a raw upstream token before doing any cryptography."""
|
||||
|
||||
REFRESH_ENVELOPE_PREFIX = "llm_refresh_"
|
||||
"""Marker prefix on every serialized REFRESH envelope. A distinct prefix keeps the two credentials
|
||||
routable without crypto and, together with the signed ``kind`` claim, stops one from being presented
|
||||
where the other is expected: a refresh envelope carries a long-lived upstream refresh token and is only
|
||||
ever presented back to the token endpoint, never forwarded upstream on a tool call."""
|
||||
|
||||
ENVELOPE_ISSUER = "litellm-mcp-bridge"
|
||||
"""``iss`` claim stamped into every envelope and required back on open."""
|
||||
|
||||
MAX_ENVELOPE_TTL_SECONDS = 3600
|
||||
"""Hard ceiling on envelope lifetime. ``exp`` is ``min(upstream expires_in, this cap)``
|
||||
"""Hard ceiling on ACCESS envelope lifetime. ``exp`` is ``min(upstream expires_in, this cap)``
|
||||
(the cap alone when the upstream omits ``expires_in``), matching the 1h lifetime of the
|
||||
BYOK session bearer this module's signing approach is borrowed from: a client-held
|
||||
credential should never outlive a bounded window even when the upstream token does."""
|
||||
|
||||
MAX_REFRESH_ENVELOPE_TTL_SECONDS = 1209600
|
||||
"""Hard ceiling on REFRESH envelope lifetime (14 days). A refresh envelope only renews the short-lived
|
||||
access envelope, and each renewal re-validates the sealed litellm key (revocation gates it) and is
|
||||
re-minted with a fresh window, so the practical bound is idle time, not a fixed session. ``exp`` is
|
||||
``min(upstream refresh_expires_in, this cap)`` (the cap alone when the upstream omits it); if the
|
||||
upstream refresh token dies first, the next renewal simply fails at the upstream and the client
|
||||
re-authenticates. The value is deliberately far shorter than a typical upstream refresh-token lifetime
|
||||
so a leaked refresh envelope is bounded even if the upstream would have honoured it for longer."""
|
||||
|
||||
MAX_ENVELOPE_BYTES = 12288
|
||||
"""Size cap on the final serialized envelope (prefix + JWT, in bytes). Upstream JWTs
|
||||
commonly run 2-4KB; base64 plus encryption overhead roughly doubles that inside the
|
||||
|
|
@ -66,21 +81,48 @@ typed error, never truncated."""
|
|||
|
||||
_ENVELOPE_JWT_ALGORITHM = "HS256"
|
||||
|
||||
EnvelopeKind = Literal["access", "refresh"]
|
||||
"""Which credential an envelope is. Stamped into the signed claims and required to match on open, so a
|
||||
signature-valid envelope of one kind cannot be replayed as the other even if its wire prefix is swapped
|
||||
(the prefix is not part of the signed payload; this claim is)."""
|
||||
|
||||
|
||||
EnvelopeSubjectType: TypeAlias = Literal["key_hash", "user_id"]
|
||||
"""Discriminator for what litellm principal the envelope binds the grant to.
|
||||
|
||||
``key_hash`` is a hashed virtual key (the scripted two-header client mints under the key it
|
||||
presents at the token endpoint); ``user_id`` is a litellm user subject (the interactive DCR
|
||||
client mints under the SSO-authenticated user, which is the only identity that browser login
|
||||
yields). Admission reloads a key record for the first and a user record for the second, then
|
||||
runs both through the same live-policy gate, so team/org/budget/revocation enforcement is
|
||||
identical either way."""
|
||||
|
||||
|
||||
class EnvelopeIdentity(BaseModel):
|
||||
"""The litellm identity the envelope binds the inner grant to.
|
||||
"""The litellm principal the envelope binds the inner grant to.
|
||||
|
||||
``key_hash`` is the hashed litellm key that authorized the mint, never a raw
|
||||
credential (and the edge rejects a bare hash presented as a bearer). Admission
|
||||
reloads the live key record by it, so the key's current team/org/object-permission
|
||||
restrictions and its revocation state are enforced at use time rather than frozen at
|
||||
mint time. ``server_id`` binds the envelope to one MCP server so it cannot be replayed
|
||||
across a server boundary.
|
||||
``subject`` is the principal identifier and ``subject_type`` says how to resolve it: a
|
||||
hashed litellm key (``key_hash``) or a litellm user id (``user_id``), never a raw
|
||||
credential (and the edge rejects a bare hash or id presented as a bearer). Admission
|
||||
reloads the live record by it, so the principal's current team/org restrictions and its
|
||||
revocation state are enforced at use time rather than frozen at mint time. ``server_id``
|
||||
binds the envelope to one MCP server so it cannot be replayed across a server boundary.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
server_id: str = Field(min_length=1)
|
||||
key_hash: str = Field(min_length=1)
|
||||
subject_type: EnvelopeSubjectType
|
||||
subject: str = Field(min_length=1)
|
||||
|
||||
|
||||
def key_hash_identity(server_id: str, key_hash: str) -> EnvelopeIdentity:
|
||||
"""The identity for the scripted client that mints under a presented virtual key."""
|
||||
return EnvelopeIdentity(server_id=server_id, subject_type="key_hash", subject=key_hash)
|
||||
|
||||
|
||||
def user_identity(server_id: str, user_id: str) -> EnvelopeIdentity:
|
||||
"""The identity for the interactive DCR client that mints under its SSO user subject."""
|
||||
return EnvelopeIdentity(server_id=server_id, subject_type="user_id", subject=user_id)
|
||||
|
||||
|
||||
class UpstreamTokenGrant(BaseModel):
|
||||
|
|
@ -99,6 +141,21 @@ class UpstreamTokenGrant(BaseModel):
|
|||
expires_in: int | None = Field(default=None, gt=0)
|
||||
|
||||
|
||||
class RefreshCredential(BaseModel):
|
||||
"""The upstream refresh grant sealed inside a refresh envelope.
|
||||
|
||||
Only the refresh token (plus the scope to re-request and the refresh token's own lifetime, when the
|
||||
upstream reports it) is sealed; the access token is never in a refresh envelope. ``refresh_token`` is
|
||||
a ``SecretStr`` so reprs never leak it, and ``expires_in`` (the refresh token's lifetime, not the
|
||||
access token's) must be positive when present.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
refresh_token: SecretStr = Field(min_length=1)
|
||||
scope: str | None = None
|
||||
expires_in: int | None = Field(default=None, gt=0)
|
||||
|
||||
|
||||
class EnvelopeKeys(BaseModel):
|
||||
"""Injected key material: the HS256 signing key and the symmetric encryption key.
|
||||
|
||||
|
|
@ -121,13 +178,21 @@ class SealedEnvelope(BaseModel):
|
|||
|
||||
|
||||
class OpenedEnvelope(BaseModel):
|
||||
"""A validated envelope: the identity it was minted for and the recovered grant."""
|
||||
"""A validated access envelope: the identity it was minted for and the recovered grant."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
identity: EnvelopeIdentity
|
||||
grant: UpstreamTokenGrant
|
||||
|
||||
|
||||
class OpenedRefreshEnvelope(BaseModel):
|
||||
"""A validated refresh envelope: the identity it was minted for and the recovered refresh grant."""
|
||||
|
||||
model_config = ConfigDict(frozen=True)
|
||||
identity: EnvelopeIdentity
|
||||
refresh: RefreshCredential
|
||||
|
||||
|
||||
class EnvelopeTooLarge(BaseModel):
|
||||
"""The serialized envelope exceeded ``MAX_ENVELOPE_BYTES``; carries sizes only."""
|
||||
|
||||
|
|
@ -199,8 +264,10 @@ class _EnvelopeClaims(BaseModel):
|
|||
iss: str
|
||||
iat: int
|
||||
exp: int
|
||||
kind: EnvelopeKind
|
||||
server_id: str = Field(min_length=1)
|
||||
key_hash: str = Field(min_length=1)
|
||||
subject_type: EnvelopeSubjectType
|
||||
subject: str = Field(min_length=1)
|
||||
grant: str = Field(min_length=1)
|
||||
|
||||
|
||||
|
|
@ -213,11 +280,25 @@ class _GrantWire(BaseModel):
|
|||
expires_in: int | None = None
|
||||
|
||||
|
||||
class _RefreshWire(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
refresh_token: str
|
||||
scope: str | None = None
|
||||
expires_in: int | None = None
|
||||
|
||||
|
||||
def is_envelope(candidate: str) -> bool:
|
||||
"""Cheap prefix check so the edge can route envelopes vs raw tokens without crypto."""
|
||||
"""Cheap prefix check for an ACCESS envelope so the edge can route envelopes vs raw tokens without
|
||||
crypto. A refresh envelope has a different prefix and is not an access envelope."""
|
||||
return candidate.startswith(ENVELOPE_PREFIX)
|
||||
|
||||
|
||||
def is_refresh_envelope(candidate: str) -> bool:
|
||||
"""Cheap prefix check for a REFRESH envelope so the token endpoint can route a refresh grant that
|
||||
carries an envelope vs a raw upstream refresh token without crypto."""
|
||||
return candidate.startswith(REFRESH_ENVELOPE_PREFIX)
|
||||
|
||||
|
||||
def mint_envelope(
|
||||
identity: EnvelopeIdentity,
|
||||
grant: UpstreamTokenGrant,
|
||||
|
|
@ -231,23 +312,15 @@ def mint_envelope(
|
|||
serialized envelope exceeds ``MAX_ENVELOPE_BYTES``.
|
||||
"""
|
||||
expires_at = now + timedelta(seconds=_envelope_ttl_seconds(grant.expires_in))
|
||||
claims = _EnvelopeClaims(
|
||||
iss=ENVELOPE_ISSUER,
|
||||
iat=int(now.timestamp()),
|
||||
exp=int(expires_at.timestamp()),
|
||||
server_id=identity.server_id,
|
||||
key_hash=identity.key_hash,
|
||||
grant=_encrypt_grant_blob(_grant_plaintext(grant), keys.encryption_key),
|
||||
return _seal(
|
||||
kind="access",
|
||||
prefix=ENVELOPE_PREFIX,
|
||||
identity=identity,
|
||||
grant_blob=_encrypt_grant_blob(_grant_plaintext(grant), keys.encryption_key),
|
||||
expires_at=expires_at,
|
||||
signing_key=keys.signing_key,
|
||||
now=now,
|
||||
)
|
||||
token = ENVELOPE_PREFIX + jwt.encode(
|
||||
claims.model_dump(),
|
||||
keys.signing_key.get_secret_value(),
|
||||
algorithm=_ENVELOPE_JWT_ALGORITHM,
|
||||
)
|
||||
size_bytes = len(token.encode("utf-8"))
|
||||
if size_bytes > MAX_ENVELOPE_BYTES:
|
||||
return EnvelopeTooLarge(size_bytes=size_bytes, max_bytes=MAX_ENVELOPE_BYTES)
|
||||
return SealedEnvelope(token=SecretStr(token), expires_at=expires_at)
|
||||
|
||||
|
||||
def open_envelope(
|
||||
|
|
@ -263,35 +336,136 @@ def open_envelope(
|
|||
re-derived, so it is stale by up to the envelope's lifetime; callers that need a
|
||||
live remaining lifetime should use ``now`` against the upstream, not this field.
|
||||
"""
|
||||
if not is_envelope(candidate):
|
||||
return NotAnEnvelope()
|
||||
# UTF-8 byte length is never below character length, so a character count already over the
|
||||
# cap rejects an oversize candidate in O(1) without encoding it; the exact byte check then
|
||||
# runs only on candidates already bounded to <= MAX_ENVELOPE_BYTES characters.
|
||||
if len(candidate) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
if len(candidate.encode("utf-8", "surrogatepass")) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
claims = _decode_claims(candidate.removeprefix(ENVELOPE_PREFIX), keys.signing_key)
|
||||
claims = _open_claims(candidate, prefix=ENVELOPE_PREFIX, expected_kind="access", keys=keys, now=now)
|
||||
if not isinstance(claims, _EnvelopeClaims):
|
||||
return claims
|
||||
if now.timestamp() >= claims.exp:
|
||||
return Expired()
|
||||
grant = _decrypt_grant(claims.grant, keys.encryption_key)
|
||||
if not isinstance(grant, UpstreamTokenGrant):
|
||||
return grant
|
||||
return OpenedEnvelope(
|
||||
identity=EnvelopeIdentity(server_id=claims.server_id, key_hash=claims.key_hash),
|
||||
identity=EnvelopeIdentity(server_id=claims.server_id, subject_type=claims.subject_type, subject=claims.subject),
|
||||
grant=grant,
|
||||
)
|
||||
|
||||
|
||||
def mint_refresh_envelope(
|
||||
identity: EnvelopeIdentity,
|
||||
refresh: RefreshCredential,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> SealedEnvelope | EnvelopeMintError:
|
||||
"""Seal ``refresh`` for ``identity`` into a long-lived, client-held refresh envelope.
|
||||
|
||||
``exp`` is ``min(refresh.expires_in, MAX_REFRESH_ENVELOPE_TTL_SECONDS)`` seconds from ``now`` (the
|
||||
cap alone when the upstream omits the refresh lifetime). Sealing a distinct ``kind="refresh"`` claim
|
||||
is what keeps a refresh envelope from ever opening as an access credential at the MCP edge. Returns
|
||||
``EnvelopeTooLarge`` when the serialized envelope exceeds ``MAX_ENVELOPE_BYTES``.
|
||||
"""
|
||||
expires_at = now + timedelta(seconds=_refresh_ttl_seconds(refresh.expires_in))
|
||||
return _seal(
|
||||
kind="refresh",
|
||||
prefix=REFRESH_ENVELOPE_PREFIX,
|
||||
identity=identity,
|
||||
grant_blob=_encrypt_grant_blob(_refresh_plaintext(refresh), keys.encryption_key),
|
||||
expires_at=expires_at,
|
||||
signing_key=keys.signing_key,
|
||||
now=now,
|
||||
)
|
||||
|
||||
|
||||
def open_refresh_envelope(
|
||||
candidate: str,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> OpenedRefreshEnvelope | EnvelopeOpenError:
|
||||
"""Validate a refresh ``candidate`` and recover the identity and inner refresh grant.
|
||||
|
||||
Total over hostile input exactly like :func:`open_envelope`: every invalid, expired, tampered,
|
||||
wrong-kind, or undecryptable candidate maps to a distinct ``EnvelopeOpenError`` variant, never a
|
||||
raise. The ``kind="refresh"`` claim is required, so an access envelope re-prefixed as a refresh one
|
||||
is rejected as ``MalformedPayload``.
|
||||
"""
|
||||
claims = _open_claims(candidate, prefix=REFRESH_ENVELOPE_PREFIX, expected_kind="refresh", keys=keys, now=now)
|
||||
if not isinstance(claims, _EnvelopeClaims):
|
||||
return claims
|
||||
refresh = _decrypt_refresh(claims.grant, keys.encryption_key)
|
||||
if not isinstance(refresh, RefreshCredential):
|
||||
return refresh
|
||||
return OpenedRefreshEnvelope(
|
||||
identity=EnvelopeIdentity(server_id=claims.server_id, subject_type=claims.subject_type, subject=claims.subject),
|
||||
refresh=refresh,
|
||||
)
|
||||
|
||||
|
||||
def _seal(
|
||||
kind: EnvelopeKind,
|
||||
prefix: str,
|
||||
identity: EnvelopeIdentity,
|
||||
grant_blob: str,
|
||||
expires_at: datetime,
|
||||
signing_key: SecretStr,
|
||||
now: datetime,
|
||||
) -> SealedEnvelope | EnvelopeTooLarge:
|
||||
"""Sign the claims for either envelope kind and enforce the size cap. Shared by both mints so the
|
||||
JWT shape, issuer, and size guard cannot drift between access and refresh envelopes."""
|
||||
claims = _EnvelopeClaims(
|
||||
iss=ENVELOPE_ISSUER,
|
||||
iat=int(now.timestamp()),
|
||||
exp=int(expires_at.timestamp()),
|
||||
kind=kind,
|
||||
server_id=identity.server_id,
|
||||
subject_type=identity.subject_type,
|
||||
subject=identity.subject,
|
||||
grant=grant_blob,
|
||||
)
|
||||
token = prefix + jwt.encode(claims.model_dump(), signing_key.get_secret_value(), algorithm=_ENVELOPE_JWT_ALGORITHM)
|
||||
size_bytes = len(token.encode("utf-8"))
|
||||
if size_bytes > MAX_ENVELOPE_BYTES:
|
||||
return EnvelopeTooLarge(size_bytes=size_bytes, max_bytes=MAX_ENVELOPE_BYTES)
|
||||
return SealedEnvelope(token=SecretStr(token), expires_at=expires_at)
|
||||
|
||||
|
||||
def _open_claims(
|
||||
candidate: str,
|
||||
prefix: str,
|
||||
expected_kind: EnvelopeKind,
|
||||
keys: EnvelopeKeys,
|
||||
now: datetime,
|
||||
) -> _EnvelopeClaims | EnvelopeOpenError:
|
||||
"""Prefix-route, size-bound, signature-verify, kind-check, and expiry-check an attacker-controlled
|
||||
candidate, shared by both openers so the security gate is identical for access and refresh. Returns
|
||||
the validated claims or a distinct ``EnvelopeOpenError``; never raises."""
|
||||
if not candidate.startswith(prefix):
|
||||
return NotAnEnvelope()
|
||||
# UTF-8 byte length is never below character length, so a character count already over the cap
|
||||
# rejects an oversize candidate in O(1) without encoding it; the exact byte check then runs only on
|
||||
# candidates already bounded to <= MAX_ENVELOPE_BYTES characters.
|
||||
if len(candidate) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
if len(candidate.encode("utf-8", "surrogatepass")) > MAX_ENVELOPE_BYTES:
|
||||
return MalformedPayload()
|
||||
claims = _decode_claims(candidate.removeprefix(prefix), keys.signing_key)
|
||||
if not isinstance(claims, _EnvelopeClaims):
|
||||
return claims
|
||||
if claims.kind != expected_kind:
|
||||
return MalformedPayload()
|
||||
if now.timestamp() >= claims.exp:
|
||||
return Expired()
|
||||
return claims
|
||||
|
||||
|
||||
def _envelope_ttl_seconds(upstream_expires_in: int | None) -> int:
|
||||
if upstream_expires_in is None:
|
||||
return MAX_ENVELOPE_TTL_SECONDS
|
||||
return min(upstream_expires_in, MAX_ENVELOPE_TTL_SECONDS)
|
||||
|
||||
|
||||
def _refresh_ttl_seconds(upstream_refresh_expires_in: int | None) -> int:
|
||||
if upstream_refresh_expires_in is None:
|
||||
return MAX_REFRESH_ENVELOPE_TTL_SECONDS
|
||||
return min(upstream_refresh_expires_in, MAX_REFRESH_ENVELOPE_TTL_SECONDS)
|
||||
|
||||
|
||||
def _grant_plaintext(grant: UpstreamTokenGrant) -> str:
|
||||
wire = _GrantWire(
|
||||
access_token=grant.access_token.get_secret_value(),
|
||||
|
|
@ -303,6 +477,15 @@ def _grant_plaintext(grant: UpstreamTokenGrant) -> str:
|
|||
return wire.model_dump_json(exclude_none=True)
|
||||
|
||||
|
||||
def _refresh_plaintext(refresh: RefreshCredential) -> str:
|
||||
wire = _RefreshWire(
|
||||
refresh_token=refresh.refresh_token.get_secret_value(),
|
||||
scope=refresh.scope,
|
||||
expires_in=refresh.expires_in,
|
||||
)
|
||||
return wire.model_dump_json(exclude_none=True)
|
||||
|
||||
|
||||
def _decode_claims(
|
||||
compact: str,
|
||||
signing_key: SecretStr,
|
||||
|
|
@ -364,3 +547,22 @@ def _decrypt_grant(
|
|||
return UpstreamTokenGrant.model_validate_json(plaintext)
|
||||
except ValidationError:
|
||||
return MalformedPayload()
|
||||
|
||||
|
||||
def _decrypt_refresh(
|
||||
blob: str,
|
||||
encryption_key: SecretStr,
|
||||
) -> RefreshCredential | DecryptFailed | MalformedPayload:
|
||||
from nacl.exceptions import CryptoError
|
||||
|
||||
try:
|
||||
plaintext = decrypt_value(
|
||||
value=base64.urlsafe_b64decode(blob),
|
||||
signing_key=encryption_key.get_secret_value(),
|
||||
)
|
||||
except (CryptoError, ValueError):
|
||||
return DecryptFailed()
|
||||
try:
|
||||
return RefreshCredential.model_validate_json(plaintext)
|
||||
except ValidationError:
|
||||
return MalformedPayload()
|
||||
|
|
|
|||
|
|
@ -3719,8 +3719,15 @@ if MCP_AVAILABLE:
|
|||
headers={"www-authenticate": upstream_www_authenticate},
|
||||
)
|
||||
|
||||
def _get_authorization_header_from_scope(scope: Scope) -> Optional[str]:
|
||||
"""First ``Authorization`` header value in the ASGI scope, or None."""
|
||||
for key, value in scope.get("headers", []):
|
||||
if key.lower() == b"authorization":
|
||||
return value.decode("latin-1")
|
||||
return None
|
||||
|
||||
def _scope_has_authorization_header(scope: Scope) -> bool:
|
||||
return any(key.lower() == b"authorization" for key, _ in scope.get("headers", []))
|
||||
return _get_authorization_header_from_scope(scope) is not None
|
||||
|
||||
def _get_forwarded_auth_from_scope(scope: Scope) -> Optional[str]:
|
||||
"""Return the upstream-bound ``Authorization`` header value, or None.
|
||||
|
|
@ -3733,17 +3740,24 @@ if MCP_AVAILABLE:
|
|||
``MCPRequestHandler.process_mcp_request``), and forwarding it upstream
|
||||
would leak the proxy key to a third-party MCP server.
|
||||
"""
|
||||
authorization = None
|
||||
has_litellm_key_header = False
|
||||
for key, value in scope.get("headers", []):
|
||||
key_lower = key.lower()
|
||||
if key_lower == b"authorization":
|
||||
authorization = value.decode("latin-1")
|
||||
elif key_lower == b"x-litellm-api-key":
|
||||
has_litellm_key_header = True
|
||||
has_litellm_key_header = any(key.lower() == b"x-litellm-api-key" for key, _ in scope.get("headers", []))
|
||||
if not has_litellm_key_header:
|
||||
return None
|
||||
return authorization
|
||||
return _get_authorization_header_from_scope(scope)
|
||||
|
||||
def _is_delegate_upstream_probe_target(server: MCPServer) -> bool:
|
||||
"""Whether ``server`` is an interactive delegate-auth server whose client-supplied
|
||||
token should be preflighted upstream.
|
||||
|
||||
Mirrors the anonymous-delegate gate in ``get_allowed_mcp_servers``: the flow is
|
||||
resolved via ``effective_oauth2_flow`` so an unstamped M2M-shape row fails closed
|
||||
(its stored client credentials drive egress; the caller's bearer is irrelevant).
|
||||
"""
|
||||
return (
|
||||
server.auth_type == MCPAuth.oauth2
|
||||
and server.delegate_auth_to_upstream is True
|
||||
and MCPServerManager.effective_oauth2_flow(server) != "client_credentials"
|
||||
)
|
||||
|
||||
async def _probe_upstream_auth(
|
||||
url: str,
|
||||
|
|
@ -3805,7 +3819,7 @@ if MCP_AVAILABLE:
|
|||
mcp_servers: Optional[List[str]],
|
||||
client_ip: Optional[str],
|
||||
) -> None:
|
||||
"""Probe pass-through upstream servers in parallel before the MCP session starts.
|
||||
"""Probe pass-through and delegate-auth upstream servers in parallel before the MCP session starts.
|
||||
|
||||
Only servers the caller's key is already authorized to reach are probed —
|
||||
the list is derived from _get_allowed_mcp_servers so that a user cannot
|
||||
|
|
@ -3813,11 +3827,42 @@ if MCP_AVAILABLE:
|
|||
|
||||
The MCP SDK commits HTTP 200 headers before invoking handlers, so a 401
|
||||
can only be returned before that point. This function raises HTTPException(401)
|
||||
with a WWW-Authenticate header if any upstream rejects the client token.
|
||||
with a WWW-Authenticate header if any upstream rejects the client token, or 403
|
||||
if the upstream accepts it but forbids the caller.
|
||||
Fails-open: network errors are logged and the request is allowed through.
|
||||
|
||||
Delegate-auth servers (``auth_type=oauth2`` + ``delegate_auth_to_upstream``)
|
||||
are probed with the caller's bare ``Authorization`` bearer. That bearer is only
|
||||
an upstream token (never a LiteLLM key) when admission took the delegate bypass,
|
||||
so the delegate target is resolved through ``get_mcp_server_by_name`` -- the same
|
||||
resolver admission used -- rather than the wider allowed-server prefix/access-group
|
||||
matching. A name that only reaches a delegate server via server_id or an access
|
||||
group would have been admitted as a real LiteLLM key, so probing it would leak that
|
||||
key upstream; requiring the admission-resolver match closes that gap. Without the
|
||||
probe a rejected token is absorbed by the tools/list handler and masked as an empty
|
||||
tool list. Gated to single-server routes so one rejected token cannot 401 a
|
||||
multi-server aggregate connect, matching the OBO preflight gating; the challenge
|
||||
echoes the requested name so aliased routes get the same resource_metadata URL as
|
||||
the tokenless preemptive challenge.
|
||||
"""
|
||||
forwarded_auth = _get_forwarded_auth_from_scope(scope)
|
||||
if not forwarded_auth:
|
||||
requested_single_target = mcp_servers[0] if mcp_servers is not None and len(mcp_servers) == 1 else None
|
||||
# The bare Authorization header (no x-litellm-api-key) is a valid upstream token
|
||||
# only when admission classified it as one, i.e. the single requested name resolves
|
||||
# to a delegate server under admission's own resolver. Resolve it the same way here
|
||||
# so a server_id- or access-group-named delegate (which admission would have treated
|
||||
# as a LiteLLM key) is never probed with that key.
|
||||
delegate_server = (
|
||||
global_mcp_server_manager.get_mcp_server_by_name(requested_single_target, client_ip=client_ip)
|
||||
if requested_single_target
|
||||
else None
|
||||
)
|
||||
delegate_auth = (
|
||||
_get_authorization_header_from_scope(scope)
|
||||
if delegate_server is not None and _is_delegate_upstream_probe_target(delegate_server)
|
||||
else None
|
||||
)
|
||||
if not forwarded_auth and not delegate_auth:
|
||||
return
|
||||
|
||||
# Use the authorized server set, not the raw user-supplied names, so that
|
||||
|
|
@ -3827,33 +3872,49 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
passthrough_servers = [
|
||||
srv
|
||||
for srv in allowed_servers
|
||||
# Restrict to genuine OAuth pass-through servers (auth_type none +
|
||||
# Authorization in extra_headers). Gateway-managed OAuth2 servers
|
||||
# must not receive the ``resource_metadata=`` challenge emitted
|
||||
# below — they require ``authorization_uri=`` pointing at the
|
||||
# gateway AS metadata. ``is_oauth_passthrough`` already requires
|
||||
# ``auth_type in (None, MCPAuth.none)``, which is mutually
|
||||
# exclusive with ``has_client_credentials`` (oauth2 + M2M flow),
|
||||
# so M2M servers are implicitly excluded here.
|
||||
if srv.is_oauth_passthrough
|
||||
]
|
||||
if not passthrough_servers:
|
||||
passthrough_targets: Tuple[Tuple[MCPServer, str, str], ...] = (
|
||||
tuple(
|
||||
(srv, forwarded_auth, srv.name)
|
||||
for srv in allowed_servers
|
||||
# Restrict to genuine OAuth pass-through servers (auth_type none +
|
||||
# Authorization in extra_headers). Gateway-managed OAuth2 servers
|
||||
# must not receive the ``resource_metadata=`` challenge emitted
|
||||
# below — they require ``authorization_uri=`` pointing at the
|
||||
# gateway AS metadata. ``is_oauth_passthrough`` already requires
|
||||
# ``auth_type in (None, MCPAuth.none)``, which is mutually
|
||||
# exclusive with ``has_client_credentials`` (oauth2 + M2M flow),
|
||||
# so M2M servers are implicitly excluded here.
|
||||
if srv.is_oauth_passthrough
|
||||
)
|
||||
if forwarded_auth
|
||||
else ()
|
||||
)
|
||||
# Probe the admission-resolved delegate server only when the caller is actually
|
||||
# authorized for it (present in the IP-filtered allowed set), keyed by server_id.
|
||||
delegate_targets: Tuple[Tuple[MCPServer, str, str], ...] = (
|
||||
tuple(
|
||||
(srv, delegate_auth, requested_single_target)
|
||||
for srv in allowed_servers
|
||||
if delegate_server is not None and srv.server_id == delegate_server.server_id
|
||||
)
|
||||
if delegate_auth and requested_single_target
|
||||
else ()
|
||||
)
|
||||
probe_targets = passthrough_targets + delegate_targets
|
||||
if not probe_targets:
|
||||
return
|
||||
|
||||
probe_results = await asyncio.gather(
|
||||
*[_probe_upstream_auth(srv.url or "", forwarded_auth) for srv in passthrough_servers]
|
||||
*[_probe_upstream_auth(srv.url or "", auth_header) for srv, auth_header, _ in probe_targets]
|
||||
)
|
||||
for srv, (probe_status, _) in zip(passthrough_servers, probe_results):
|
||||
for (srv, _, challenge_server_name), (probe_status, _) in zip(probe_targets, probe_results):
|
||||
if probe_status == 401:
|
||||
# Token is missing or expired: keep pass-through clients on the
|
||||
# protected-resource discovery flow so they re-authorize against
|
||||
# the upstream IdP metadata proxied by LiteLLM.
|
||||
www_authenticate = _get_passthrough_www_authenticate(
|
||||
scope=scope,
|
||||
server_name=srv.name,
|
||||
server_name=challenge_server_name,
|
||||
invalid_token=True,
|
||||
)
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -1118,6 +1118,7 @@ class GenerateKeyRequest(KeyRequestBase):
|
|||
class GenerateKeyResponse(KeyRequestBase):
|
||||
key: str # type: ignore
|
||||
key_name: Optional[str] = None
|
||||
key_type: str | None = None
|
||||
expires: Optional[datetime] = None
|
||||
user_id: Optional[str] = None
|
||||
token_id: Optional[str] = None
|
||||
|
|
@ -2421,6 +2422,16 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
"is active as a reminder that hard enforcement is relaxed."
|
||||
),
|
||||
)
|
||||
skip_user_budget_on_team_key: bool | None = Field(
|
||||
None,
|
||||
description=(
|
||||
"If True, restores the legacy behavior where a user's personal "
|
||||
"max_budget is NOT enforced when their key belongs to a team; only "
|
||||
"the team (and team-member) budgets apply. Defaults to False, meaning "
|
||||
"the user's personal max_budget is always enforced regardless of "
|
||||
"whether the key belongs to a team (see GitHub issue #12905)."
|
||||
),
|
||||
)
|
||||
user_url_validation: Optional[bool] = Field(
|
||||
None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -626,26 +626,29 @@ async def common_checks(
|
|||
)
|
||||
|
||||
async def _user_max_budget_check() -> None:
|
||||
# 4.1 personal budget, if personal key
|
||||
if (
|
||||
(team_object is None or team_object.team_id is None)
|
||||
and user_object is not None
|
||||
and user_object.max_budget is not None
|
||||
):
|
||||
from litellm.proxy.proxy_server import get_current_spend
|
||||
if user_object is None or user_object.max_budget is None:
|
||||
return
|
||||
skip_for_team = (
|
||||
general_settings.get("skip_user_budget_on_team_key") is True
|
||||
and team_object is not None
|
||||
and team_object.team_id is not None
|
||||
)
|
||||
if skip_for_team:
|
||||
return
|
||||
from litellm.proxy.proxy_server import get_current_spend
|
||||
|
||||
user_budget = user_object.max_budget
|
||||
user_spend = await get_current_spend(
|
||||
counter_key=f"spend:user:{user_object.user_id}",
|
||||
fallback_spend=user_object.spend or 0.0,
|
||||
user_budget = user_object.max_budget
|
||||
user_spend = await get_current_spend(
|
||||
counter_key=f"spend:user:{user_object.user_id}",
|
||||
fallback_spend=user_object.spend or 0.0,
|
||||
max_budget=user_budget,
|
||||
)
|
||||
if math.isfinite(user_budget) and user_spend >= user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=user_spend,
|
||||
max_budget=user_budget,
|
||||
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
|
||||
)
|
||||
if math.isfinite(user_budget) and user_spend >= user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=user_spend,
|
||||
max_budget=user_budget,
|
||||
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
|
||||
)
|
||||
|
||||
# Each scope reads a distinct counter key with no cross-scope ordering
|
||||
# dependency, so the per-scope Redis-first reads run concurrently instead
|
||||
|
|
|
|||
|
|
@ -2442,6 +2442,7 @@ async def _reserve_budget_after_common_checks(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
end_user_id=end_user_id,
|
||||
end_user_object=end_user_object,
|
||||
skip_user_budget_on_team_key=general_settings.get("skip_user_budget_on_team_key") is True,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
# Bounds the __cause__/__context__ walk in is_database_service_unavailable_error_in_chain.
|
||||
# Real exception chains are a few links deep; the cap also makes the walk cycle-safe.
|
||||
_MAX_EXCEPTION_CHAIN_DEPTH = 20
|
||||
|
||||
|
||||
class PrismaDBExceptionHandler:
|
||||
"""
|
||||
|
|
@ -218,6 +222,32 @@ class PrismaDBExceptionHandler:
|
|||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def is_database_service_unavailable_error_in_chain(e: BaseException) -> bool:
|
||||
"""Like ``is_database_service_unavailable_error`` but also walks the
|
||||
``__cause__`` / ``__context__`` chain.
|
||||
|
||||
``is_database_service_unavailable_error`` classifies a single exception
|
||||
by type, which a caller that catches a raw DB failure and re-raises a
|
||||
domain exception of a different type defeats. ``get_user_object`` in
|
||||
``litellm/proxy/auth/auth_checks.py`` is the concrete case: it wraps
|
||||
every DB error, a genuine outage included, in a bare ``ValueError``
|
||||
whose original error survives only as ``__context__``. A type check on
|
||||
the ``ValueError`` misses the outage, so the caller would mistake an
|
||||
infrastructure fault for an auth failure. Walking the chain recovers the
|
||||
real signal, which is the PEP 3134 way to inspect a wrapped cause.
|
||||
|
||||
The walk is depth-bounded, which also makes it cycle-safe.
|
||||
"""
|
||||
current: BaseException | None = e
|
||||
for _ in range(_MAX_EXCEPTION_CHAIN_DEPTH):
|
||||
if not isinstance(current, Exception):
|
||||
return False
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(current):
|
||||
return True
|
||||
current = current.__cause__ or current.__context__
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def handle_db_exception(e: Exception):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -34,6 +34,12 @@ def is_text_content_call_type(call_type: str) -> bool:
|
|||
|
||||
TEXT_PART_TYPES: FrozenSet[str] = frozenset({"text", "input_text", "output_text"})
|
||||
|
||||
# Responses-API item types whose ``output`` field carries user/tool text
|
||||
# that guardrails should inspect. ``function_call_output`` is the
|
||||
# built-in shape; ``custom_tool_call_output`` is the custom-tool
|
||||
# counterpart (see ``ChatCompletionCustomToolCallOutput``).
|
||||
_OUTPUT_ITEM_TYPES: frozenset[str] = frozenset({"function_call_output", "custom_tool_call_output"})
|
||||
|
||||
|
||||
def _iter_text_parts_in_content(content: Any) -> Iterator[str]:
|
||||
"""Yield text fragments from a ``message.content`` value (string or
|
||||
|
|
@ -72,7 +78,7 @@ def _coerce_input_to_messages(input_value: Any) -> List[Dict[str, Any]]:
|
|||
messages.append({"role": item.get("role") or "user", "content": [item]})
|
||||
elif "content" in item:
|
||||
messages.append({"role": item.get("role") or "user", "content": item["content"]})
|
||||
elif item.get("type") == "function_call_output" and "output" in item:
|
||||
elif item.get("type") in _OUTPUT_ITEM_TYPES and "output" in item:
|
||||
messages.append({"role": item.get("role") or "tool", "content": item["output"]})
|
||||
return messages
|
||||
|
||||
|
|
@ -157,7 +163,7 @@ def walk_user_text(data: Dict[str, Any], visit: Callable[[str], str]) -> int:
|
|||
input_value[idx] = {**item, "text": visit(item["text"])}
|
||||
elif "content" in item:
|
||||
item["content"] = _rewrite_content(item["content"])
|
||||
elif item.get("type") == "function_call_output" and "output" in item:
|
||||
elif item.get("type") in _OUTPUT_ITEM_TYPES and "output" in item:
|
||||
item["output"] = _rewrite_content(item["output"])
|
||||
return visited
|
||||
|
||||
|
|
|
|||
|
|
@ -772,7 +772,13 @@ class LassoGuardrail(CustomGuardrail):
|
|||
data: Request data (used for conversation_id generation and tools extraction)
|
||||
cache: Cache instance for storing conversation_id (optional for post-call)
|
||||
"""
|
||||
payload: Dict[str, Any] = {"messages": messages, "messageType": message_type}
|
||||
payload: Dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"messageType": message_type,
|
||||
# Drives the "Used By" badge on Lasso Application API Keys: every call from this
|
||||
# integration is attributed as "litellm" on the keys list.
|
||||
"source": {"type": "litellm"},
|
||||
}
|
||||
|
||||
# Add optional parameters if available
|
||||
if self.user_id:
|
||||
|
|
|
|||
|
|
@ -28,11 +28,20 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import (
|
|||
)
|
||||
from litellm.proxy.utils import ProxyUpdateSpend
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
StandardLoggingPayload,
|
||||
StandardLoggingPayloadErrorInformation,
|
||||
)
|
||||
from litellm.utils import get_end_user_id_for_cost_tracking
|
||||
|
||||
_PASS_THROUGH_CALL_TYPES: frozenset[str] = frozenset(
|
||||
{
|
||||
CallTypes.pass_through.value,
|
||||
CallTypes.llm_passthrough_route.value,
|
||||
CallTypes.allm_passthrough_route.value,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _ProxyDBLogger(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -219,11 +228,13 @@ class _ProxyDBLogger(CustomLogger):
|
|||
verbose_proxy_logger.debug(
|
||||
f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}"
|
||||
)
|
||||
call_type: Optional[str] = kwargs.get("call_type")
|
||||
if _should_track_cost_callback(
|
||||
user_api_key=user_api_key,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
end_user_id=end_user_id,
|
||||
call_type=call_type,
|
||||
):
|
||||
## UPDATE DATABASE
|
||||
await _update_database_and_spend_counters(
|
||||
|
|
@ -412,9 +423,15 @@ def _should_track_cost_callback(
|
|||
user_id: Optional[str],
|
||||
team_id: Optional[str],
|
||||
end_user_id: Optional[str],
|
||||
call_type: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if the cost callback should be tracked based on the kwargs
|
||||
|
||||
Pass-through endpoints can be configured with ``auth=false``, which leaves
|
||||
the request with no key/user/team/end-user to attribute spend to. Those
|
||||
requests still forward real provider traffic that operators expect to see
|
||||
in request/usage logs, so they are tracked even when unauthenticated.
|
||||
"""
|
||||
|
||||
# don't run track cost callback if user opted into disabling spend
|
||||
|
|
@ -423,7 +440,7 @@ def _should_track_cost_callback(
|
|||
|
||||
if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None:
|
||||
return True
|
||||
return False
|
||||
return call_type in _PASS_THROUGH_CALL_TYPES
|
||||
|
||||
|
||||
def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]:
|
||||
|
|
|
|||
|
|
@ -468,7 +468,10 @@ def handle_key_type(data: GenerateKeyRequest, data_json: dict) -> dict:
|
|||
Handle the key type.
|
||||
"""
|
||||
key_type = data.key_type
|
||||
data_json.pop("key_type", None)
|
||||
if key_type is None:
|
||||
data_json.pop("key_type", None)
|
||||
return data_json
|
||||
data_json["key_type"] = key_type.value
|
||||
if key_type == LiteLLMKeyType.LLM_API:
|
||||
data_json["allowed_routes"] = ["llm_api_routes"]
|
||||
elif key_type == LiteLLMKeyType.MANAGEMENT:
|
||||
|
|
@ -3566,6 +3569,7 @@ async def generate_key_helper_fn(
|
|||
created_by: Optional[str] = None,
|
||||
updated_by: Optional[str] = None,
|
||||
allowed_routes: Optional[list] = None,
|
||||
key_type: str | None = None,
|
||||
sso_user_id: Optional[str] = None,
|
||||
object_permission_id: Optional[str] = None, # object_permission_id <-> LiteLLM_ObjectPermissionTable
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None,
|
||||
|
|
@ -3706,6 +3710,7 @@ async def generate_key_helper_fn(
|
|||
"created_by": created_by,
|
||||
"updated_by": updated_by,
|
||||
"allowed_routes": allowed_routes or [],
|
||||
"key_type": key_type,
|
||||
"object_permission_id": object_permission_id,
|
||||
"router_settings": router_settings_json,
|
||||
"access_group_ids": access_group_ids or [],
|
||||
|
|
|
|||
|
|
@ -3550,7 +3550,7 @@ async def team_info(
|
|||
try:
|
||||
team_info: Optional[BaseModel] = await TeamRepository(prisma_client).table.find_unique(
|
||||
where={"team_id": team_id},
|
||||
include={"object_permission": True},
|
||||
include={"litellm_model_table": True, "object_permission": True},
|
||||
)
|
||||
if team_info is None:
|
||||
raise Exception
|
||||
|
|
|
|||
|
|
@ -4034,30 +4034,41 @@ class MicrosoftSSOHandler:
|
|||
base_url = MicrosoftSSOHandler.get_graph_api_base_url()
|
||||
# Endpoint to get app role assignments for the given service principal
|
||||
endpoint = f"/servicePrincipals/{service_principal_id}/appRoleAssignedTo"
|
||||
url = base_url + endpoint
|
||||
next_link: str | None = base_url + endpoint
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
response = await async_client.get(url, headers=headers)
|
||||
response_json = response.json()
|
||||
verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}")
|
||||
group_ids: List[str] = []
|
||||
service_principal_teams: List[MicrosoftServicePrincipalTeam] = []
|
||||
page_count = 0
|
||||
|
||||
for _object in response_json.get("value", []):
|
||||
if _object.get("principalType") == "Group":
|
||||
# Append the group ID to the list
|
||||
group_ids.append(_object.get("principalId"))
|
||||
# Append the service principal team to the list
|
||||
service_principal_teams.append(
|
||||
MicrosoftServicePrincipalTeam(
|
||||
principalDisplayName=_object.get("principalDisplayName"),
|
||||
principalId=_object.get("principalId"),
|
||||
while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
|
||||
response = await async_client.get(next_link, headers=headers)
|
||||
response_json = response.json()
|
||||
verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}")
|
||||
|
||||
for _object in response_json.get("value", []):
|
||||
if _object.get("principalType") == "Group":
|
||||
# Append the group ID to the list
|
||||
group_ids.append(_object.get("principalId"))
|
||||
# Append the service principal team to the list
|
||||
service_principal_teams.append(
|
||||
MicrosoftServicePrincipalTeam(
|
||||
principalDisplayName=_object.get("principalDisplayName"),
|
||||
principalId=_object.get("principalId"),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
next_link = response_json.get("@odata.nextLink")
|
||||
page_count += 1
|
||||
|
||||
if next_link is not None and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Reached maximum page limit of {MicrosoftSSOHandler.MAX_GRAPH_API_PAGES}. Some service principal group assignments may not be included."
|
||||
)
|
||||
|
||||
return group_ids, service_principal_teams
|
||||
|
||||
|
|
|
|||
|
|
@ -14805,6 +14805,7 @@ async def get_config_list(
|
|||
"forward_client_headers_to_llm_api": {"type": "Boolean"},
|
||||
"mcp_required_fields": {"type": "List"},
|
||||
"cancel_on_disconnect": {"type": "Boolean"},
|
||||
"skip_user_budget_on_team_key": {"type": "Boolean"},
|
||||
}
|
||||
|
||||
return_val = []
|
||||
|
|
|
|||
|
|
@ -422,6 +422,7 @@ model LiteLLM_VerificationToken {
|
|||
budget_reset_at DateTime?
|
||||
allowed_cache_controls String[] @default([])
|
||||
allowed_routes String[] @default([])
|
||||
key_type String?
|
||||
policies String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
model_spend Json @default("{}")
|
||||
|
|
@ -516,6 +517,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
budget_reset_at DateTime?
|
||||
allowed_cache_controls String[] @default([])
|
||||
allowed_routes String[] @default([])
|
||||
key_type String?
|
||||
policies String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
model_spend Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -123,6 +123,7 @@ async def reserve_budget_for_request(
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
end_user_id: Optional[str] = None,
|
||||
end_user_object: Optional[Any] = None,
|
||||
skip_user_budget_on_team_key: bool = False,
|
||||
) -> Optional[dict]:
|
||||
if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
|
||||
return None
|
||||
|
|
@ -141,6 +142,7 @@ async def reserve_budget_for_request(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
end_user_id=end_user_id,
|
||||
end_user_object=end_user_object,
|
||||
skip_user_budget_on_team_key=skip_user_budget_on_team_key,
|
||||
)
|
||||
if not counters:
|
||||
return None
|
||||
|
|
@ -296,6 +298,7 @@ async def _get_budget_counters(
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
end_user_id: Optional[str] = None,
|
||||
end_user_object: Optional[Any] = None,
|
||||
skip_user_budget_on_team_key: bool = False,
|
||||
) -> List[_BudgetCounter]:
|
||||
counters: List[_BudgetCounter] = []
|
||||
|
||||
|
|
@ -344,8 +347,9 @@ async def _get_budget_counters(
|
|||
)
|
||||
)
|
||||
|
||||
is_team_key = team_object is not None and team_object.team_id is not None
|
||||
if (
|
||||
(team_object is None or team_object.team_id is None)
|
||||
not (is_team_key and skip_user_budget_on_team_key)
|
||||
and user_object is not None
|
||||
and user_object.user_id is not None
|
||||
and user_object.max_budget is not None
|
||||
|
|
|
|||
|
|
@ -261,7 +261,7 @@ async def aresponses_api_with_mcp(
|
|||
pre_processed_mcp_tools=original_mcp_tools,
|
||||
)
|
||||
|
||||
return LiteLLM_Proxy_MCP_Handler._create_mcp_streaming_response(
|
||||
mcp_streaming_response = LiteLLM_Proxy_MCP_Handler._create_mcp_streaming_response(
|
||||
input=input,
|
||||
model=model,
|
||||
all_tools=all_tools,
|
||||
|
|
@ -272,6 +272,10 @@ async def aresponses_api_with_mcp(
|
|||
tool_server_map=tool_server_map,
|
||||
**kwargs,
|
||||
)
|
||||
await mcp_streaming_response._create_initial_response_iterator()
|
||||
if mcp_streaming_response._initial_creation_error is not None:
|
||||
raise mcp_streaming_response._initial_creation_error
|
||||
return mcp_streaming_response
|
||||
|
||||
# Determine if we should auto-execute tools
|
||||
should_auto_execute = bool(mcp_tools_with_litellm_proxy) and LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from litellm._uuid import uuid
|
|||
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
|
||||
from litellm.types.llms.openai import (
|
||||
BaseLiteLLMOpenAIResponseObject,
|
||||
ErrorEvent,
|
||||
ErrorEventError,
|
||||
MCPCallArgumentsDeltaEvent,
|
||||
MCPCallArgumentsDoneEvent,
|
||||
MCPCallCompletedEvent,
|
||||
|
|
@ -316,6 +318,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
# Cache the response ID to ensure consistency across all events
|
||||
self._cached_response_id: Optional[str] = None
|
||||
|
||||
self._initial_creation_error: Exception | None = None
|
||||
self._stream_error: Exception | None = None
|
||||
self._error_event_emitted = False
|
||||
self._last_sequence_number = 0
|
||||
|
||||
def _extract_mcp_headers_from_params(self) -> None:
|
||||
"""Extract MCP headers from original request params to pass to tool calls"""
|
||||
from typing import Dict, Optional
|
||||
|
|
@ -380,10 +387,31 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
return LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(self.mcp_tools_with_litellm_proxy)
|
||||
|
||||
def _make_stream_error_event(self) -> ResponsesAPIStreamingResponse:
|
||||
err = self._stream_error
|
||||
status_code = getattr(err, "status_code", None)
|
||||
return ErrorEvent(
|
||||
type=ResponsesAPIStreamEvents.ERROR,
|
||||
sequence_number=self._last_sequence_number + 1,
|
||||
error=ErrorEventError(
|
||||
type="mcp_gateway_error",
|
||||
code=str(status_code) if status_code is not None else "internal_error",
|
||||
message=str(err) if err is not None else "MCP gateway stream failed",
|
||||
param=None,
|
||||
),
|
||||
)
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> ResponsesAPIStreamingResponse:
|
||||
chunk = await self._anext_impl()
|
||||
sequence_number = getattr(chunk, "sequence_number", None)
|
||||
if isinstance(sequence_number, int) and sequence_number > self._last_sequence_number:
|
||||
self._last_sequence_number = sequence_number
|
||||
return chunk
|
||||
|
||||
async def _anext_impl(self) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
Phase-based streaming:
|
||||
1. initial_response - Stream the first LLM response (includes response.created, response.in_progress, response.output_item.added)
|
||||
|
|
@ -438,10 +466,16 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.phase = "continue_initial_response"
|
||||
return await self.__anext__()
|
||||
self.phase = "finished"
|
||||
if self._stream_error is not None and not self._error_event_emitted:
|
||||
self._error_event_emitted = True
|
||||
return self._make_stream_error_event()
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Phase 6: Finished
|
||||
if self.phase == "finished":
|
||||
if self._stream_error is not None and not self._error_event_emitted:
|
||||
self._error_event_emitted = True
|
||||
return self._make_stream_error_event()
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Should not reach here
|
||||
|
|
@ -530,6 +564,12 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
chunk = await cast(Any, self.base_iterator).__anext__() # type: ignore[attr-defined]
|
||||
|
||||
if self._cached_response_id is None and hasattr(chunk, "response"):
|
||||
new_response = getattr(chunk, "response", None)
|
||||
new_response_id = getattr(new_response, "id", None) if new_response is not None else None
|
||||
if new_response_id:
|
||||
self._cached_response_id = new_response_id
|
||||
|
||||
# Ensure response ID consistency - update chunk if needed
|
||||
if self._cached_response_id and hasattr(chunk, "response"):
|
||||
response_obj = getattr(chunk, "response", None)
|
||||
|
|
@ -589,6 +629,8 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
traceback.print_exc()
|
||||
self.base_iterator = None
|
||||
self._initial_creation_error = e
|
||||
self._stream_error = e
|
||||
# Don't set phase to "finished" here — let __anext__ emit any
|
||||
# pre-generated MCP discovery events before ending the iteration.
|
||||
|
||||
|
|
@ -761,6 +803,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
if hasattr(follow_up_response, "__aiter__"):
|
||||
self.base_iterator = follow_up_response
|
||||
self.collected_response = None
|
||||
self._cached_response_id = None
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error creating follow-up iterator: {e}")
|
||||
|
|
@ -768,6 +811,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
traceback.print_exc()
|
||||
self.base_iterator = None
|
||||
self._stream_error = e
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
|
|
|||
|
|
@ -7605,13 +7605,13 @@ class Router:
|
|||
model_name = entry.get("model_name") if isinstance(entry, dict) else entry.model_name
|
||||
if not model_name or not lp:
|
||||
continue
|
||||
if model_name in self.adaptive_routers:
|
||||
continue
|
||||
deployment = Deployment(
|
||||
model_name=model_name,
|
||||
litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)),
|
||||
model_info=(entry.get("model_info") if isinstance(entry, dict) else entry.model_info),
|
||||
)
|
||||
if model_name in self.adaptive_routers:
|
||||
continue
|
||||
self.init_adaptive_router_deployment(deployment=deployment)
|
||||
|
||||
for model_name, complexity_router in self.complexity_routers.items():
|
||||
|
|
@ -10707,56 +10707,39 @@ class Router:
|
|||
if self.routing_plugins:
|
||||
await self._run_routing_plugins(model=model, request_kwargs=request_kwargs, messages=messages)
|
||||
|
||||
#########################################################
|
||||
# Check if any auto-router should be used
|
||||
#########################################################
|
||||
if model in self.auto_routers:
|
||||
return await self.auto_routers[model].async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
router_strategy = (
|
||||
self.auto_routers.get(model)
|
||||
or self.complexity_routers.get(model)
|
||||
or self.adaptive_routers.get(model)
|
||||
or self.quality_routers.get(model)
|
||||
)
|
||||
if router_strategy is None:
|
||||
return None
|
||||
|
||||
#########################################################
|
||||
# Check if any complexity-router should be used
|
||||
#########################################################
|
||||
if model in self.complexity_routers:
|
||||
return await self.complexity_routers[model].async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
pre_routing_hook_response = await router_strategy.async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Check if an adaptive-router should be used
|
||||
#########################################################
|
||||
adaptive_router = self.adaptive_routers.get(model)
|
||||
if adaptive_router is not None:
|
||||
return await adaptive_router.async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
# `model` (the alias, e.g. "smart-router") is never the deployment actually
|
||||
# called - apply the alias's own litellm_params (besides `model` itself,
|
||||
# which is just the alias marker) to the request, since the tier/route
|
||||
# deployment the hook selected won't have them. Router-only fields
|
||||
# (tpm, rpm, weight, complexity_router_config, ...) are excluded from the
|
||||
# actual outbound LLM call downstream by litellm.types.utils.all_litellm_params,
|
||||
# not here.
|
||||
if pre_routing_hook_response is not None:
|
||||
alias_index = self.model_name_to_deployment_indices.get(model, [])
|
||||
if alias_index:
|
||||
alias_litellm_params = self.model_list[alias_index[0]].get("litellm_params", {})
|
||||
for key, value in alias_litellm_params.items():
|
||||
if key != "model" and value is not None:
|
||||
request_kwargs.setdefault(key, value)
|
||||
|
||||
#########################################################
|
||||
# Check if any quality-router should be used
|
||||
#########################################################
|
||||
if model in self.quality_routers:
|
||||
return await self.quality_routers[model].async_pre_routing_hook(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
|
||||
return None
|
||||
return pre_routing_hook_response
|
||||
|
||||
def get_available_deployment(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from __future__ import annotations
|
|||
import asyncio
|
||||
import random
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Literal, Union, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -809,6 +809,44 @@ class ComplexityRouter(CustomLogger):
|
|||
|
||||
return user_message, system_prompt
|
||||
|
||||
@staticmethod
|
||||
def _iter_metadata_dicts(request_kwargs: dict) -> list[dict]:
|
||||
"""Metadata may land on `metadata` or `litellm_metadata` depending on the
|
||||
endpoint, mirroring DeploymentAffinityCheck's precedence."""
|
||||
return [
|
||||
metadata
|
||||
for metadata_key in ("litellm_metadata", "metadata")
|
||||
if isinstance(metadata := request_kwargs.get(metadata_key), dict)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id_from_request_kwargs(request_kwargs: dict) -> str | None:
|
||||
"""Resolve a client-supplied session_id."""
|
||||
for metadata in ComplexityRouter._iter_metadata_dicts(request_kwargs):
|
||||
session_id = metadata.get("session_id")
|
||||
if session_id is not None:
|
||||
return str(session_id)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _get_user_api_key_hash_from_request_kwargs(request_kwargs: dict) -> str | None:
|
||||
"""Resolve the proxy-derived API key hash, the same trust boundary
|
||||
DeploymentAffinityCheck uses for its own key-based affinity (not the
|
||||
client-supplied OpenAI `user` param, which isn't authenticated)."""
|
||||
for metadata in ComplexityRouter._iter_metadata_dicts(request_kwargs):
|
||||
user_key = metadata.get("user_api_key_hash")
|
||||
if user_key is not None:
|
||||
return str(user_key)
|
||||
return None
|
||||
|
||||
def _get_session_affinity_cache_key(self, session_id: str, request_kwargs: dict) -> str:
|
||||
# Namespace by the caller's API key hash so two different callers reusing the
|
||||
# same client-supplied session_id can't poison each other's routing pin. Falls
|
||||
# back to "unscoped" only when there's no authenticated caller to scope by
|
||||
# (e.g. direct Router usage without the proxy layer).
|
||||
caller_scope = self._get_user_api_key_hash_from_request_kwargs(request_kwargs) or "unscoped"
|
||||
return f"complexity_router_session_affinity:v1:{self.model_name}:{caller_scope}:{session_id}"
|
||||
|
||||
async def async_pre_routing_hook(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -816,10 +854,70 @@ class ComplexityRouter(CustomLogger):
|
|||
messages: list[dict[str, Any]] | None = None,
|
||||
input: Union[str, list] | None = None,
|
||||
specific_deployment: bool | None = False,
|
||||
) -> Optional[PreRoutingHookResponse]:
|
||||
) -> PreRoutingHookResponse | None:
|
||||
"""
|
||||
Pre-routing hook called before the routing decision.
|
||||
|
||||
When `session_affinity` is enabled and a session_id is resolvable on the request,
|
||||
pins the model chosen on the session's first turn and reuses it for every later
|
||||
turn, skipping classification entirely. Otherwise delegates to `_classify_and_route`.
|
||||
"""
|
||||
from litellm.types.router import PreRoutingHookResponse
|
||||
|
||||
session_id = self._get_session_id_from_request_kwargs(request_kwargs) if self.config.session_affinity else None
|
||||
cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None
|
||||
|
||||
if cache_key is not None:
|
||||
pinned_model = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
|
||||
if isinstance(pinned_model, str):
|
||||
# Refresh the TTL on every hit so an active session doesn't lose its
|
||||
# pin mid-conversation just because it outlives the original write.
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=pinned_model,
|
||||
ttl=self.config.session_affinity_ttl_seconds,
|
||||
)
|
||||
if self.config.adaptive:
|
||||
from litellm.router_strategy.adaptive_router.config import (
|
||||
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
|
||||
)
|
||||
|
||||
kwargs_metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(kwargs_metadata, dict):
|
||||
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = pinned_model
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause=session_affinity_pin, routed_model={pinned_model}"
|
||||
)
|
||||
has_original_messages = messages is not None and len(messages) > 0
|
||||
return PreRoutingHookResponse(
|
||||
model=pinned_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
response = await self._classify_and_route(
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
if cache_key is not None and response is not None:
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=response.model,
|
||||
ttl=self.config.session_affinity_ttl_seconds,
|
||||
)
|
||||
return response
|
||||
|
||||
async def _classify_and_route(
|
||||
self,
|
||||
model: str,
|
||||
request_kwargs: dict,
|
||||
messages: list[dict[str, Any]] | None = None,
|
||||
input: Union[str, list] | None = None,
|
||||
specific_deployment: bool | None = False,
|
||||
) -> PreRoutingHookResponse | None:
|
||||
"""
|
||||
Classifies the request by complexity and returns the appropriate model.
|
||||
Supports chat completions (messages), Responses API (input), and other
|
||||
formats via the guardrail translation handler dispatch.
|
||||
|
|
|
|||
|
|
@ -361,6 +361,20 @@ class ComplexityRouterConfig(BaseModel):
|
|||
description="Minimum cosine similarity for a semantic keyword match",
|
||||
)
|
||||
|
||||
# Session affinity: pin the first turn's routed model for the rest of the session
|
||||
session_affinity: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"When True and a session_id is resolvable on the request, pin the model chosen on the "
|
||||
"session's first turn and reuse it for every later turn, skipping re-classification."
|
||||
),
|
||||
)
|
||||
session_affinity_ttl_seconds: int = Field(
|
||||
default=3600,
|
||||
gt=0,
|
||||
description="TTL for the session affinity pin; refreshed on every cache hit",
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="allow") # Allow additional fields
|
||||
|
||||
@field_validator("tiers", mode="before")
|
||||
|
|
|
|||
|
|
@ -213,6 +213,8 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_input_audio_tokens_metric",
|
||||
"litellm_output_reasoning_tokens_metric",
|
||||
"litellm_output_audio_tokens_metric",
|
||||
"litellm_video_duration_seconds_metric",
|
||||
"litellm_images_generated_metric",
|
||||
"litellm_deployment_successful_fallbacks",
|
||||
"litellm_deployment_failed_fallbacks",
|
||||
"litellm_remaining_team_budget_metric",
|
||||
|
|
@ -506,6 +508,9 @@ class PrometheusMetricLabels:
|
|||
litellm_output_reasoning_tokens_metric = litellm_output_tokens_metric
|
||||
litellm_output_audio_tokens_metric = litellm_output_tokens_metric
|
||||
|
||||
litellm_video_duration_seconds_metric = litellm_output_tokens_metric
|
||||
litellm_images_generated_metric = litellm_output_tokens_metric
|
||||
|
||||
litellm_deployment_state = [
|
||||
UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value,
|
||||
UserAPIKeyLabelNames.MODEL_ID.value,
|
||||
|
|
@ -717,6 +722,8 @@ class PrometheusMetricLabels:
|
|||
"litellm_input_tokens_metric",
|
||||
"litellm_total_tokens_metric",
|
||||
"litellm_output_tokens_metric",
|
||||
"litellm_video_duration_seconds_metric",
|
||||
"litellm_images_generated_metric",
|
||||
}
|
||||
)
|
||||
# Managed batch metrics
|
||||
|
|
|
|||
|
|
@ -3210,6 +3210,16 @@ all_litellm_params = (
|
|||
"_litellm_tpm_reserved_model",
|
||||
"_litellm_tpm_reserved_scopes",
|
||||
"_litellm_tpm_reservation_released",
|
||||
"auto_router_config_path",
|
||||
"auto_router_config",
|
||||
"auto_router_default_model",
|
||||
"auto_router_embedding_model",
|
||||
"complexity_router_config",
|
||||
"complexity_router_default_model",
|
||||
"adaptive_router_config",
|
||||
"adaptive_router_default_model",
|
||||
"quality_router_config",
|
||||
"quality_router_default_model",
|
||||
]
|
||||
+ list(StandardCallbackDynamicParams.__annotations__.keys())
|
||||
+ list(CustomPricingLiteLLMParams.model_fields.keys())
|
||||
|
|
|
|||
|
|
@ -11331,6 +11331,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
|
|
@ -11362,6 +11363,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
|
|
@ -11424,6 +11426,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_max_reasoning_effort": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -11479,6 +11482,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
|
|
@ -11506,6 +11510,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
|
|
@ -11559,6 +11564,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_output_config": true
|
||||
|
|
@ -11586,6 +11592,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_output_config": true
|
||||
|
|
@ -11614,6 +11621,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"provider_specific_entry": {
|
||||
|
|
@ -11648,6 +11656,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"provider_specific_entry": {
|
||||
|
|
@ -11682,6 +11691,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -11718,6 +11728,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
@ -11788,6 +11799,7 @@
|
|||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_native_structured_output": true,
|
||||
"supports_sampling_params": false,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm"
|
||||
version = "1.93.0"
|
||||
version = "1.94.0"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10, <3.14"
|
||||
|
|
@ -62,8 +62,8 @@ proxy = [
|
|||
"azure-identity>=1.25.2,<2.0",
|
||||
"azure-storage-blob>=12.28.0,<13.0",
|
||||
"mcp>=1.26.0,<2.0",
|
||||
"litellm-proxy-extras==0.4.76",
|
||||
"litellm-enterprise==0.1.49",
|
||||
"litellm-proxy-extras==0.4.77",
|
||||
"litellm-enterprise==0.1.50",
|
||||
"RestrictedPython>=8.1,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
"polars>=1.38.1,<2.0",
|
||||
|
|
@ -205,7 +205,7 @@ ci = [
|
|||
# protobuf, Pillow is a compiled C extension).
|
||||
"tenacity==8.5.0",
|
||||
"google-generativeai==0.8.6",
|
||||
"Pillow==12.2.0",
|
||||
"Pillow==12.3.0",
|
||||
# Azure batch E2E tests still import psycopg2 directly.
|
||||
"psycopg2-binary==2.9.11",
|
||||
"pytest-codspeed==4.3.0",
|
||||
|
|
@ -264,6 +264,8 @@ constraint-dependencies = [
|
|||
"aiohttp>=3.14.1,<4.0",
|
||||
"packaging>=24.0",
|
||||
"soupsieve>=2.8.4",
|
||||
"httplib2>=0.32.0",
|
||||
"setuptools>=83.0.0",
|
||||
]
|
||||
override-dependencies = [
|
||||
# a2a-sdk 1.x requires packaging>=24.0; lunary 1.4.x still caps at <24.0.
|
||||
|
|
@ -284,7 +286,7 @@ members = ["enterprise", "litellm-proxy-extras"]
|
|||
profile = "black"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.93.0"
|
||||
version = "1.94.0"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -422,6 +422,7 @@ model LiteLLM_VerificationToken {
|
|||
budget_reset_at DateTime?
|
||||
allowed_cache_controls String[] @default([])
|
||||
allowed_routes String[] @default([])
|
||||
key_type String?
|
||||
policies String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
model_spend Json @default("{}")
|
||||
|
|
@ -516,6 +517,7 @@ model LiteLLM_DeletedVerificationToken {
|
|||
budget_reset_at DateTime?
|
||||
allowed_cache_controls String[] @default([])
|
||||
allowed_routes String[] @default([])
|
||||
key_type String?
|
||||
policies String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
model_spend Json @default("{}")
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
|
|||
- `embeddings/` - the `/embeddings` endpoint across providers
|
||||
- `batches/` - the `/batches` endpoint (placeholder until the first test lands)
|
||||
- `realtime/` - realtime websocket sessions, including the pipecat audio path
|
||||
- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window) and `spend_tracking/` (spend logging and cost attribution on `/spend/*`)
|
||||
- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`)
|
||||
- `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials; also the dashboard UI behavior on top of them, driven through the proxy-served UI at /ui with playwright (optional dep behind importorskip)
|
||||
- `logging/` - logging-integration delivery (datadog and friends)
|
||||
- `security/` - secret handling and log-leak protection
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Callable
|
||||
|
||||
import pytest
|
||||
|
|
@ -49,7 +50,7 @@ from e2e_http import (
|
|||
unwrap,
|
||||
)
|
||||
from lifecycle import ResourceManager
|
||||
from models import KeyGenerateBody, SpendLogRow, SpendLogsParams
|
||||
from models import KeyGenerateBody, SpendLogRow
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
|
@ -349,6 +350,10 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
|
|||
the file-read path fires while the batch itself is not blocked.
|
||||
``resources.key()`` cannot set limits, so the key is minted on the gateway
|
||||
directly and its delete deferred.
|
||||
|
||||
Snapshots read /spend/logs/v2 over a bounded window around the test instead
|
||||
of the unpaginated /spend/logs whole-table read, which grows with the
|
||||
environment and OOMed the e2e runner on stage.
|
||||
"""
|
||||
user_id = f"e2e-batch-rl-{unique_marker()}"
|
||||
key = client.gateway.generate_key(
|
||||
|
|
@ -356,8 +361,13 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
|
|||
)
|
||||
resources.defer(lambda: client.gateway.delete_key(key))
|
||||
|
||||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
window_end = window_start + timedelta(hours=2)
|
||||
before = frozenset(
|
||||
row.request_id for row in unattributed_rows(client.gateway.spend_logs(SpendLogsParams()))
|
||||
row.request_id
|
||||
for row in unattributed_rows(
|
||||
client.gateway.spend_logs_window(start=window_start, end=window_end)
|
||||
)
|
||||
)
|
||||
|
||||
file = unwrap(
|
||||
|
|
@ -379,7 +389,9 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
|
|||
|
||||
new_orphans = [
|
||||
row
|
||||
for row in unattributed_rows(client.gateway.spend_logs(SpendLogsParams()))
|
||||
for row in unattributed_rows(
|
||||
client.gateway.spend_logs_window(start=window_start, end=window_end)
|
||||
)
|
||||
if row.request_id not in before
|
||||
]
|
||||
assert not new_orphans, (
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@
|
|||
- {id: logging.datadog.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/datadog/datadog.py", rationale: "Powers dashboards/alerts; cardinality regressions common"}
|
||||
- {id: logging.datadog.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions], source: "integrations/datadog/datadog.py", rationale: "Failure metrics for alerting/SLO"}
|
||||
- {id: logging.prometheus.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/prometheus.py", rationale: "Standard OSS metrics; per-key cardinality (existing e2e)"}
|
||||
- {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"}
|
||||
- {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"}
|
||||
- {id: logging.otel.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions, messages], source: "integrations/otel/logger.py", rationale: "Error spans for observability continuity"}
|
||||
- {id: logging.braintrust.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/braintrust_logging.py", rationale: "Evals platform spend"}
|
||||
- {id: logging.langsmith.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langsmith.py", rationale: "LangChain ecosystem"}
|
||||
|
|
|
|||
|
|
@ -18,6 +18,12 @@ configs:
|
|||
type: redis
|
||||
host: redis
|
||||
port: 6379
|
||||
# OTEL v2 trace destination for the logging suite's trace-completeness
|
||||
# tests: the arize_phoenix preset is OTLP with a configurable endpoint
|
||||
# (PHOENIX_COLLECTOR_HTTP_ENDPOINT below points it at the jaeger service),
|
||||
# so gen-AI spans export through a preset-owned provider - the code path
|
||||
# where trace splits actually happen - with no cloud credentials needed.
|
||||
callbacks: ["arize_phoenix"]
|
||||
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
|
|
@ -68,9 +74,14 @@ services:
|
|||
condition: service_healthy
|
||||
redis:
|
||||
condition: service_healthy
|
||||
jaeger:
|
||||
condition: service_healthy
|
||||
env_file: .env
|
||||
environment:
|
||||
LITELLM_MASTER_KEY: sk-1234
|
||||
LITELLM_OTEL_V2: "true"
|
||||
PHOENIX_COLLECTOR_HTTP_ENDPOINT: http://jaeger:4318/v1/traces
|
||||
PHOENIX_API_KEY: local-jaeger-noauth
|
||||
DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm
|
||||
UI_USERNAME: admin
|
||||
UI_PASSWORD: sk-1234
|
||||
|
|
@ -114,3 +125,15 @@ services:
|
|||
interval: 3s
|
||||
timeout: 3s
|
||||
retries: 20
|
||||
|
||||
# throwaway OTEL trace destination (OTLP ingest on 4318 inside the network,
|
||||
# query API on host 16686 for test read-back; see E2E_OTEL_QUERY_URL)
|
||||
jaeger:
|
||||
image: jaegertracing/all-in-one:1.62.0
|
||||
ports:
|
||||
- "16686:16686"
|
||||
healthcheck:
|
||||
test: ["CMD", "wget", "-qO-", "http://localhost:14269/"]
|
||||
interval: 3s
|
||||
timeout: 3s
|
||||
retries: 20
|
||||
|
|
|
|||
|
|
@ -24,6 +24,14 @@ CONTROL_PLANE_BASE_URL = os.environ.get(
|
|||
UI_USERNAME = os.environ.get("E2E_UI_USERNAME", "admin")
|
||||
UI_PASSWORD = os.environ.get("E2E_UI_PASSWORD", MASTER_KEY)
|
||||
|
||||
CHEAP_ANTHROPIC_MODEL = os.environ.get("E2E_CHEAP_ANTHROPIC_MODEL", "claude-haiku-4-5")
|
||||
CHEAP_OPENAI_MODEL = os.environ.get("E2E_CHEAP_OPENAI_MODEL", "gpt-5.5")
|
||||
|
||||
# Jaeger query API of the compose stack's OTEL trace destination (the `jaeger`
|
||||
# service in docker-compose.yml maps it to host 16686). Trace-completeness tests
|
||||
# read exported spans back through it.
|
||||
OTEL_QUERY_URL = os.environ.get("E2E_OTEL_QUERY_URL", "http://localhost:16686").rstrip("/")
|
||||
|
||||
# Writes on the proxy are eventually consistent (e.g. spend rows flush on
|
||||
# proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once.
|
||||
POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120"))
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import time
|
|||
import warnings
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
|
||||
from e2e_http import (
|
||||
NoBody,
|
||||
|
|
@ -50,6 +51,8 @@ from models import (
|
|||
OcrResponse,
|
||||
SpendLogRow,
|
||||
SpendLogs,
|
||||
SpendLogsPage,
|
||||
SpendLogsPageParams,
|
||||
SpendLogsParams,
|
||||
)
|
||||
from e2e_config import (
|
||||
|
|
@ -255,6 +258,28 @@ class Gateway:
|
|||
case _:
|
||||
return []
|
||||
|
||||
def spend_logs_window(self, *, start: datetime, end: datetime) -> list[SpendLogRow]:
|
||||
def fetch(page: int) -> SpendLogsPage:
|
||||
return unwrap(
|
||||
self.transport.get(
|
||||
"/spend/logs/v2",
|
||||
headers=self.transport.master,
|
||||
params=SpendLogsPageParams(
|
||||
start_date=start.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
end_date=end.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
page=page,
|
||||
page_size=100,
|
||||
),
|
||||
response_type=SpendLogsPage,
|
||||
)
|
||||
)
|
||||
|
||||
first = fetch(1)
|
||||
return [
|
||||
*first.data,
|
||||
*(row for page in range(2, first.total_pages + 1) for row in fetch(page).data),
|
||||
]
|
||||
|
||||
def poll_logs_for_key(
|
||||
self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None
|
||||
) -> list[SpendLogRow]:
|
||||
|
|
|
|||
|
|
@ -109,14 +109,16 @@ class StreamingResponse(BaseModel):
|
|||
"""Raw outcome for calls whose body is provider-native or streamed: status, the
|
||||
x-litellm-call-id header, the x-litellm-response-cost header (StandardLogging
|
||||
response_cost), the content-type (which tells streaming `text/event-stream` from
|
||||
non-streaming `application/json`), and the body. SpendLogs.request_id is the
|
||||
completion body id, not call_id. Used by passthrough and streaming, where one
|
||||
validated JSON model does not fit."""
|
||||
non-streaming `application/json`), the response headers (lowercased names, e.g.
|
||||
the x-ratelimit-* pacing headers and retry-after on a 429), and the body.
|
||||
SpendLogs.request_id is the completion body id, not call_id. Used by passthrough
|
||||
and streaming, where one validated JSON model does not fit."""
|
||||
|
||||
status_code: int
|
||||
call_id: str | None = None # x-litellm-call-id header
|
||||
response_cost: float | None = None # x-litellm-response-cost header
|
||||
content_type: str | None = None
|
||||
headers: dict[str, str] = {}
|
||||
body: str
|
||||
chunks: int = 0 # streamed events (0 for non-streaming)
|
||||
|
||||
|
|
@ -276,12 +278,14 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon
|
|||
call_id = _hdr(resp, "x-litellm-call-id")
|
||||
response_cost = _parse_response_cost(resp)
|
||||
content_type = _hdr(resp, "content-type")
|
||||
headers = {name.lower(): value for name, value in resp.headers.items()}
|
||||
if not stream or not (200 <= resp.status_code < 300):
|
||||
return StreamingResponse(
|
||||
status_code=resp.status_code,
|
||||
call_id=call_id,
|
||||
response_cost=response_cost,
|
||||
content_type=content_type,
|
||||
headers=headers,
|
||||
body=resp.text,
|
||||
)
|
||||
lines = cast("Iterator[bytes]", resp.iter_lines())
|
||||
|
|
@ -291,6 +295,7 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon
|
|||
call_id=call_id,
|
||||
response_cost=response_cost,
|
||||
content_type=content_type,
|
||||
headers=headers,
|
||||
body="<streamed>",
|
||||
chunks=chunks,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import os
|
|||
import pytest
|
||||
|
||||
from logging_client import LangfuseCreds, LoggingClient, build_logging_client, load_langfuse_creds
|
||||
from otel_client import OtelReader, build_otel_reader
|
||||
|
||||
|
||||
def pytest_configure(config: pytest.Config) -> None:
|
||||
|
|
@ -28,6 +29,12 @@ def client() -> LoggingClient:
|
|||
return build_logging_client()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def otel_reader() -> OtelReader:
|
||||
"""Read-back client for the compose stack's Jaeger trace destination."""
|
||||
return build_otel_reader()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def datadog_creds() -> None:
|
||||
"""Require Datadog shipping credentials. Hard-fail when absent; never skip."""
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ from e2e_http import (
|
|||
unwrap,
|
||||
)
|
||||
from models import (
|
||||
AnthropicMessagesBody,
|
||||
ChatBody,
|
||||
ChatMessage,
|
||||
ChatResponse,
|
||||
|
|
@ -75,6 +76,14 @@ WEATHER_TOOL = ChatTool(
|
|||
)
|
||||
|
||||
|
||||
class ResponsesRequestBody(BaseModel):
|
||||
"""OpenAI Responses API /v1/responses request (non-streaming)."""
|
||||
|
||||
model: str
|
||||
input: str
|
||||
max_output_tokens: int
|
||||
|
||||
|
||||
class TeamCallbackBody(BaseModel):
|
||||
callback_name: Literal["langfuse_otel", "langfuse", "langsmith", "gcs"]
|
||||
callback_type: Literal["success", "failure", "success_and_failure"]
|
||||
|
|
@ -455,6 +464,32 @@ class LoggingClient:
|
|||
json=body,
|
||||
)
|
||||
|
||||
def messages_raw(self, key: str, model: str, text: str, *, max_tokens: int = 16) -> StreamingResponse:
|
||||
"""Non-streaming POST /v1/messages (Anthropic-native body): raw outcome
|
||||
judged by status/body/headers, for tests that need x-litellm-call-id."""
|
||||
return self.gateway.transport.send(
|
||||
"/v1/messages",
|
||||
headers=self.gateway.transport.bearer(key),
|
||||
json=AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=max_tokens,
|
||||
messages=[ChatMessage(role="user", content=text)],
|
||||
),
|
||||
)
|
||||
|
||||
def responses_raw(
|
||||
self, key: str, model: str, text: str, *, max_output_tokens: int = 64
|
||||
) -> StreamingResponse:
|
||||
"""Non-streaming POST /v1/responses (OpenAI Responses API): raw outcome
|
||||
judged by status/body/headers, for tests that need x-litellm-call-id.
|
||||
max_output_tokens caps reasoning-model output cost; a capped response is
|
||||
still a 200 and still exports the trace."""
|
||||
return self.gateway.transport.send(
|
||||
"/v1/responses",
|
||||
headers=self.gateway.transport.bearer(key),
|
||||
json=ResponsesRequestBody(model=model, input=text, max_output_tokens=max_output_tokens),
|
||||
)
|
||||
|
||||
def scrape_metrics(self) -> str:
|
||||
return self.gateway.probe("/metrics", params=NoBody()).body
|
||||
|
||||
|
|
|
|||
138
tests/e2e/logging/otel_client.py
Normal file
138
tests/e2e/logging/otel_client.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
"""Jaeger read-back for the OTEL trace-completeness tests: typed models over the
|
||||
Jaeger query API (the destination's own API - completeness is judged on what the
|
||||
backend actually holds, never on "export succeeded" proxy-side).
|
||||
|
||||
Traces are fetched server-side by the ``litellm.call_id`` tag the gen-AI span
|
||||
carries (the request's x-litellm-call-id response header), so read-back is
|
||||
immune to the query page filling up with unrelated traffic (background jobs,
|
||||
other suites sharing the stack). Jaeger returns every span of a matching trace,
|
||||
so the completeness assertions see the whole tree. A failed query is a hard
|
||||
failure, never an empty result - an unreachable destination must not read as
|
||||
"the trace never arrived".
|
||||
|
||||
External reads go through ``e2e_http`` (the only module allowed to call
|
||||
``requests.*``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from e2e_config import OTEL_QUERY_URL, POLL_INTERVAL, POLL_TIMEOUT
|
||||
from e2e_http import URL, NoBody, Success, get
|
||||
|
||||
#: OTEL resource service.name the proxy exports under (OTEL_SERVICE_NAME default).
|
||||
JAEGER_SERVICE = "litellm"
|
||||
#: Span tag carrying the request's x-litellm-call-id (stamped on the gen-AI span).
|
||||
CALL_ID_TAG = "litellm.call_id"
|
||||
|
||||
|
||||
class JaegerTag(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
key: str
|
||||
value: str | int | float | bool | None = None
|
||||
|
||||
|
||||
class JaegerReference(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", populate_by_name=True)
|
||||
|
||||
ref_type: str = Field(alias="refType")
|
||||
trace_id: str = Field(alias="traceID")
|
||||
span_id: str = Field(alias="spanID")
|
||||
|
||||
|
||||
class JaegerSpan(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", populate_by_name=True)
|
||||
|
||||
span_id: str = Field(alias="spanID")
|
||||
operation_name: str = Field(alias="operationName")
|
||||
start_time: int = Field(default=0, alias="startTime")
|
||||
references: list[JaegerReference] = []
|
||||
tags: list[JaegerTag] = []
|
||||
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
for tag in self.tags:
|
||||
if tag.key == "span.kind":
|
||||
return str(tag.value)
|
||||
return ""
|
||||
|
||||
|
||||
class JaegerTrace(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", populate_by_name=True)
|
||||
|
||||
trace_id: str = Field(alias="traceID")
|
||||
spans: list[JaegerSpan] = []
|
||||
|
||||
def span_names(self) -> list[str]:
|
||||
return sorted(span.operation_name for span in self.spans)
|
||||
|
||||
|
||||
class JaegerTracesPage(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
data: list[JaegerTrace] = []
|
||||
|
||||
|
||||
class _TracesQuery(BaseModel):
|
||||
service: str
|
||||
tags: str
|
||||
limit: int = 20
|
||||
lookback: str = "1h"
|
||||
|
||||
|
||||
def _settled(trace: JaegerTrace, names: set[str], prefixes: set[str]) -> bool:
|
||||
present = set(trace.span_names())
|
||||
return names.issubset(present) and all(
|
||||
any(name.startswith(prefix) for name in present) for prefix in prefixes
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OtelReader:
|
||||
query_url: str
|
||||
|
||||
def traces_for_call(self, call_id: str) -> list[JaegerTrace]:
|
||||
"""Every trace holding a span tagged with this call id. Jaeger matches
|
||||
spans server-side and returns their full traces; more than one hit for
|
||||
one call IS the split-trace bug, so this never collapses to one."""
|
||||
result = get(
|
||||
URL(f"{self.query_url}/api/traces"),
|
||||
headers=NoBody(),
|
||||
params=_TracesQuery(service=JAEGER_SERVICE, tags=json.dumps({CALL_ID_TAG: call_id})),
|
||||
response_type=JaegerTracesPage,
|
||||
timeout=30.0,
|
||||
)
|
||||
match result:
|
||||
case Success(data=page):
|
||||
return page.data
|
||||
case failure:
|
||||
pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}")
|
||||
|
||||
def poll_traces_for_call(
|
||||
self, *, call_id: str, settled_names: set[str], settled_prefixes: set[str]
|
||||
) -> list[JaegerTrace]:
|
||||
"""Poll until exactly one trace holds the call and it carries every span
|
||||
name in ``settled_names`` plus at least one name per prefix in
|
||||
``settled_prefixes`` (spans flush in batches, the cost write lands after
|
||||
the response), then return the hits. At the deadline the last hits are
|
||||
returned as-is so the caller's assertions report the real final state -
|
||||
on a split trace this never settles and the orphan comes back."""
|
||||
deadline = time.monotonic() + POLL_TIMEOUT
|
||||
hits: list[JaegerTrace] = []
|
||||
while time.monotonic() < deadline:
|
||||
hits = self.traces_for_call(call_id)
|
||||
if len(hits) == 1 and _settled(hits[0], settled_names, settled_prefixes):
|
||||
return hits
|
||||
time.sleep(POLL_INTERVAL)
|
||||
return hits
|
||||
|
||||
|
||||
def build_otel_reader() -> OtelReader:
|
||||
return OtelReader(query_url=OTEL_QUERY_URL)
|
||||
259
tests/e2e/logging/test_otel_trace_e2e.py
Normal file
259
tests/e2e/logging/test_otel_trace_e2e.py
Normal file
|
|
@ -0,0 +1,259 @@
|
|||
"""Live e2e: OTEL trace completeness on the admin-owned destination (LIT-3787).
|
||||
|
||||
Covers logging.otel.success.exports_metric: a successful non-streaming call must
|
||||
land at the OTEL destination as ONE connected trace - a single root SERVER span
|
||||
with the auth phase, db lookups, and cost write under it, and the gen-AI CLIENT
|
||||
span parented into the same tree. The regression this pins: the proxy publishing
|
||||
the global TracerProvider before callbacks init made server spans export through
|
||||
a different provider than the preset's gen-AI spans, so the destination received
|
||||
the gen-AI span alone, dangling (fixed in #30590; verified failing at its parent
|
||||
commit 1bd603d1ac).
|
||||
|
||||
Both halves of the contract are asserted: the recorded state (the proxy reports
|
||||
the OTEL v2 logger active via /health/readiness/details) and the enforced
|
||||
behavior (the complete span tree at the destination, read back through the
|
||||
destination's own query API - never proxy-side "export succeeded" logs).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, unique_marker
|
||||
from e2e_http import NoBody, StreamingResponse, require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from logging_client import LoggingClient
|
||||
from otel_client import JaegerTrace, OtelReader
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = CHEAP_ANTHROPIC_MODEL
|
||||
COST_SPAN = "batch_write_to_db _PROXY_track_cost_callback"
|
||||
DB_SPAN_PREFIX = "postgres "
|
||||
#: The active OTEL v2 logger's name in /health/readiness/details success_callbacks.
|
||||
OTEL_V2_LOGGER_NAME = "OpenTelemetryV2"
|
||||
|
||||
|
||||
class _ReadinessDetails(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
success_callbacks: list[str] = []
|
||||
|
||||
|
||||
def _assert_otel_destination_configured(client: LoggingClient) -> None:
|
||||
"""Recorded state: the proxy reports the OTEL v2 logger among its active
|
||||
callbacks, so a missing/failed destination config fails here, before any
|
||||
traffic-based assertion can time out confusingly."""
|
||||
result = client.gateway.probe("/health/readiness/details", params=NoBody())
|
||||
assert result.status_code == 200, (
|
||||
f"/health/readiness/details must answer 200, got {result.status_code}: {result.body[:300]}"
|
||||
)
|
||||
details = _ReadinessDetails.model_validate_json(result.body)
|
||||
assert OTEL_V2_LOGGER_NAME in details.success_callbacks, (
|
||||
f"the proxy must report the {OTEL_V2_LOGGER_NAME} callback active "
|
||||
f"(LITELLM_OTEL_V2 + arize_phoenix preset in the compose config); got: {details.success_callbacks}"
|
||||
)
|
||||
|
||||
|
||||
def _first_ok(client: LoggingClient, send: Callable[[], StreamingResponse]) -> StreamingResponse:
|
||||
"""First successful call on a fresh key. A fresh key may briefly 401 until
|
||||
the data plane's auth cache picks it up, so retry on 401 to a deadline; a
|
||||
401 is rejected before the LLM call so it exports no gen-AI span and cannot
|
||||
contaminate the trace assertions. Any other failure is behavior under test
|
||||
and fails hard."""
|
||||
deadline = time.monotonic() + client.gateway.poll_timeout
|
||||
while True:
|
||||
outcome = send()
|
||||
if outcome.ok:
|
||||
return outcome
|
||||
if outcome.status_code != 401 or time.monotonic() >= deadline:
|
||||
require_successful_call(outcome)
|
||||
time.sleep(client.gateway.poll_interval)
|
||||
|
||||
|
||||
def _parent_ids(span_id: str, trace: JaegerTrace) -> list[str]:
|
||||
span = next(s for s in trace.spans if s.span_id == span_id)
|
||||
return [ref.span_id for ref in span.references if ref.ref_type == "CHILD_OF"]
|
||||
|
||||
|
||||
def _chain_reaches(span_id: str, root_id: str, trace: JaegerTrace) -> bool:
|
||||
"""Walk parent references (within the trace) from span_id up to root_id."""
|
||||
seen: set[str] = set()
|
||||
in_trace = {s.span_id for s in trace.spans}
|
||||
current = span_id
|
||||
while current not in seen:
|
||||
if current == root_id:
|
||||
return True
|
||||
seen.add(current)
|
||||
parents = [p for p in _parent_ids(current, trace) if p in in_trace]
|
||||
if not parents:
|
||||
return False
|
||||
current = parents[0]
|
||||
return False
|
||||
|
||||
|
||||
def _assert_complete_trace(hits: list[JaegerTrace], *, route: str, genai_span: str) -> None:
|
||||
"""The enforced behavior: the destination holds exactly one trace for the
|
||||
call, rooted at the SERVER span, with auth/db/cost children and the gen-AI
|
||||
span all connected into that one tree - no dangling parent references."""
|
||||
assert hits, (
|
||||
"no trace for this call arrived at the destination within the deadline "
|
||||
"(nothing tagged with its call id was found)"
|
||||
)
|
||||
assert len(hits) == 1, (
|
||||
f"expected exactly ONE trace for the call, got {len(hits)}: "
|
||||
f"{[(t.trace_id, t.span_names()) for t in hits]} - more than one trace for "
|
||||
"one call is the split-trace bug (gen-AI span exported away from its root)"
|
||||
)
|
||||
trace = hits[0]
|
||||
names = trace.span_names()
|
||||
in_trace = {span.span_id for span in trace.spans}
|
||||
|
||||
dangling = [
|
||||
span.operation_name
|
||||
for span in trace.spans
|
||||
if span.references and not any(ref.span_id in in_trace for ref in span.references)
|
||||
]
|
||||
assert not dangling, (
|
||||
f"span(s) {dangling} reference a parent that never reached the destination "
|
||||
f"(orphaned trace); spans present: {names}"
|
||||
)
|
||||
|
||||
roots = [span for span in trace.spans if not span.references]
|
||||
assert len(roots) == 1, f"expected exactly one root span, got {[s.operation_name for s in roots]}; spans: {names}"
|
||||
root = roots[0]
|
||||
assert root.operation_name == f"POST {route}", (
|
||||
f"the root must be the SERVER span 'POST {route}', got {root.operation_name!r}"
|
||||
)
|
||||
assert root.kind == "server", f"the root span must have kind=server, got {root.kind!r}"
|
||||
|
||||
assert f"auth {route}" in names, f"auth phase span 'auth {route}' missing; spans: {names}"
|
||||
assert any(name.startswith(DB_SPAN_PREFIX) for name in names), (
|
||||
f"no db ('{DB_SPAN_PREFIX}*') span in the trace; spans: {names}"
|
||||
)
|
||||
assert COST_SPAN in names, f"cost write span {COST_SPAN!r} missing; spans: {names}"
|
||||
|
||||
genai = next((span for span in trace.spans if span.operation_name == genai_span), None)
|
||||
assert genai is not None, f"gen-AI span {genai_span!r} missing; spans: {names}"
|
||||
assert genai.kind == "client", f"gen-AI span must have kind=client, got {genai.kind!r}"
|
||||
assert _chain_reaches(genai.span_id, root.span_id, trace), (
|
||||
f"gen-AI span {genai_span!r} is in the trace but its parent chain does not "
|
||||
f"reach the root SERVER span; spans: {names}"
|
||||
)
|
||||
|
||||
|
||||
def _settled_names(*, route: str, genai_span: str) -> set[str]:
|
||||
return {f"POST {route}", f"auth {route}", COST_SPAN, genai_span}
|
||||
|
||||
|
||||
class TestOtelTraceCompleteness:
|
||||
@pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["chat_completions"])
|
||||
def test_chat_completions_exports_complete_trace(
|
||||
self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager
|
||||
) -> None:
|
||||
"""This test verifies that a successful non-streaming
|
||||
/chat/completions request produces one complete OTEL trace.
|
||||
|
||||
The trace should have a single server root span for the incoming request, with
|
||||
the authentication, database, and cost-recording work beneath it. The span for
|
||||
the actual model call must also belong to that same trace, rather than being
|
||||
exported separately with a missing parent.
|
||||
|
||||
This matters because a split trace is easy to miss: all of the spans may still
|
||||
arrive, but the model call appears without the surrounding request context.
|
||||
That makes it difficult to understand where time was spent, connect the model
|
||||
cost to the original request, or investigate a slow or failed call.
|
||||
|
||||
/chat/completions is the main OpenAI-compatible route used by most customers,
|
||||
so it is important that trace parenting works correctly on this path.
|
||||
"""
|
||||
route = "/chat/completions"
|
||||
_assert_otel_destination_configured(client)
|
||||
|
||||
key = client.key_with_alias(f"otel-trace-chat-{unique_marker()}", models=[MODEL])
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
marker = unique_marker()
|
||||
outcome = _first_ok(
|
||||
client, lambda: client.chat_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16)
|
||||
)
|
||||
assert outcome.call_id is not None, "success response must carry x-litellm-call-id"
|
||||
|
||||
hits = otel_reader.poll_traces_for_call(
|
||||
call_id=outcome.call_id,
|
||||
settled_names=_settled_names(route=route, genai_span=f"chat {MODEL}"),
|
||||
settled_prefixes={DB_SPAN_PREFIX},
|
||||
)
|
||||
_assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}")
|
||||
|
||||
@pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["messages"])
|
||||
def test_messages_exports_complete_trace(
|
||||
self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager
|
||||
) -> None:
|
||||
"""This test verifies that one successful non-streaming /v1/messages request
|
||||
produces exactly one complete OTEL trace.
|
||||
|
||||
The trace must have a single root span named "POST /v1/messages". The
|
||||
authentication, database, cost-writing, and model-call spans must all belong to
|
||||
the same trace and have valid parent relationships leading back to that root.
|
||||
|
||||
The model-call span is expected to be named "chat <model>". The test fails if
|
||||
the request is split across multiple traces, if any span references a missing
|
||||
parent, or if the model-call span cannot be connected back to the root."""
|
||||
route = "/v1/messages"
|
||||
_assert_otel_destination_configured(client)
|
||||
|
||||
key = client.key_with_alias(f"otel-trace-messages-{unique_marker()}", models=[MODEL])
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
marker = unique_marker()
|
||||
outcome = _first_ok(
|
||||
client, lambda: client.messages_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16)
|
||||
)
|
||||
assert outcome.call_id is not None, "success response must carry x-litellm-call-id"
|
||||
|
||||
hits = otel_reader.poll_traces_for_call(
|
||||
call_id=outcome.call_id,
|
||||
settled_names=_settled_names(route=route, genai_span=f"chat {MODEL}"),
|
||||
settled_prefixes={DB_SPAN_PREFIX},
|
||||
)
|
||||
_assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}")
|
||||
|
||||
@pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["responses"])
|
||||
def test_responses_exports_complete_trace(
|
||||
self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager
|
||||
) -> None:
|
||||
"""This test verifies that one successful non-streaming /v1/responses request
|
||||
produces exactly one complete OTEL trace.
|
||||
|
||||
The trace must have a single root span named "POST /v1/responses". The
|
||||
authentication, database, cost-writing, and model-call spans must all belong to
|
||||
the same trace and have valid parent relationships leading back to that root.
|
||||
|
||||
The model-call span is expected to be named "chat <model>". The test fails if
|
||||
the request is split across multiple traces, if any span references a missing
|
||||
parent, or if the model-call span cannot be connected back to the root."""
|
||||
route = "/v1/responses"
|
||||
_assert_otel_destination_configured(client)
|
||||
|
||||
key = client.key_with_alias(f"otel-trace-responses-{unique_marker()}", models=[CHEAP_OPENAI_MODEL])
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
marker = unique_marker()
|
||||
outcome = _first_ok(
|
||||
client,
|
||||
lambda: client.responses_raw(key, CHEAP_OPENAI_MODEL, f"reply with one word {marker}"),
|
||||
)
|
||||
assert outcome.call_id is not None, "success response must carry x-litellm-call-id"
|
||||
|
||||
genai_span = f"chat {CHEAP_OPENAI_MODEL}"
|
||||
hits = otel_reader.poll_traces_for_call(
|
||||
call_id=outcome.call_id,
|
||||
settled_names=_settled_names(route=route, genai_span=genai_span),
|
||||
settled_prefixes={DB_SPAN_PREFIX},
|
||||
)
|
||||
_assert_complete_trace(hits, route=route, genai_span=genai_span)
|
||||
|
|
@ -111,7 +111,7 @@ class TestKeyRoutes:
|
|||
key = _generate_key(
|
||||
client,
|
||||
resources,
|
||||
KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=424242),
|
||||
KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=424242, rpm_limit=424243),
|
||||
)
|
||||
|
||||
info = client.gateway.key_info(key)
|
||||
|
|
@ -122,6 +122,9 @@ class TestKeyRoutes:
|
|||
assert info.tpm_limit == 424242, (
|
||||
f"/key/info reports tpm_limit {info.tpm_limit}, configured 424242"
|
||||
)
|
||||
assert info.rpm_limit == 424243, (
|
||||
f"/key/info reports rpm_limit {info.rpm_limit}, configured 424243"
|
||||
)
|
||||
|
||||
_poll_chat_ok(client, key, "gemini-2.5-flash")
|
||||
_assert_model_denied(
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from __future__ import annotations
|
|||
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, RootModel
|
||||
from pydantic import BaseModel, ConfigDict, RootModel, model_validator
|
||||
|
||||
# ---------- keys ----------
|
||||
|
||||
|
|
@ -82,6 +82,7 @@ class KeyInfo(BaseModel):
|
|||
key_alias: str | None = None
|
||||
models: list[str] = []
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
team_id: str | None = None
|
||||
spend: float | None = None
|
||||
max_budget: float | None = None
|
||||
|
|
@ -255,6 +256,16 @@ class SpendLogsParams(BaseModel):
|
|||
request_id: str | None = None
|
||||
api_key: str | None = None
|
||||
|
||||
@model_validator(mode="after")
|
||||
def require_filter(self) -> SpendLogsParams:
|
||||
if self.request_id is None and self.api_key is None:
|
||||
raise ValueError(
|
||||
"unfiltered /spend/logs returns the entire spend table and OOMs the "
|
||||
"runner on long-lived environments; filter by request_id or api_key, "
|
||||
"or use Gateway.spend_logs_window for a bounded /spend/logs/v2 read"
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
class SpendLogsPageParams(BaseModel):
|
||||
"""Query for /spend/logs/v2, which requires an explicit date window and
|
||||
|
|
|
|||
15
tests/e2e/quota_management/ratelimit/conftest.py
Normal file
15
tests/e2e/quota_management/ratelimit/conftest.py
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
"""Quota-management suite's `client` fixture.
|
||||
|
||||
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
|
||||
live in the parent tests/e2e/conftest.py. QuotaClient holds the shared Gateway,
|
||||
so the `resources` fixture cleans up keys through it.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from quota_client import QuotaClient, build_client
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def client() -> QuotaClient:
|
||||
return build_client()
|
||||
32
tests/e2e/quota_management/ratelimit/quota_client.py
Normal file
32
tests/e2e/quota_management/ratelimit/quota_client.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
"""Client for the quota-management suite: the shared Gateway plus raw chat
|
||||
calls judged by HTTP status, body, and headers (a rate-limit block is a 429
|
||||
whose body and retry-after header carry the contract, not a typed success
|
||||
model)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from e2e_gateway import Gateway, build_gateway
|
||||
from e2e_http import StreamingResponse
|
||||
from models import ChatBody, ChatMessage
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class QuotaClient:
|
||||
gateway: Gateway
|
||||
|
||||
def chat(self, key: str, model: str, content: str, *, max_tokens: int = 16) -> StreamingResponse:
|
||||
return self.gateway.transport.send(
|
||||
"/chat/completions",
|
||||
headers=self.gateway.transport.bearer(key),
|
||||
json=ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=content)],
|
||||
max_tokens=max_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def build_client() -> QuotaClient:
|
||||
return QuotaClient(gateway=build_gateway())
|
||||
254
tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py
Normal file
254
tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py
Normal file
|
|
@ -0,0 +1,254 @@
|
|||
"""Live e2e: key-level rpm/tpm rate limits on the gateway.
|
||||
|
||||
Covers quota_management.ratelimit.*: a key generated with rpm_limit/tpm_limit gets a
|
||||
429 once the limit is crossed inside one window (blocks_over_limit), serves
|
||||
again once the window rolls and no sooner (resets_after_window), and successful responses
|
||||
report x-ratelimit-* limit/remaining headers so clients can pace
|
||||
(headers_report_remaining). Each test asserts both halves of the contract: the
|
||||
recorded state (/key/info echoes the configured limit) and the enforced
|
||||
behavior (the 429, the recovery, or the headers on live traffic).
|
||||
|
||||
The v3 limiter counts a request against the rpm budget at the pre-call hook,
|
||||
before model routing, so every call that clears auth consumes budget whether or
|
||||
not it ultimately succeeds. The tpm budget is reserved pre-call from an estimate
|
||||
(message chars // 4 + max_tokens) and reconciled to the body's actual
|
||||
usage.total_tokens after the call, so a block may legitimately fire before the
|
||||
actual spend crosses the limit; the tpm test asserts the exact contract on both
|
||||
sides (a 429 only once the blocked call's reservation exceeds the remaining
|
||||
budget, and no later than the first call after actual spend reaches the limit).
|
||||
All calls of one test must land inside a single window
|
||||
(LITELLM_RATE_LIMIT_WINDOW_SIZE, 60s default), which real chat latency
|
||||
comfortably allows.
|
||||
|
||||
The window opens at the pre-call hook of the first counted call, which happens
|
||||
after that call is sent, so the send timestamp of the winning first call is a
|
||||
lower bound on the window start. The reset test uses it to reject an early
|
||||
reset: recovery must not arrive before the full window has elapsed from that
|
||||
send, less a small tolerance for the limiter's integer-second window
|
||||
arithmetic.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker
|
||||
from e2e_http import StreamingResponse, require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import KeyGenerateBody
|
||||
from quota_client import QuotaClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
MODEL = CHEAP_ANTHROPIC_MODEL
|
||||
TPM_LIMIT = 60
|
||||
CHAT_MAX_TOKENS = 16
|
||||
RESERVATION_CHARS_PER_TOKEN = 4
|
||||
WINDOW_SECONDS = 60
|
||||
RESET_TOLERANCE_SECONDS = 5
|
||||
LAST_CALL_LATENCY_MARGIN_SECONDS = 10
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _FirstOk:
|
||||
sent_at: float
|
||||
response: StreamingResponse
|
||||
|
||||
|
||||
class _ChatUsage(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
total_tokens: int
|
||||
|
||||
|
||||
class _ChatBodyWithUsage(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
usage: _ChatUsage
|
||||
|
||||
|
||||
def _total_tokens(outcome: StreamingResponse) -> int:
|
||||
try:
|
||||
return _ChatBodyWithUsage.model_validate_json(outcome.body).usage.total_tokens
|
||||
except ValidationError:
|
||||
pytest.fail(f"successful chat body must report usage.total_tokens, got: {outcome.body[:300]}")
|
||||
|
||||
|
||||
def _reserved_tokens(content: str) -> int:
|
||||
return max(1, len(content) // RESERVATION_CHARS_PER_TOKEN) + CHAT_MAX_TOKENS
|
||||
|
||||
|
||||
def _remaining_from_429(body: str) -> int:
|
||||
found = re.search(r"Remaining: (\d+)", body)
|
||||
if found is None:
|
||||
pytest.fail(f"429 body must report the remaining budget, got: {body[:300]}")
|
||||
return int(found.group(1))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BlockedByReservation:
|
||||
outcome: StreamingResponse
|
||||
content: str
|
||||
spent: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CrossedLimit:
|
||||
spent: int
|
||||
|
||||
|
||||
def _spend_until_blocked_or_crossed(
|
||||
client: QuotaClient, key: str, first: _FirstOk
|
||||
) -> _BlockedByReservation | _CrossedLimit:
|
||||
"""Drive chat traffic, summing each body's actual usage.total_tokens, until
|
||||
the limiter blocks (which the reservation may do before the actual spend
|
||||
crosses the limit) or the actual spend reaches the limit."""
|
||||
window_deadline = first.sent_at + WINDOW_SECONDS - LAST_CALL_LATENCY_MARGIN_SECONDS
|
||||
spent = _total_tokens(first.response)
|
||||
while spent < TPM_LIMIT:
|
||||
assert time.monotonic() < window_deadline, (
|
||||
f"spent only {spent} of {TPM_LIMIT} tokens before the {WINDOW_SECONDS}s window could roll; "
|
||||
"the exact-crossing assertion needs every call inside one window"
|
||||
)
|
||||
content = f"reply with one word {unique_marker()}"
|
||||
outcome = client.chat(key, MODEL, content, max_tokens=CHAT_MAX_TOKENS)
|
||||
if outcome.status_code == 429:
|
||||
return _BlockedByReservation(outcome=outcome, content=content, spent=spent)
|
||||
require_successful_call(outcome)
|
||||
spent += _total_tokens(outcome)
|
||||
return _CrossedLimit(spent=spent)
|
||||
|
||||
|
||||
def _limited_key(
|
||||
client: QuotaClient,
|
||||
resources: ResourceManager,
|
||||
*,
|
||||
rpm_limit: int | None = None,
|
||||
tpm_limit: int | None = None,
|
||||
) -> str:
|
||||
key = client.gateway.generate_key(KeyGenerateBody(models=[MODEL], rpm_limit=rpm_limit, tpm_limit=tpm_limit))
|
||||
resources.defer(lambda: client.gateway.delete_key(key))
|
||||
return key
|
||||
|
||||
|
||||
def _chat(client: QuotaClient, key: str) -> StreamingResponse:
|
||||
return client.chat(key, MODEL, f"reply with one word {unique_marker()}")
|
||||
|
||||
|
||||
def _first_ok(client: QuotaClient, key: str) -> _FirstOk:
|
||||
"""First successful call on a fresh key, which opens the rate-limit window;
|
||||
`sent_at` is captured just before the winning send, so the window opened no
|
||||
earlier than it. A fresh key may briefly 401 until the data plane's auth
|
||||
cache picks it up, so retry on 401 to a deadline; a 401 never reaches the
|
||||
rate limiter, so only the successful call consumes budget. Any other failure
|
||||
is behavior under test and fails hard."""
|
||||
deadline = time.monotonic() + client.gateway.poll_timeout
|
||||
while True:
|
||||
sent_at = time.monotonic()
|
||||
outcome = _chat(client, key)
|
||||
if outcome.ok:
|
||||
return _FirstOk(sent_at=sent_at, response=outcome)
|
||||
if outcome.status_code != 401 or time.monotonic() >= deadline:
|
||||
require_successful_call(outcome)
|
||||
time.sleep(client.gateway.poll_interval)
|
||||
|
||||
|
||||
def _assert_rate_limited(outcome: StreamingResponse, limit_type: str) -> None:
|
||||
assert outcome.status_code == 429, (
|
||||
f"expected a 429 {limit_type} rate-limit block, got {outcome.status_code}: {outcome.body[:300]}"
|
||||
)
|
||||
assert "Rate limit exceeded for api_key" in outcome.body, (
|
||||
f"429 body must name the api_key scope, got: {outcome.body[:300]}"
|
||||
)
|
||||
assert f"Limit type: {limit_type}" in outcome.body, (
|
||||
f"429 body must carry 'Limit type: {limit_type}', got: {outcome.body[:300]}"
|
||||
)
|
||||
retry_after = outcome.headers.get("retry-after")
|
||||
assert retry_after is not None and retry_after.isdigit() and int(retry_after) > 0, (
|
||||
f"429 must carry a positive integer retry-after header, got {retry_after!r}"
|
||||
)
|
||||
|
||||
|
||||
class TestKeyRateLimits:
|
||||
@pytest.mark.covers("quota_management.ratelimit.rpm.blocks_over_limit")
|
||||
def test_rpm_limit_blocks_over_limit(self, client: QuotaClient, resources: ResourceManager) -> None:
|
||||
key = _limited_key(client, resources, rpm_limit=3)
|
||||
info = client.gateway.key_info(key)
|
||||
assert info.rpm_limit == 3, f"/key/info reports rpm_limit {info.rpm_limit}, configured 3"
|
||||
|
||||
_ = _first_ok(client, key)
|
||||
for _ in range(2):
|
||||
require_successful_call(_chat(client, key))
|
||||
|
||||
_assert_rate_limited(_chat(client, key), "requests")
|
||||
|
||||
@pytest.mark.covers("quota_management.ratelimit.tpm.blocks_over_limit")
|
||||
def test_tpm_limit_blocks_over_limit(self, client: QuotaClient, resources: ResourceManager) -> None:
|
||||
key = _limited_key(client, resources, tpm_limit=TPM_LIMIT)
|
||||
info = client.gateway.key_info(key)
|
||||
assert info.tpm_limit == TPM_LIMIT, f"/key/info reports tpm_limit {info.tpm_limit}, configured {TPM_LIMIT}"
|
||||
|
||||
first = _first_ok(client, key)
|
||||
match _spend_until_blocked_or_crossed(client, key, first):
|
||||
case _BlockedByReservation(outcome=outcome, content=content, spent=spent):
|
||||
_assert_rate_limited(outcome, "tokens")
|
||||
remaining = _remaining_from_429(outcome.body)
|
||||
reserved = _reserved_tokens(content)
|
||||
assert reserved > remaining, (
|
||||
f"blocked while the call still fit: {remaining} of {TPM_LIMIT} tokens remained but the call "
|
||||
f"reserved only {reserved} ({spent} actual tokens spent so far)"
|
||||
)
|
||||
case _CrossedLimit():
|
||||
_assert_rate_limited(_chat(client, key), "tokens")
|
||||
|
||||
@pytest.mark.covers("quota_management.ratelimit.rpm.resets_after_window")
|
||||
def test_rpm_limit_resets_after_window(self, client: QuotaClient, resources: ResourceManager) -> None:
|
||||
key = _limited_key(client, resources, rpm_limit=1)
|
||||
|
||||
first = _first_ok(client, key)
|
||||
_assert_rate_limited(_chat(client, key), "requests")
|
||||
|
||||
deadline = time.monotonic() + client.gateway.poll_timeout
|
||||
while time.monotonic() < deadline:
|
||||
attempt_sent_at = time.monotonic()
|
||||
outcome = _chat(client, key)
|
||||
if outcome.ok:
|
||||
window_age = attempt_sent_at - first.sent_at
|
||||
assert window_age >= WINDOW_SECONDS - RESET_TOLERANCE_SECONDS, (
|
||||
f"the key recovered {window_age:.1f}s after the window opened, before the "
|
||||
f"{WINDOW_SECONDS}s window (less {RESET_TOLERANCE_SECONDS}s tolerance) elapsed; "
|
||||
"the limiter reset early instead of after the window"
|
||||
)
|
||||
return
|
||||
assert outcome.status_code == 429, (
|
||||
f"while the window drains only 429s are acceptable, got {outcome.status_code}: {outcome.body[:300]}"
|
||||
)
|
||||
time.sleep(client.gateway.poll_interval)
|
||||
pytest.fail("a blocked key never recovered after the rate-limit window elapsed")
|
||||
|
||||
@pytest.mark.covers("quota_management.ratelimit.rpm.headers_report_remaining")
|
||||
def test_headers_report_limit_and_remaining(self, client: QuotaClient, resources: ResourceManager) -> None:
|
||||
key = _limited_key(client, resources, rpm_limit=5, tpm_limit=100000)
|
||||
|
||||
first = _first_ok(client, key).response
|
||||
assert first.headers.get("x-ratelimit-api_key-limit-requests") == "5", (
|
||||
f"success response must report the key's request limit, headers: "
|
||||
f"{ {k: v for k, v in first.headers.items() if 'ratelimit' in k} }"
|
||||
)
|
||||
assert first.headers.get("x-ratelimit-api_key-remaining-requests") == str(5 - 1), (
|
||||
f"first call against rpm_limit=5 must leave {5 - 1} remaining, got "
|
||||
f"{first.headers.get('x-ratelimit-api_key-remaining-requests')!r}"
|
||||
)
|
||||
assert first.headers.get("x-ratelimit-api_key-limit-tokens") == "100000", (
|
||||
f"success response must report the key's token limit, got "
|
||||
f"{first.headers.get('x-ratelimit-api_key-limit-tokens')!r}"
|
||||
)
|
||||
remaining_tokens = first.headers.get("x-ratelimit-api_key-remaining-tokens")
|
||||
assert remaining_tokens is not None and remaining_tokens.isdigit() and int(remaining_tokens) < 100000, (
|
||||
f"one call must leave remaining tokens reported and below the limit, got {remaining_tokens!r}"
|
||||
)
|
||||
|
|
@ -1,17 +1,23 @@
|
|||
"""Unit coverage for the Gateway model-management surface (create_model /
|
||||
delete_model).
|
||||
delete_model) and the bounded spend read-back (spend_logs_window).
|
||||
|
||||
The batches conftest and several llm_translation tests register deployments at
|
||||
runtime through gateway.create_model; when that method went missing, every batch
|
||||
test errored at fixture setup (AttributeError) before a single request reached
|
||||
the proxy. This pins the surface with a typed fake Transport so a rename or
|
||||
signature drift fails here instead of in a live stage run.
|
||||
|
||||
spend_logs_window exists because the unpaginated /spend/logs whole-table read
|
||||
grew past the e2e runner's memory limit on stage and OOMKilled every run; these
|
||||
tests pin its /spend/logs/v2 pagination and that SpendLogsParams can no longer
|
||||
express the unfiltered read.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from batches.batch_client import BatchClient
|
||||
from e2e_gateway import Gateway
|
||||
|
|
@ -30,6 +36,9 @@ from models import (
|
|||
ModelNewBody,
|
||||
ModelNewResponse,
|
||||
ModelsListResponse,
|
||||
SpendLogsPage,
|
||||
SpendLogsPageParams,
|
||||
SpendLogsParams,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -46,6 +55,8 @@ class _RecordingTransport:
|
|||
servable_after_gets: int = 0
|
||||
models_error: UnknownApiError | None = None
|
||||
model_gets: int = 0
|
||||
spend_total: int = 0
|
||||
spend_gets: list[SpendLogsPageParams] = field(default_factory=list)
|
||||
_created: list[str] = field(default_factory=list)
|
||||
|
||||
def post[R: BaseModel](
|
||||
|
|
@ -91,6 +102,22 @@ class _RecordingTransport:
|
|||
return Success(
|
||||
data=response_type.model_validate({"data": [{"id": name} for name in visible]})
|
||||
)
|
||||
if path == "/spend/logs/v2" and response_type is SpendLogsPage:
|
||||
assert isinstance(params, SpendLogsPageParams)
|
||||
self.spend_gets.append(params)
|
||||
offset = (params.page - 1) * params.page_size
|
||||
count = min(params.page_size, max(self.spend_total - offset, 0))
|
||||
return Success(
|
||||
data=response_type.model_validate(
|
||||
{
|
||||
"data": [{"request_id": f"req-{offset + i}"} for i in range(count)],
|
||||
"total": self.spend_total,
|
||||
"page": params.page,
|
||||
"page_size": params.page_size,
|
||||
"total_pages": (self.spend_total + params.page_size - 1) // params.page_size,
|
||||
}
|
||||
)
|
||||
)
|
||||
raise AssertionError(f"unexpected get: {path}")
|
||||
|
||||
def delete[R: BaseModel](
|
||||
|
|
@ -202,3 +229,45 @@ def test_gateway_delete_model_posts_the_model_id() -> None:
|
|||
assert path == "/model/delete"
|
||||
assert isinstance(body, ModelDeleteBody)
|
||||
assert body.id == "registered-id"
|
||||
|
||||
|
||||
WINDOW_START = datetime(2026, 7, 14, 12, 0, 0, tzinfo=timezone.utc)
|
||||
WINDOW_END = datetime(2026, 7, 14, 14, 0, 0, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
def test_gateway_spend_logs_window_pages_through_every_row_in_the_window() -> None:
|
||||
transport = _RecordingTransport(spend_total=250)
|
||||
gateway = Gateway(transport=transport)
|
||||
|
||||
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
|
||||
|
||||
assert len(rows) == 250
|
||||
assert len({row.request_id for row in rows}) == 250
|
||||
assert [params.page for params in transport.spend_gets] == [1, 2, 3]
|
||||
assert all(params.start_date == "2026-07-14 12:00:00" for params in transport.spend_gets)
|
||||
assert all(params.end_date == "2026-07-14 14:00:00" for params in transport.spend_gets)
|
||||
|
||||
|
||||
def test_gateway_spend_logs_window_stops_at_an_exact_page_boundary() -> None:
|
||||
transport = _RecordingTransport(spend_total=200)
|
||||
gateway = Gateway(transport=transport)
|
||||
|
||||
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
|
||||
|
||||
assert len(rows) == 200
|
||||
assert [params.page for params in transport.spend_gets] == [1, 2]
|
||||
|
||||
|
||||
def test_gateway_spend_logs_window_returns_empty_for_an_empty_window() -> None:
|
||||
transport = _RecordingTransport(spend_total=0)
|
||||
gateway = Gateway(transport=transport)
|
||||
|
||||
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
|
||||
|
||||
assert rows == []
|
||||
assert [params.page for params in transport.spend_gets] == [1]
|
||||
|
||||
|
||||
def test_spend_logs_params_rejects_the_unfiltered_whole_table_read() -> None:
|
||||
with pytest.raises(ValidationError, match="spend_logs_window"):
|
||||
SpendLogsParams()
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import io
|
|||
import os
|
||||
import sys
|
||||
|
||||
from typing import Optional
|
||||
from typing import Optional, Union
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
|
|
@ -12,6 +12,7 @@ import json
|
|||
import logging
|
||||
import time
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from datetime import datetime
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -20,17 +21,24 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.responses.main import mock_responses_api_response
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.types.utils import (
|
||||
ModelResponse,
|
||||
ResponsesAPIResponse,
|
||||
StandardLoggingPayload,
|
||||
TextCompletionResponse,
|
||||
)
|
||||
|
||||
|
||||
class TestCustomLogger(CustomLogger):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.logged_standard_logging_payload: Optional[StandardLoggingPayload] = None
|
||||
self.response_obj: Optional[Union[ModelResponse, TextCompletionResponse, ResponsesAPIResponse]] = None
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
standard_logging_payload = kwargs.get("standard_logging_object", None)
|
||||
self.logged_standard_logging_payload = standard_logging_payload
|
||||
self.response_obj = response_obj
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -108,6 +116,78 @@ async def test_dynamic_turn_off_message_logging_overrides_global_off(dynamic_tur
|
|||
assert standard_logging_payload["messages"][0]["content"] == expected_message_content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redaction_with_custom_logger_streaming():
|
||||
"""Test redaction of responses for custom logger callbacks"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
class LoggingWithoutSyncSuccessHandler(Logging):
|
||||
def success_handler(self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs):
|
||||
pass
|
||||
|
||||
litellm.turn_off_message_logging = True
|
||||
test_custom_logger = TestCustomLogger()
|
||||
|
||||
try:
|
||||
litellm_logging_obj = LoggingWithoutSyncSuccessHandler(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=True,
|
||||
call_type="acompletion",
|
||||
litellm_call_id="1234",
|
||||
start_time=datetime.now(),
|
||||
function_id="1234",
|
||||
dynamic_async_success_callbacks=[test_custom_logger],
|
||||
)
|
||||
|
||||
response = await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="hello",
|
||||
stream=True,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
# Consume the stream to trigger logging
|
||||
chunks = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
async_complete_streaming_response = test_custom_logger.response_obj
|
||||
assert async_complete_streaming_response is not None
|
||||
assert async_complete_streaming_response.choices[0].message.content == "redacted-by-litellm"
|
||||
finally:
|
||||
litellm.turn_off_message_logging = False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_redaction_scoped_to_opted_out_logger():
|
||||
"""One logger opting out of message logging must not blank the response for other loggers"""
|
||||
litellm.turn_off_message_logging = False
|
||||
opted_out_logger = TestCustomLogger(message_logging=False)
|
||||
compliant_logger = TestCustomLogger()
|
||||
litellm.callbacks = [opted_out_logger, compliant_logger]
|
||||
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
mock_response="hello",
|
||||
stream=True,
|
||||
)
|
||||
async for _ in response:
|
||||
pass
|
||||
|
||||
await asyncio.sleep(1)
|
||||
assert opted_out_logger.response_obj is not None
|
||||
assert opted_out_logger.response_obj.choices[0].message.content == "redacted-by-litellm"
|
||||
assert compliant_logger.response_obj is not None
|
||||
assert compliant_logger.response_obj.choices[0].message.content == "hello"
|
||||
finally:
|
||||
litellm.callbacks = []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redaction_responses_api():
|
||||
"""Test redaction with ResponsesAPIResponse format"""
|
||||
|
|
|
|||
|
|
@ -219,8 +219,8 @@ async def test_aaauser_personal_budgets(key_ownership):
|
|||
"""
|
||||
Set a personal budget on a user
|
||||
|
||||
- have it only apply when key belongs to user -> raises BudgetExceededError
|
||||
- if key belongs to team, have key respect team budget -> allows call to go through
|
||||
User budget is enforced regardless of key ownership (personal or team).
|
||||
Both cases should raise BudgetExceededError when the user is over budget.
|
||||
"""
|
||||
import asyncio
|
||||
import time
|
||||
|
|
@ -229,7 +229,12 @@ async def test_aaauser_personal_budgets(key_ownership):
|
|||
from starlette.datastructures import URL
|
||||
import litellm
|
||||
|
||||
from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTable,
|
||||
ProxyErrorTypes,
|
||||
ProxyException,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.proxy_server import hash_token, user_api_key_cache
|
||||
|
||||
|
|
@ -273,14 +278,9 @@ async def test_aaauser_personal_budgets(key_ownership):
|
|||
== valid_token
|
||||
)
|
||||
|
||||
try:
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(request=request, api_key="Bearer " + user_key)
|
||||
|
||||
if key_ownership == "user_key":
|
||||
pytest.fail("Expected this call to fail. User is over limit.")
|
||||
except Exception:
|
||||
if key_ownership == "team_key":
|
||||
pytest.fail("Expected this call to work. Key is below team budget.")
|
||||
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1614,3 +1614,105 @@ class TestGuardrailInterventionClassification:
|
|||
|
||||
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert slg["guardrail_status"] == "guardrail_intervened"
|
||||
|
||||
|
||||
class _ApplyStyleGuardrail(CustomGuardrail):
|
||||
"""Overrides only apply_guardrail, like openai_moderation; async_pre_call_hook stays the CustomLogger no-op."""
|
||||
|
||||
def __init__(self, block: bool):
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
super().__init__(
|
||||
guardrail_name="apply-style-guardrail",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=False,
|
||||
)
|
||||
self.block = block
|
||||
self.apply_called = False
|
||||
self.seen_texts = None
|
||||
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
from fastapi import HTTPException
|
||||
|
||||
self.apply_called = True
|
||||
self.seen_texts = inputs.get("texts")
|
||||
if self.block:
|
||||
raise HTTPException(status_code=400, detail={"error": "Violated moderation policy"})
|
||||
return inputs
|
||||
|
||||
|
||||
class TestApplyGuardrailStyleDeploymentDispatch:
|
||||
"""LIT-4217 regression: model-level guardrails that implement only the
|
||||
unified apply_guardrail interface must execute in
|
||||
async_pre_call_deployment_hook instead of silently hitting the
|
||||
async_pre_call_hook no-op."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("call_type", [CallTypes.completion, CallTypes.acompletion])
|
||||
async def test_blocks_when_requested_via_model_level_guardrails(self, call_type):
|
||||
from fastapi import HTTPException
|
||||
|
||||
guardrail = _ApplyStyleGuardrail(block=True)
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "flagged content"}],
|
||||
"guardrails": ["apply-style-guardrail"],
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
await guardrail.async_pre_call_deployment_hook(kwargs, call_type)
|
||||
|
||||
assert guardrail.apply_called is True
|
||||
assert guardrail.seen_texts == ["flagged content"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_path_runs_guardrail_and_strips_dispatch_key(self):
|
||||
guardrail = _ApplyStyleGuardrail(block=False)
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"guardrails": ["apply-style-guardrail"],
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
result = await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
|
||||
|
||||
assert guardrail.apply_called is True
|
||||
assert result is not None
|
||||
assert "guardrail_to_apply" not in result
|
||||
assert result["messages"] == [{"role": "user", "content": "hello"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_when_not_requested(self):
|
||||
guardrail = _ApplyStyleGuardrail(block=True)
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"guardrails": ["some-other-guardrail"],
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
result = await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
|
||||
|
||||
assert guardrail.apply_called is False
|
||||
assert result is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fails_closed_when_proxy_extras_missing(self):
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
guardrail = _ApplyStyleGuardrail(block=True)
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "flagged content"}],
|
||||
"guardrails": ["apply-style-guardrail"],
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
with patch.dict(sys.modules, {"litellm.proxy.utils": None}):
|
||||
with pytest.raises(ImportError, match="litellm\\[proxy\\]"):
|
||||
await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
|
||||
|
||||
assert guardrail.apply_called is False
|
||||
|
|
|
|||
|
|
@ -0,0 +1,180 @@
|
|||
"""
|
||||
Unit tests for the video-seconds and images-generated Prometheus counters (LIT-4254).
|
||||
|
||||
Video providers report ``duration_seconds`` inside the usage object that lands
|
||||
on ``standard_logging_payload["metadata"]["usage_object"]``; image generation
|
||||
calls report ``output_image_count`` there. Both counters are sparse: only
|
||||
incremented when the value is present and > 0.
|
||||
"""
|
||||
|
||||
from typing import get_args
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.types.integrations.prometheus import (
|
||||
DEFINED_PROMETHEUS_METRICS,
|
||||
PrometheusMetricLabels,
|
||||
UserAPIKeyLabelValues,
|
||||
)
|
||||
|
||||
MEDIA_GENERATION_METRICS = [
|
||||
"litellm_video_duration_seconds_metric",
|
||||
"litellm_images_generated_metric",
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_enum_values():
|
||||
return UserAPIKeyLabelValues(
|
||||
end_user="test-end-user",
|
||||
hashed_api_key="test-key-hash",
|
||||
api_key_alias="test-key-alias",
|
||||
team="test-team",
|
||||
team_alias="test-team-alias",
|
||||
user="test-user",
|
||||
model="sora-2",
|
||||
)
|
||||
|
||||
|
||||
def _make_mock_logger():
|
||||
logger = MagicMock()
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
setattr(logger, name, MagicMock())
|
||||
logger.get_labels_for_metric = MagicMock(
|
||||
return_value=[
|
||||
"model",
|
||||
"hashed_api_key",
|
||||
"api_key_alias",
|
||||
"team",
|
||||
"team_alias",
|
||||
"end_user",
|
||||
"user",
|
||||
]
|
||||
)
|
||||
return logger
|
||||
|
||||
|
||||
class TestMediaGenerationMetricsRegistration:
|
||||
def test_metrics_in_defined_prometheus_metrics(self):
|
||||
defined = get_args(DEFINED_PROMETHEUS_METRICS)
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
assert name in defined, f"{name} missing from DEFINED_PROMETHEUS_METRICS"
|
||||
|
||||
def test_metric_labels_defined(self):
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
assert hasattr(PrometheusMetricLabels, name), f"{name} missing from PrometheusMetricLabels"
|
||||
|
||||
def test_metrics_share_output_token_label_set(self):
|
||||
assert (
|
||||
PrometheusMetricLabels.litellm_video_duration_seconds_metric
|
||||
== PrometheusMetricLabels.litellm_output_tokens_metric
|
||||
)
|
||||
assert (
|
||||
PrometheusMetricLabels.litellm_images_generated_metric
|
||||
== PrometheusMetricLabels.litellm_output_tokens_metric
|
||||
)
|
||||
|
||||
def test_runtime_label_set_matches_output_tokens_metric(self):
|
||||
"""Full parity with litellm_output_tokens_metric, including the org labels
|
||||
appended via _org_label_metrics, so existing token dashboards can be cloned."""
|
||||
expected = PrometheusMetricLabels.get_labels("litellm_output_tokens_metric")
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
assert PrometheusMetricLabels.get_labels(name) == expected
|
||||
|
||||
|
||||
class TestIncrementMediaGenerationMetrics:
|
||||
def test_video_duration_incremented(self, sample_enum_values):
|
||||
logger = _make_mock_logger()
|
||||
payload = {"metadata": {"usage_object": {"duration_seconds": 8.0}}}
|
||||
|
||||
PrometheusLogger._increment_media_generation_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
logger.litellm_video_duration_seconds_metric.labels().inc.assert_called_once_with(8.0)
|
||||
logger.litellm_images_generated_metric.labels.assert_not_called()
|
||||
|
||||
def test_image_count_incremented(self, sample_enum_values):
|
||||
logger = _make_mock_logger()
|
||||
payload = {
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"prompt_tokens": 18,
|
||||
"completion_tokens": 391,
|
||||
"total_tokens": 409,
|
||||
"output_image_count": 2,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
PrometheusLogger._increment_media_generation_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
logger.litellm_images_generated_metric.labels().inc.assert_called_once_with(2.0)
|
||||
logger.litellm_video_duration_seconds_metric.labels.assert_not_called()
|
||||
|
||||
def test_token_only_usage_is_a_noop(self, sample_enum_values):
|
||||
logger = _make_mock_logger()
|
||||
payload = {
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 30,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
PrometheusLogger._increment_media_generation_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
getattr(logger, name).labels.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize("bad_value", [0, 0.0, None, -4.0, "4", True])
|
||||
def test_non_positive_or_non_numeric_values_are_ignored(self, sample_enum_values, bad_value):
|
||||
logger = _make_mock_logger()
|
||||
payload = {
|
||||
"metadata": {
|
||||
"usage_object": {
|
||||
"duration_seconds": bad_value,
|
||||
"output_image_count": bad_value,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
PrometheusLogger._increment_media_generation_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
getattr(logger, name).labels.assert_not_called()
|
||||
|
||||
def test_missing_usage_object_is_a_noop(self, sample_enum_values):
|
||||
logger = _make_mock_logger()
|
||||
|
||||
for payload in ({"metadata": {}}, {"metadata": None}, {"metadata": {"usage_object": "redacted"}}):
|
||||
PrometheusLogger._increment_media_generation_metrics(
|
||||
logger,
|
||||
standard_logging_payload=payload,
|
||||
enum_values=sample_enum_values,
|
||||
)
|
||||
|
||||
for name in MEDIA_GENERATION_METRICS:
|
||||
getattr(logger, name).labels.assert_not_called()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
|
@ -326,3 +326,148 @@ async def test_should_leave_rate_limit_labels_blank_for_non_rate_limit_failure()
|
|||
assert isinstance(enum_values, UserAPIKeyLabelValues)
|
||||
assert enum_values.rate_limit_category is None
|
||||
assert enum_values.rate_limit_type is None
|
||||
|
||||
|
||||
def _logger_with_mock_virtual_key_gauges() -> PrometheusLogger:
|
||||
with patch(
|
||||
"litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None
|
||||
):
|
||||
logger = PrometheusLogger()
|
||||
logger.litellm_remaining_api_key_requests_for_model = MagicMock()
|
||||
logger.litellm_remaining_api_key_tokens_for_model = MagicMock()
|
||||
logger.get_labels_for_metric = MagicMock(return_value=[])
|
||||
return logger
|
||||
|
||||
|
||||
def _kwargs_with_v3_rate_limit_headers(additional_headers: dict) -> dict:
|
||||
return {
|
||||
"litellm_params": {"metadata": {"model_group": "gpt-4o-mini"}},
|
||||
"standard_logging_object": {
|
||||
"metadata": {},
|
||||
"hidden_params": {"additional_headers": additional_headers},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _set_virtual_key_metrics(logger: PrometheusLogger, kwargs: dict) -> None:
|
||||
logger._set_virtual_key_rate_limit_metrics(
|
||||
user_api_key="test-hash",
|
||||
user_api_key_alias="test-alias",
|
||||
kwargs=kwargs,
|
||||
metadata=kwargs["litellm_params"]["metadata"],
|
||||
model_id="model-123",
|
||||
)
|
||||
|
||||
|
||||
def test_should_read_v3_remaining_headers_when_metadata_keys_absent():
|
||||
"""
|
||||
Regression for LIT-2577: the default v3 rate limiter writes remaining
|
||||
per-(key, model) values into
|
||||
``standard_logging_object.hidden_params.additional_headers`` as
|
||||
``x-ratelimit-model_per_key-remaining-{requests,tokens}`` and never sets
|
||||
the legacy ``litellm-key-remaining-*`` metadata keys, so the gauges were
|
||||
pinned to ``sys.maxsize``.
|
||||
"""
|
||||
logger = _logger_with_mock_virtual_key_gauges()
|
||||
kwargs = _kwargs_with_v3_rate_limit_headers(
|
||||
{
|
||||
"x-ratelimit-model_per_key-remaining-requests": 42,
|
||||
"x-ratelimit-model_per_key-remaining-tokens": 900,
|
||||
"x-ratelimit-model_per_key-limit-requests": 100,
|
||||
"x-ratelimit-model_per_key-limit-tokens": 1000,
|
||||
}
|
||||
)
|
||||
|
||||
_set_virtual_key_metrics(logger, kwargs)
|
||||
|
||||
logger.litellm_remaining_api_key_requests_for_model.labels.return_value.set.assert_called_once_with(
|
||||
42
|
||||
)
|
||||
logger.litellm_remaining_api_key_tokens_for_model.labels.return_value.set.assert_called_once_with(
|
||||
900
|
||||
)
|
||||
|
||||
|
||||
def test_should_prefer_legacy_metadata_keys_over_v3_headers():
|
||||
logger = _logger_with_mock_virtual_key_gauges()
|
||||
kwargs = _kwargs_with_v3_rate_limit_headers(
|
||||
{
|
||||
"x-ratelimit-model_per_key-remaining-requests": 42,
|
||||
"x-ratelimit-model_per_key-remaining-tokens": 900,
|
||||
}
|
||||
)
|
||||
kwargs["litellm_params"]["metadata"].update(
|
||||
{
|
||||
"litellm-key-remaining-requests-gpt-4o-mini": 3,
|
||||
"litellm-key-remaining-tokens-gpt-4o-mini": 200,
|
||||
}
|
||||
)
|
||||
|
||||
_set_virtual_key_metrics(logger, kwargs)
|
||||
|
||||
logger.litellm_remaining_api_key_requests_for_model.labels.return_value.set.assert_called_once_with(
|
||||
3
|
||||
)
|
||||
logger.litellm_remaining_api_key_tokens_for_model.labels.return_value.set.assert_called_once_with(
|
||||
200
|
||||
)
|
||||
|
||||
|
||||
def test_should_treat_zero_v3_remaining_as_zero():
|
||||
logger = _logger_with_mock_virtual_key_gauges()
|
||||
kwargs = _kwargs_with_v3_rate_limit_headers(
|
||||
{
|
||||
"x-ratelimit-model_per_key-remaining-requests": 0,
|
||||
"x-ratelimit-model_per_key-remaining-tokens": 0,
|
||||
}
|
||||
)
|
||||
|
||||
_set_virtual_key_metrics(logger, kwargs)
|
||||
|
||||
logger.litellm_remaining_api_key_requests_for_model.labels.return_value.set.assert_called_once_with(
|
||||
0
|
||||
)
|
||||
logger.litellm_remaining_api_key_tokens_for_model.labels.return_value.set.assert_called_once_with(
|
||||
0
|
||||
)
|
||||
|
||||
|
||||
def test_should_keep_maxsize_sentinel_when_no_rate_limit_source_present():
|
||||
import sys
|
||||
|
||||
logger = _logger_with_mock_virtual_key_gauges()
|
||||
kwargs = {
|
||||
"litellm_params": {"metadata": {"model_group": "gpt-4o-mini"}},
|
||||
"standard_logging_object": {"metadata": {}, "hidden_params": {}},
|
||||
}
|
||||
|
||||
_set_virtual_key_metrics(logger, kwargs)
|
||||
|
||||
logger.litellm_remaining_api_key_requests_for_model.labels.return_value.set.assert_called_once_with(
|
||||
sys.maxsize
|
||||
)
|
||||
logger.litellm_remaining_api_key_tokens_for_model.labels.return_value.set.assert_called_once_with(
|
||||
sys.maxsize
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad_value", ["not-a-number", None, True])
|
||||
def test_should_ignore_non_int_v3_header_values(bad_value):
|
||||
import sys
|
||||
|
||||
logger = _logger_with_mock_virtual_key_gauges()
|
||||
kwargs = _kwargs_with_v3_rate_limit_headers(
|
||||
{
|
||||
"x-ratelimit-model_per_key-remaining-requests": bad_value,
|
||||
"x-ratelimit-model_per_key-remaining-tokens": bad_value,
|
||||
}
|
||||
)
|
||||
|
||||
_set_virtual_key_metrics(logger, kwargs)
|
||||
|
||||
logger.litellm_remaining_api_key_requests_for_model.labels.return_value.set.assert_called_once_with(
|
||||
sys.maxsize
|
||||
)
|
||||
logger.litellm_remaining_api_key_tokens_for_model.labels.return_value.set.assert_called_once_with(
|
||||
sys.maxsize
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3707,3 +3707,69 @@ def test_set_cost_breakdown_stores_reasoning_cost():
|
|||
cost_for_built_in_tools_cost_usd_dollar=0.0,
|
||||
)
|
||||
assert "reasoning_cost" not in no_reasoning.cost_breakdown
|
||||
|
||||
|
||||
def _build_payload_for_media_response(logging_obj, init_response_obj, kwargs=None):
|
||||
import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload,
|
||||
)
|
||||
|
||||
now = datetime.datetime.now()
|
||||
return get_standard_logging_object_payload(
|
||||
kwargs=kwargs or {"litellm_call_id": "media-call-id", "model": "test-model", "messages": []},
|
||||
init_response_obj=init_response_obj,
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
logging_obj=logging_obj,
|
||||
status="success",
|
||||
)
|
||||
|
||||
|
||||
def test_image_response_sets_output_image_count_on_usage_object(logging_obj):
|
||||
"""Generated-image count must land on metadata.usage_object for callbacks (e.g. Prometheus)."""
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
response = ImageResponse(created=1, data=[{"url": "https://img/1"}, {"url": "https://img/2"}])
|
||||
|
||||
payload = _build_payload_for_media_response(logging_obj, response)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["metadata"]["usage_object"]["output_image_count"] == 2
|
||||
|
||||
|
||||
def test_output_image_count_survives_message_redaction(logging_obj, monkeypatch):
|
||||
"""Redaction replaces the ImageResponse body, so the count must be captured pre-redaction."""
|
||||
import litellm
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
monkeypatch.setattr(litellm, "turn_off_message_logging", True)
|
||||
response = ImageResponse(created=1, data=[{"url": "https://img/1"}])
|
||||
|
||||
payload = _build_payload_for_media_response(logging_obj, response)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["response"] == {"text": "redacted-by-litellm"}
|
||||
assert payload["metadata"]["usage_object"]["output_image_count"] == 1
|
||||
|
||||
|
||||
def test_non_image_response_has_no_output_image_count(logging_obj):
|
||||
payload = _build_payload_for_media_response(
|
||||
logging_obj, {"id": "chatcmpl-1", "usage": {"prompt_tokens": 1, "completion_tokens": 2}}
|
||||
)
|
||||
|
||||
assert payload is not None
|
||||
assert "output_image_count" not in payload["metadata"]["usage_object"]
|
||||
|
||||
|
||||
def test_zero_token_video_usage_preserves_duration_seconds(logging_obj):
|
||||
"""Video usage bills by duration; the payload must keep duration_seconds even with zero tokens."""
|
||||
payload = _build_payload_for_media_response(
|
||||
logging_obj, {"id": "video-1", "usage": {"duration_seconds": 4.0}}
|
||||
)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["metadata"]["usage_object"]["duration_seconds"] == 4.0
|
||||
assert payload["total_tokens"] == 0
|
||||
assert payload["completion_tokens"] == 0
|
||||
|
|
|
|||
|
|
@ -10,9 +10,11 @@ from types import SimpleNamespace
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.redact_messages import (
|
||||
_redact_responses_api_output,
|
||||
perform_redaction,
|
||||
redact_streaming_responses_for_custom_logger,
|
||||
should_redact_message_logging,
|
||||
)
|
||||
from litellm.responses.main import mock_responses_api_response
|
||||
|
|
@ -442,3 +444,109 @@ class TestPerformRedaction:
|
|||
assert "vertex_ai_url_context_metadata" not in hidden_params
|
||||
assert "vertex_ai_safety_ratings" not in hidden_params
|
||||
assert "vertex_ai_citation_metadata" not in hidden_params
|
||||
|
||||
def test_redact_async_complete_streaming_response(self):
|
||||
"""Test that async_complete_streaming_response is properly redacted."""
|
||||
response_obj = litellm.ModelResponse(
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
message=litellm.Message(content="secret content", role="assistant")
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
model_call_details = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"prompt": "hi",
|
||||
"input": "hi",
|
||||
"stream": True,
|
||||
"async_complete_streaming_response": response_obj,
|
||||
}
|
||||
|
||||
perform_redaction(model_call_details, result=None)
|
||||
|
||||
redacted_response = model_call_details["async_complete_streaming_response"]
|
||||
assert redacted_response.choices[0].message.content == "redacted-by-litellm"
|
||||
|
||||
def test_redact_complete_streaming_response(self):
|
||||
"""Test that complete_streaming_response is properly redacted."""
|
||||
response_obj = litellm.ModelResponse(
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
message=litellm.Message(content="secret content", role="assistant")
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
model_call_details = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"prompt": "hi",
|
||||
"input": "hi",
|
||||
"stream": True,
|
||||
"complete_streaming_response": response_obj,
|
||||
}
|
||||
|
||||
perform_redaction(model_call_details, result=None)
|
||||
|
||||
redacted_response = model_call_details["complete_streaming_response"]
|
||||
assert redacted_response.choices[0].message.content == "redacted-by-litellm"
|
||||
|
||||
def test_streaming_responses_untouched_when_disabled(self):
|
||||
response_obj = litellm.ModelResponse(
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
message=litellm.Message(content="secret content", role="assistant")
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
model_call_details = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"prompt": "hi",
|
||||
"input": "hi",
|
||||
"stream": True,
|
||||
"async_complete_streaming_response": response_obj,
|
||||
}
|
||||
|
||||
perform_redaction(model_call_details, result=None, redact_streaming_responses=False)
|
||||
|
||||
assert response_obj.choices[0].message.content == "secret content"
|
||||
|
||||
|
||||
class TestRedactStreamingResponsesForCustomLogger:
|
||||
def _model_call_details(self):
|
||||
response_obj = litellm.ModelResponse(
|
||||
choices=[
|
||||
litellm.Choices(
|
||||
message=litellm.Message(content="secret content", role="assistant")
|
||||
)
|
||||
]
|
||||
)
|
||||
return {
|
||||
"stream": True,
|
||||
"async_complete_streaming_response": response_obj,
|
||||
}, response_obj
|
||||
|
||||
def test_opted_out_logger_gets_redacted_copy(self):
|
||||
model_call_details, response_obj = self._model_call_details()
|
||||
opted_out_logger = CustomLogger(message_logging=False)
|
||||
|
||||
redacted_details = redact_streaming_responses_for_custom_logger(
|
||||
model_call_details=model_call_details, custom_logger=opted_out_logger
|
||||
)
|
||||
|
||||
redacted_response = redacted_details["async_complete_streaming_response"]
|
||||
assert redacted_response.choices[0].message.content == "redacted-by-litellm"
|
||||
assert response_obj.choices[0].message.content == "secret content"
|
||||
assert model_call_details["async_complete_streaming_response"] is response_obj
|
||||
|
||||
def test_compliant_logger_gets_shared_response(self):
|
||||
model_call_details, response_obj = self._model_call_details()
|
||||
compliant_logger = CustomLogger()
|
||||
|
||||
result_details = redact_streaming_responses_for_custom_logger(
|
||||
model_call_details=model_call_details, custom_logger=compliant_logger
|
||||
)
|
||||
|
||||
assert result_details is model_call_details
|
||||
assert response_obj.choices[0].message.content == "secret content"
|
||||
|
|
|
|||
|
|
@ -957,15 +957,15 @@ def test_anthropic_structured_output_beta_header():
|
|||
@pytest.mark.parametrize(
|
||||
"model_name",
|
||||
[
|
||||
"claude-opus-4-6-20250918",
|
||||
"claude-opus-4.6-20250918",
|
||||
"claude-opus-4-8",
|
||||
"claude-opus-4-6-20260205",
|
||||
"claude-opus-4-5-20251101",
|
||||
"claude-opus-4.5-20251101",
|
||||
],
|
||||
)
|
||||
def test_opus_uses_native_structured_output(model_name):
|
||||
"""
|
||||
Test that Opus 4.5 and 4.6 models use native Anthropic structured outputs
|
||||
Test that supported Opus models use native Anthropic structured outputs
|
||||
(output_format) rather than the tool-based workaround.
|
||||
"""
|
||||
config = AnthropicConfig()
|
||||
|
|
@ -1005,6 +1005,43 @@ def test_opus_uses_native_structured_output(model_name):
|
|||
assert optional_params.get("json_mode") is True
|
||||
|
||||
|
||||
def test_native_structured_output_uses_bundled_capability_when_remote_map_lags(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
model = "claude-opus-4-8"
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
{model: {"supports_response_schema": True}},
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
try:
|
||||
optional_params = AnthropicConfig().map_openai_params(
|
||||
non_default_params={
|
||||
"response_format": {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "answer",
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"properties": {"answer": {"type": "string"}},
|
||||
"required": ["answer"],
|
||||
},
|
||||
},
|
||||
}
|
||||
},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
finally:
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
assert "output_format" in optional_params
|
||||
assert "tools" not in optional_params
|
||||
|
||||
|
||||
def test_non_structured_output_model_uses_tool_workaround():
|
||||
"""
|
||||
Test that models NOT in the native structured output list still use the
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import pytest
|
|||
|
||||
from litellm.constants import (
|
||||
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
|
||||
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
|
||||
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
|
||||
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
|
||||
)
|
||||
|
|
@ -174,6 +175,81 @@ def test_unrecognized_effort_raises_clean_400():
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
def test_pinned_temperature_dropped_when_adaptive_downgraded_to_enabled():
|
||||
"""Regression (#33203): Claude Code's safety classifier sends adaptive thinking +
|
||||
temperature=0 to Haiku 4.5. The adaptive interface is downgraded to legacy enabled
|
||||
thinking, but Anthropic rejects "temperature may only be set to 1 when thinking is
|
||||
enabled". The pinned temperature must be dropped so the request succeeds while the
|
||||
downgraded thinking is preserved."""
|
||||
params = _claude_code_payload(effort="medium")
|
||||
params["temperature"] = 0
|
||||
result = _transform("claude-haiku-4-5", params)
|
||||
|
||||
assert result["thinking"] == {
|
||||
"type": "enabled",
|
||||
"budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
|
||||
}
|
||||
assert "temperature" not in result
|
||||
|
||||
|
||||
def test_temperature_one_preserved_with_enabled_thinking():
|
||||
"""temperature=1 is compatible with extended thinking, so it must be kept."""
|
||||
params = _claude_code_payload(effort="medium")
|
||||
params["temperature"] = 1
|
||||
result = _transform("claude-haiku-4-5", params)
|
||||
|
||||
assert result["thinking"]["type"] == "enabled"
|
||||
assert result["temperature"] == 1
|
||||
|
||||
|
||||
def test_pinned_temperature_preserved_when_thinking_dropped():
|
||||
"""When thinking is dropped entirely (non-reasoning model), there is no thinking
|
||||
conflict, so a pinned temperature must survive untouched."""
|
||||
params = _claude_code_payload(effort="medium")
|
||||
params["temperature"] = 0
|
||||
result = _transform("claude-3-5-haiku-latest", params)
|
||||
|
||||
assert "thinking" not in result
|
||||
assert result["temperature"] == 0
|
||||
|
||||
|
||||
def test_pinned_temperature_preserved_for_adaptive_model():
|
||||
"""Adaptive models (4.6+) own the thinking/temperature relationship natively, so
|
||||
the passthrough must not strip a pinned temperature for them."""
|
||||
params = _claude_code_payload(effort="high")
|
||||
params["temperature"] = 0
|
||||
result = _transform("claude-sonnet-4-6", params)
|
||||
|
||||
assert result["thinking"] == {"type": "adaptive"}
|
||||
assert result["temperature"] == 0
|
||||
|
||||
|
||||
def test_pinned_temperature_dropped_for_opus_4_5_effort():
|
||||
"""Opus 4.5 keeps native output_config.effort (extended thinking), which is equally
|
||||
incompatible with a pinned non-1 temperature, so the temperature must be dropped."""
|
||||
params = _claude_code_payload(effort="medium")
|
||||
params["temperature"] = 0
|
||||
result = _transform("claude-opus-4-5", params)
|
||||
|
||||
assert result["output_config"] == {"effort": "medium"}
|
||||
assert "temperature" not in result
|
||||
|
||||
|
||||
def test_reasoning_effort_with_pinned_temperature_drops_temperature():
|
||||
"""The reasoning_effort alias synthesizes legacy enabled thinking on a non-adaptive
|
||||
model; a co-pinned non-1 temperature must be dropped to avoid the Anthropic 400."""
|
||||
result = _transform(
|
||||
"claude-haiku-4-5",
|
||||
{"max_tokens": 8192, "reasoning_effort": "low", "temperature": 0},
|
||||
)
|
||||
|
||||
assert result["thinking"] == {
|
||||
"type": "enabled",
|
||||
"budget_tokens": DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
|
||||
}
|
||||
assert "temperature" not in result
|
||||
|
||||
|
||||
def test_non_adaptive_request_without_effort_is_untouched():
|
||||
"""A non-adaptive model receiving a request with no adaptive interface (no
|
||||
effort, no adaptive thinking) must pass through untouched."""
|
||||
|
|
|
|||
|
|
@ -99,17 +99,11 @@ class TestContextManagementConversion:
|
|||
}
|
||||
)
|
||||
kwargs = _ADAPTER.translate_request(req)
|
||||
assert kwargs["context_management"] == [
|
||||
{"type": "compaction", "compact_threshold": 100000}
|
||||
]
|
||||
assert kwargs["context_management"] == [{"type": "compaction", "compact_threshold": 100000}]
|
||||
|
||||
def test_translate_request_drops_anthropic_only_context_management(self):
|
||||
"""context_management with only unknown edit types is omitted from kwargs."""
|
||||
req = _make_request(
|
||||
context_management={
|
||||
"edits": [{"type": "clear_thinking_20251015", "keep": "all"}]
|
||||
}
|
||||
)
|
||||
req = _make_request(context_management={"edits": [{"type": "clear_thinking_20251015", "keep": "all"}]})
|
||||
kwargs = _ADAPTER.translate_request(req)
|
||||
assert "context_management" not in kwargs
|
||||
|
||||
|
|
@ -134,9 +128,7 @@ class TestOutputConfigStructuredOutput:
|
|||
|
||||
def test_output_config_format_json_schema_converted(self):
|
||||
"""output_config.format.json_schema is converted to OpenAI text.format."""
|
||||
req = _make_request(
|
||||
output_config={"format": {"type": "json_schema", "schema": self._SCHEMA}}
|
||||
)
|
||||
req = _make_request(output_config={"format": {"type": "json_schema", "schema": self._SCHEMA}})
|
||||
kwargs = _ADAPTER.translate_request(req)
|
||||
assert "text" in kwargs
|
||||
fmt = kwargs["text"]["format"]
|
||||
|
|
@ -153,9 +145,7 @@ class TestOutputConfigStructuredOutput:
|
|||
|
||||
def test_output_format_still_works(self):
|
||||
"""The original output_format field still takes precedence when present."""
|
||||
req = _make_request(
|
||||
output_format={"type": "json_schema", "schema": self._SCHEMA}
|
||||
)
|
||||
req = _make_request(output_format={"type": "json_schema", "schema": self._SCHEMA})
|
||||
kwargs = _ADAPTER.translate_request(req)
|
||||
assert "text" in kwargs
|
||||
assert kwargs["text"]["format"]["type"] == "json_schema"
|
||||
|
|
@ -250,9 +240,7 @@ class TestTranslateMessagesToResponsesInput:
|
|||
]
|
||||
result = _translate_messages(messages)
|
||||
assert len(result) == 1
|
||||
assert result[0]["content"] == [
|
||||
{"type": "input_image", "image_url": "data:image/png;base64,abc123"}
|
||||
]
|
||||
assert result[0]["content"] == [{"type": "input_image", "image_url": "data:image/png;base64,abc123"}]
|
||||
|
||||
def test_user_url_image(self):
|
||||
"""User message with URL image source becomes input_image with the URL."""
|
||||
|
|
@ -268,9 +256,7 @@ class TestTranslateMessagesToResponsesInput:
|
|||
}
|
||||
]
|
||||
result = _translate_messages(messages)
|
||||
assert result[0]["content"] == [
|
||||
{"type": "input_image", "image_url": "https://example.com/img.jpg"}
|
||||
]
|
||||
assert result[0]["content"] == [{"type": "input_image", "image_url": "https://example.com/img.jpg"}]
|
||||
|
||||
def test_user_base64_image_empty_data_skipped(self):
|
||||
"""Base64 image with empty data is skipped (no URL can be formed)."""
|
||||
|
|
@ -341,9 +327,7 @@ class TestTranslateMessagesToResponsesInput:
|
|||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": "call_null", "content": None}
|
||||
],
|
||||
"content": [{"type": "tool_result", "tool_use_id": "call_null", "content": None}],
|
||||
}
|
||||
]
|
||||
result = _translate_messages(messages)
|
||||
|
|
@ -370,9 +354,7 @@ class TestTranslateMessagesToResponsesInput:
|
|||
}
|
||||
]
|
||||
result = _translate_messages(messages)
|
||||
assert result[0]["content"] == [
|
||||
{"type": "output_text", "text": "Here is the answer."}
|
||||
]
|
||||
assert result[0]["content"] == [{"type": "output_text", "text": "Here is the answer."}]
|
||||
|
||||
def test_assistant_tool_use_becomes_function_call(self):
|
||||
"""Assistant tool_use block becomes a top-level function_call item."""
|
||||
|
|
@ -404,15 +386,11 @@ class TestTranslateMessagesToResponsesInput:
|
|||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": "Let me reason step by step."}
|
||||
],
|
||||
"content": [{"type": "thinking", "thinking": "Let me reason step by step."}],
|
||||
}
|
||||
]
|
||||
result = _translate_messages(messages)
|
||||
assert result[0]["content"] == [
|
||||
{"type": "output_text", "text": "Let me reason step by step."}
|
||||
]
|
||||
assert result[0]["content"] == [{"type": "output_text", "text": "Let me reason step by step."}]
|
||||
|
||||
def test_assistant_empty_thinking_block_skipped(self):
|
||||
"""Assistant thinking block with empty thinking text is skipped."""
|
||||
|
|
@ -584,28 +562,27 @@ class TestTranslateToolsToResponsesAPI:
|
|||
|
||||
|
||||
class TestTranslateToolChoiceToResponsesAPI:
|
||||
"""Anthropic tool_choice -> Responses API tool_choice."""
|
||||
"""Anthropic tool_choice -> Responses API tool_choice.
|
||||
|
||||
def test_auto_maps_to_auto(self):
|
||||
assert _ADAPTER.translate_tool_choice_to_responses_api({"type": "auto"}) == {
|
||||
"type": "auto"
|
||||
}
|
||||
The Responses API's tool_choice schema (openai.types.responses.tool_choice_options)
|
||||
is a bare Literal["none", "auto", "required"] for these simple cases - not an
|
||||
object like {"type": "auto"}. Sending the object shape to an OpenAI-compatible
|
||||
server gets rejected with a pydantic validation error.
|
||||
"""
|
||||
|
||||
def test_any_maps_to_required(self):
|
||||
assert _ADAPTER.translate_tool_choice_to_responses_api({"type": "any"}) == {
|
||||
"type": "required"
|
||||
}
|
||||
def test_auto_maps_to_bare_string_auto(self):
|
||||
assert _ADAPTER.translate_tool_choice_to_responses_api({"type": "auto"}) == "auto"
|
||||
|
||||
def test_any_maps_to_bare_string_required(self):
|
||||
assert _ADAPTER.translate_tool_choice_to_responses_api({"type": "any"}) == "required"
|
||||
|
||||
def test_none_maps_to_bare_string_none(self):
|
||||
assert _ADAPTER.translate_tool_choice_to_responses_api({"type": "none"}) == "none"
|
||||
|
||||
def test_specific_tool_maps_to_function(self):
|
||||
result = _ADAPTER.translate_tool_choice_to_responses_api(
|
||||
{"type": "tool", "name": "get_weather"}
|
||||
)
|
||||
result = _ADAPTER.translate_tool_choice_to_responses_api({"type": "tool", "name": "get_weather"})
|
||||
assert result == {"type": "function", "name": "get_weather"}
|
||||
|
||||
def test_unknown_type_defaults_to_auto(self):
|
||||
result = _ADAPTER.translate_tool_choice_to_responses_api({"type": "none"})
|
||||
assert result == {"type": "auto"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# translate_thinking_to_reasoning
|
||||
|
|
@ -616,17 +593,13 @@ class TestTranslateThinkingToReasoning:
|
|||
"""Anthropic thinking param -> Responses API reasoning param."""
|
||||
|
||||
def test_budget_high_effort(self):
|
||||
result = _ADAPTER.translate_thinking_to_reasoning(
|
||||
{"type": "enabled", "budget_tokens": 10000}
|
||||
)
|
||||
result = _ADAPTER.translate_thinking_to_reasoning({"type": "enabled", "budget_tokens": 10000})
|
||||
# Default (reasoning_auto_summary=False): only effort, no summary
|
||||
assert result == {"effort": "high"}
|
||||
assert result is not None and "summary" not in result
|
||||
|
||||
def test_budget_above_threshold_high_effort(self):
|
||||
result = _ADAPTER.translate_thinking_to_reasoning(
|
||||
{"type": "enabled", "budget_tokens": 50000}
|
||||
)
|
||||
result = _ADAPTER.translate_thinking_to_reasoning({"type": "enabled", "budget_tokens": 50000})
|
||||
assert result is not None
|
||||
assert result["effort"] == "high"
|
||||
assert "summary" not in result
|
||||
|
|
@ -652,9 +625,7 @@ class TestTranslateThinkingToReasoning:
|
|||
assert result is not None and "summary" not in result
|
||||
|
||||
def test_budget_minimal_effort(self):
|
||||
result = _ADAPTER.translate_thinking_to_reasoning(
|
||||
{"type": "enabled", "budget_tokens": 500}
|
||||
)
|
||||
result = _ADAPTER.translate_thinking_to_reasoning({"type": "enabled", "budget_tokens": 500})
|
||||
assert result == {"effort": "minimal"}
|
||||
assert result is not None and "summary" not in result
|
||||
|
||||
|
|
@ -707,9 +678,7 @@ class TestTranslateThinkingToReasoning:
|
|||
original = litellm.reasoning_auto_summary
|
||||
try:
|
||||
litellm.reasoning_auto_summary = True
|
||||
result = _ADAPTER.translate_thinking_to_reasoning(
|
||||
{"type": "enabled", "budget_tokens": 10000}
|
||||
)
|
||||
result = _ADAPTER.translate_thinking_to_reasoning({"type": "enabled", "budget_tokens": 10000})
|
||||
assert result == {"effort": "high", "summary": "detailed"}
|
||||
finally:
|
||||
litellm.reasoning_auto_summary = original
|
||||
|
|
@ -789,11 +758,7 @@ class TestTranslateRequestBroaderCoverage:
|
|||
assert kwargs["top_p"] == 0.9
|
||||
|
||||
def test_tools_translated(self):
|
||||
req = _make_request(
|
||||
tools=[
|
||||
{"name": "calculator", "description": "Does math.", "input_schema": {}}
|
||||
]
|
||||
)
|
||||
req = _make_request(tools=[{"name": "calculator", "description": "Does math.", "input_schema": {}}])
|
||||
kwargs = _ADAPTER.translate_request(req)
|
||||
assert len(kwargs["tools"]) == 1
|
||||
assert kwargs["tools"][0]["name"] == "calculator"
|
||||
|
|
@ -929,9 +894,7 @@ class TestTranslateResponse:
|
|||
|
||||
def test_multiple_text_parts(self):
|
||||
"""Multiple output_text parts become multiple text content blocks."""
|
||||
response = _make_mock_response(
|
||||
output=[_make_output_message(["Part 1", "Part 2"])]
|
||||
)
|
||||
response = _make_mock_response(output=[_make_output_message(["Part 1", "Part 2"])])
|
||||
result: Any = _ADAPTER.translate_response(response)
|
||||
assert len(result["content"]) == 2
|
||||
assert result["content"][0]["text"] == "Part 1"
|
||||
|
|
|
|||
|
|
@ -158,6 +158,54 @@ class TestClaudePlatformActionsCovered:
|
|||
)
|
||||
|
||||
|
||||
class TestBedrockMantleActionsCovered:
|
||||
"""LIT-3859: bedrock_mantle inference authorizes against the
|
||||
``bedrock-mantle`` action namespace, so the session-policy ceiling
|
||||
must include it or every Mantle request via OIDC/WIF auth denies
|
||||
with "no session policy allows the bedrock-mantle:CreateInference
|
||||
action" even when the role's identity policy grants it."""
|
||||
|
||||
def test_bedrock_mantle_create_inference_present(self):
|
||||
policy = _captured_policy()
|
||||
all_actions: set = set()
|
||||
for stmt in policy["Statement"]:
|
||||
stmt_actions = stmt.get("Action")
|
||||
if isinstance(stmt_actions, str):
|
||||
all_actions.add(stmt_actions)
|
||||
elif isinstance(stmt_actions, list):
|
||||
all_actions.update(stmt_actions)
|
||||
assert "bedrock-mantle:CreateInference" in all_actions, (
|
||||
"bedrock-mantle:CreateInference missing from session policy — "
|
||||
"bedrock_mantle/* requests will 403 on OIDC/WIF auth"
|
||||
)
|
||||
|
||||
def test_bedrock_mantle_statement_allows(self):
|
||||
policy = _captured_policy()
|
||||
stmt = _statement_by_sid(policy, "BedrockMantleLiteLLM")
|
||||
assert stmt["Effect"] == "Allow"
|
||||
assert stmt["Resource"] == "*"
|
||||
|
||||
def test_no_bedrock_mantle_wildcard(self):
|
||||
policy = _captured_policy()
|
||||
stmt = _statement_by_sid(policy, "BedrockMantleLiteLLM")
|
||||
actions = stmt["Action"]
|
||||
if isinstance(actions, str):
|
||||
actions = [actions]
|
||||
assert "bedrock-mantle:*" not in actions, (
|
||||
"session policy must not grant bedrock-mantle:* — "
|
||||
"the ceiling should match the documented action set"
|
||||
)
|
||||
|
||||
def test_bedrock_mantle_statement_carries_secure_transport_condition(self):
|
||||
policy = _captured_policy()
|
||||
stmt = _statement_by_sid(policy, "BedrockMantleLiteLLM")
|
||||
cond = stmt.get("Condition") or {}
|
||||
assert cond.get("Bool", {}).get("aws:SecureTransport") == "true", (
|
||||
"BedrockMantleLiteLLM must require aws:SecureTransport=true "
|
||||
"to keep parity with the bedrock statement"
|
||||
)
|
||||
|
||||
|
||||
def _make_jwt(payload: dict) -> str:
|
||||
def _segment(data: dict) -> str:
|
||||
return base64.urlsafe_b64encode(json.dumps(data).encode()).rstrip(b"=").decode()
|
||||
|
|
|
|||
|
|
@ -44,6 +44,60 @@ class TestOpenAIResponsesAPIConfig:
|
|||
# The function should return the params unchanged
|
||||
assert result == test_params
|
||||
|
||||
@pytest.mark.parametrize("max_output_tokens", [1, 15])
|
||||
def test_map_openai_params_clamps_max_output_tokens_below_minimum(self, max_output_tokens):
|
||||
"""OpenAI's Responses API rejects max_output_tokens < 16.
|
||||
|
||||
Claude Code (via the Anthropic Messages -> Responses adapter) sends a
|
||||
max_tokens=1 warmup probe when running `/model`, which produced:
|
||||
"Invalid 'max_output_tokens': integer below minimum value.
|
||||
Expected a value >= 16, but got 1 instead."
|
||||
Clamp anything below the minimum up to 16 instead of erroring.
|
||||
"""
|
||||
result = self.config.map_openai_params(
|
||||
response_api_optional_params={"max_output_tokens": max_output_tokens},
|
||||
model=self.model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["max_output_tokens"] == 16
|
||||
|
||||
def test_map_openai_params_preserves_max_output_tokens_at_or_above_minimum(self):
|
||||
"""Values already >= 16 must pass through untouched."""
|
||||
result = self.config.map_openai_params(
|
||||
response_api_optional_params={"max_output_tokens": 256},
|
||||
model=self.model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["max_output_tokens"] == 256
|
||||
|
||||
def test_map_openai_params_leaves_max_output_tokens_absent(self):
|
||||
"""A request without max_output_tokens must not gain the key."""
|
||||
result = self.config.map_openai_params(
|
||||
response_api_optional_params={"input": "hi"},
|
||||
model=self.model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "max_output_tokens" not in result
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value, expected",
|
||||
[
|
||||
(1, 16),
|
||||
(15, 16),
|
||||
(16, 16),
|
||||
(17, 17),
|
||||
(256, 256),
|
||||
(None, None),
|
||||
],
|
||||
)
|
||||
def test_enforce_min_max_output_tokens(self, value, expected):
|
||||
"""Below the minimum clamps to 16; the boundary, larger values, and None
|
||||
are returned unchanged so no previously-valid request regresses."""
|
||||
assert self.config._enforce_min_max_output_tokens(value) == expected
|
||||
|
||||
def validate_responses_api_request_params(self, params, expected_fields):
|
||||
"""
|
||||
Validate that the params dict has the expected structure of ResponsesAPIRequestParams
|
||||
|
|
|
|||
|
|
@ -4910,6 +4910,7 @@ class TestMCPDcrBridgeDelegateAdmission:
|
|||
cls,
|
||||
*,
|
||||
key_hash=None,
|
||||
user_id=None,
|
||||
server_id="bridge-server-id",
|
||||
access_token="inner-upstream-access-token",
|
||||
token_type="Bearer",
|
||||
|
|
@ -4921,17 +4922,23 @@ class TestMCPDcrBridgeDelegateAdmission:
|
|||
envelope_keys_from_master_key,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
EnvelopeIdentity,
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
key_hash_identity,
|
||||
mint_envelope,
|
||||
user_identity,
|
||||
)
|
||||
from pydantic import SecretStr
|
||||
|
||||
identity = (
|
||||
user_identity(server_id=server_id, user_id=user_id)
|
||||
if user_id is not None
|
||||
else key_hash_identity(server_id=server_id, key_hash=key_hash or cls._KEY_HASH)
|
||||
)
|
||||
keys = envelope_keys_from_master_key(master_key or cls._MASTER_KEY)
|
||||
now = minted_at or datetime.now(timezone.utc)
|
||||
sealed = mint_envelope(
|
||||
identity=EnvelopeIdentity(server_id=server_id, key_hash=key_hash or cls._KEY_HASH),
|
||||
identity=identity,
|
||||
grant=UpstreamTokenGrant(
|
||||
access_token=SecretStr(access_token),
|
||||
token_type=token_type,
|
||||
|
|
@ -4999,6 +5006,38 @@ class TestMCPDcrBridgeDelegateAdmission:
|
|||
stack.enter_context(patcher)
|
||||
yield get_key_object
|
||||
|
||||
@staticmethod
|
||||
@contextlib.contextmanager
|
||||
def _patch_user_reload(*, return_value=None, side_effect=None):
|
||||
"""Patch the user-subject reload path an interactively-minted envelope takes: the
|
||||
``get_user_object`` lookup ``_reload_admitted_user`` runs (which also drives the SCIM gate),
|
||||
plus the ``prisma_client`` / ``user_api_key_cache`` globals. The centralized gate's own
|
||||
fetches fail-safe to None under the MagicMock prisma, so an unblocked user admits. Yields the
|
||||
``get_user_object`` mock so a caller can assert the sealed user_id was the reload key."""
|
||||
get_user_object = AsyncMock(return_value=return_value, side_effect=side_effect)
|
||||
with (
|
||||
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
|
||||
):
|
||||
yield get_user_object
|
||||
|
||||
@staticmethod
|
||||
def _wrapped_user_lookup_error(original: BaseException) -> ValueError:
|
||||
"""Reproduce get_user_object's real exception contract (litellm/proxy/auth/auth_checks.py): it
|
||||
catches every DB failure in a broad ``except`` and re-raises a bare ``ValueError``, so the
|
||||
original error (a missing-user Exception or a real outage) survives only as ``__context__``.
|
||||
Injecting a raw ConnectionError/Exception instead would exercise a shape production never
|
||||
produces and let a chain-blind outage classifier pass. That wrapping fidelity is itself pinned by
|
||||
test_get_user_object_wraps_db_outage_as_valueerror_preserving_context in test_auth_checks."""
|
||||
try:
|
||||
raise original
|
||||
except BaseException:
|
||||
try:
|
||||
raise ValueError(f"User doesn't exist in db. Got error - {original}")
|
||||
except ValueError as wrapped:
|
||||
return wrapped
|
||||
|
||||
@staticmethod
|
||||
def _mcp_request(path="/mcp/bridge_delegate_server"):
|
||||
"""A minimal ``Request`` for direct ``_admit_dcr_bridge_delegate`` calls, mirroring how
|
||||
|
|
@ -5060,6 +5099,155 @@ class TestMCPDcrBridgeDelegateAdmission:
|
|||
"bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"}
|
||||
}
|
||||
|
||||
async def test_user_subject_envelope_admits_under_the_reloaded_user(self):
|
||||
"""An interactively-minted (user_id) envelope admits under the reloaded USER, not a key: the
|
||||
reload is keyed by the sealed user_id, the admitted auth carries that user_id, the raw-key
|
||||
pipeline is never invoked, and the inner upstream token is injected for egress. This is the
|
||||
interactive-DCR admission the whole flow exists for."""
|
||||
envelope = self._mint_bridge_envelope(user_id="sso-user-7")
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bridge_delegate_server",
|
||||
"headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))],
|
||||
}
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_auth,
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
|
||||
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
|
||||
self._patch_user_reload(
|
||||
return_value=MagicMock(
|
||||
user_id="sso-user-7",
|
||||
metadata={"scim_active": True},
|
||||
user_role=None,
|
||||
object_permission=None,
|
||||
object_permission_id=None,
|
||||
)
|
||||
) as get_user_object,
|
||||
):
|
||||
mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server()
|
||||
(auth_result, _h, _s, mcp_server_auth_headers, _o, _r) = await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
||||
assert get_user_object.await_args.kwargs["user_id"] == "sso-user-7"
|
||||
assert auth_result.user_id == "sso-user-7"
|
||||
mock_auth.assert_not_called()
|
||||
assert mcp_server_auth_headers == {
|
||||
"bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"}
|
||||
}
|
||||
|
||||
async def test_user_subject_envelope_carries_the_users_mcp_object_permission(self):
|
||||
"""The admitted user's own MCP object permission rides on the returned auth so the shared
|
||||
get_allowed_mcp_servers grants the user their litellm-granted servers, rather than admitting a
|
||||
bare user with no MCP access. Regression for the signed-in SSO client getting zero tools because
|
||||
the reload dropped the user's object permission."""
|
||||
object_permission = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="op-user-7", mcp_servers=["bridge_delegate_server"]
|
||||
)
|
||||
envelope = self._mint_bridge_envelope(user_id="sso-user-7")
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bridge_delegate_server",
|
||||
"headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))],
|
||||
}
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
|
||||
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
|
||||
self._patch_user_reload(
|
||||
return_value=MagicMock(
|
||||
user_id="sso-user-7",
|
||||
metadata={"scim_active": True},
|
||||
user_role=None,
|
||||
object_permission=object_permission,
|
||||
object_permission_id="op-user-7",
|
||||
)
|
||||
),
|
||||
):
|
||||
mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server()
|
||||
(auth_result, _h, _s, _headers, _o, _r) = await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
||||
assert auth_result.object_permission is not None
|
||||
assert auth_result.object_permission.mcp_servers == ["bridge_delegate_server"]
|
||||
|
||||
async def test_user_subject_envelope_missing_user_fails_closed_401(self):
|
||||
"""A user_id envelope whose user has since been deleted must fail closed with a 401, not a 500.
|
||||
get_user_object catches the missing row and re-raises a bare ValueError (it does not return None
|
||||
on the production path), so the reload must fail closed rather than let it propagate as an opaque
|
||||
500, and must not mistake the wrapped ValueError for a DB outage. Regression for the missing-user
|
||||
path surfacing as a 500."""
|
||||
envelope = self._mint_bridge_envelope(user_id="ghost-user")
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bridge_delegate_server",
|
||||
"headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))],
|
||||
}
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
|
||||
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
|
||||
self._patch_user_reload(side_effect=self._wrapped_user_lookup_error(Exception())),
|
||||
):
|
||||
mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
|
||||
async def test_user_subject_envelope_db_outage_is_retryable_503(self):
|
||||
"""A transient database outage while reloading the envelope's user is a retryable 503, not an
|
||||
opaque 500, matching the key path's contract so an interactive DCR client retries instead of
|
||||
treating a live identity as invalid. get_user_object wraps the outage in a bare ValueError, so this
|
||||
exercises the chain-aware classifier; a raw ConnectionError would falsely pass even the old
|
||||
chain-blind check because it is an OSError. Regression for the user reload dropping the 503 arm."""
|
||||
envelope = self._mint_bridge_envelope(user_id="sso-user-7")
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bridge_delegate_server",
|
||||
"headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))],
|
||||
}
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
|
||||
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
|
||||
self._patch_user_reload(
|
||||
side_effect=self._wrapped_user_lookup_error(ConnectionError("auth database unreachable"))
|
||||
),
|
||||
):
|
||||
mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
||||
assert exc_info.value.status_code == 503
|
||||
|
||||
async def test_user_subject_envelope_scim_deactivated_user_fails_closed_401(self):
|
||||
"""SCIM-deactivating the envelope's user revokes it immediately: the reloaded user carries
|
||||
scim_active False, so admission 401s rather than letting an offboarded user keep tool access
|
||||
until the envelope expires."""
|
||||
envelope = self._mint_bridge_envelope(user_id="offboarded-user")
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp/bridge_delegate_server",
|
||||
"headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))],
|
||||
}
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
|
||||
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
|
||||
self._patch_user_reload(return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False})),
|
||||
):
|
||||
mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server()
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
|
||||
async def test_revoked_key_envelope_fails_closed_401(self):
|
||||
"""An envelope whose key has since been deleted must fail closed: ``get_key_object`` raises
|
||||
for the missing row, so admission 401s instead of admitting the caller as an unrestricted
|
||||
|
|
|
|||
|
|
@ -0,0 +1,147 @@
|
|||
"""Classification matrix for upstream OAuth/DCR rejections: who is blamed depends only on the §5.2
|
||||
code and whose credentials the gateway presented, never on the upstream's HTTP status."""
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.faults.classify import (
|
||||
classify_upstream_dcr_rejection,
|
||||
classify_upstream_token_rejection,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.types import (
|
||||
CallerRejected,
|
||||
GatewayRejected,
|
||||
UpstreamProtocolFault,
|
||||
UpstreamReportedFault,
|
||||
)
|
||||
|
||||
|
||||
def _response(status_code: int, *, json_body: object = None, text_body: str = "", headers: dict = None) -> httpx.Response:
|
||||
request = httpx.Request("POST", "https://idp.example.com/token")
|
||||
if json_body is not None:
|
||||
return httpx.Response(status_code, json=json_body, request=request)
|
||||
return httpx.Response(status_code, text=text_body, headers=headers or {}, request=request)
|
||||
|
||||
|
||||
def test_caller_fault_code_classifies_as_caller_rejected_regardless_of_status():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(500, json_body={"error": "invalid_grant", "error_description": "Code expired."}),
|
||||
credential_source="gateway_stored",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, CallerRejected)
|
||||
assert fault.code == "invalid_grant"
|
||||
assert fault.description == "Code expired."
|
||||
|
||||
|
||||
def test_credential_code_with_gateway_stored_credentials_indicts_gateway():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(401, json_body={"error": "invalid_client", "error_description": "not found"}),
|
||||
credential_source="gateway_stored",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, GatewayRejected)
|
||||
assert fault.code == "invalid_client"
|
||||
|
||||
|
||||
def test_credential_code_with_caller_supplied_credentials_stays_caller_fault():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(401, json_body={"error": "invalid_client"}),
|
||||
credential_source="caller_supplied",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, CallerRejected)
|
||||
assert fault.code == "invalid_client"
|
||||
|
||||
|
||||
def test_unknown_code_relays_as_caller_rejected():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(400, json_body={"error": "slow_down", "error_description": "Polling too fast."}),
|
||||
credential_source="gateway_stored",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, CallerRejected)
|
||||
assert fault.code == "slow_down"
|
||||
|
||||
|
||||
def test_body_without_error_field_is_protocol_fault():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(404, text_body="<html>not here</html>"),
|
||||
credential_source="gateway_stored",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, UpstreamProtocolFault)
|
||||
assert fault.note == "upstream token endpoint returned HTTP 404"
|
||||
|
||||
|
||||
def test_unreadable_body_is_protocol_fault_not_exception():
|
||||
unreadable = httpx.Response(
|
||||
400,
|
||||
stream=httpx.ByteStream(b"\x1f\x8bnot-gzip"),
|
||||
headers={"content-encoding": "gzip"},
|
||||
request=httpx.Request("POST", "https://idp.example.com/token"),
|
||||
)
|
||||
fault = classify_upstream_token_rejection(unreadable, credential_source="gateway_stored", log_context="srv")
|
||||
assert isinstance(fault, UpstreamProtocolFault)
|
||||
|
||||
|
||||
def test_wire_fields_are_bounded():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(400, json_body={"error": "invalid_request", "error_description": "x" * 5000}),
|
||||
credential_source="gateway_stored",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, CallerRejected)
|
||||
assert len(fault.description) == 500
|
||||
|
||||
|
||||
def test_dcr_rejection_with_rfc7591_code_is_caller_rejected():
|
||||
fault = classify_upstream_dcr_rejection(
|
||||
_response(400, json_body={"error": "invalid_redirect_uri", "error_description": "not allowed"}),
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, CallerRejected)
|
||||
assert fault.code == "invalid_redirect_uri"
|
||||
|
||||
|
||||
def test_dcr_rejection_without_code_is_protocol_fault():
|
||||
fault = classify_upstream_dcr_rejection(_response(500, text_body="<html>trace</html>"), log_context="srv")
|
||||
assert isinstance(fault, UpstreamProtocolFault)
|
||||
assert fault.note == "upstream registration failed with HTTP 500"
|
||||
|
||||
|
||||
def test_upstream_self_blame_codes_stay_upstream_faults():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(400, json_body={"error": "server_error", "error_description": "boom"}),
|
||||
credential_source="caller_supplied",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, UpstreamReportedFault)
|
||||
assert fault.code == "server_error"
|
||||
|
||||
|
||||
def test_temporarily_unavailable_is_upstream_fault():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(503, json_body={"error": "temporarily_unavailable"}),
|
||||
credential_source="gateway_stored",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, UpstreamReportedFault)
|
||||
assert fault.code == "temporarily_unavailable"
|
||||
|
||||
|
||||
def test_invalid_target_is_gateway_fault_even_with_caller_credentials():
|
||||
fault = classify_upstream_token_rejection(
|
||||
_response(400, json_body={"error": "invalid_target"}),
|
||||
credential_source="caller_supplied",
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, GatewayRejected)
|
||||
assert fault.code == "invalid_target"
|
||||
|
||||
|
||||
def test_dcr_server_error_code_is_not_blamed_on_caller():
|
||||
fault = classify_upstream_dcr_rejection(
|
||||
_response(500, json_body={"error": "server_error"}),
|
||||
log_context="srv",
|
||||
)
|
||||
assert isinstance(fault, UpstreamReportedFault)
|
||||
|
|
@ -0,0 +1,96 @@
|
|||
"""Rendering contract: status, wire code, and prose all derive from the fault tag, so a caller-fault
|
||||
code can never ship on a server-fault status and gateway-side faults never carry provider prose."""
|
||||
|
||||
import json
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.faults.render_oauth import (
|
||||
dcr_fault_detail,
|
||||
render_token_fault,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.types import (
|
||||
CallerRejected,
|
||||
GatewayRejected,
|
||||
UpstreamProtocolFault,
|
||||
UpstreamReportedFault,
|
||||
)
|
||||
|
||||
|
||||
def test_caller_rejected_renders_code_derived_status():
|
||||
response = render_token_fault(CallerRejected(code="invalid_grant", description="Code expired."))
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body) == {"error": "invalid_grant", "error_description": "Code expired."}
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
|
||||
|
||||
def test_caller_rejected_invalid_client_renders_401():
|
||||
response = render_token_fault(CallerRejected(code="invalid_client"))
|
||||
assert response.status_code == 401
|
||||
assert json.loads(response.body) == {"error": "invalid_client"}
|
||||
|
||||
|
||||
def test_caller_rejected_includes_error_uri_only_when_present():
|
||||
response = render_token_fault(
|
||||
CallerRejected(code="invalid_scope", description="bad scope", error_uri="https://idp.example.com/e")
|
||||
)
|
||||
assert json.loads(response.body) == {
|
||||
"error": "invalid_scope",
|
||||
"error_description": "bad scope",
|
||||
"error_uri": "https://idp.example.com/e",
|
||||
}
|
||||
|
||||
|
||||
def test_gateway_rejected_renders_502_with_gateway_prose():
|
||||
response = render_token_fault(GatewayRejected(code="invalid_client"))
|
||||
assert response.status_code == 502
|
||||
body = json.loads(response.body)
|
||||
assert body["error"] == "server_error"
|
||||
assert "invalid_client" in body["error_description"]
|
||||
assert "client_id and client_secret" in body["error_description"]
|
||||
|
||||
|
||||
def test_gateway_invalid_target_prose_names_resource_indicators():
|
||||
response = render_token_fault(GatewayRejected(code="invalid_target"))
|
||||
body = json.loads(response.body)
|
||||
assert response.status_code == 502
|
||||
assert "RFC 8707" in body["error_description"]
|
||||
|
||||
|
||||
def test_protocol_fault_renders_502_note():
|
||||
response = render_token_fault(UpstreamProtocolFault(note="upstream token endpoint returned HTTP 503"))
|
||||
assert response.status_code == 502
|
||||
assert json.loads(response.body) == {
|
||||
"error": "server_error",
|
||||
"error_description": "upstream token endpoint returned HTTP 503",
|
||||
}
|
||||
|
||||
|
||||
def test_dcr_caller_rejection_is_400_per_rfc7591_regardless_of_upstream_status():
|
||||
status_code, detail = dcr_fault_detail(CallerRejected(code="invalid_client_metadata", description="bad grant types"))
|
||||
assert status_code == 400
|
||||
assert detail == "invalid_client_metadata: bad grant types"
|
||||
|
||||
|
||||
def test_dcr_protocol_fault_is_502():
|
||||
status_code, detail = dcr_fault_detail(UpstreamProtocolFault(note="upstream registration failed with HTTP 500"))
|
||||
assert status_code == 502
|
||||
assert detail == "upstream registration failed with HTTP 500"
|
||||
|
||||
|
||||
def test_upstream_reported_server_error_renders_502_with_matching_code():
|
||||
response = render_token_fault(UpstreamReportedFault(code="server_error"))
|
||||
assert response.status_code == 502
|
||||
assert json.loads(response.body)["error"] == "server_error"
|
||||
|
||||
|
||||
def test_upstream_reported_temporarily_unavailable_renders_503_with_matching_code():
|
||||
response = render_token_fault(UpstreamReportedFault(code="temporarily_unavailable"))
|
||||
assert response.status_code == 503
|
||||
body = json.loads(response.body)
|
||||
assert body["error"] == "temporarily_unavailable"
|
||||
assert "retry" in body["error_description"]
|
||||
|
||||
|
||||
def test_dcr_upstream_reported_fault_maps_to_5xx():
|
||||
status_code, detail = dcr_fault_detail(UpstreamReportedFault(code="server_error"))
|
||||
assert status_code == 502
|
||||
assert "internal error" in detail
|
||||
|
|
@ -15,10 +15,14 @@ from pydantic import SecretStr
|
|||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import (
|
||||
BridgeEnvelopeAdmitted,
|
||||
BridgeEnvelopeInvalid,
|
||||
BridgeRefreshInvalid,
|
||||
BridgeRefreshOpened,
|
||||
NotBridgeEnvelope,
|
||||
build_bridge_refresh_token_response,
|
||||
build_bridge_token_response,
|
||||
envelope_keys_from_master_key,
|
||||
is_bridge_envelope_shaped,
|
||||
open_bridge_refresh_envelope,
|
||||
resolve_bridge_envelope,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
|
||||
|
|
@ -26,15 +30,17 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import
|
|||
EnvelopeIdentity,
|
||||
EnvelopeKeys,
|
||||
EnvelopeTooLarge,
|
||||
RefreshCredential,
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
key_hash_identity,
|
||||
mint_envelope,
|
||||
)
|
||||
|
||||
_NOW = datetime(2026, 7, 9, 12, 0, 0, tzinfo=timezone.utc)
|
||||
_MASTER_KEY = "sk-master-key-for-derivation-tests-0123456789"
|
||||
_ACCESS_TOKEN = "upstream-access-token-do-not-leak-8f14e45fceea"
|
||||
_IDENTITY = EnvelopeIdentity(server_id="srv-456", key_hash="hashed-key-123")
|
||||
_IDENTITY = key_hash_identity(server_id="srv-456", key_hash="hashed-key-123")
|
||||
_SERVER_ID = _IDENTITY.server_id
|
||||
|
||||
|
||||
|
|
@ -48,6 +54,76 @@ def _sealed_token(keys: EnvelopeKeys, now: datetime = _NOW, identity: EnvelopeId
|
|||
return sealed.token.get_secret_value()
|
||||
|
||||
|
||||
_UPSTREAM_REFRESH = "upstream-refresh-do-not-leak-9b2c"
|
||||
|
||||
|
||||
def _sealed_refresh(keys: EnvelopeKeys, now: datetime = _NOW, identity: EnvelopeIdentity = _IDENTITY) -> str:
|
||||
sealed = build_bridge_refresh_token_response(
|
||||
identity, RefreshCredential(refresh_token=SecretStr(_UPSTREAM_REFRESH)), keys, now
|
||||
)
|
||||
assert isinstance(sealed, SealedEnvelope)
|
||||
return sealed.token.get_secret_value()
|
||||
|
||||
|
||||
def test_open_bridge_refresh_envelope_round_trips_identity_and_refresh():
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
result = open_bridge_refresh_envelope(_sealed_refresh(keys), keys, _NOW, _SERVER_ID)
|
||||
assert isinstance(result, BridgeRefreshOpened)
|
||||
assert result.identity == _IDENTITY
|
||||
assert result.refresh.refresh_token.get_secret_value() == _UPSTREAM_REFRESH
|
||||
|
||||
|
||||
def test_open_bridge_refresh_envelope_strips_bearer_scheme():
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
result = open_bridge_refresh_envelope(f"Bearer {_sealed_refresh(keys)}", keys, _NOW, _SERVER_ID)
|
||||
assert isinstance(result, BridgeRefreshOpened)
|
||||
|
||||
|
||||
def test_open_bridge_refresh_envelope_rejects_wrong_server():
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
result = open_bridge_refresh_envelope(_sealed_refresh(keys), keys, _NOW, "a-different-server")
|
||||
assert isinstance(result, BridgeRefreshInvalid)
|
||||
|
||||
|
||||
def test_open_bridge_refresh_envelope_rejects_non_refresh_bearers():
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
# an access envelope is not a refresh envelope; a raw upstream refresh token is not one either
|
||||
assert isinstance(open_bridge_refresh_envelope(_sealed_token(keys), keys, _NOW, _SERVER_ID), BridgeRefreshInvalid)
|
||||
assert isinstance(open_bridge_refresh_envelope("raw-refresh-token", keys, _NOW, _SERVER_ID), BridgeRefreshInvalid)
|
||||
|
||||
|
||||
def test_open_bridge_refresh_envelope_rejects_under_wrong_master_key():
|
||||
minted = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
other = envelope_keys_from_master_key(_MASTER_KEY + "-rotated")
|
||||
result = open_bridge_refresh_envelope(_sealed_refresh(minted), other, _NOW, _SERVER_ID)
|
||||
assert isinstance(result, BridgeRefreshInvalid)
|
||||
|
||||
|
||||
def test_refresh_envelope_is_never_admitted_at_the_tool_call_edge():
|
||||
"""A refresh envelope must never authenticate a tool call. The admission edge engages the bridge arm
|
||||
for it (is_bridge_envelope_shaped is true for either envelope kind), and the consumer rejects it as
|
||||
BridgeEnvelopeInvalid, which admission fails closed (401): a refresh credential is only ever
|
||||
presented back to the token endpoint."""
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
refresh = _sealed_refresh(keys)
|
||||
assert is_bridge_envelope_shaped(refresh) is True
|
||||
assert is_bridge_envelope_shaped(f"Bearer {refresh}") is True
|
||||
result = resolve_bridge_envelope(refresh, keys, _NOW, _SERVER_ID)
|
||||
assert isinstance(result, BridgeEnvelopeInvalid)
|
||||
|
||||
|
||||
def test_refresh_jwt_wearing_the_access_prefix_is_rejected_at_the_edge():
|
||||
"""Belt-and-suspenders against a swapped wire prefix: a refresh JWT re-prefixed as an access envelope
|
||||
opens far enough to hit the signed kind claim, which rejects it, so admission fails closed rather
|
||||
than forwarding a refresh credential's contents upstream."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import REFRESH_ENVELOPE_PREFIX
|
||||
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
swapped = ENVELOPE_PREFIX + _sealed_refresh(keys).removeprefix(REFRESH_ENVELOPE_PREFIX)
|
||||
result = resolve_bridge_envelope(swapped, keys, _NOW, _SERVER_ID)
|
||||
assert isinstance(result, BridgeEnvelopeInvalid)
|
||||
|
||||
|
||||
def test_key_derivation_is_deterministic():
|
||||
assert envelope_keys_from_master_key(_MASTER_KEY) == envelope_keys_from_master_key(_MASTER_KEY)
|
||||
|
||||
|
|
@ -138,7 +214,7 @@ def test_resolve_envelope_minted_for_another_server_is_invalid():
|
|||
captured or misrouted envelope cannot forward one server's upstream credential to
|
||||
another. The valid access token stays sealed; the mismatch alone fails the resolve."""
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
other_server_identity = EnvelopeIdentity(server_id="srv-OTHER", key_hash=_IDENTITY.key_hash)
|
||||
other_server_identity = key_hash_identity(server_id="srv-OTHER", key_hash=_IDENTITY.subject)
|
||||
token = _sealed_token(keys, identity=other_server_identity)
|
||||
result = resolve_bridge_envelope(token, keys, _NOW, _SERVER_ID)
|
||||
assert isinstance(result, BridgeEnvelopeInvalid)
|
||||
|
|
@ -155,7 +231,7 @@ def test_resolve_non_ascii_server_id_stays_total_and_does_not_raise():
|
|||
unicode server_id); it stays total and returns a typed result. A matching non-ASCII id admits,
|
||||
a mismatching one is BridgeEnvelopeInvalid, and neither raises."""
|
||||
keys = envelope_keys_from_master_key(_MASTER_KEY)
|
||||
unicode_identity = EnvelopeIdentity(server_id="srv-café", key_hash=_IDENTITY.key_hash)
|
||||
unicode_identity = key_hash_identity(server_id="srv-café", key_hash=_IDENTITY.subject)
|
||||
token = _sealed_token(keys, identity=unicode_identity)
|
||||
assert isinstance(resolve_bridge_envelope(token, keys, _NOW, "srv-café"), BridgeEnvelopeAdmitted)
|
||||
assert isinstance(resolve_bridge_envelope(token, keys, _NOW, "srv-cafe"), BridgeEnvelopeInvalid)
|
||||
|
|
|
|||
|
|
@ -24,6 +24,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import
|
|||
ENVELOPE_PREFIX,
|
||||
MAX_ENVELOPE_BYTES,
|
||||
MAX_ENVELOPE_TTL_SECONDS,
|
||||
MAX_REFRESH_ENVELOPE_TTL_SECONDS,
|
||||
REFRESH_ENVELOPE_PREFIX,
|
||||
BadSignature,
|
||||
DecryptFailed,
|
||||
EnvelopeIdentity,
|
||||
|
|
@ -33,11 +35,18 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import
|
|||
MalformedPayload,
|
||||
NotAnEnvelope,
|
||||
OpenedEnvelope,
|
||||
OpenedRefreshEnvelope,
|
||||
RefreshCredential,
|
||||
SealedEnvelope,
|
||||
UpstreamTokenGrant,
|
||||
is_envelope,
|
||||
is_refresh_envelope,
|
||||
key_hash_identity,
|
||||
mint_envelope,
|
||||
mint_refresh_envelope,
|
||||
open_envelope,
|
||||
open_refresh_envelope,
|
||||
user_identity,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value, encrypt_value
|
||||
|
||||
|
|
@ -51,7 +60,7 @@ _WRONG_SIGNING = EnvelopeKeys(signing_key=SecretStr(_OTHER_SIGNING_KEY), encrypt
|
|||
_WRONG_ENCRYPTION = EnvelopeKeys(signing_key=SecretStr(_SIGNING_KEY), encryption_key=SecretStr(_OTHER_ENCRYPTION_KEY))
|
||||
_ACCESS_TOKEN = "upstream-access-token-do-not-leak-8f14e45fceea"
|
||||
_REFRESH_TOKEN = "upstream-refresh-token-do-not-leak-1d0aa4b7"
|
||||
_IDENTITY = EnvelopeIdentity(server_id="srv-456", key_hash="hashed-key-123")
|
||||
_IDENTITY = key_hash_identity(server_id="srv-456", key_hash="hashed-key-123")
|
||||
|
||||
|
||||
def _full_grant() -> UpstreamTokenGrant:
|
||||
|
|
@ -137,17 +146,99 @@ def test_minimal_grant_round_trips_without_none_leakage_into_claims():
|
|||
def test_claim_layout_and_no_plaintext_token_in_envelope():
|
||||
token = _sealed_token(_full_grant())
|
||||
claims = _unverified_claims(token)
|
||||
assert set(claims) == {"iss", "iat", "exp", "server_id", "key_hash", "grant"}
|
||||
assert set(claims) == {"iss", "iat", "exp", "kind", "server_id", "subject_type", "subject", "grant"}
|
||||
assert claims["iss"] == ENVELOPE_ISSUER
|
||||
assert claims["iat"] == int(_NOW.timestamp())
|
||||
assert claims["exp"] == int(_NOW.timestamp()) + 600
|
||||
assert claims["kind"] == "access"
|
||||
assert claims["server_id"] == "srv-456"
|
||||
assert claims["key_hash"] == "hashed-key-123"
|
||||
assert claims["subject_type"] == "key_hash"
|
||||
assert claims["subject"] == "hashed-key-123"
|
||||
assert _ACCESS_TOKEN not in token
|
||||
assert _ACCESS_TOKEN not in json.dumps(claims)
|
||||
assert _REFRESH_TOKEN not in json.dumps(claims)
|
||||
|
||||
|
||||
def _refresh_credential() -> RefreshCredential:
|
||||
return RefreshCredential(refresh_token=SecretStr(_REFRESH_TOKEN), scope="read:tools", expires_in=None)
|
||||
|
||||
|
||||
def _sealed_refresh_token(refresh: RefreshCredential | None = None, keys: EnvelopeKeys = _KEYS) -> str:
|
||||
sealed = mint_refresh_envelope(_IDENTITY, refresh or _refresh_credential(), keys, _NOW)
|
||||
assert isinstance(sealed, SealedEnvelope)
|
||||
return sealed.token.get_secret_value()
|
||||
|
||||
|
||||
def test_refresh_envelope_round_trips_identity_and_refresh_token():
|
||||
token = _sealed_refresh_token()
|
||||
assert is_refresh_envelope(token)
|
||||
assert not is_envelope(token)
|
||||
opened = open_refresh_envelope(token, _KEYS, _NOW)
|
||||
assert isinstance(opened, OpenedRefreshEnvelope)
|
||||
assert opened.identity == _IDENTITY
|
||||
assert opened.refresh.refresh_token.get_secret_value() == _REFRESH_TOKEN
|
||||
assert opened.refresh.scope == "read:tools"
|
||||
|
||||
|
||||
def test_refresh_envelope_ttl_is_min_of_upstream_refresh_lifetime_and_cap():
|
||||
short = mint_refresh_envelope(
|
||||
_IDENTITY, RefreshCredential(refresh_token=SecretStr("r"), expires_in=120), _KEYS, _NOW
|
||||
)
|
||||
assert isinstance(short, SealedEnvelope)
|
||||
assert short.expires_at == _NOW + timedelta(seconds=120)
|
||||
capped = mint_refresh_envelope(
|
||||
_IDENTITY,
|
||||
RefreshCredential(refresh_token=SecretStr("r"), expires_in=MAX_REFRESH_ENVELOPE_TTL_SECONDS + 86400),
|
||||
_KEYS,
|
||||
_NOW,
|
||||
)
|
||||
assert isinstance(capped, SealedEnvelope)
|
||||
assert capped.expires_at == _NOW + timedelta(seconds=MAX_REFRESH_ENVELOPE_TTL_SECONDS)
|
||||
default = mint_refresh_envelope(_IDENTITY, RefreshCredential(refresh_token=SecretStr("r")), _KEYS, _NOW)
|
||||
assert isinstance(default, SealedEnvelope)
|
||||
assert default.expires_at == _NOW + timedelta(seconds=MAX_REFRESH_ENVELOPE_TTL_SECONDS)
|
||||
|
||||
|
||||
def test_access_and_refresh_envelopes_do_not_cross_open():
|
||||
access = _sealed_token(_full_grant())
|
||||
refresh = _sealed_refresh_token()
|
||||
# each opener rejects the other kind's prefix outright
|
||||
assert isinstance(open_refresh_envelope(access, _KEYS, _NOW), NotAnEnvelope)
|
||||
assert isinstance(open_envelope(refresh, _KEYS, _NOW), NotAnEnvelope)
|
||||
|
||||
|
||||
def test_prefix_swap_is_rejected_by_the_signed_kind_claim():
|
||||
# the wire prefix is not signed, so swap it; the signed kind claim must still reject the cross-use
|
||||
refresh = _sealed_refresh_token()
|
||||
swapped_to_access = ENVELOPE_PREFIX + refresh.removeprefix(REFRESH_ENVELOPE_PREFIX)
|
||||
assert isinstance(open_envelope(swapped_to_access, _KEYS, _NOW), MalformedPayload)
|
||||
access = _sealed_token(_full_grant())
|
||||
swapped_to_refresh = REFRESH_ENVELOPE_PREFIX + access.removeprefix(ENVELOPE_PREFIX)
|
||||
assert isinstance(open_refresh_envelope(swapped_to_refresh, _KEYS, _NOW), MalformedPayload)
|
||||
|
||||
|
||||
def test_refresh_envelope_total_over_hostile_input():
|
||||
token = _sealed_refresh_token()
|
||||
# expired against the injected clock
|
||||
assert isinstance(
|
||||
open_refresh_envelope(token, _KEYS, _NOW + timedelta(seconds=MAX_REFRESH_ENVELOPE_TTL_SECONDS)), Expired
|
||||
)
|
||||
# wrong signing key
|
||||
assert isinstance(open_refresh_envelope(token, _WRONG_SIGNING, _NOW), BadSignature)
|
||||
# right signature, wrong encryption key
|
||||
assert isinstance(open_refresh_envelope(token, _WRONG_ENCRYPTION, _NOW), DecryptFailed)
|
||||
# not an envelope at all
|
||||
assert isinstance(open_refresh_envelope("raw-upstream-refresh-token", _KEYS, _NOW), NotAnEnvelope)
|
||||
|
||||
|
||||
def test_refresh_envelope_never_leaks_the_refresh_token_in_plaintext():
|
||||
token = _sealed_refresh_token()
|
||||
assert _REFRESH_TOKEN not in token
|
||||
claims = jwt.decode(token.removeprefix(REFRESH_ENVELOPE_PREFIX), options={"verify_signature": False})
|
||||
assert claims["kind"] == "refresh"
|
||||
assert _REFRESH_TOKEN not in json.dumps(claims)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"expires_in, expected_ttl",
|
||||
[
|
||||
|
|
@ -226,11 +317,11 @@ def test_wrong_issuer_is_malformed_payload():
|
|||
|
||||
def test_missing_identity_claim_is_malformed_payload():
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _forge({key: value for key, value in claims.items() if key != "key_hash"})
|
||||
forged = _forge({key: value for key, value in claims.items() if key != "subject"})
|
||||
assert isinstance(open_envelope(forged, _KEYS, _NOW), MalformedPayload)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("identity_claim", ["server_id", "key_hash"])
|
||||
@pytest.mark.parametrize("identity_claim", ["server_id", "subject"])
|
||||
def test_signed_empty_identity_claim_is_malformed_payload_not_a_raise(identity_claim):
|
||||
claims = _unverified_claims(_sealed_token(_full_grant()))
|
||||
forged = _forge({**claims, identity_claim: ""})
|
||||
|
|
@ -463,9 +554,11 @@ def test_non_positive_expires_in_is_rejected_at_construction_without_leaking():
|
|||
|
||||
def test_empty_identity_and_key_fields_are_rejected_at_construction():
|
||||
with pytest.raises(ValidationError):
|
||||
EnvelopeIdentity(server_id="", key_hash="hashed-key-123")
|
||||
EnvelopeIdentity(server_id="", subject_type="key_hash", subject="hashed-key-123")
|
||||
with pytest.raises(ValidationError):
|
||||
EnvelopeIdentity(server_id="srv-456", key_hash="")
|
||||
EnvelopeIdentity(server_id="srv-456", subject_type="key_hash", subject="")
|
||||
with pytest.raises(ValidationError):
|
||||
EnvelopeIdentity(server_id="srv-456", subject_type="not-a-subject-type", subject="x")
|
||||
with pytest.raises(ValidationError):
|
||||
EnvelopeKeys(signing_key=SecretStr(""), encryption_key=SecretStr(_ENCRYPTION_KEY))
|
||||
with pytest.raises(ValidationError):
|
||||
|
|
@ -474,6 +567,20 @@ def test_empty_identity_and_key_fields_are_rejected_at_construction():
|
|||
UpstreamTokenGrant(access_token=SecretStr(""), token_type="Bearer")
|
||||
|
||||
|
||||
def test_user_subject_identity_round_trips():
|
||||
"""The user_id subject variant seals and opens with its discriminator intact, so the edge can
|
||||
tell an interactively-minted (user) envelope from a scripted (key_hash) one and reload the right
|
||||
kind of record."""
|
||||
identity = user_identity(server_id="srv-456", user_id="user-42")
|
||||
sealed = mint_envelope(identity, _full_grant(), _KEYS, _NOW)
|
||||
assert isinstance(sealed, SealedEnvelope)
|
||||
opened = open_envelope(sealed.token.get_secret_value(), _KEYS, _NOW)
|
||||
assert isinstance(opened, OpenedEnvelope)
|
||||
assert opened.identity.server_id == "srv-456"
|
||||
assert opened.identity.subject_type == "user_id"
|
||||
assert opened.identity.subject == "user-42"
|
||||
|
||||
|
||||
def test_public_models_are_frozen():
|
||||
sealed = mint_envelope(_IDENTITY, _full_grant(), _KEYS, _NOW)
|
||||
assert isinstance(sealed, SealedEnvelope)
|
||||
|
|
@ -484,4 +591,4 @@ def test_public_models_are_frozen():
|
|||
with pytest.raises(ValidationError):
|
||||
opened.grant = _minimal_grant()
|
||||
with pytest.raises(ValidationError):
|
||||
_IDENTITY.key_hash = "someone-elses-hash"
|
||||
_IDENTITY.subject = "someone-elses-hash"
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue