mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
chore: merge litellm_internal_staging into litellm_add_milvus_grpc_transport
This commit is contained in:
commit
6af361c4bb
166 changed files with 9221 additions and 1012 deletions
31
.github/workflows/test-terraform-modules.yml
vendored
31
.github/workflows/test-terraform-modules.yml
vendored
|
|
@ -4,6 +4,7 @@ on:
|
|||
push:
|
||||
paths:
|
||||
- "terraform/litellm/aws/**"
|
||||
- "terraform/litellm/gcp/**"
|
||||
- ".github/workflows/test-terraform-modules.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -13,6 +14,7 @@ on:
|
|||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/litellm/aws/**"
|
||||
- "terraform/litellm/gcp/**"
|
||||
- ".github/workflows/test-terraform-modules.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -52,3 +54,32 @@ jobs:
|
|||
# Plan-only, mock_provider-backed: no AWS credentials, no API calls.
|
||||
- name: test
|
||||
run: terraform test
|
||||
|
||||
gcp-module:
|
||||
name: fmt, validate, test (gcp)
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
defaults:
|
||||
run:
|
||||
working-directory: terraform/litellm/gcp
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- uses: hashicorp/setup-terraform@b9cd54a3c349d3f38e8881555d616ced269862dd # v3.1.2
|
||||
with:
|
||||
terraform_version: 1.13.3
|
||||
terraform_wrapper: false
|
||||
|
||||
- name: fmt
|
||||
run: terraform fmt -recursive -check -diff
|
||||
|
||||
- name: init
|
||||
run: terraform init -backend=false -input=false
|
||||
|
||||
- name: validate
|
||||
run: terraform validate
|
||||
|
||||
- name: test
|
||||
run: terraform test
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 13424
|
||||
"limit": 13422
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2148
|
||||
"limit": 2137
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 3364
|
||||
"limit": 3363
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5570
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15252
|
||||
"limit": 15246
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 42776
|
||||
"limit": 42551
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38207
|
||||
"limit": 38191
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19578
|
||||
"limit": 19576
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29817
|
||||
"limit": 29814
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 110
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 4
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 811
|
||||
"limit": 810
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -433,9 +433,9 @@ _default_detect_secrets_config = {
|
|||
"name": "ZendeskSecretKeyDetector",
|
||||
"path": _custom_plugins_path + "/zendesk_secret_key.py",
|
||||
},
|
||||
{"name": "Base64HighEntropyString", "limit": 3.0},
|
||||
{"name": "Base64HighEntropyString", "limit": 4.5},
|
||||
{"name": "HexHighEntropyString", "limit": 3.0},
|
||||
]
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -466,16 +466,19 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
|
|||
|
||||
os.remove(temp_file.name)
|
||||
|
||||
detected_secrets = []
|
||||
for file in secrets.files:
|
||||
for found_secret in secrets[file]:
|
||||
if found_secret.secret_value is None:
|
||||
continue
|
||||
detected_secrets.append(
|
||||
{"type": found_secret.type, "value": found_secret.secret_value}
|
||||
)
|
||||
|
||||
return detected_secrets
|
||||
return [
|
||||
{"type": found_secret.type, "value": found_secret.secret_value}
|
||||
for file in sorted(secrets.files)
|
||||
for found_secret in sorted(
|
||||
secrets[file],
|
||||
key=lambda secret: (
|
||||
-len(secret.secret_value or ""),
|
||||
secret.type,
|
||||
secret.secret_value or "",
|
||||
),
|
||||
)
|
||||
if found_secret.secret_value is not None
|
||||
]
|
||||
|
||||
def redact_text(self, text: str, source: str = "message") -> str:
|
||||
"""Replace every detected secret in ``text`` with ``[REDACTED]`` and
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ This plugin searches for OpenAI API Keys.
|
|||
"""
|
||||
|
||||
import re
|
||||
from collections.abc import Generator
|
||||
|
||||
from detect_secrets.plugins.base import RegexBasedDetector
|
||||
|
||||
|
|
@ -16,4 +17,16 @@ class OpenAIApiKeyDetector(RegexBasedDetector):
|
|||
|
||||
@property
|
||||
def denylist(self) -> list[re.Pattern]:
|
||||
return [re.compile(r"""(sk-[a-zA-Z0-9]{5,})""")]
|
||||
return [
|
||||
re.compile(
|
||||
r"((?:(?<![a-zA-Z0-9])|(?<=%[0-9A-Fa-f]{2}))"
|
||||
r"sk[-_]"
|
||||
r"[a-zA-Z0-9_-]{5,}"
|
||||
r"(?![a-zA-Z0-9_-]))"
|
||||
)
|
||||
]
|
||||
|
||||
def analyze_string(self, string: str) -> Generator[str, None, None]:
|
||||
# the digit check lives outside the regex: a lookahead re-scans the token
|
||||
# from every `sk` inside it, which is quadratic on `-sk-sk-sk-...` input
|
||||
yield from (match for match in super().analyze_string(string) if re.search(r"[0-9]", match))
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.64"
|
||||
version = "0.1.65"
|
||||
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.64"
|
||||
version = "0.1.65"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.93"
|
||||
version = "0.4.94"
|
||||
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.93"
|
||||
version = "0.4.94"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
|
|||
"router_general_settings",
|
||||
"ignore_invalid_deployments",
|
||||
"fallback_access_check",
|
||||
"heuristic_v2_router_limit",
|
||||
"auto_router_capability_limit",
|
||||
}
|
||||
)
|
||||
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
|
|
@ -89,6 +89,7 @@ LITELLM_MAX_STREAMING_DURATION_SECONDS: Final = (
|
|||
# Data URIs exceeding this are replaced with a size placeholder.
|
||||
# Set to 0 to disable truncation.
|
||||
MAX_BASE64_LENGTH_FOR_LOGGING: Final = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 64))
|
||||
BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS: Final = 256 * 1024
|
||||
REDACTED_BY_LITELLM: Final = "redacted-by-litellm"
|
||||
# in-memory stand-in handed to provider converters for redacted arguments; never stored
|
||||
REDACTED_TOOL_CALL_ARGUMENTS_PLACEHOLDER: Final = "{}"
|
||||
|
|
@ -215,6 +216,9 @@ MAX_CALLBACKS: Final = get_env_int("LITELLM_MAX_CALLBACKS", 100)
|
|||
# so the deployment-level hook does not re-run them for the same request
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY: Final = "_pre_call_executed_guardrails"
|
||||
|
||||
# Attribute stamped on log_guardrail_information wrappers so __init_subclass__ does not wrap them again
|
||||
LOGS_GUARDRAIL_INFORMATION_MARKER: Final = "_litellm_logs_guardrail_information"
|
||||
|
||||
# Generic fallback for unknown models
|
||||
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET: Final = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128)
|
||||
|
|
|
|||
|
|
@ -97,7 +97,11 @@ class CloudZeroStreamer:
|
|||
continue
|
||||
|
||||
# Convert lists back to DataFrames
|
||||
return {date_key: pl.DataFrame(records) for date_key, records in daily_batches.items() if records}
|
||||
return {
|
||||
date_key: pl.DataFrame(records, infer_schema_length=None)
|
||||
for date_key, records in daily_batches.items()
|
||||
if records
|
||||
}
|
||||
|
||||
def _parse_and_convert_timestamp(self, timestamp_str: str) -> datetime:
|
||||
"""Parse timestamp string and convert to UTC."""
|
||||
|
|
|
|||
|
|
@ -95,7 +95,7 @@ class CBFTransformer:
|
|||
if len(cbf_data) > 0:
|
||||
console.print(f"[green]✓ Successfully transformed {len(cbf_data):,} records[/green]")
|
||||
|
||||
return pl.DataFrame(cbf_data)
|
||||
return pl.DataFrame(cbf_data, infer_schema_length=None)
|
||||
|
||||
def _create_cbf_record(self, row: dict[str, object]) -> CBFRecord:
|
||||
"""Create a single CBF record from LiteLLM daily spend row."""
|
||||
|
|
|
|||
|
|
@ -46,6 +46,7 @@ dc: Final = DualCache()
|
|||
|
||||
from litellm.constants import (
|
||||
GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS,
|
||||
LOGS_GUARDRAIL_INFORMATION_MARKER,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
)
|
||||
from litellm.exceptions import (
|
||||
|
|
@ -151,6 +152,13 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
records_own_guardrail_information: ClassVar[bool] = False
|
||||
|
||||
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
|
||||
super().__init_subclass__(**kwargs)
|
||||
own_apply_guardrail: Final = cls.__dict__.get("apply_guardrail")
|
||||
if own_apply_guardrail is None or LOGS_GUARDRAIL_INFORMATION_MARKER in vars(own_apply_guardrail):
|
||||
return
|
||||
cls.apply_guardrail = log_guardrail_information(own_apply_guardrail)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: str | None = None,
|
||||
|
|
@ -940,6 +948,23 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
return False
|
||||
|
||||
def _suppressed_by_auto_router_compression(self) -> bool:
|
||||
"""True when an auto router's own compression policy suppresses this guardrail.
|
||||
|
||||
Reads request-scoped state set by `arm_pre_call`, never request metadata. The
|
||||
caller controls metadata, and metadata reaches spend logs the caller can read,
|
||||
so a suppression list carried there would be one a request could replay to
|
||||
switch off a PII or content-filter guardrail for itself.
|
||||
"""
|
||||
name: Final = self.guardrail_name
|
||||
if not name:
|
||||
return False
|
||||
from litellm.proxy.guardrails.auto_router_compression import (
|
||||
suppressed_compression_guardrails,
|
||||
)
|
||||
|
||||
return name in suppressed_compression_guardrails()
|
||||
|
||||
def should_run_guardrail(
|
||||
self,
|
||||
data,
|
||||
|
|
@ -948,6 +973,9 @@ class CustomGuardrail(CustomLogger):
|
|||
"""
|
||||
Returns True if the guardrail should be run on the event_type
|
||||
"""
|
||||
if self._suppressed_by_auto_router_compression():
|
||||
return False
|
||||
|
||||
requested_guardrails: Final = self.get_guardrail_from_metadata(data)
|
||||
disable_global_guardrail: Final = self.get_disable_global_guardrail(data)
|
||||
opted_out_global_guardrails: Final = self.get_opted_out_global_guardrails_from_metadata(data)
|
||||
|
|
@ -1559,4 +1587,5 @@ def log_guardrail_information(func):
|
|||
return async_wrapper(*args, **kwargs)
|
||||
return sync_wrapper(*args, **kwargs)
|
||||
|
||||
vars(wrapper)[LOGS_GUARDRAIL_INFORMATION_MARKER] = True # rebind-ok: stamps the wrapper this call just built
|
||||
return wrapper
|
||||
|
|
|
|||
|
|
@ -61,9 +61,9 @@ _MAX_CONCURRENT_SHADOW_TASKS: Final = 16
|
|||
_MAX_JUDGE_RESPONSE_CHARS: Final = 8_000
|
||||
_MAX_JUDGE_PROMPT_CHARS: Final = 24_000
|
||||
|
||||
# The judge answers with a small JSON object; a tighter budget truncates the JSON
|
||||
# mid-object and the attempt is lost to an error row.
|
||||
JUDGE_MAX_OUTPUT_TOKENS: Final = 1500
|
||||
# Covers the judge's reasoning tokens as well as its small JSON answer: a judge deployment
|
||||
# carrying an elevated reasoning_effort spends a tight cap before it ever answers.
|
||||
JUDGE_MAX_OUTPUT_TOKENS: Final = 4096
|
||||
|
||||
_MAX_ERROR_CHARS: Final = 500
|
||||
|
||||
|
|
@ -419,6 +419,20 @@ def _failure_detail(e: BaseException) -> str:
|
|||
return f"{type(e).__name__}{location}: {e}"
|
||||
|
||||
|
||||
def _judge_reply_shape(response: object) -> str:
|
||||
"""How an unparseable judge reply was shaped. The parser's own message cannot separate a
|
||||
judge that answered with nothing from one truncated mid-object, and those want opposite
|
||||
fixes. Shape only, never the reply text: the judge quotes the sampled turns it compares,
|
||||
and no attempt row carries sampled content today."""
|
||||
read: Final = _chat_message_reader(response)
|
||||
if read is None:
|
||||
return "unreadable judge reply"
|
||||
content: Final = read("content")
|
||||
served: Final = str(_field_reader(response)("model") or "unknown")
|
||||
body: Final = f"{len(str(content))} chars" if content else "no content"
|
||||
return f"finish_reason={_chat_finish_reason(response)}, content={body}, model={served}"
|
||||
|
||||
|
||||
def _call_cost(response: object) -> float:
|
||||
"""Price one eval-arm call with the figure the spend pipeline bills: the router client
|
||||
stamps _hidden_params.response_cost from the deployment's own pricing, which the public
|
||||
|
|
@ -1266,7 +1280,9 @@ class ShadowEvalLogger(CustomLogger):
|
|||
verdict: Final = PairwiseVerdict.model_validate(parse_json_verdict(raw))
|
||||
except Exception as e: # noqa: BLE001 # malformed verdicts become error rows
|
||||
verbose_logger.debug("shadow_eval: unparseable judge verdict: %s", e)
|
||||
return _CallFailure(f"unparseable judge verdict: {e}", cost=_call_cost(response))
|
||||
return _CallFailure(
|
||||
f"unparseable judge verdict: {e}; {_judge_reply_shape(response)}", cost=_call_cost(response)
|
||||
)
|
||||
return _JudgeVerdict(
|
||||
preference=_unmask_preference(verdict.preference, real_is_a),
|
||||
confidence=max(0.0, min(1.0, verdict.confidence)),
|
||||
|
|
|
|||
|
|
@ -78,7 +78,10 @@ from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
|||
from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import (
|
||||
InteractionsUsageObjectTransformation,
|
||||
)
|
||||
from litellm.litellm_core_utils.logging_utils import truncate_base64_in_messages
|
||||
from litellm.litellm_core_utils.logging_utils import (
|
||||
truncate_base64_in_messages,
|
||||
truncate_base64_in_messages_async,
|
||||
)
|
||||
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,
|
||||
|
|
@ -538,6 +541,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.standard_built_in_tools_params: StandardBuiltInToolsParams = (
|
||||
self.initialize_standard_built_in_tools_params(kwargs)
|
||||
)
|
||||
self.truncated_messages_for_logging: str | list | dict | None = None # mutable-ok: logged messages shape
|
||||
## TIME TO FIRST TOKEN LOGGING ##
|
||||
self.completion_start_time: datetime.datetime | None = None
|
||||
self._llm_caching_handler: LLMCachingHandler | None = None
|
||||
|
|
@ -1820,6 +1824,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
and litellm_params.get(CallTypes.aanthropic_messages.value, False) is not True
|
||||
and litellm_params.get(CallTypes.agenerate_content.value, False) is not True
|
||||
and litellm_params.get(CallTypes.agenerate_content_stream.value, False) is not True
|
||||
and litellm_params.get(CallTypes.arealtime.value, False) is not True
|
||||
)
|
||||
|
||||
def _is_assembled_stream_success(self, result=None) -> bool:
|
||||
|
|
@ -1913,7 +1918,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
two paths cannot mutate it at the same time. ``prefer_async_handlers`` only
|
||||
bypasses the sync-SDK-only shortcut (e.g. ``async for`` on a stream from
|
||||
``completion()``); legacy string callbacks still run via
|
||||
``executor.submit(failure_handler)`` when configured.
|
||||
``executor.submit(failure_handler)`` when configured, and still get submitted
|
||||
when the awaiting task is cancelled (e.g. the event loop shuts down right after
|
||||
the request failed).
|
||||
"""
|
||||
litellm_params: Final = self.model_call_details.get("litellm_params", {}) or {}
|
||||
sync_sdk: Final = self._is_sync_litellm_request(litellm_params)
|
||||
|
|
@ -1922,12 +1929,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.failure_handler(exception, traceback_exception)
|
||||
return
|
||||
|
||||
await self.async_failure_handler(exception, traceback_exception)
|
||||
|
||||
if not self._should_run_sync_failure_callbacks_for_async_calls():
|
||||
return
|
||||
|
||||
executor.submit(self.failure_handler, exception, traceback_exception)
|
||||
try:
|
||||
await self.async_failure_handler(exception, traceback_exception)
|
||||
finally:
|
||||
if self._should_run_sync_failure_callbacks_for_async_calls():
|
||||
executor.submit(self.failure_handler, exception, traceback_exception)
|
||||
|
||||
def should_run_logging(
|
||||
self,
|
||||
|
|
@ -2932,6 +2938,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result._hidden_params["batch_failed_requests"] = batch_result.failed_requests # pyright: ignore[reportPrivateUsage] # rebind-ok: same pattern as above
|
||||
result.usage = batch_result.usage
|
||||
|
||||
self.truncated_messages_for_logging = await truncate_base64_in_messages_async(
|
||||
StandardLoggingPayloadSetup.append_system_prompt_messages(
|
||||
kwargs=self.model_call_details, messages=self.model_call_details.get("messages")
|
||||
)
|
||||
)
|
||||
start_time, end_time, result = self._success_handler_helper_fn(
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
@ -3224,8 +3235,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details = {}
|
||||
|
||||
if (
|
||||
self.model_call_details.get("log_event_type") == "failed_api_call"
|
||||
and self.model_call_details.get("exception") is exception
|
||||
self.model_call_details.get("exception") is exception
|
||||
and self.model_call_details.get("standard_logging_object") is not None
|
||||
):
|
||||
return start_time, self.model_call_details["end_time"]
|
||||
|
|
@ -6201,9 +6211,13 @@ def get_standard_logging_object_payload(
|
|||
model_id=_model_id,
|
||||
requester_ip_address=clean_metadata.get("requester_ip_address", None),
|
||||
user_agent=clean_metadata.get("user_agent", None),
|
||||
messages=truncate_base64_in_messages(
|
||||
StandardLoggingPayloadSetup.append_system_prompt_messages(
|
||||
kwargs=kwargs, messages=kwargs.get("messages")
|
||||
messages=(
|
||||
logging_obj.truncated_messages_for_logging
|
||||
if logging_obj.truncated_messages_for_logging is not None
|
||||
else truncate_base64_in_messages(
|
||||
StandardLoggingPayloadSetup.append_system_prompt_messages(
|
||||
kwargs=kwargs, messages=kwargs.get("messages")
|
||||
)
|
||||
)
|
||||
),
|
||||
response=final_response_obj,
|
||||
|
|
|
|||
|
|
@ -3,12 +3,15 @@ import functools
|
|||
import inspect
|
||||
import re
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MAX_BASE64_LENGTH_FOR_LOGGING
|
||||
from litellm.constants import (
|
||||
BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS,
|
||||
MAX_BASE64_LENGTH_FOR_LOGGING,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -141,6 +144,39 @@ def truncate_base64_in_messages(
|
|||
return messages
|
||||
|
||||
|
||||
_StringTree = str | Sequence["_StringTree"] | Mapping[str, "_StringTree"] | None
|
||||
|
||||
|
||||
def _iter_string_leaves(value: _StringTree) -> Iterator[str]:
|
||||
stack: Final[list[_StringTree]] = [value] # mutable-ok: explicit stack, recursive functions are banned in litellm/
|
||||
while stack:
|
||||
match stack.pop():
|
||||
case str() as text:
|
||||
yield text
|
||||
case Mapping() as mapping:
|
||||
stack.extend(mapping.values())
|
||||
case Sequence() as items:
|
||||
stack.extend(items)
|
||||
case None:
|
||||
pass
|
||||
|
||||
|
||||
async def truncate_base64_in_messages_async(
|
||||
messages: str | list | dict | None, # mutable-ok: same contract as truncate_base64_in_messages
|
||||
) -> str | list | dict | None: # mutable-ok: same contract as truncate_base64_in_messages
|
||||
"""
|
||||
Same result as truncate_base64_in_messages, but payloads whose string content
|
||||
reaches BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS are scanned in a worker
|
||||
thread so the regex pass over multi-MB base64 images does not block the event loop.
|
||||
"""
|
||||
if messages is None or MAX_BASE64_LENGTH_FOR_LOGGING <= 0:
|
||||
return messages
|
||||
total_chars: Final = sum(len(leaf) for leaf in _iter_string_leaves(messages))
|
||||
if total_chars < BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS:
|
||||
return truncate_base64_in_messages(messages)
|
||||
return await asyncio.to_thread(truncate_base64_in_messages, messages)
|
||||
|
||||
|
||||
# Global service logger instance to avoid recreating it
|
||||
_service_logger = None
|
||||
|
||||
|
|
|
|||
|
|
@ -29,3 +29,11 @@ def websocket_close_reason(message: str, fallback: str) -> str:
|
|||
if len(encoded) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES:
|
||||
return message
|
||||
return encoded[:WEBSOCKET_CLOSE_REASON_MAX_BYTES].decode("utf-8", errors="ignore")
|
||||
|
||||
|
||||
def client_close_code(upstream_code: int) -> int:
|
||||
from websockets.frames import EXTERNAL_CLOSE_CODES, CloseCode
|
||||
|
||||
if upstream_code in EXTERNAL_CLOSE_CODES or 3000 <= upstream_code < 5000:
|
||||
return upstream_code
|
||||
return int(CloseCode.INTERNAL_ERROR)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,15 @@
|
|||
import asyncio
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, cast
|
||||
import traceback
|
||||
from collections.abc import Coroutine, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, TypedDict, cast
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -19,9 +22,11 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.realtime import ALL_DELTA_TYPES
|
||||
|
||||
from .litellm_logging import Logging as LiteLLMLogging
|
||||
from .realtime_errors import client_close_code, realtime_error_event, websocket_close_reason
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
|
|
@ -30,8 +35,30 @@ else:
|
|||
CLIENT_CONNECTION_CLASS = Any
|
||||
|
||||
|
||||
class _ClientWebSocketExceptions(Protocol):
|
||||
ConnectionClosed: type[Exception]
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BackendClose:
|
||||
code: int
|
||||
reason: str
|
||||
|
||||
@property
|
||||
def message(self) -> str:
|
||||
if not self.reason:
|
||||
return f"upstream websocket closed with code {self.code}"
|
||||
return f"upstream websocket closed with code {self.code}: {self.reason}"
|
||||
|
||||
|
||||
class ClientLoopExit(Enum):
|
||||
CLIENT_DISCONNECTED = auto()
|
||||
BACKEND_CLOSED = auto()
|
||||
|
||||
|
||||
def backend_close_from(error: "ConnectionClosed") -> BackendClose:
|
||||
if error.rcvd is None:
|
||||
return BackendClose(code=1006, reason=str(error))
|
||||
return BackendClose(code=error.rcvd.code, reason=error.rcvd.reason)
|
||||
|
||||
|
||||
class _ASGIScope(TypedDict, total=False):
|
||||
|
|
@ -69,10 +96,13 @@ class _ScopedWebSocket(Protocol):
|
|||
|
||||
|
||||
class _ClientWebSocket(_ScopedWebSocket, Protocol):
|
||||
exceptions: _ClientWebSocketExceptions
|
||||
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
async def receive_text(self) -> str: ...
|
||||
async def close(self, code: int = 1000, reason: str | None = None) -> None: ...
|
||||
|
||||
|
||||
class _LoggingWorker(Protocol):
|
||||
def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None: ...
|
||||
|
||||
|
||||
def _decode_json_object(payload: str) -> Mapping[str, object]:
|
||||
|
|
@ -108,11 +138,14 @@ class RealTimeStreaming:
|
|||
backend_uses_beta_protocol: bool | None = None,
|
||||
force_transcription_model: str | None = None,
|
||||
event_normalizer: RealtimeEventNormalizer | None = None,
|
||||
logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER,
|
||||
):
|
||||
self.websocket: _ClientWebSocket = websocket
|
||||
self.backend_ws = backend_ws
|
||||
self.logging_obj = logging_obj
|
||||
self._logging_worker = logging_worker
|
||||
self.messages: list[OpenAIRealtimeEvents] = []
|
||||
self._backend_sent_frames: bool = False
|
||||
self.input_message: dict = {}
|
||||
self.input_messages: list[dict[str, str]] = []
|
||||
self.session_tools: list[dict] = []
|
||||
|
|
@ -388,9 +421,10 @@ class RealTimeStreaming:
|
|||
# Route through the bounded logging worker (per-coroutine timeout +
|
||||
# concurrency cap) instead of a bare create_task, so a slow callback
|
||||
# can't leave suspended tasks pinning each call's response in memory.
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True)
|
||||
)
|
||||
self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
|
||||
async def _send_to_backend(self, message: str) -> bool:
|
||||
"""Send a message to the backend WebSocket.
|
||||
|
|
@ -1035,60 +1069,84 @@ class RealTimeStreaming:
|
|||
return True
|
||||
return False
|
||||
|
||||
async def backend_to_client_send_messages(self):
|
||||
async def _relay_backend_messages(self) -> NoReturn:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv(decode=False)
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv()
|
||||
self._backend_sent_frames = True
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
try:
|
||||
raw_response = raw_response.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
verbose_logger.warning("Received non-UTF-8 binary frame from backend, skipping.")
|
||||
continue
|
||||
|
||||
if self.provider_config:
|
||||
try:
|
||||
await self._handle_provider_config_message(raw_response)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error processing backend message, skipping: %s", e)
|
||||
continue
|
||||
else:
|
||||
event = self._parse_backend_event(raw_response)
|
||||
if event is None:
|
||||
await self.websocket.send_text(raw_response)
|
||||
continue
|
||||
|
||||
if self._should_drop_event_from_client(event):
|
||||
continue
|
||||
|
||||
if await self._handle_raw_backend_message(event, raw_response):
|
||||
continue
|
||||
|
||||
event = self._normalize_event_for_ga_client(event)
|
||||
self.store_message(event)
|
||||
|
||||
if not self._client_wants_beta:
|
||||
await self.websocket.send_text(json.dumps(event))
|
||||
continue
|
||||
|
||||
translated = self._translate_event_to_beta(event)
|
||||
if translated is None:
|
||||
continue
|
||||
await self.websocket.send_text(json.dumps(translated))
|
||||
|
||||
async def backend_to_client_send_messages(self) -> BackendClose:
|
||||
import websockets
|
||||
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv(decode=False)
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv()
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
try:
|
||||
raw_response = raw_response.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
verbose_logger.warning("Received non-UTF-8 binary frame from backend, skipping.")
|
||||
continue
|
||||
|
||||
if self.provider_config:
|
||||
try:
|
||||
await self._handle_provider_config_message(raw_response)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error processing backend message, skipping: %s", e)
|
||||
continue
|
||||
else:
|
||||
event = self._parse_backend_event(raw_response)
|
||||
if event is None:
|
||||
await self.websocket.send_text(raw_response)
|
||||
continue
|
||||
|
||||
if self._should_drop_event_from_client(event):
|
||||
continue
|
||||
|
||||
if await self._handle_raw_backend_message(event, raw_response):
|
||||
continue
|
||||
|
||||
event = self._normalize_event_for_ga_client(event)
|
||||
self.store_message(event)
|
||||
|
||||
if not self._client_wants_beta:
|
||||
await self.websocket.send_text(json.dumps(event))
|
||||
continue
|
||||
|
||||
translated = self._translate_event_to_beta(event)
|
||||
if translated is None:
|
||||
continue
|
||||
await self.websocket.send_text(json.dumps(translated))
|
||||
|
||||
await self._relay_backend_messages()
|
||||
except websockets.exceptions.ConnectionClosed as e:
|
||||
verbose_logger.exception("Connection closed in backend to client send messages - %s", e)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in backend to client send messages: %s", e)
|
||||
finally:
|
||||
close: Final = backend_close_from(e)
|
||||
self._flush_unbilled_transcription_usage()
|
||||
if self._backend_refused_session(close):
|
||||
await self.log_backend_refusal(e)
|
||||
else:
|
||||
await self.log_messages()
|
||||
return close
|
||||
except asyncio.CancelledError:
|
||||
self._flush_unbilled_transcription_usage()
|
||||
await self.log_messages()
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in backend to client send messages: %s", e)
|
||||
self._flush_unbilled_transcription_usage()
|
||||
await self.log_messages()
|
||||
return BackendClose(code=1011, reason="proxy failed while relaying the upstream websocket")
|
||||
|
||||
def _backend_refused_session(self, close: BackendClose) -> bool:
|
||||
return close.code != 1000 and not self._backend_sent_frames
|
||||
|
||||
async def log_backend_refusal(self, error: Exception) -> None:
|
||||
if not self.logging_obj:
|
||||
return
|
||||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_failure_handlers(error, traceback.format_exc(), prefer_async_handlers=True)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _detect_beta_header(websocket: _ScopedWebSocket) -> bool:
|
||||
|
|
@ -1243,11 +1301,22 @@ class RealTimeStreaming:
|
|||
item["content"] = new_content
|
||||
return item
|
||||
|
||||
async def client_ack_messages(self):
|
||||
async def _receive_client_message(self) -> str | None:
|
||||
try:
|
||||
return await self.websocket.receive_text()
|
||||
except Exception as e: # noqa: BLE001 # whatever the client socket raises, the client is gone
|
||||
verbose_logger.debug("Client disconnected: %s", e)
|
||||
return None
|
||||
|
||||
async def client_ack_messages(self) -> ClientLoopExit:
|
||||
import websockets
|
||||
|
||||
client_event: _ClientEventFrame
|
||||
try:
|
||||
while True:
|
||||
message = await self.websocket.receive_text()
|
||||
message = await self._receive_client_message()
|
||||
if message is None:
|
||||
return ClientLoopExit.CLIENT_DISCONNECTED
|
||||
|
||||
## GUARDRAIL: intercept conversation.item.create for text-based injection.
|
||||
guardrail_turn_detection_injected = False
|
||||
|
|
@ -1481,23 +1550,38 @@ class RealTimeStreaming:
|
|||
if guardrail_turn_detection_injected and sent:
|
||||
self._guardrail_turn_detection_update_sent = True
|
||||
|
||||
except websockets.exceptions.ConnectionClosed as e:
|
||||
verbose_logger.debug("Backend closed while forwarding a client message: %s", e)
|
||||
return ClientLoopExit.BACKEND_CLOSED
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Error in client ack messages: %s", e)
|
||||
return ClientLoopExit.CLIENT_DISCONNECTED
|
||||
|
||||
async def bidirectional_forward(self):
|
||||
async def bidirectional_forward(self) -> None:
|
||||
forward_task: Final = asyncio.create_task(self.backend_to_client_send_messages())
|
||||
client_task: Final = asyncio.create_task(self.client_ack_messages())
|
||||
try:
|
||||
await self.client_ack_messages()
|
||||
except self.websocket.exceptions.ConnectionClosed:
|
||||
verbose_logger.debug("Connection closed")
|
||||
forward_task.cancel()
|
||||
await asyncio.wait((forward_task, client_task), return_when=asyncio.FIRST_COMPLETED)
|
||||
if client_task.done() and client_task.result() is ClientLoopExit.CLIENT_DISCONNECTED:
|
||||
return
|
||||
await self._close_client(await forward_task)
|
||||
finally:
|
||||
if not forward_task.done():
|
||||
forward_task.cancel()
|
||||
try:
|
||||
await forward_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
forward_task.cancel()
|
||||
client_task.cancel()
|
||||
await asyncio.gather(forward_task, client_task, return_exceptions=True)
|
||||
|
||||
async def _close_client(self, close: BackendClose) -> None:
|
||||
redacted_message: Final = redact_internal_details_from_client_message(close.message)
|
||||
redacted_reason: Final = redact_internal_details_from_client_message(close.reason)
|
||||
try:
|
||||
if close.code != 1000:
|
||||
await self.websocket.send_text(realtime_error_event(redacted_message, error_type="server_error"))
|
||||
await self.websocket.close(
|
||||
code=client_close_code(close.code),
|
||||
reason=websocket_close_reason(redacted_reason, fallback=redacted_message),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # the client may already be gone; the session is over either way
|
||||
verbose_logger.debug("Could not relay the upstream close to the client: %s", e)
|
||||
|
||||
|
||||
def client_sent_openai_beta_realtime_header(websocket: _ScopedWebSocket) -> bool:
|
||||
|
|
|
|||
|
|
@ -5,6 +5,20 @@ from typing import Final
|
|||
from fastapi import HTTPException
|
||||
|
||||
|
||||
class MCPServerURLCredentialsError(HTTPException):
|
||||
"""A fixed, sanitized URL-credential migration error safe for operator previews."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__(
|
||||
status_code=500,
|
||||
detail=(
|
||||
"misconfigured: auth_type none cannot be used with credentials embedded in the upstream URL; "
|
||||
"remove them from the URL and configure Basic Auth with auth_type: basic and "
|
||||
"auth_value: username:password"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class MCPUpstreamAuthError(Exception):
|
||||
"""Raised when an upstream MCP server returns an authentication failure
|
||||
(typically HTTP 401) and the gateway should surface it transparently to
|
||||
|
|
|
|||
|
|
@ -159,6 +159,9 @@ from litellm.proxy._types import (
|
|||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
|
||||
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
||||
id_jag_assertion_capture_gap_at_startup,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, get_server_root_path
|
||||
from litellm.repositories.table_repositories import MCPServerRepository
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
|
@ -1382,6 +1385,20 @@ def _warn_internal_delegate_pkce_if_applicable(server: MCPServer, *, source: str
|
|||
)
|
||||
|
||||
|
||||
def _warn_config_id_jag_server_outruns_sso(server: MCPServer) -> None:
|
||||
if server.auth_type != MCPAuth.oauth2_id_jag:
|
||||
return
|
||||
gap: Final = id_jag_assertion_capture_gap_at_startup()
|
||||
if gap is None:
|
||||
return
|
||||
verbose_logger.warning(
|
||||
"MCP server %r (id=%s, source=config) is declared with auth_type=oauth2_id_jag, but %s.",
|
||||
get_server_prefix(server),
|
||||
server.server_id,
|
||||
gap,
|
||||
)
|
||||
|
||||
|
||||
def _deserialize_json_dict(data: str | _StringMap | None) -> dict[str, str] | None:
|
||||
"""
|
||||
Deserialize optional JSON mappings stored in the database.
|
||||
|
|
@ -2393,6 +2410,7 @@ class MCPServerManager:
|
|||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
_warn_config_id_jag_server_outruns_sso(new_server)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
self._set_oauth_discovery_deferred(
|
||||
server_id,
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from pydantic import SecretStr
|
|||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.experimental_mcp_client.client import strip_auth_scheme, to_basic_credentials
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import MCPServerURLCredentialsError
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
|
|
@ -293,6 +294,8 @@ def raise_public(error: CredError) -> NoReturn:
|
|||
)
|
||||
case "misconfigured":
|
||||
raise HTTPException(status_code=500, detail=error.summary)
|
||||
case "url_credentials_not_allowed":
|
||||
raise MCPServerURLCredentialsError()
|
||||
case "upstream_unavailable":
|
||||
raise HTTPException(status_code=503, detail=error.summary)
|
||||
case "unsupported_mode":
|
||||
|
|
|
|||
|
|
@ -31,11 +31,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
|
|||
TokenCacheBackend,
|
||||
TokenStoreUnavailable,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_distributed_lock import (
|
||||
RedisDistributedLock,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
|
||||
RedisRefreshCoordinator,
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.runtime_refresh_coordinator import (
|
||||
runtime_refresh_coordinator,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_cache_codec import (
|
||||
OAuthTokenCacheCodec,
|
||||
|
|
@ -131,23 +128,17 @@ def _runtime_backend_and_coordinator() -> tuple[TokenCacheBackend | None, Refres
|
|||
)
|
||||
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415
|
||||
|
||||
redis_cache: Final = user_api_key_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
coordinator: Final = runtime_refresh_coordinator()
|
||||
if coordinator is None:
|
||||
return None, None, False
|
||||
codec: Final = OAuthTokenCacheCodec(
|
||||
encrypt_value_helper,
|
||||
lambda blob: decrypt_value_helper(blob, "mcp_per_user_token", exception_type="debug"),
|
||||
)
|
||||
# user_api_key_cache satisfies the AsyncCache slice (DualCache types ttl via **kwargs) and the
|
||||
# Redis client from init_async_client() is partially typed - both are untyped-boundary casts.
|
||||
# user_api_key_cache satisfies the AsyncCache slice (DualCache types ttl via **kwargs) - an
|
||||
# untyped-boundary cast.
|
||||
cache: Final[AsyncCache] = user_api_key_cache # pyright: ignore
|
||||
redis_client: Final = redis_cache.init_async_client() # pyright: ignore
|
||||
lock: Final = RedisDistributedLock(
|
||||
redis_client, # pyright: ignore
|
||||
namespace_key=redis_cache.check_and_fix_namespace,
|
||||
)
|
||||
backend: Final = DualCacheTokenCacheBackend(cache, codec)
|
||||
coordinator: Final = RedisRefreshCoordinator(lock)
|
||||
return backend, coordinator, True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -46,11 +46,13 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
|||
Ok,
|
||||
Result,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_refresher import (
|
||||
default_sso_assertion_store,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
AssertionStoreUnavailable,
|
||||
DbSSOAssertionStore,
|
||||
SSOAssertionStore,
|
||||
SSOIdentityAssertion,
|
||||
assertion_expired,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint import (
|
||||
ExchangedToken,
|
||||
|
|
@ -129,12 +131,12 @@ class UpstreamCredentialProvider:
|
|||
self._token_endpoint: TokenEndpointClient = token_endpoint or TokenEndpointClient()
|
||||
self._exchanged_tokens: ExchangedTokenCache = exchanged_tokens or ExchangedTokenCache()
|
||||
self._client_credentials_source = client_credentials_source or ClientCredentialsTokenSource()
|
||||
self._sso_assertion_store: SSOAssertionStore = sso_assertion_store or DbSSOAssertionStore()
|
||||
self._sso_assertion_store: SSOAssertionStore = sso_assertion_store or default_sso_assertion_store()
|
||||
|
||||
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Result[httpx.Auth, CredError]:
|
||||
match server.config:
|
||||
case NoneConfig():
|
||||
return Ok(NoOpAuth())
|
||||
return self._none(server)
|
||||
case ApiKeyConfig() as config:
|
||||
return self._api_key(config)
|
||||
case PassthroughConfig():
|
||||
|
|
@ -151,6 +153,15 @@ class UpstreamCredentialProvider:
|
|||
return _not_implemented(AuthSpecKind.aws_sigv4)
|
||||
assert_never(server.config)
|
||||
|
||||
def _none(self, server: ServerSpec) -> Result[httpx.Auth, CredError]:
|
||||
try:
|
||||
resource: Final = httpx.URL(server.resource)
|
||||
except httpx.InvalidURL:
|
||||
return Ok(NoOpAuth())
|
||||
if resource.userinfo:
|
||||
return Error(CredError.of_url_credentials_not_allowed())
|
||||
return Ok(NoOpAuth())
|
||||
|
||||
async def has_user_token(self, subject: Subject, server: ServerSpec) -> bool:
|
||||
"""Whether a usable per-user token exists for this server (the preemptive 401's check).
|
||||
|
||||
|
|
@ -237,7 +248,7 @@ class UpstreamCredentialProvider:
|
|||
"Sign in through LiteLLM SSO so the gateway captures one."
|
||||
)
|
||||
)
|
||||
if _assertion_expired(assertion, datetime.now(timezone.utc)):
|
||||
if assertion_expired(assertion, datetime.now(timezone.utc)):
|
||||
return Error(
|
||||
CredError.of_precondition_required(
|
||||
"The stored IdP identity assertion for this user has expired. Sign in through "
|
||||
|
|
@ -396,19 +407,6 @@ def _id_jag_slot_key(subject: Subject, server: ServerSpec) -> str:
|
|||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
|
||||
def _assertion_expired(assertion: SSOIdentityAssertion, now: datetime) -> bool:
|
||||
"""Whether the stored assertion's ``exp`` has passed. An assertion carrying no expiry is
|
||||
treated as usable and left for the IdP to reject, since the store records what the id_token
|
||||
claimed rather than imposing a lifetime of its own. A naive ``expires_at`` is read as UTC so a
|
||||
stored value that lost its offset compares instead of raising.
|
||||
"""
|
||||
expires_at: Final = assertion.expires_at
|
||||
if expires_at is None:
|
||||
return False
|
||||
normalized: Final = expires_at if expires_at.tzinfo is not None else expires_at.replace(tzinfo=timezone.utc)
|
||||
return normalized <= now
|
||||
|
||||
|
||||
def _id_jag_fingerprint(subject_token: str, server_id: str, config: IdJagConfig) -> str:
|
||||
"""What the cached leg-2 bearer was minted from: the subject token, the server, and the config.
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,41 @@
|
|||
"""The runtime ``RefreshCoordinator``: cross-replica single-flight when Redis is wired.
|
||||
|
||||
Builds ``RedisRefreshCoordinator`` over the proxy's shared Redis so one refresh runs per key
|
||||
across the fleet, or returns ``None`` when Redis is absent so the caller keeps the foundation's
|
||||
in-process default (correct for a single replica). The proxy globals it reads are not ready at
|
||||
import time, so this is called per composition rather than held as module state.
|
||||
|
||||
Shared by every credential arm that renews a stored grant: a rotating refresh token must be
|
||||
redeemed once across all workers, so each arm electing its own winner with its own lock shape
|
||||
would be a bug waiting to differ.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
RefreshCoordinator,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_distributed_lock import (
|
||||
RedisDistributedLock,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.redis_refresh_coordinator import (
|
||||
RedisRefreshCoordinator,
|
||||
)
|
||||
|
||||
|
||||
def runtime_refresh_coordinator() -> RefreshCoordinator | None:
|
||||
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # runtime global
|
||||
|
||||
redis_cache: Final = user_api_key_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
return None
|
||||
# The Redis client from init_async_client() is only partially typed; the lock validates every
|
||||
# reply it depends on, so the untyped boundary is contained here.
|
||||
redis_client: Final = redis_cache.init_async_client() # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # litellm redis wrapper is untyped
|
||||
lock: Final = RedisDistributedLock(
|
||||
redis_client, # pyright: ignore[reportArgumentType,reportUnknownArgumentType] # litellm redis wrapper is untyped
|
||||
namespace_key=redis_cache.check_and_fix_namespace,
|
||||
)
|
||||
return RedisRefreshCoordinator(lock)
|
||||
|
|
@ -0,0 +1,469 @@
|
|||
"""Renew the stored SSO identity assertion so an ID-JAG agent outlives one id_token.
|
||||
|
||||
The ``oauth2_id_jag`` arm asserts the id_token captured at the user's last interactive sign-in, so
|
||||
without renewal an agent holding a brokered LiteLLM key can act for that user only until that token's
|
||||
``exp``, typically an hour, and the sole recovery is another interactive login. The assertion already
|
||||
carries the IdP refresh token beside it; this module is what redeems it.
|
||||
|
||||
``RefreshingSSOAssertionStore`` wraps any ``SSOAssertionStore`` and satisfies the same protocol, so
|
||||
the egress arm is unchanged: it still reads one assertion and still judges expiry itself. Renewal is
|
||||
lazy (only a read that finds a near-expiry assertion triggers one, so IdP traffic tracks actual use,
|
||||
not the size of the user table) and single-flighted per user through the same ``RefreshCoordinator``
|
||||
the ``authorization_code`` arm uses, because an IdP that rotates refresh tokens treats two concurrent
|
||||
redemptions of one token as replay and can revoke the whole grant chain.
|
||||
|
||||
The refresh is redeemed against the generic-OIDC client the login itself used
|
||||
(``GENERIC_TOKEN_ENDPOINT`` / ``GENERIC_CLIENT_ID`` / ``GENERIC_CLIENT_SECRET``, which the proxy
|
||||
reconciles from the stored SSO row into the process environment at startup), authenticated the way
|
||||
that login authenticated: the non-PKCE path always sends HTTP Basic, while the PKCE path sends the
|
||||
credentials in the body when ``GENERIC_INCLUDE_CLIENT_ID`` is set, and an IdP application may accept
|
||||
only one of the two. An assertion can only exist if that client minted it, so no other client could
|
||||
redeem its refresh token, and no other method is known to be accepted. A deployment whose
|
||||
``GENERIC_SCOPE`` omits ``offline_access`` captures no refresh token at all, which is why that miss
|
||||
logs the scope by name rather than failing silently.
|
||||
|
||||
Failures are values internally (``Result[_, RefreshFailure]``). At the store boundary they collapse
|
||||
onto the protocol's existing two-outcome contract: a refusal returns the expired assertion unchanged
|
||||
so the reader's own guard challenges the user to sign in again, while a transient IdP failure raises
|
||||
``AssertionStoreUnavailable`` so the reader answers 503 instead of blaming the user for an outage.
|
||||
|
||||
One ambiguity remains under Redis-coordinated renewal across replicas. A cross-replica loser that
|
||||
finds the row still expiring after the holder finished cannot tell a refused refresh from a renewal
|
||||
that could not be recorded. Redeeming itself could consume a refresh token the holder may already
|
||||
have rotated, so it answers retryable 503 rather than guessing a sign-in challenge. The next
|
||||
uncontended read settles the outcome itself: a refusal challenges, and a successful refresh persists.
|
||||
If the holder rotated the token but its write failed, that rotation is lost and the next uncontended
|
||||
read's refusal challenges, which is the only honest answer because the rotated token was never
|
||||
recorded. On the refusal path, the loser pays for one retry before that challenge.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final, Literal, Protocol
|
||||
|
||||
import httpx
|
||||
from pydantic import SecretStr, TypeAdapter, ValidationError
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import Timeout
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
|
||||
build_token_endpoint_client_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
InProcessRefreshCoordinator,
|
||||
RefreshCoordinator,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
||||
Error,
|
||||
Ok,
|
||||
Result,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.runtime_refresh_coordinator import (
|
||||
runtime_refresh_coordinator,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
AssertionStoreUnavailable,
|
||||
DbSSOAssertionStore,
|
||||
SSOAssertionStore,
|
||||
SSOIdentityAssertion,
|
||||
assertion_expired,
|
||||
assertion_from_sso_login,
|
||||
fetch_sso_identity_assertion,
|
||||
persist_sso_identity_assertion,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.mcp import MCPTokenEndpointAuthMethod
|
||||
|
||||
_BODY_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(dict[str, object])
|
||||
|
||||
_REFRESH_GRANT_TYPE: Final = "refresh_token"
|
||||
# The lock namespace for the one assertion row a user has; the sibling arm keys the same lock by
|
||||
# server_id, and no server_id can collide with this literal.
|
||||
_SINGLE_FLIGHT_KEY: Final = "sso_identity_assertion"
|
||||
# Renew this far ahead of ``exp`` so a token that would die between resolution and the second leg of
|
||||
# the exchange is replaced first. Matches the sibling per-user token store's skew.
|
||||
_DEFAULT_EXPIRY_SKEW_SECONDS: Final = 60.0
|
||||
|
||||
|
||||
class AssertionRead(Protocol):
|
||||
"""Reads the user's stored assertion row."""
|
||||
|
||||
async def __call__(self, user_id: str) -> SSOIdentityAssertion | None: ...
|
||||
|
||||
|
||||
class AssertionWrite(Protocol):
|
||||
"""Replaces the user's stored assertion row."""
|
||||
|
||||
async def __call__(self, user_id: str, assertion: SSOIdentityAssertion) -> None: ...
|
||||
|
||||
|
||||
class CoordinatorFactory(Protocol):
|
||||
"""Builds the cross-replica coordinator, or ``None`` when there is no shared lock to build on."""
|
||||
|
||||
def __call__(self) -> RefreshCoordinator | None: ...
|
||||
|
||||
|
||||
class FormPost(Protocol):
|
||||
"""POSTs an OAuth form and hands back the raw response."""
|
||||
|
||||
async def __call__(
|
||||
self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
|
||||
) -> httpx.Response | None: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SSOClientConfig:
|
||||
"""The generic-OIDC client credentials a refresh_token grant has to authenticate as, and how."""
|
||||
|
||||
token_endpoint: str
|
||||
client_id: str
|
||||
client_secret: SecretStr
|
||||
auth_method: MCPTokenEndpointAuthMethod
|
||||
|
||||
|
||||
def sso_client_config(env: Mapping[str, str]) -> SSOClientConfig | None:
|
||||
"""The configured generic-OIDC client, or ``None`` when the deployment has none.
|
||||
|
||||
Read from the process environment because that is where the login path reads it
|
||||
(``_setup_generic_sso_env_vars``) and where the proxy materializes the stored ``sso_config`` row
|
||||
at startup, so this resolves to the same client that minted the assertion. ``None`` is an
|
||||
ordinary state, not an error: a deployment signing in through a provider that captures no
|
||||
assertion has nothing here to renew, and a client with no secret is not a confidential client
|
||||
that could redeem one.
|
||||
|
||||
``auth_method`` is derived from the same ``GENERIC_INCLUDE_CLIENT_ID`` the login reads, because
|
||||
the two login paths do not agree: the non-PKCE path always authenticates with HTTP Basic, while
|
||||
the PKCE path puts the credentials in the body when that flag is set. Both capture assertions, so
|
||||
a constant here would authenticate the renewal differently from the sign-in that produced the
|
||||
refresh token and 401 against an IdP application registered for only one of the two.
|
||||
"""
|
||||
token_endpoint: Final = env.get("GENERIC_TOKEN_ENDPOINT")
|
||||
client_id: Final = env.get("GENERIC_CLIENT_ID")
|
||||
client_secret: Final = env.get("GENERIC_CLIENT_SECRET")
|
||||
if not token_endpoint or not client_id or not client_secret:
|
||||
return None
|
||||
includes_client_id: Final = env.get("GENERIC_INCLUDE_CLIENT_ID", "false").lower() == "true"
|
||||
return SSOClientConfig(
|
||||
token_endpoint=token_endpoint,
|
||||
client_id=client_id,
|
||||
client_secret=SecretStr(client_secret),
|
||||
auth_method="client_secret_post" if includes_client_id else "client_secret_basic",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RefreshFailure:
|
||||
"""Why a renewal produced nothing, split by what the caller can do about it.
|
||||
|
||||
``rejected`` is settled: this refresh token will never work again, so the user has to sign in.
|
||||
``unavailable`` is transient: the same attempt may succeed in a minute, so telling the user to
|
||||
sign in again would be a lie about whose problem it is. Both arms carry the same payload, so
|
||||
this is a ``Literal`` discriminant rather than a ``tagged_union``; consumers still ``match`` on
|
||||
``kind`` with an ``assert_never`` tail.
|
||||
"""
|
||||
|
||||
kind: Literal["rejected", "unavailable"]
|
||||
detail: str
|
||||
|
||||
@staticmethod
|
||||
def of_rejected(detail: str) -> RefreshFailure:
|
||||
return RefreshFailure(kind="rejected", detail=detail)
|
||||
|
||||
@staticmethod
|
||||
def of_unavailable(detail: str) -> RefreshFailure:
|
||||
return RefreshFailure(kind="unavailable", detail=detail)
|
||||
|
||||
|
||||
class TokenEndpointTransport(Protocol):
|
||||
"""One form POST to the IdP token endpoint, with the refusal/outage split preserved.
|
||||
|
||||
That split is the whole reason this is not the resolver's ``TokenEndpointClient``: that
|
||||
collaborator maps every non-2xx to ``upstream_unavailable``, which is right for an exchange leg
|
||||
and wrong here, where a 400 ``invalid_grant`` means the stored refresh token is dead and the user
|
||||
must act.
|
||||
"""
|
||||
|
||||
async def post(
|
||||
self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
|
||||
) -> Result[Mapping[str, object], RefreshFailure]: ...
|
||||
|
||||
|
||||
async def post_form(url: str, form: Mapping[str, str], headers: Mapping[str, str]) -> httpx.Response | None:
|
||||
# litellm's httpx handler is only partially typed; nothing but the response object crosses back,
|
||||
# and the transport below validates its body, so the untyped boundary is contained here.
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped
|
||||
return await client.post(url, data=form, headers=headers) # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType,reportReturnType,reportArgumentType] # litellm http handler is untyped and its stub narrows data=/headers= to dict, which httpx itself does not require
|
||||
|
||||
|
||||
class HttpxTokenEndpointTransport:
|
||||
"""The live transport. 4xx is the IdP refusing this grant; anything else is an outage.
|
||||
|
||||
The POST itself is injected so that split, which decides whether the user is challenged or told
|
||||
to wait, is testable without a live IdP.
|
||||
"""
|
||||
|
||||
def __init__(self, post: FormPost = post_form) -> None:
|
||||
self._post = post
|
||||
|
||||
async def post(
|
||||
self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
|
||||
) -> Result[Mapping[str, object], RefreshFailure]:
|
||||
try:
|
||||
response: Final = await self._post(url, form, headers)
|
||||
if response is None:
|
||||
return Error(RefreshFailure.of_unavailable("the IdP token endpoint returned no response"))
|
||||
response.raise_for_status()
|
||||
body: Final = _BODY_ADAPTER.validate_python(response.json()) # pyright: ignore[reportAny] # untyped JSON; the adapter is the type gate
|
||||
except httpx.HTTPStatusError as exc:
|
||||
status: Final = exc.response.status_code
|
||||
if 400 <= status < 500:
|
||||
return Error(RefreshFailure.of_rejected(f"the IdP refused the refresh with status {status}"))
|
||||
return Error(RefreshFailure.of_unavailable(f"the IdP token endpoint answered with status {status}"))
|
||||
except (httpx.RequestError, Timeout) as exc:
|
||||
return Error(RefreshFailure.of_unavailable(f"the IdP token endpoint is unreachable ({type(exc).__name__})"))
|
||||
except json.JSONDecodeError:
|
||||
return Error(RefreshFailure.of_unavailable("the IdP token endpoint returned a non-JSON response"))
|
||||
except ValidationError:
|
||||
return Error(RefreshFailure.of_unavailable("the IdP token endpoint returned a non-object response"))
|
||||
return Ok(body)
|
||||
|
||||
|
||||
class SSOAssertionRefresher:
|
||||
"""Redeems the stored refresh token for a current id_token and writes the rotation back.
|
||||
|
||||
Collaborators are injected so the orchestration, the untyped response parsing and the
|
||||
write-back race are all testable without an IdP or a database.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transport: TokenEndpointTransport,
|
||||
*,
|
||||
client_config: Callable[[], SSOClientConfig | None] = lambda: sso_client_config(os.environ),
|
||||
read: AssertionRead = fetch_sso_identity_assertion,
|
||||
write: AssertionWrite = persist_sso_identity_assertion,
|
||||
) -> None:
|
||||
self._transport = transport
|
||||
self._client_config = client_config
|
||||
self._read = read
|
||||
self._write = write
|
||||
|
||||
async def refresh(
|
||||
self, user_id: str, assertion: SSOIdentityAssertion
|
||||
) -> Result[SSOIdentityAssertion, RefreshFailure]:
|
||||
if assertion.refresh_token is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"ID-JAG: the stored IdP identity assertion for user_id=%s has expired and no refresh token was "
|
||||
"captured with it, so it cannot be renewed without another interactive sign-in. Add "
|
||||
"'offline_access' to GENERIC_SCOPE so the SSO login captures one.",
|
||||
user_id,
|
||||
)
|
||||
return Error(RefreshFailure.of_rejected("no refresh token was captured at sign-in"))
|
||||
config: Final = self._client_config()
|
||||
if config is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"ID-JAG: the stored IdP identity assertion for user_id=%s has expired and cannot be renewed "
|
||||
"because the generic SSO client is not configured (GENERIC_TOKEN_ENDPOINT, GENERIC_CLIENT_ID, "
|
||||
"GENERIC_CLIENT_SECRET).",
|
||||
user_id,
|
||||
)
|
||||
return Error(RefreshFailure.of_rejected("the generic SSO client is not configured"))
|
||||
|
||||
carried_refresh_token: Final = assertion.refresh_token.get_secret_value()
|
||||
# Whichever method the SSO login used for this client, since that is the one the IdP
|
||||
# application is known to accept: an assertion only exists to renew because a sign-in already
|
||||
# authenticated this client that way.
|
||||
client_auth: Final = build_token_endpoint_client_auth(
|
||||
auth_method=config.auth_method,
|
||||
client_id=config.client_id,
|
||||
client_secret=config.client_secret.get_secret_value(),
|
||||
)
|
||||
form: Final = { # mutable-ok: the RFC 6749 form body is a wire format the HTTP client takes as a mapping
|
||||
"grant_type": _REFRESH_GRANT_TYPE,
|
||||
"refresh_token": carried_refresh_token,
|
||||
**client_auth.body,
|
||||
}
|
||||
match await self._transport.post(config.token_endpoint, form, client_auth.headers):
|
||||
case Error(failure):
|
||||
return Error(failure)
|
||||
case Ok(body):
|
||||
return await self._renewed_from(user_id, assertion, body, carried_refresh_token)
|
||||
|
||||
async def _renewed_from(
|
||||
self,
|
||||
user_id: str,
|
||||
previous: SSOIdentityAssertion,
|
||||
body: Mapping[str, object],
|
||||
carried_refresh_token: str,
|
||||
) -> Result[SSOIdentityAssertion, RefreshFailure]:
|
||||
"""The renewed assertion, built by the same validator the login path uses.
|
||||
|
||||
A rotated refresh token replaces the stored one; an omitted one carries forward, since an
|
||||
IdP that does not rotate expects the original to keep working.
|
||||
"""
|
||||
rotated: Final = body.get("refresh_token")
|
||||
renewed: Final = assertion_from_sso_login(
|
||||
body.get("id_token"),
|
||||
rotated if isinstance(rotated, str) and rotated else carried_refresh_token,
|
||||
)
|
||||
if renewed is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"ID-JAG: the IdP accepted the refresh for user_id=%s but returned no usable id_token, so there "
|
||||
"is nothing to assert upstream. The SSO client's grant needs the 'openid' scope for the token "
|
||||
"endpoint to return one on a refresh.",
|
||||
user_id,
|
||||
)
|
||||
return Error(RefreshFailure.of_rejected("the IdP's refresh response carried no usable id_token"))
|
||||
failure: Final = await self._store_renewal(user_id, previous, renewed)
|
||||
if failure is not None:
|
||||
return Error(failure)
|
||||
return Ok(renewed)
|
||||
|
||||
async def _store_renewal(
|
||||
self, user_id: str, previous: SSOIdentityAssertion, renewed: SSOIdentityAssertion
|
||||
) -> RefreshFailure | None:
|
||||
"""Write the renewal back, unless the row moved on while this renewal was in flight.
|
||||
|
||||
The row is one per user and last-write-wins, so an interactive sign-in landing mid-renewal
|
||||
would otherwise be overwritten with a refresh token the IdP has already rotated away, costing
|
||||
that user a sign-in later. Comparing against the id_token this renewal started from is what
|
||||
detects that; skipping is safe because the newer row is the one the reader wants anyway.
|
||||
|
||||
A failed write is transient, not settled. The store, not this return value, is what every
|
||||
caller reads, so a renewal that could not be recorded is a renewal nobody will see; saying so
|
||||
keeps a database problem answering 503 rather than telling the user to sign in again over it.
|
||||
"""
|
||||
try:
|
||||
current: Final = await self._read(user_id)
|
||||
if current is not None and current.id_token.get_secret_value() != previous.id_token.get_secret_value():
|
||||
verbose_proxy_logger.info(
|
||||
"ID-JAG: a newer IdP identity assertion for user_id=%s was stored while this renewal was in "
|
||||
"flight; keeping the stored one.",
|
||||
user_id,
|
||||
)
|
||||
return None
|
||||
await self._write(user_id, renewed)
|
||||
except Exception as exc: # noqa: BLE001 # any storage failure is transient here, never the user's fault
|
||||
verbose_proxy_logger.warning(
|
||||
"ID-JAG: could not persist the renewed IdP identity assertion for user_id=%s, so the rotated "
|
||||
"refresh token is lost and this user will have to sign in again once the renewed token expires: %s",
|
||||
user_id,
|
||||
exc,
|
||||
)
|
||||
return RefreshFailure.of_unavailable("the renewed IdP identity assertion could not be persisted")
|
||||
return None
|
||||
|
||||
|
||||
class RefreshingSSOAssertionStore:
|
||||
"""An ``SSOAssertionStore`` that renews a near-expiry assertion before handing it back.
|
||||
|
||||
Reads the inner store; an assertion still comfortably inside its lifetime is returned untouched,
|
||||
so the common path costs exactly what it did before. Otherwise one renewal runs per user through
|
||||
the injected ``RefreshCoordinator`` and every caller then re-reads the inner store, which is the
|
||||
authority: the winner's write is what they all observe, and a renewal the write-back guard
|
||||
skipped yields the newer assertion that displaced it rather than a private copy.
|
||||
|
||||
A refusal leaves the expired assertion in place for the reader's own guard to reject, so the user
|
||||
sees the same sign-in-again challenge as before this store existed. A transient IdP failure
|
||||
raises ``AssertionStoreUnavailable``, the protocol's existing signal for "this is not the user's
|
||||
fault"; concurrent in-process callers share that outcome, while a cross-replica loser answers 503
|
||||
when its re-read still finds the row expiring. On the refusal path that costs the loser one retry,
|
||||
which then challenges. If the holder rotated the token but its write failed, the rotation is lost
|
||||
and the next uncontended read's refusal challenges, the only honest answer because that token was
|
||||
never recorded.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inner: SSOAssertionStore,
|
||||
refresher: SSOAssertionRefresher,
|
||||
*,
|
||||
fresh_read: AssertionRead,
|
||||
coordinator_factory: CoordinatorFactory = runtime_refresh_coordinator,
|
||||
expiry_skew_seconds: float = _DEFAULT_EXPIRY_SKEW_SECONDS,
|
||||
clock: Callable[[], datetime] = lambda: datetime.now(timezone.utc),
|
||||
) -> None:
|
||||
self._inner = inner
|
||||
self._refresher = refresher
|
||||
self._fresh_read = fresh_read
|
||||
self._coordinator_factory = coordinator_factory
|
||||
self._in_process_coordinator = InProcessRefreshCoordinator()
|
||||
self._distributed_coordinator: RefreshCoordinator | None = None
|
||||
self._skew = timedelta(seconds=expiry_skew_seconds)
|
||||
self._clock = clock
|
||||
|
||||
async def fetch(self, user_id: str) -> SSOIdentityAssertion | None:
|
||||
assertion: Final = await self._inner.fetch(user_id)
|
||||
if not self._expiring(assertion):
|
||||
return assertion
|
||||
await self._coordinator().run(
|
||||
user_id,
|
||||
_SINGLE_FLIGHT_KEY,
|
||||
refresh=lambda: self._renew(user_id),
|
||||
reread=lambda: self._reread_renewed(user_id),
|
||||
)
|
||||
return await self._fresh_read(user_id)
|
||||
|
||||
def _expiring(self, assertion: SSOIdentityAssertion | None) -> bool:
|
||||
return assertion is not None and assertion_expired(assertion, self._clock() + self._skew)
|
||||
|
||||
def _coordinator(self) -> RefreshCoordinator:
|
||||
"""The cross-replica coordinator once Redis is reachable, else the in-process one.
|
||||
|
||||
Built on first use and kept, because the proxy's Redis client is not wired at import time;
|
||||
retried while it is absent so a proxy that gains Redis later stops electing per-worker.
|
||||
"""
|
||||
if self._distributed_coordinator is None:
|
||||
self._distributed_coordinator = self._coordinator_factory()
|
||||
return self._distributed_coordinator or self._in_process_coordinator
|
||||
|
||||
async def _renew(self, user_id: str) -> None:
|
||||
"""The elected renewal, judged from a fresh read so a rotation another replica just landed is
|
||||
never redeemed again. Returns nothing: the inner store, not this return value, is what every
|
||||
caller reads afterwards, so the winner and the losers cannot disagree."""
|
||||
latest: Final = await self._fresh_read(user_id)
|
||||
if latest is None or not self._expiring(latest):
|
||||
return
|
||||
match await self._refresher.refresh(user_id, latest):
|
||||
case Ok(_):
|
||||
return
|
||||
case Error(failure):
|
||||
match failure.kind:
|
||||
case "rejected":
|
||||
return
|
||||
case "unavailable":
|
||||
raise AssertionStoreUnavailable(failure.detail)
|
||||
assert_never(failure.kind)
|
||||
|
||||
async def _reread_renewed(self, user_id: str) -> None:
|
||||
"""A loser cannot distinguish refusal from an unrecorded renewal without risking token replay.
|
||||
|
||||
It answers retryable 503 instead of guessing a sign-in challenge; the retry runs uncontended
|
||||
and settles the outcome itself.
|
||||
"""
|
||||
latest: Final = await self._fresh_read(user_id)
|
||||
if self._expiring(latest):
|
||||
raise AssertionStoreUnavailable(
|
||||
f"the IdP identity assertion for user_id={user_id} was being renewed by another replica "
|
||||
"and is not yet current; retry shortly"
|
||||
)
|
||||
|
||||
|
||||
def default_sso_assertion_store() -> SSOAssertionStore:
|
||||
"""The live read seam for the ``id_jag`` arm: the stored assertion, renewed when it is stale."""
|
||||
db_store: Final = DbSSOAssertionStore()
|
||||
fresh_read: Final = db_store.fetch_uncached
|
||||
return RefreshingSSOAssertionStore(
|
||||
db_store,
|
||||
SSOAssertionRefresher(HttpxTokenEndpointTransport(), read=fresh_read),
|
||||
fresh_read=fresh_read,
|
||||
)
|
||||
|
|
@ -127,6 +127,23 @@ def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIden
|
|||
)
|
||||
|
||||
|
||||
def assertion_expired(assertion: SSOIdentityAssertion, now: datetime) -> bool:
|
||||
"""Whether the assertion's ``exp`` has passed at ``now``. An assertion carrying no expiry is
|
||||
treated as usable and left for the IdP to reject, since the store records what the id_token
|
||||
claimed rather than imposing a lifetime of its own. A naive ``expires_at`` is read as UTC so a
|
||||
stored value that lost its offset compares instead of raising.
|
||||
|
||||
Lives beside the model rather than in either reader so the egress guard and the renewal
|
||||
trigger judge the same field the same way; passing a ``now`` in the future is how a caller
|
||||
asks "is this about to expire" without a second, driftable predicate.
|
||||
"""
|
||||
expires_at: Final = assertion.expires_at
|
||||
if expires_at is None:
|
||||
return False
|
||||
normalized: Final = expires_at if expires_at.tzinfo is not None else expires_at.replace(tzinfo=timezone.utc)
|
||||
return normalized <= now
|
||||
|
||||
|
||||
async def ema_assertion_retention_enabled() -> bool:
|
||||
"""Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only
|
||||
retains bearer material while an EMA upstream exists to spend it on. Judged against the two
|
||||
|
|
@ -146,7 +163,9 @@ async def ema_assertion_retention_enabled() -> bool:
|
|||
return True
|
||||
if prisma_client is None:
|
||||
return False
|
||||
row = await prisma_client.db.litellm_mcpservertable.find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value})
|
||||
row: Final = await prisma_client.db.litellm_mcpservertable.find_first(
|
||||
where={"auth_type": MCPAuth.oauth2_id_jag.value}
|
||||
)
|
||||
return row is not None
|
||||
|
||||
|
||||
|
|
@ -158,7 +177,7 @@ async def persist_sso_identity_assertion(
|
|||
|
||||
if prisma_client is None:
|
||||
return
|
||||
payload: Final[dict[str, str]] = {
|
||||
payload: Final = {
|
||||
"id_token": assertion.id_token.get_secret_value(),
|
||||
**({"refresh_token": assertion.refresh_token.get_secret_value()} if assertion.refresh_token else {}),
|
||||
**({"issuer": assertion.issuer} if assertion.issuer else {}),
|
||||
|
|
@ -220,11 +239,13 @@ async def fetch_sso_identity_assertion(
|
|||
|
||||
|
||||
class AssertionStoreUnavailable(Exception):
|
||||
"""Raised by ``fetch`` when the backing store is unreachable (e.g. the DB is down).
|
||||
"""Raised by ``fetch`` when the assertion cannot be read for a transient reason: the DB is
|
||||
down, or the IdP behind a renewing store could not be reached.
|
||||
|
||||
Distinct from returning ``None`` for "this user has no captured assertion": an outage must not
|
||||
read as a definite absence, which would tell the user to sign in again over a transient failure,
|
||||
and it must not escape as an unhandled error on the egress or retry path. Mirrors
|
||||
and it must not escape as an unhandled error on the egress or retry path. The message names the
|
||||
real component for the operator log; callers get the reader's generic 503. Mirrors
|
||||
``TokenStoreUnavailable`` on the sibling per-user OAuth store.
|
||||
"""
|
||||
|
||||
|
|
@ -257,6 +278,12 @@ class DbSSOAssertionStore:
|
|||
except Exception as exc: # noqa: BLE001 # any driver/storage failure is an outage, not an absence
|
||||
raise AssertionStoreUnavailable(str(exc)) from exc
|
||||
|
||||
async def fetch_uncached(self, user_id: str) -> SSOIdentityAssertion | None:
|
||||
try:
|
||||
return await _read_assertion_from_db(user_id)
|
||||
except Exception as exc: # noqa: BLE001 # any driver/storage failure is an outage, not an absence
|
||||
raise AssertionStoreUnavailable(str(exc)) from exc
|
||||
|
||||
|
||||
async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None:
|
||||
"""Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation,
|
||||
|
|
@ -280,7 +307,9 @@ async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient,
|
|||
row.user_id,
|
||||
)
|
||||
return False
|
||||
re_encrypted = _STR_ADAPTER.validate_python(encrypt_value_helper(plaintext, new_encryption_key=new_master_key))
|
||||
re_encrypted: Final = _STR_ADAPTER.validate_python(
|
||||
encrypt_value_helper(plaintext, new_encryption_key=new_master_key)
|
||||
)
|
||||
await prisma_client.db.litellm_ssoidentityassertion.update(
|
||||
where={"user_id": row.user_id},
|
||||
data={"assertion_b64": re_encrypted},
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ class CredError:
|
|||
tag: Literal[
|
||||
"unauthorized",
|
||||
"misconfigured",
|
||||
"url_credentials_not_allowed",
|
||||
"upstream_unavailable",
|
||||
"unsupported_mode",
|
||||
"precondition_required",
|
||||
|
|
@ -103,6 +104,7 @@ class CredError:
|
|||
|
||||
unauthorized: Unauthorized = case() # no usable credential for this (subject, server) -> 401 challenge
|
||||
misconfigured: str = case() # the declared mode is missing required config -> 5xx (operator)
|
||||
url_credentials_not_allowed: None = case()
|
||||
upstream_unavailable: str = case() # the IdP / token endpoint could not be reached -> 503
|
||||
unsupported_mode: str = case() # a raw mode string did not parse into AuthSpecKind (boundary)
|
||||
precondition_required: str = case() # a required per-user value (e.g. an env var) has not been provided -> 412
|
||||
|
|
@ -129,6 +131,10 @@ class CredError:
|
|||
def of_misconfigured(detail: str) -> CredError:
|
||||
return CredError(misconfigured=detail)
|
||||
|
||||
@staticmethod
|
||||
def of_url_credentials_not_allowed() -> CredError:
|
||||
return CredError(url_credentials_not_allowed=None)
|
||||
|
||||
@staticmethod
|
||||
def of_upstream_unavailable(detail: str) -> CredError:
|
||||
return CredError(upstream_unavailable=detail)
|
||||
|
|
@ -154,6 +160,12 @@ class CredError:
|
|||
return f"unauthorized: {self.unauthorized.detail}"
|
||||
case "misconfigured":
|
||||
return f"misconfigured: {self.misconfigured}"
|
||||
case "url_credentials_not_allowed":
|
||||
return (
|
||||
"misconfigured: auth_type none cannot be used with credentials embedded in the upstream URL; "
|
||||
"remove them from the URL and configure Basic Auth with auth_type: basic and "
|
||||
"auth_value: username:password"
|
||||
)
|
||||
case "upstream_unavailable":
|
||||
return f"upstream unavailable: {self.upstream_unavailable}"
|
||||
case "unsupported_mode":
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm.exceptions import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPServerListError,
|
||||
MCPServerURLCredentialsError,
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
||||
|
|
@ -75,6 +76,8 @@ _MCP_GUARDRAIL_REJECTIONS: Final = (
|
|||
|
||||
|
||||
def _connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str:
|
||||
if isinstance(exc, MCPServerURLCredentialsError):
|
||||
return str(exc.detail)
|
||||
if isinstance(exc, TimeoutError):
|
||||
return (
|
||||
f"Failed to connect to MCP server: no response from {url or 'the server'} "
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ if TYPE_CHECKING:
|
|||
|
||||
|
||||
AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router"
|
||||
HEURISTIC_V2_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
|
||||
AUTO_ROUTER_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
|
||||
|
||||
|
||||
class LicenseCheck:
|
||||
|
|
@ -153,11 +153,12 @@ class LicenseCheck:
|
|||
return False
|
||||
return team_count > _max_teams_in_license
|
||||
|
||||
def heuristic_v2_router_limit(self) -> int | None:
|
||||
def auto_router_capability_limit(self) -> int | None:
|
||||
"""
|
||||
How many heuristic_v2 auto-routers this proxy may hold: unlimited (None) only when the
|
||||
signed license lists the auto_router feature, otherwise one. A license verified through
|
||||
the API carries no feature list, so it does not lift the limit either.
|
||||
How many auto-routers may claim each licensed capability (heuristic_v2, operator-defined
|
||||
tier_definitions): unlimited (None) only when the signed license lists the auto_router
|
||||
feature, otherwise one per capability. A license verified through the API carries no
|
||||
feature list, so it does not lift the limit either.
|
||||
"""
|
||||
if self.airgapped_license_data is None:
|
||||
return 1
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from litellm.proxy.common_utils.sse_keepalive import (
|
|||
wrap_sse_stream_with_keepalive_pings,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
|
||||
from litellm.router import Router
|
||||
|
|
@ -2004,6 +2005,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
trust_client_model_info=False,
|
||||
)
|
||||
|
||||
# An auto router with its own compression policy is authoritative for this
|
||||
# request: suppress every other compression guardrail and arm whichever one
|
||||
# the policy names for the model call, before those guardrails get a chance
|
||||
# to run below.
|
||||
await _arm_auto_router_compression(data=self.data, llm_router=llm_router)
|
||||
|
||||
self.data = await proxy_logging_obj.pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=self.data,
|
||||
|
|
|
|||
|
|
@ -199,6 +199,12 @@ class PrismaDBExceptionHandler:
|
|||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def is_prisma_error(e: Exception) -> bool:
|
||||
import prisma
|
||||
|
||||
return isinstance(e, _exception_types(prisma.errors.PrismaError))
|
||||
|
||||
@staticmethod
|
||||
def is_deadlock_error(e: Exception) -> bool:
|
||||
"""True iff ``e`` is a Postgres deadlock (P2034 / 40P01) surfaced through prisma."""
|
||||
|
|
|
|||
267
litellm/proxy/guardrails/auto_router_compression.py
Normal file
267
litellm/proxy/guardrails/auto_router_compression.py
Normal file
|
|
@ -0,0 +1,267 @@
|
|||
"""
|
||||
Decouples prompt compression between an auto router's routing decision and the model
|
||||
it routes to, via ``auto_router_routing_compression`` / ``auto_router_model_compression``
|
||||
on the marker deployment: a guardrail name, or ``"none"``.
|
||||
|
||||
Neither key set inherits today's behaviour. Either key set makes the auto router
|
||||
authoritative and suppresses every other compression guardrail for that request.
|
||||
"""
|
||||
|
||||
import contextvars
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.router_utils.auto_router_model_naming import AUTO_ROUTER_MODEL_PREFIX
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.router import Router
|
||||
|
||||
COMPRESSION_GUARDRAIL_PROVIDERS: Final = frozenset({"headroom", "compresr"})
|
||||
_NO_COMPRESSION: Final = "none"
|
||||
|
||||
# A ContextVar, not metadata: metadata reaches spend logs the caller can read, and a
|
||||
# suppression list they can read is one they can replay to disable any guardrail.
|
||||
_suppressed_compression_guardrails: Final[contextvars.ContextVar[frozenset[str]]] = contextvars.ContextVar(
|
||||
"litellm_auto_router_suppressed_compression_guardrails", default=frozenset()
|
||||
)
|
||||
|
||||
|
||||
def suppressed_compression_guardrails() -> frozenset[str]:
|
||||
"""Names of the compression guardrails this request's auto router suppresses."""
|
||||
return _suppressed_compression_guardrails.get()
|
||||
|
||||
|
||||
# Only the proxy calls `arm_pre_call`, so on the SDK path nothing arms and nothing
|
||||
# compresses; the router must not assume the model hop already ran.
|
||||
_model_hop_armed: Final[contextvars.ContextVar[bool]] = contextvars.ContextVar(
|
||||
"litellm_auto_router_model_hop_armed", default=False
|
||||
)
|
||||
|
||||
|
||||
def model_hop_compression_armed() -> bool:
|
||||
"""True when this request's model-side compression guardrail was actually armed."""
|
||||
return _model_hop_armed.get()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AutoRouterCompressionPolicy:
|
||||
"""An auto router's compression choice for each hop. ``None`` means no compression."""
|
||||
|
||||
routing: str | None
|
||||
model: str | None
|
||||
|
||||
@property
|
||||
def is_same(self) -> bool:
|
||||
return self.routing == self.model
|
||||
|
||||
|
||||
def _normalized_compression_choice(raw: object) -> str | None:
|
||||
if not isinstance(raw, str) or not raw:
|
||||
return None
|
||||
return None if raw.strip().lower() == _NO_COMPRESSION else raw
|
||||
|
||||
|
||||
def policy_from_litellm_params(litellm_params: Mapping[str, object]) -> AutoRouterCompressionPolicy | None:
|
||||
raw_routing: Final = litellm_params.get("auto_router_routing_compression")
|
||||
raw_model: Final = litellm_params.get("auto_router_model_compression")
|
||||
if raw_routing is None and raw_model is None:
|
||||
return None
|
||||
return AutoRouterCompressionPolicy(
|
||||
routing=_normalized_compression_choice(raw_routing),
|
||||
model=_normalized_compression_choice(raw_model),
|
||||
)
|
||||
|
||||
|
||||
def policy_for_model(
|
||||
llm_router: "Router | None",
|
||||
model_alias: str,
|
||||
team_id: str | None,
|
||||
request_tags: Sequence[str],
|
||||
) -> AutoRouterCompressionPolicy | None:
|
||||
"""The compression policy of the auto router marker `model_alias` resolves to.
|
||||
|
||||
Pre-call arming and the routing hook both resolve through here, so an alias with
|
||||
several tag-scoped markers cannot suppress under one and then route under another.
|
||||
"""
|
||||
if llm_router is None:
|
||||
return None
|
||||
deployments: Final = llm_router.get_model_list(model_name=model_alias, team_id=team_id) or ()
|
||||
markers: Final = tuple(
|
||||
litellm_params
|
||||
for deployment in deployments
|
||||
if isinstance(litellm_params := deployment.get("litellm_params"), Mapping) # pyright: ignore[reportUnnecessaryIsInstance] # filters out non-Mapping
|
||||
and str(litellm_params.get("model", "")).startswith(AUTO_ROUTER_MODEL_PREFIX)
|
||||
)
|
||||
requested: Final = frozenset(request_tags)
|
||||
tag_matched: Final = tuple(
|
||||
params for params in markers if (tags := params.get("tags")) and requested.issuperset(frozenset(tags))
|
||||
)
|
||||
# Untagged only: a marker scoped to tags this request lacks describes other traffic.
|
||||
untagged: Final = tuple(params for params in markers if not params.get("tags"))
|
||||
# Lazy, so the first marker carrying a policy wins and the rest are never read.
|
||||
candidates: Final = (policy_from_litellm_params(params) for params in (*tag_matched, *untagged))
|
||||
return next((policy for policy in candidates if policy is not None), None)
|
||||
|
||||
|
||||
def team_id_from_request(request_kwargs: Mapping[str, object]) -> str | None:
|
||||
"""The caller's team id, from whichever metadata bucket this surface writes to."""
|
||||
for meta_key in ("metadata", "litellm_metadata"):
|
||||
meta = request_kwargs.get(meta_key)
|
||||
if isinstance(meta, Mapping):
|
||||
team_id = meta.get("user_api_key_team_id")
|
||||
if isinstance(team_id, str):
|
||||
return team_id
|
||||
return None
|
||||
|
||||
|
||||
def _compression_guardrail_classes() -> tuple[type, ...]:
|
||||
"""The registered guardrail classes whose provider compresses prompts."""
|
||||
from litellm.proxy.guardrails.guardrail_registry import guardrail_class_registry
|
||||
|
||||
return tuple(cls for name, cls in guardrail_class_registry.items() if name in COMPRESSION_GUARDRAIL_PROVIDERS)
|
||||
|
||||
|
||||
def is_compression_guardrail(guardrail: object) -> bool:
|
||||
"""Whether `guardrail` is an instance of a compression guardrail provider.
|
||||
|
||||
Both hops validate through here: the policy fields are operator-supplied names, and
|
||||
an unvalidated one would get handed the conversation and invoked.
|
||||
"""
|
||||
classes: Final = _compression_guardrail_classes()
|
||||
return bool(classes) and isinstance(guardrail, classes)
|
||||
|
||||
|
||||
def _active_compression_guardrails() -> tuple["CustomGuardrail", ...]:
|
||||
"""Every currently-active guardrail whose type is a compression guardrail."""
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
||||
if not _compression_guardrail_classes():
|
||||
return ()
|
||||
active: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomGuardrail)
|
||||
return tuple(cb for cb in active if is_compression_guardrail(cb) and cb.guardrail_name)
|
||||
|
||||
|
||||
async def arm_pre_call(
|
||||
data: dict[str, object], # mutable-ok: arms the live request dict in place
|
||||
llm_router: "Router | None",
|
||||
) -> None:
|
||||
"""Apply an auto router's compression policy, if any, before guardrails run.
|
||||
|
||||
Suppresses every other compression guardrail and re-enables the model-side
|
||||
guardrail the policy names (if any) even when it isn't ``default_on``.
|
||||
"""
|
||||
_suppressed_compression_guardrails.set(frozenset())
|
||||
_model_hop_armed.set(False)
|
||||
if llm_router is None:
|
||||
return
|
||||
|
||||
model_alias: Final = data.get("model")
|
||||
if not isinstance(model_alias, str) or not model_alias:
|
||||
return
|
||||
|
||||
from litellm.router_strategy.tag_based_routing import (
|
||||
_get_tags_from_request_kwargs, # pyright: ignore[reportPrivateUsage] # used in router.py and budget_limiter.py too
|
||||
)
|
||||
|
||||
policy: Final = policy_for_model(
|
||||
llm_router=llm_router,
|
||||
model_alias=model_alias,
|
||||
team_id=team_id_from_request(data),
|
||||
request_tags=_get_tags_from_request_kwargs(data),
|
||||
)
|
||||
if policy is None:
|
||||
return
|
||||
|
||||
_suppressed_compression_guardrails.set(
|
||||
frozenset(
|
||||
name
|
||||
for guardrail in _active_compression_guardrails()
|
||||
if (name := guardrail.guardrail_name) and name != policy.model
|
||||
)
|
||||
)
|
||||
|
||||
# Arming adds the name to `metadata["guardrails"]`, which runs it even if not default_on.
|
||||
armed_model_hop: Final = policy.model is not None and any(
|
||||
guardrail.guardrail_name == policy.model for guardrail in _active_compression_guardrails()
|
||||
)
|
||||
if policy.model is not None and not armed_model_hop:
|
||||
verbose_proxy_logger.warning(
|
||||
"AutoRouter compression: '%s' is not an active compression guardrail; the model hop is uncompressed",
|
||||
policy.model,
|
||||
)
|
||||
|
||||
if armed_model_hop:
|
||||
_model_hop_armed.set(True)
|
||||
_, metadata = get_or_create_metadata_bucket(data)
|
||||
requested: Final = metadata.get("guardrails")
|
||||
existing: Final = tuple(requested) if isinstance(requested, (list, tuple)) else ()
|
||||
if policy.model not in existing:
|
||||
# A list: litellm_pre_call_utils isinstance-checks this key and drops a tuple.
|
||||
metadata["guardrails"] = [*existing, policy.model] # mutable-ok: this key's contract is a list
|
||||
|
||||
|
||||
def _as_routing_messages(
|
||||
messages: Iterable[Mapping[str, object]],
|
||||
) -> list[dict[str, object]]: # mutable-ok: shape fixed by the pre-routing hook protocol
|
||||
"""A fresh, independently mutable copy, the shape the pre-routing hook takes."""
|
||||
return [dict(message) for message in messages] # mutable-ok: shape fixed by the pre-routing hook protocol
|
||||
|
||||
|
||||
async def messages_for_routing(
|
||||
policy: AutoRouterCompressionPolicy | None,
|
||||
# list[dict], not Sequence[Mapping]: fixed by the async_pre_routing_hook protocol.
|
||||
messages: list[dict[str, object]] | None, # mutable-ok: shape fixed by the pre-routing hook protocol
|
||||
request_kwargs: Mapping[str, object],
|
||||
) -> list[dict[str, object]] | None: # mutable-ok: shape fixed by the pre-routing hook protocol
|
||||
"""Messages to use for a routing decision, per `policy.routing`. None means the
|
||||
caller should route on whatever it already has.
|
||||
|
||||
Reads the live messages, never a pre-guardrail copy: this compresses through a real
|
||||
guardrail that POSTs the text out, so routing on a pre-masking snapshot would leak
|
||||
what the masking guardrail stripped. When the model hop already compressed and the
|
||||
hops differ, routing therefore reads the compressed text rather than the original.
|
||||
"""
|
||||
if policy is None or policy.routing is None:
|
||||
return None
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
from litellm.proxy.common_utils.registry_read_through import (
|
||||
get_initialized_guardrail_with_read_through,
|
||||
)
|
||||
|
||||
guardrail: Final = await get_initialized_guardrail_with_read_through(policy.routing)
|
||||
if guardrail is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"AutoRouter compression: guardrail '%s' not found; routing on uncompressed messages", policy.routing
|
||||
)
|
||||
return _as_routing_messages(messages)
|
||||
|
||||
if not is_compression_guardrail(guardrail):
|
||||
verbose_proxy_logger.warning(
|
||||
"AutoRouter compression: guardrail '%s' is not a compression guardrail; routing on uncompressed messages",
|
||||
policy.routing,
|
||||
)
|
||||
return _as_routing_messages(messages)
|
||||
|
||||
inputs: Final[GenericGuardrailAPIInputs] = {
|
||||
"structured_messages": _as_routing_messages(messages) # pyright: ignore[reportAssignmentType] # plain dicts, not AllMessageValues; see headroom.py's own use of this shape
|
||||
}
|
||||
model: Final = request_kwargs.get("model")
|
||||
# Throwaway: apply_guardrail writes stats here, so routing never double-counts into
|
||||
# extract_compression_saved_tokens.
|
||||
stats_sink: Final = {"messages": messages, "model": model} # mutable-ok: apply_guardrail writes its stats here
|
||||
result: Final = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=stats_sink,
|
||||
input_type="request",
|
||||
)
|
||||
compressed: Final = result.get("structured_messages")
|
||||
return compressed if isinstance(compressed, list) else _as_routing_messages(messages)
|
||||
|
|
@ -64,6 +64,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
||||
id_jag_assertion_capture_gap,
|
||||
)
|
||||
from litellm.proxy.management_helpers.audit_logs import (
|
||||
get_audit_log_changed_by,
|
||||
is_audit_logging_enabled,
|
||||
|
|
@ -272,6 +275,22 @@ if MCP_AVAILABLE:
|
|||
_validate_mcp_server_name_fields(payload)
|
||||
_validate_upstream_token_header(payload)
|
||||
|
||||
def warn_if_id_jag_server_outruns_sso(server_id: str | None, auth_type: MCPAuth | str | None) -> None:
|
||||
"""Registering an ``oauth2_id_jag`` server under an SSO provider that captures no IdP
|
||||
identity assertion is a dead configuration: nothing here fails, and then every ID-JAG call
|
||||
fails for every user with a message that only ever tells them to sign in again. Say it once,
|
||||
at the moment the admin can still act on it."""
|
||||
if auth_type != MCPAuth.oauth2_id_jag:
|
||||
return
|
||||
gap = id_jag_assertion_capture_gap()
|
||||
if gap is None:
|
||||
return
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s is registered with auth_type=oauth2_id_jag, but %s.",
|
||||
server_id,
|
||||
gap,
|
||||
)
|
||||
|
||||
def stamp_omitted_oauth2_flow(payload: NewMCPServerRequest) -> None:
|
||||
"""Fallback only: fill in oauth2_flow when an oauth2 create omits it.
|
||||
|
||||
|
|
@ -1623,6 +1642,8 @@ if MCP_AVAILABLE:
|
|||
detail={"error": f"Error creating mcp server: {e}"},
|
||||
)
|
||||
|
||||
warn_if_id_jag_server_outruns_sso(new_mcp_server.server_id, new_mcp_server.auth_type)
|
||||
|
||||
# Registry refresh is best-effort: the row is already committed, so a
|
||||
# failure here (e.g. an unrelated malformed row in the table) must not
|
||||
# surface as a 500 and orphan the created server, which would push the
|
||||
|
|
@ -2726,6 +2747,7 @@ if MCP_AVAILABLE:
|
|||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail={"error": f"MCP Server not found, passed server_id={payload.server_id}"},
|
||||
)
|
||||
warn_if_id_jag_server_outruns_sso(mcp_server_record_updated.server_id, mcp_server_record_updated.auth_type)
|
||||
await global_mcp_server_manager.update_server(mcp_server_record_updated)
|
||||
|
||||
# Ensure registry is up to date by reloading from database
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ from litellm.proxy._types import (
|
|||
TeamModelDeleteRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY
|
||||
from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_REMEDY
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import (
|
||||
coordination_redis_cache,
|
||||
|
|
@ -98,11 +98,13 @@ from litellm.router_strategy.complexity_router import (
|
|||
normalize_classification_prompt,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
GATED_AUTO_ROUTER_CAPABILITIES,
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
capability_limit_violation,
|
||||
carries_complexity_router_settings,
|
||||
count_heuristic_v2_routers,
|
||||
heuristic_v2_limit_violation,
|
||||
uses_heuristic_v2_classifier,
|
||||
count_capability_routers,
|
||||
gated_capability_of,
|
||||
is_complexity_router_model,
|
||||
validate_complexity_router_config_placement,
|
||||
validate_complexity_router_config_write,
|
||||
validate_strategy_router_model_write,
|
||||
|
|
@ -237,11 +239,13 @@ def _strategy_router_write_violation(
|
|||
An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is
|
||||
the discriminator the router loads it by; a write that mangles it makes the
|
||||
router drop the deployment silently under ``ignore_invalid_deployments``.
|
||||
Only writes that supply ``litellm_params.model`` are judged on the naming
|
||||
contract, against the merged (stored + incoming) params, so partial patches
|
||||
and restores of an already-corrupted row stay legal. A config is judged only
|
||||
when the write carries one, for the same reason: a rename must not be held
|
||||
hostage by a stored config it does not touch. Returns the violation, or None.
|
||||
A patch adding auto-router settings is judged against the effective model,
|
||||
decrypting the stored model when the patch omits it, so a regular deployment
|
||||
cannot claim a strategy-router configuration. Unrelated partial patches and
|
||||
restores that do not touch strategy-router settings stay legal. A config is
|
||||
judged only when the write carries one, for the same reason: a rename must
|
||||
not be held hostage by a stored config it does not touch. Returns the
|
||||
violation, or None.
|
||||
"""
|
||||
if incoming_params is None:
|
||||
return None
|
||||
|
|
@ -256,14 +260,18 @@ def _strategy_router_write_violation(
|
|||
for source in (incoming_params, existing_params)
|
||||
if source is not None and getattr(source, field, None) is not None
|
||||
)
|
||||
# Scope reads the incoming model because the stored one is encrypted at rest.
|
||||
if carries_complexity_router_settings(incoming_params.model, present_fields):
|
||||
effective_params: Final = _effective_complexity_router_params(incoming_params, existing_params)
|
||||
effective_model: Final = effective_params.get("model")
|
||||
if carries_complexity_router_settings(
|
||||
effective_model if isinstance(effective_model, str) else None, present_fields
|
||||
):
|
||||
placement_violation: Final = validate_complexity_router_config_placement(incoming_params.model_extra)
|
||||
if placement_violation is not None:
|
||||
return placement_violation
|
||||
if incoming_params.model is None:
|
||||
return None
|
||||
return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields)
|
||||
return validate_strategy_router_model_write(
|
||||
model=effective_model if isinstance(effective_model, str) else "",
|
||||
present_fields=present_fields,
|
||||
)
|
||||
|
||||
|
||||
def _raise_on_strategy_router_write_violation(
|
||||
|
|
@ -281,14 +289,23 @@ def _raise_on_strategy_router_write_violation(
|
|||
)
|
||||
|
||||
|
||||
HEURISTIC_V2_SLOT_LOCK_KEY: Final = 5_872_301
|
||||
_HEURISTIC_V2_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)"
|
||||
_HEURISTIC_V2_DB_ROWS_SQL: Final = """
|
||||
SELECT count(*)::int AS held FROM "LiteLLM_ProxyModelTable"
|
||||
AUTO_ROUTER_CAPABILITY_SLOT_LOCK_KEY: Final = 5_872_301
|
||||
_CAPABILITY_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)"
|
||||
_STORED_LITELLM_PARAMS_SQL: Final = (
|
||||
"(CASE jsonb_typeof(litellm_params) WHEN 'string' THEN (litellm_params #>> '{}')::jsonb ELSE litellm_params END)"
|
||||
)
|
||||
_STORED_COMPLEXITY_CONFIG_SQL: Final = f"{_STORED_LITELLM_PARAMS_SQL} -> 'complexity_router_config'"
|
||||
_CAPABILITY_DB_ROWS_SQL: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
capability.key: f"""
|
||||
SELECT {_STORED_LITELLM_PARAMS_SQL} ->> 'model' AS model
|
||||
FROM "LiteLLM_ProxyModelTable"
|
||||
WHERE model_id <> $1
|
||||
AND (CASE jsonb_typeof(litellm_params) WHEN 'string' THEN (litellm_params #>> '{}')::jsonb ELSE litellm_params END)
|
||||
-> 'complexity_router_config' ->> 'classifier_type' = 'heuristic_v2'
|
||||
AND ({capability.sql_config_predicate.format(config=_STORED_COMPLEXITY_CONFIG_SQL)})
|
||||
"""
|
||||
for capability in GATED_AUTO_ROUTER_CAPABILITIES
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _effective_complexity_router_config(
|
||||
|
|
@ -301,13 +318,44 @@ def _effective_complexity_router_config(
|
|||
return existing_params.complexity_router_config
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _heuristic_v2_slot(
|
||||
prisma_client: PrismaClient, *, effective_config: object, model_id: str | None
|
||||
) -> AsyncGenerator[_ProxyModelTable, None]:
|
||||
"""Hand out the model table to write through while the row's claim on a heuristic_v2 slot is settled.
|
||||
def _effective_model(
|
||||
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
|
||||
) -> str | None:
|
||||
"""The model a write leaves on the row, decrypting an existing value only when the patch omits it."""
|
||||
incoming: Final = None if incoming_params is None else incoming_params.model
|
||||
if incoming is not None:
|
||||
return incoming
|
||||
existing: Final = None if existing_params is None else existing_params.model
|
||||
if existing is None:
|
||||
return None
|
||||
decrypted: Final = decrypt_value_helper(
|
||||
value=existing,
|
||||
key="model",
|
||||
exception_type="debug",
|
||||
return_original_value=True,
|
||||
)
|
||||
return decrypted if isinstance(decrypted, str) else None
|
||||
|
||||
A write that leaves the row on classifier_type heuristic_v2 under a limited license runs
|
||||
|
||||
def _effective_complexity_router_params(
|
||||
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
|
||||
) -> Mapping[str, object]:
|
||||
"""The model and complexity config a write leaves, for placement and capability decisions."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
"model": _effective_model(incoming_params, existing_params),
|
||||
"complexity_router_config": _effective_complexity_router_config(incoming_params, existing_params),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _auto_router_capability_slot(
|
||||
prisma_client: PrismaClient, *, effective_params: Mapping[str, object], model_id: str | None
|
||||
) -> AsyncGenerator[_ProxyModelTable, None]:
|
||||
"""Hand out the model table to write through while the row's claim on a licensed capability is settled.
|
||||
|
||||
A write that leaves the row claiming a licensed capability under a limited license runs
|
||||
inside one transaction that takes an advisory lock in its own statement before counting
|
||||
(a statement's snapshot predates anything it locks), so pods cannot both pass the count:
|
||||
the DB rows (any pod, either JSON shape) plus this proxy's config.yaml routers are judged
|
||||
|
|
@ -321,21 +369,37 @@ async def _heuristic_v2_slot(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import _license_check, llm_router
|
||||
|
||||
limit: Final = _license_check.heuristic_v2_router_limit()
|
||||
if limit is None or not uses_heuristic_v2_classifier(effective_config):
|
||||
limit: Final = _license_check.auto_router_capability_limit()
|
||||
capability: Final = gated_capability_of(effective_params)
|
||||
if limit is None or capability is None:
|
||||
yield _proxy_model_table(prisma_client)
|
||||
return
|
||||
async with prisma_client.db.tx() as tx_ctx:
|
||||
tables: Final[_TxModelTables] = tx_ctx
|
||||
await tx_ctx.query_raw(_HEURISTIC_V2_LOCK_SQL, HEURISTIC_V2_SLOT_LOCK_KEY)
|
||||
rows: Sequence[Mapping[str, object]] = await tx_ctx.query_raw(_HEURISTIC_V2_DB_ROWS_SQL, model_id or "")
|
||||
db_held: Final = rows[0].get("held") if rows else 0
|
||||
await tx_ctx.query_raw(_CAPABILITY_LOCK_SQL, AUTO_ROUTER_CAPABILITY_SLOT_LOCK_KEY)
|
||||
rows: Sequence[Mapping[str, object]] = await tx_ctx.query_raw(
|
||||
_CAPABILITY_DB_ROWS_SQL[capability.key], model_id or ""
|
||||
)
|
||||
db_held: Final = sum(
|
||||
1
|
||||
for row in rows
|
||||
for stored_model in (row.get("model"),)
|
||||
if isinstance(stored_model, str)
|
||||
and is_complexity_router_model(
|
||||
decrypt_value_helper(
|
||||
value=stored_model,
|
||||
key="model",
|
||||
exception_type="debug",
|
||||
return_original_value=True,
|
||||
)
|
||||
)
|
||||
)
|
||||
config_rows: Final = () if llm_router is None else tuple(llm_router.config_deployments())
|
||||
held: Final = (db_held if isinstance(db_held, int) else 0) + count_heuristic_v2_routers(config_rows)
|
||||
violation: Final = heuristic_v2_limit_violation(held=held + 1, limit=limit)
|
||||
held: Final = db_held + count_capability_routers(config_rows, capability=capability)
|
||||
violation: Final = capability_limit_violation(capability=capability, held=held + 1, limit=limit)
|
||||
if violation is not None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail=f"{violation} {HEURISTIC_V2_LICENSE_REMEDY}"
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail=f"{violation} {AUTO_ROUTER_LICENSE_REMEDY}"
|
||||
)
|
||||
yield tables.litellm_proxymodeltable
|
||||
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
|
||||
|
|
@ -791,6 +855,9 @@ async def patch_model(
|
|||
existing_params=db_model.litellm_params,
|
||||
)
|
||||
|
||||
effective_params: Final = _effective_complexity_router_params(
|
||||
patch_data.litellm_params, db_model.litellm_params
|
||||
)
|
||||
requested_model_name: Final = patch_data.model_name
|
||||
stored_model_name: str | None = None
|
||||
|
||||
|
|
@ -799,11 +866,9 @@ async def patch_model(
|
|||
stored_model_name = update_data.get("model_name")
|
||||
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
update_data["updated_at"] = cast(str, get_utc_datetime())
|
||||
async with _heuristic_v2_slot(
|
||||
async with _auto_router_capability_slot(
|
||||
prisma_client,
|
||||
effective_config=_effective_complexity_router_config(
|
||||
patch_data.litellm_params, db_model.litellm_params
|
||||
),
|
||||
effective_params=effective_params,
|
||||
model_id=model_id,
|
||||
) as table:
|
||||
return await table.update(where={"model_id": model_id}, data=update_data)
|
||||
|
|
@ -1959,9 +2024,12 @@ async def add_new_model(
|
|||
model_params=priced_model_params,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
slot=_heuristic_v2_slot(
|
||||
slot=_auto_router_capability_slot(
|
||||
prisma_client,
|
||||
effective_config=priced_model_params.litellm_params.complexity_router_config,
|
||||
effective_params=_effective_complexity_router_params(
|
||||
priced_model_params.litellm_params,
|
||||
None,
|
||||
),
|
||||
model_id=priced_model_params.model_info.id,
|
||||
),
|
||||
)
|
||||
|
|
@ -2110,6 +2178,9 @@ async def update_model(
|
|||
incoming_params=model_params.litellm_params,
|
||||
existing_params=deployment.litellm_params,
|
||||
)
|
||||
effective_params: Final = _effective_complexity_router_params(
|
||||
model_params.litellm_params, deployment.litellm_params
|
||||
)
|
||||
|
||||
# update DB
|
||||
if store_model_in_db is True:
|
||||
|
|
@ -2147,11 +2218,9 @@ async def update_model(
|
|||
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
|
||||
**({} if renamed_to is None else {"model_name": renamed_to}),
|
||||
}
|
||||
async with _heuristic_v2_slot(
|
||||
async with _auto_router_capability_slot(
|
||||
prisma_client,
|
||||
effective_config=_effective_complexity_router_config(
|
||||
model_params.litellm_params, deployment.litellm_params
|
||||
),
|
||||
effective_params=effective_params,
|
||||
model_id=_model_id,
|
||||
) as table:
|
||||
model_response: Final = await table.update(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,81 @@
|
|||
"""Whether the SSO provider the login callback dispatches to can capture an IdP identity assertion.
|
||||
|
||||
An ``oauth2_id_jag`` MCP server spends the ``id_token`` captured at SSO login as its RFC 8693
|
||||
subject token. Only the generic OIDC login path reaches a token response the gateway retains one
|
||||
from, so a deployment whose SSO runs through Google, Microsoft or SAML never stores an assertion
|
||||
and every store-sourced ID-JAG exchange fails for every user, however many times they sign in.
|
||||
Neither side can see that alone: the MCP registration knows nothing about SSO and the login knows
|
||||
nothing about MCP. This module is the one shared answer both warn from.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from enum import Enum
|
||||
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
|
||||
|
||||
_GENERIC_OIDC_REMEDY = (
|
||||
"Point SSO at the generic OIDC provider (GENERIC_CLIENT_ID), the one login path whose token "
|
||||
"response the gateway retains an id_token from"
|
||||
)
|
||||
|
||||
|
||||
class ActiveSSOProvider(str, Enum):
|
||||
google = "google"
|
||||
microsoft = "microsoft"
|
||||
generic = "generic"
|
||||
saml = "saml"
|
||||
none = "none"
|
||||
|
||||
|
||||
def active_sso_provider() -> ActiveSSOProvider:
|
||||
"""The provider the SSO callback will dispatch to.
|
||||
|
||||
Mirrors the callback's precedence rather than reporting everything configured: an environment
|
||||
carrying both GOOGLE_CLIENT_ID and GENERIC_CLIENT_ID runs the Google branch, so it must report
|
||||
Google. Presence is judged the way the callback judges it, so a client id set to the empty
|
||||
string still selects that branch here.
|
||||
"""
|
||||
if os.getenv("GOOGLE_CLIENT_ID") is not None:
|
||||
return ActiveSSOProvider.google
|
||||
if os.getenv("MICROSOFT_CLIENT_ID") is not None:
|
||||
return ActiveSSOProvider.microsoft
|
||||
if os.getenv("GENERIC_CLIENT_ID") is not None:
|
||||
return ActiveSSOProvider.generic
|
||||
if SAMLAuthHandler.is_saml_configured():
|
||||
return ActiveSSOProvider.saml
|
||||
return ActiveSSOProvider.none
|
||||
|
||||
|
||||
def id_jag_assertion_capture_gap() -> str | None:
|
||||
"""Why ID-JAG cannot work under the active SSO provider, phrased for an operator reading a log,
|
||||
or ``None`` when that provider does capture an assertion."""
|
||||
provider = active_sso_provider()
|
||||
match provider:
|
||||
case ActiveSSOProvider.generic:
|
||||
return None
|
||||
case ActiveSSOProvider.none:
|
||||
return (
|
||||
"no SSO provider is configured, so no IdP identity assertion is ever captured and "
|
||||
f"ID-JAG credential resolution fails for every user. {_GENERIC_OIDC_REMEDY}"
|
||||
)
|
||||
case ActiveSSOProvider.google | ActiveSSOProvider.microsoft | ActiveSSOProvider.saml:
|
||||
return (
|
||||
f"the active SSO provider ({provider.value}) has no identity-assertion capture path, so no "
|
||||
"IdP id_token is ever stored and ID-JAG credential resolution fails for every user no matter "
|
||||
f"how often they sign in. {_GENERIC_OIDC_REMEDY}"
|
||||
)
|
||||
case _:
|
||||
assert_never(provider)
|
||||
|
||||
|
||||
def id_jag_assertion_capture_gap_at_startup() -> str | None:
|
||||
"""Config load runs before SSO settings stored in the database are reconciled into the process
|
||||
environment, so an unresolved provider at that point is not yet a gap; the SSO callback reports it
|
||||
once a login happens."""
|
||||
if active_sso_provider() is ActiveSSOProvider.none:
|
||||
return None
|
||||
return id_jag_assertion_capture_gap()
|
||||
|
|
@ -16,7 +16,7 @@ import json
|
|||
import os
|
||||
import re
|
||||
import secrets
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from copy import deepcopy
|
||||
from html import escape
|
||||
from types import MappingProxyType
|
||||
|
|
@ -29,6 +29,7 @@ from typing import (
|
|||
NoReturn,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
Union,
|
||||
cast,
|
||||
overload,
|
||||
|
|
@ -70,6 +71,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
SSOIdentityAssertion,
|
||||
assertion_from_sso_login,
|
||||
ema_assertion_retention_enabled,
|
||||
retain_sso_identity_assertion_for_ema,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -105,6 +107,9 @@ from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
|||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
||||
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
|
||||
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
||||
id_jag_assertion_capture_gap,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
|
||||
from litellm.proxy.management_endpoints.sso_helper_utils import (
|
||||
check_is_admin_only_access,
|
||||
|
|
@ -1677,6 +1682,46 @@ async def get_generic_sso_response(
|
|||
return result or {}, received_response, access_token_payload, sso_assertion
|
||||
|
||||
|
||||
RetentionCheck: TypeAlias = Callable[[], Awaitable[bool]] # mutable-ok: Callable parameter syntax
|
||||
|
||||
|
||||
async def warn_if_id_jag_assertion_uncaptured(
|
||||
assertion: SSOIdentityAssertion | None, *, retention_enabled: RetentionCheck | None = None
|
||||
) -> None:
|
||||
"""Say, at the one moment it is knowable, that this login gave an ``oauth2_id_jag`` server
|
||||
nothing to spend. Without it the operator only ever sees the per-request failure, which cannot
|
||||
tell a user who has never signed in from a provider that will never capture. Kept strictly
|
||||
diagnostic: a store outage is swallowed, since a login must not fail over a log line."""
|
||||
if assertion is not None:
|
||||
return
|
||||
try:
|
||||
check: Final = retention_enabled if retention_enabled is not None else ema_assertion_retention_enabled
|
||||
if not await check():
|
||||
return
|
||||
except Exception as exc: # noqa: BLE001 # diagnostics must never break the login
|
||||
verbose_proxy_logger.debug("Could not check for oauth2_id_jag MCP servers after SSO login: %s", exc)
|
||||
return
|
||||
gap: Final = id_jag_assertion_capture_gap()
|
||||
verbose_proxy_logger.warning(
|
||||
"SSO login captured no IdP identity assertion while an oauth2_id_jag MCP server is registered: %s",
|
||||
gap if gap is not None else "the identity provider's token response carried no usable id_token",
|
||||
)
|
||||
|
||||
|
||||
async def warn_if_id_jag_capture_gap(*, retention_enabled: RetentionCheck | None = None) -> None:
|
||||
gap: Final = id_jag_assertion_capture_gap()
|
||||
if gap is None:
|
||||
return
|
||||
try:
|
||||
check: Final = retention_enabled if retention_enabled is not None else ema_assertion_retention_enabled
|
||||
if not await check():
|
||||
return
|
||||
except Exception as exc: # noqa: BLE001 # diagnostics must never break the page they annotate
|
||||
verbose_proxy_logger.debug("Could not check for oauth2_id_jag MCP servers: %s", exc)
|
||||
return
|
||||
verbose_proxy_logger.warning("SSO debug callback ran with an oauth2_id_jag capture gap: %s", gap)
|
||||
|
||||
|
||||
async def create_team_member_add_task(team_id, user_info):
|
||||
"""Create a task for adding a member to a team."""
|
||||
try:
|
||||
|
|
@ -2269,6 +2314,7 @@ async def _complete_cli_sso_callback_session(
|
|||
raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO")
|
||||
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion)
|
||||
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
|
||||
|
||||
teams: list[str] = []
|
||||
if hasattr(user_info, "teams") and user_info.teams:
|
||||
|
|
@ -3599,6 +3645,7 @@ class SSOAuthenticationHandler:
|
|||
|
||||
if isinstance(user_id, str) and user_id:
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
|
||||
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
|
||||
|
||||
disabled_non_admin_personal_key_creation: Final = get_disabled_non_admin_personal_key_creation()
|
||||
litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/")
|
||||
|
|
@ -4733,6 +4780,7 @@ async def debug_sso_callback(request: Request):
|
|||
safe_raw_claims: Final = {k: v for k, v in (received_response or {}).items() if k not in _OAUTH_TOKEN_FIELDS}
|
||||
safe_access_token_claims = {k: v for k, v in (access_token_payload or {}).items() if k not in _OAUTH_TOKEN_FIELDS}
|
||||
|
||||
await warn_if_id_jag_capture_gap()
|
||||
sso_payload: Final = {
|
||||
"parsed_by_proxy": filtered_result,
|
||||
"raw_claims": safe_raw_claims,
|
||||
|
|
|
|||
|
|
@ -118,10 +118,11 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
get_hidden_params_dict,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
GATED_AUTO_ROUTER_CAPABILITIES,
|
||||
STRATEGY_ROUTER_PARAM_FIELDS,
|
||||
capability_limit_violation,
|
||||
carries_complexity_router_settings,
|
||||
count_heuristic_v2_routers,
|
||||
heuristic_v2_limit_violation,
|
||||
count_capability_routers,
|
||||
validate_complexity_router_config_placement,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -303,7 +304,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
)
|
||||
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY, LicenseCheck
|
||||
from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_REMEDY, LicenseCheck
|
||||
from litellm.proxy.auth.model_checks import (
|
||||
expand_wildcard_deployments_for_model_info,
|
||||
get_all_fallbacks,
|
||||
|
|
@ -4340,17 +4341,28 @@ def validate_deployment_complexity_router_placement(model: Mapping[str, object])
|
|||
raise ValueError(f"model {model.get('model_name', '')!r}: {violation}")
|
||||
|
||||
|
||||
def validate_heuristic_v2_router_limit(model_list: Sequence[Mapping[str, object]], *, limit: int | None) -> None:
|
||||
def validate_auto_router_capability_limits(model_list: Sequence[Mapping[str, object]], *, limit: int | None) -> None:
|
||||
"""
|
||||
Refuse to start when config.yaml defines more heuristic_v2 auto-routers than the license allows.
|
||||
Refuse to start when config.yaml defines more auto-routers claiming a licensed capability than allowed.
|
||||
|
||||
Checked here rather than left to router registration for the same reason as the two
|
||||
validators above: the proxy builds its router with `ignore_invalid_deployments=True`, so
|
||||
the router's own refusal would turn the extra router into a silently missing model.
|
||||
"""
|
||||
violation: Final = heuristic_v2_limit_violation(held=count_heuristic_v2_routers(model_list), limit=limit)
|
||||
if violation is not None:
|
||||
raise ValueError(f"config.yaml model_list: {violation} {HEURISTIC_V2_LICENSE_REMEDY}")
|
||||
violations: Final = tuple(
|
||||
message
|
||||
for capability in GATED_AUTO_ROUTER_CAPABILITIES
|
||||
if (
|
||||
message := capability_limit_violation(
|
||||
capability=capability,
|
||||
held=count_capability_routers(model_list, capability=capability),
|
||||
limit=limit,
|
||||
)
|
||||
)
|
||||
is not None
|
||||
)
|
||||
if violations:
|
||||
raise ValueError(f"config.yaml model_list: {' '.join(violations)} {AUTO_ROUTER_LICENSE_REMEDY}")
|
||||
|
||||
|
||||
def pin_complexity_router_model_id(model: dict) -> None: # mutable-ok: out-param, model_info is stamped in place
|
||||
|
|
@ -5758,7 +5770,7 @@ class ProxyConfig:
|
|||
model_list: Final = config.get("model_list", None)
|
||||
if model_list:
|
||||
router_params["model_list"] = model_list
|
||||
validate_heuristic_v2_router_limit(model_list, limit=_license_check.heuristic_v2_router_limit())
|
||||
validate_auto_router_capability_limits(model_list, limit=_license_check.auto_router_capability_limit())
|
||||
print( # noqa: T201
|
||||
"\033[32mLiteLLM: Proxy initialized with Config, Set models:\033[0m"
|
||||
)
|
||||
|
|
@ -5848,7 +5860,7 @@ class ProxyConfig:
|
|||
),
|
||||
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
|
||||
auto_router_capability_limit=_license_check.auto_router_capability_limit,
|
||||
)
|
||||
|
||||
if redis_usage_cache is not None and router.cache.redis_cache is None:
|
||||
|
|
@ -6309,7 +6321,7 @@ class ProxyConfig:
|
|||
search_tools=search_tools,
|
||||
ignore_invalid_deployments=True,
|
||||
fallback_access_check=router_fallback_access_check,
|
||||
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
|
||||
auto_router_capability_limit=_license_check.auto_router_capability_limit,
|
||||
)
|
||||
verbose_proxy_logger.debug("updated llm_router: %s", llm_router)
|
||||
else:
|
||||
|
|
@ -11458,6 +11470,37 @@ def _realtime_query_params_template(model: str | None, intent: str | None) -> tu
|
|||
return tuple(params)
|
||||
|
||||
|
||||
async def _release_realtime_budget_reservation(user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
release_or_invalidate_budget_reservation,
|
||||
)
|
||||
|
||||
await release_or_invalidate_budget_reservation(
|
||||
budget_reservation=user_api_key_dict.budget_reservation,
|
||||
)
|
||||
|
||||
|
||||
async def _reject_realtime_session(
|
||||
websocket: WebSocket,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
*,
|
||||
code: int,
|
||||
reason: str,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
try:
|
||||
if error_message is not None:
|
||||
try:
|
||||
await websocket.send_text(
|
||||
json.dumps({"type": "error", "error": {"type": "guardrail_error", "message": error_message}})
|
||||
)
|
||||
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
|
||||
verbose_proxy_logger.debug("Could not send realtime pre-call error event to client; closing anyway")
|
||||
await websocket.close(code=code, reason=reason)
|
||||
finally:
|
||||
await _release_realtime_budget_reservation(user_api_key_dict)
|
||||
|
||||
|
||||
@app.websocket("/openai/v1/realtime")
|
||||
@app.websocket("/v1/realtime")
|
||||
@app.websocket("/realtime")
|
||||
|
|
@ -11483,7 +11526,9 @@ async def realtime_websocket_endpoint(
|
|||
if intent == "transcription":
|
||||
route_model = "gpt-realtime-whisper"
|
||||
else:
|
||||
await websocket.close(code=1008, reason="model query parameter is required")
|
||||
await _reject_realtime_session(
|
||||
websocket, user_api_key_dict, code=1008, reason="model query parameter is required"
|
||||
)
|
||||
return
|
||||
assert route_model is not None
|
||||
try:
|
||||
|
|
@ -11494,7 +11539,7 @@ async def realtime_websocket_endpoint(
|
|||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyException as e:
|
||||
await websocket.close(code=1008, reason=e.message[:120])
|
||||
await _reject_realtime_session(websocket, user_api_key_dict, code=1008, reason=e.message[:120])
|
||||
return
|
||||
await websocket.accept(**accept_kwargs)
|
||||
|
||||
|
|
@ -11553,21 +11598,9 @@ async def realtime_websocket_endpoint(
|
|||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Realtime pre-call error")
|
||||
try:
|
||||
await websocket.send_text(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "guardrail_error",
|
||||
"message": str(e),
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
await websocket.close(code=1011, reason="Pre-call error")
|
||||
await _reject_realtime_session(
|
||||
websocket, user_api_key_dict, code=1011, reason="Pre-call error", error_message=str(e)
|
||||
)
|
||||
return
|
||||
|
||||
# Phase 2: route to upstream LLM.
|
||||
|
|
@ -11597,6 +11630,13 @@ async def realtime_websocket_endpoint(
|
|||
)
|
||||
except Exception: # noqa: BLE001 # the lower layer may have closed the socket already; closing twice is not an error
|
||||
verbose_proxy_logger.debug("Could not close realtime client websocket; it is already gone")
|
||||
finally:
|
||||
from litellm.litellm_core_utils.realtime_streaming import (
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
)
|
||||
|
||||
if not litellm_logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY):
|
||||
await _release_realtime_budget_reservation(user_api_key_dict)
|
||||
|
||||
|
||||
######################################################################
|
||||
|
|
@ -13501,6 +13541,27 @@ def _is_auto_router_model(model: Mapping[str, object]) -> bool:
|
|||
return isinstance(litellm_model, str) and litellm_model.startswith("auto_router/")
|
||||
|
||||
|
||||
def _model_in_access_group(model: Mapping[str, object], access_group: str) -> bool:
|
||||
model_info: Final = model.get("model_info")
|
||||
if not isinstance(model_info, Mapping):
|
||||
return False
|
||||
access_groups: Final = model_info.get("access_groups")
|
||||
return isinstance(access_groups, (list, tuple)) and access_group in access_groups
|
||||
|
||||
|
||||
def _matches_model_info_filters(
|
||||
model: Mapping[str, object],
|
||||
exclude_auto_routers: bool | None,
|
||||
access_group: str | None,
|
||||
wildcard_only: bool | None,
|
||||
) -> bool:
|
||||
if exclude_auto_routers is True and _is_auto_router_model(model):
|
||||
return False
|
||||
if isinstance(access_group, str) and not _model_in_access_group(model, access_group):
|
||||
return False
|
||||
return wildcard_only is not True or "*" in str(model.get("model_name") or "")
|
||||
|
||||
|
||||
def _paginate_models_response(
|
||||
all_models: list[dict[str, Any]],
|
||||
page: int,
|
||||
|
|
@ -13811,6 +13872,14 @@ async def model_info_v2(
|
|||
"existing callers are unaffected"
|
||||
),
|
||||
),
|
||||
access_group: str | None = fastapi.Query(
|
||||
None,
|
||||
description="Only return deployments whose `model_info.access_groups` contains this access group",
|
||||
),
|
||||
wildcard_only: bool | None = fastapi.Query(
|
||||
False,
|
||||
description="Only return wildcard deployments, i.e. those whose `model_name` contains `*`",
|
||||
),
|
||||
):
|
||||
"""
|
||||
Paginated model metadata for proxy deployments (pricing, provider, team access).
|
||||
|
|
@ -13828,6 +13897,8 @@ async def model_info_v2(
|
|||
modelId: Return a single deployment by LiteLLM model id.
|
||||
teamId: Filter to models with direct access or team membership for this team id.
|
||||
sortBy / sortOrder: Sort by model_name, created_at, updated_at, costs, or status.
|
||||
access_group: Only return deployments in this model access group.
|
||||
wildcard_only: Only return deployments whose `model_name` contains `*`.
|
||||
|
||||
Example request:
|
||||
```
|
||||
|
|
@ -13981,8 +14052,9 @@ async def model_info_v2(
|
|||
|
||||
# `is True` because direct-call tests bypass FastAPI, so the Query default arrives as a
|
||||
# truthy sentinel object rather than False.
|
||||
if exclude_auto_routers is True:
|
||||
all_models = [m for m in all_models if not _is_auto_router_model(m)]
|
||||
all_models = [
|
||||
m for m in all_models if _matches_model_info_filters(m, exclude_auto_routers, access_group, wildcard_only)
|
||||
]
|
||||
|
||||
# Update total count to include agents
|
||||
search_total_count = len(all_models)
|
||||
|
|
|
|||
|
|
@ -373,6 +373,33 @@ async def invalidate_budget_reservation_counters(
|
|||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
|
||||
|
||||
async def release_or_invalidate_budget_reservation(
|
||||
budget_reservation: dict | None, # mutable-ok: stamps finalized on the caller's shared reservation dict
|
||||
) -> None:
|
||||
"""Reconcile a still-open reservation on a terminal path that settles no cost.
|
||||
|
||||
A failed or upstream-refused request never runs the success cost callback, so
|
||||
its pre-call reservation stays open and keeps the spend counter pinned above
|
||||
real spend until the counter's TTL expires, 429ing later requests on the same
|
||||
key. Release it to zero; if the release itself fails (e.g. the counter store is
|
||||
unreachable) drop the reserved counters directly and mark the reservation
|
||||
finalized so nothing reprocesses it. Idempotent: the finalized guard makes a
|
||||
second call a no-op once success or failure handling already reconciled.
|
||||
"""
|
||||
if budget_reservation is None or budget_reservation.get("finalized") is True:
|
||||
return
|
||||
try:
|
||||
await asyncio.shield(release_budget_reservation(budget_reservation=budget_reservation))
|
||||
except Exception: # noqa: BLE001 # a cleanup failure must not pin the counter; drop it directly instead
|
||||
verbose_proxy_logger.exception("Failed to release budget reservation; invalidating counters")
|
||||
try:
|
||||
await invalidate_budget_reservation_counters(budget_reservation=budget_reservation)
|
||||
except Exception: # noqa: BLE001 # nothing left to try; the finalized stamp below keeps it from being reprocessed
|
||||
verbose_proxy_logger.exception("Failed to invalidate budget reservation counters after release failed")
|
||||
finally:
|
||||
budget_reservation["finalized"] = True
|
||||
|
||||
|
||||
async def _get_budget_counters(
|
||||
request_body: dict,
|
||||
valid_token: UserAPIKeyAuth,
|
||||
|
|
|
|||
|
|
@ -6386,10 +6386,17 @@ class ProxyUpdateSpend:
|
|||
)
|
||||
break
|
||||
except Exception as e:
|
||||
if not PrismaDBExceptionHandler.is_database_transport_error(e):
|
||||
if not _is_transient_spend_log_write_error(e):
|
||||
if PrismaDBExceptionHandler.is_prisma_error(e):
|
||||
await enqueue_spend_logs(prisma_client, logs_to_process, at_head=True)
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking - DB error writing spend logs, requeued %d rows for the next flush. error=%s",
|
||||
len(logs_to_process),
|
||||
str(e),
|
||||
)
|
||||
raise
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking - DB connection error writing spend logs, retry %d/%d. logs_count=%d, error=%s",
|
||||
"Spend tracking - transient DB error writing spend logs, retry %d/%d. logs_count=%d, error=%s",
|
||||
i + 1,
|
||||
n_retry_times,
|
||||
len(logs_to_process),
|
||||
|
|
@ -6732,6 +6739,10 @@ async def _monitor_spend_logs_queue(
|
|||
MAX_SPEND_LOG_ISOLATION_FAILURES_PER_BATCH: Final = 256
|
||||
|
||||
|
||||
def _is_transient_spend_log_write_error(e: Exception) -> bool:
|
||||
return PrismaDBExceptionHandler.is_database_transport_error(e) or PrismaDBExceptionHandler.is_deadlock_error(e)
|
||||
|
||||
|
||||
async def _create_spend_logs_with_poison_isolation(
|
||||
repo: SpendLogsRepository,
|
||||
rows: Sequence[Mapping[str, object]],
|
||||
|
|
@ -6767,6 +6778,8 @@ async def _create_spend_logs_with_poison_isolation(
|
|||
raise
|
||||
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
|
||||
raise
|
||||
if PrismaDBExceptionHandler.is_deadlock_error(e):
|
||||
raise
|
||||
budget_left: Final = max(failure_budget - 1, 0)
|
||||
if len(rows) == 1:
|
||||
request_id: Final = rows[0].get("request_id")
|
||||
|
|
|
|||
|
|
@ -27,7 +27,7 @@ from litellm.types.realtime import (
|
|||
RealtimeTranscriptionSessionRequest,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.types.utils import CallTypes, LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
from ..litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
|
|
@ -360,7 +360,7 @@ async def _arealtime(
|
|||
user: Final = kwargs.get("user", None)
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
litellm_params_dict: Final = {**get_litellm_params(**kwargs), CallTypes.arealtime.value: True}
|
||||
|
||||
model, _custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import contextvars
|
||||
from collections.abc import Coroutine, Generator, Iterable, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
|
||||
|
|
@ -23,6 +24,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.openai_like.responses.transformation import OpenAILikeResponsesConfig
|
||||
from litellm.responses.litellm_completion_transformation.handler import (
|
||||
LiteLLMCompletionTransformationHandler,
|
||||
)
|
||||
|
|
@ -403,8 +405,40 @@ def _bridges_to_chat_completions(
|
|||
return responses_api_provider_config is None or use_chat_completions_api is True
|
||||
|
||||
|
||||
def _deployment_passes_through_responses(model_info: object) -> bool:
|
||||
"""Whether ``model_info.supported_endpoints`` opts the deployment into native ``{api_base}/responses``."""
|
||||
if not isinstance(model_info, dict):
|
||||
return False
|
||||
supported_endpoints: Final = model_info.get("supported_endpoints")
|
||||
return isinstance(supported_endpoints, (list, tuple)) and "/v1/responses" in supported_endpoints
|
||||
|
||||
|
||||
def _deployment_model_info_after_prompt_swap(
|
||||
requested_provider: str | None, resolved_provider: str | None, model_info: object
|
||||
) -> object:
|
||||
"""Deployment metadata only describes the upstream while the prompt manager keeps its provider."""
|
||||
return model_info if resolved_provider == requested_provider else None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _AsyncPromptManagementOutcome:
|
||||
merged_optional_params: Mapping[str, object]
|
||||
deployment_model_info: object
|
||||
|
||||
|
||||
def _resolve_responses_api_provider_config(
|
||||
model: str, custom_llm_provider: str, model_info: object
|
||||
) -> BaseResponsesAPIConfig | None:
|
||||
provider_config: Final = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model, provider=custom_llm_provider
|
||||
)
|
||||
if provider_config is not None or not _deployment_passes_through_responses(model_info):
|
||||
return provider_config
|
||||
return OpenAILikeResponsesConfig()
|
||||
|
||||
|
||||
def _will_bridge_to_chat_completions(
|
||||
model: str, custom_llm_provider: str | None, use_chat_completions_api: bool
|
||||
model: str, custom_llm_provider: str | None, use_chat_completions_api: bool, model_info: object
|
||||
) -> bool:
|
||||
"""``_bridges_to_chat_completions`` for callers running before the provider config is resolved.
|
||||
|
||||
|
|
@ -418,9 +452,7 @@ def _will_bridge_to_chat_completions(
|
|||
if custom_llm_provider is None:
|
||||
return True
|
||||
return _bridges_to_chat_completions(
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=normalized_model[0], provider=custom_llm_provider
|
||||
),
|
||||
_resolve_responses_api_provider_config(normalized_model[0], custom_llm_provider, model_info),
|
||||
use_chat_completions_api or normalized_model[1],
|
||||
)
|
||||
|
||||
|
|
@ -527,7 +559,10 @@ async def aresponses(
|
|||
with _prompt_management_sees_a_provisional_message_list(
|
||||
kwargs,
|
||||
bridged=_will_bridge_to_chat_completions(
|
||||
model, custom_llm_provider, bool(kwargs.get("use_chat_completions_api"))
|
||||
model,
|
||||
custom_llm_provider,
|
||||
bool(kwargs.get("use_chat_completions_api")),
|
||||
kwargs.get("model_info"),
|
||||
),
|
||||
):
|
||||
(
|
||||
|
|
@ -552,6 +587,7 @@ async def aresponses(
|
|||
merged_input=merged_input,
|
||||
),
|
||||
)
|
||||
requested_provider: Final = custom_llm_provider
|
||||
if model != original_model:
|
||||
custom_llm_provider = _resolve_prompt_swapped_provider(
|
||||
original_model=original_model,
|
||||
|
|
@ -561,7 +597,12 @@ async def aresponses(
|
|||
prompt_id=prompt_id,
|
||||
)
|
||||
kwargs.pop("prompt_id", None)
|
||||
kwargs["_async_prompt_merged_params"] = merged_optional_params
|
||||
kwargs["_async_prompt_merged_params"] = _AsyncPromptManagementOutcome(
|
||||
merged_optional_params=merged_optional_params,
|
||||
deployment_model_info=_deployment_model_info_after_prompt_swap(
|
||||
requested_provider, custom_llm_provider, kwargs.get("model_info")
|
||||
),
|
||||
)
|
||||
|
||||
func: Final = partial(
|
||||
responses,
|
||||
|
|
@ -666,12 +707,14 @@ def _apply_prompt_management_to_responses_call(
|
|||
kwargs: dict[str, Any],
|
||||
local_vars: dict[str, object],
|
||||
use_chat_completions_api: bool,
|
||||
) -> tuple[str | ResponseInputParam, str, str | None]:
|
||||
async_merged: Final[Mapping[str, object] | None] = kwargs.pop("_async_prompt_merged_params", None)
|
||||
if async_merged is not None:
|
||||
for key, value in async_merged.items():
|
||||
) -> tuple[str | ResponseInputParam, str, str | None, object]:
|
||||
"""Returns the prompt-managed input, model and provider, plus the deployment metadata that still
|
||||
describes the upstream (``None`` once the prompt manager moved the request to another provider)."""
|
||||
async_outcome: Final[_AsyncPromptManagementOutcome | None] = kwargs.pop("_async_prompt_merged_params", None)
|
||||
if async_outcome is not None:
|
||||
for key, value in async_outcome.merged_optional_params.items():
|
||||
local_vars[key] = value
|
||||
return input, model, custom_llm_provider
|
||||
return input, model, custom_llm_provider, async_outcome.deployment_model_info
|
||||
|
||||
prompt_id: Final = cast(str | None, kwargs.get("prompt_id", None))
|
||||
prompt_variables: Final = cast(dict | None, kwargs.get("prompt_variables", None))
|
||||
|
|
@ -684,7 +727,9 @@ def _apply_prompt_management_to_responses_call(
|
|||
):
|
||||
with _prompt_management_sees_a_provisional_message_list(
|
||||
kwargs,
|
||||
bridged=_will_bridge_to_chat_completions(model, custom_llm_provider, use_chat_completions_api),
|
||||
bridged=_will_bridge_to_chat_completions(
|
||||
model, custom_llm_provider, use_chat_completions_api, kwargs.get("model_info")
|
||||
),
|
||||
):
|
||||
(
|
||||
model,
|
||||
|
|
@ -710,19 +755,28 @@ def _apply_prompt_management_to_responses_call(
|
|||
)
|
||||
local_vars["input"] = input
|
||||
local_vars["model"] = model
|
||||
if model != original_model:
|
||||
custom_llm_provider = _resolve_prompt_swapped_provider(
|
||||
resolved_provider: Final = (
|
||||
custom_llm_provider
|
||||
if model == original_model
|
||||
else _resolve_prompt_swapped_provider(
|
||||
original_model=original_model,
|
||||
swapped_model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs=kwargs,
|
||||
prompt_id=prompt_id,
|
||||
)
|
||||
local_vars["custom_llm_provider"] = custom_llm_provider
|
||||
)
|
||||
local_vars["custom_llm_provider"] = resolved_provider
|
||||
for key, value in merged_optional_params.items():
|
||||
local_vars[key] = value
|
||||
return (
|
||||
input,
|
||||
model,
|
||||
resolved_provider,
|
||||
_deployment_model_info_after_prompt_swap(custom_llm_provider, resolved_provider, kwargs.get("model_info")),
|
||||
)
|
||||
|
||||
return input, model, custom_llm_provider
|
||||
return input, model, custom_llm_provider, kwargs.get("model_info")
|
||||
|
||||
|
||||
# Opt-in via model id (mirrors the `responses/` prefix pattern on chat completions).
|
||||
|
|
@ -1052,7 +1106,7 @@ def responses(
|
|||
)
|
||||
local_vars["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
input, model, custom_llm_provider = _apply_prompt_management_to_responses_call(
|
||||
input, model, custom_llm_provider, deployment_model_info = _apply_prompt_management_to_responses_call(
|
||||
input=input,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
@ -1123,9 +1177,8 @@ def responses(
|
|||
if custom_llm_provider is None:
|
||||
responses_api_provider_config = None
|
||||
else:
|
||||
responses_api_provider_config = ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
responses_api_provider_config = _resolve_responses_api_provider_config(
|
||||
model, custom_llm_provider, deployment_model_info
|
||||
)
|
||||
|
||||
local_vars.update(kwargs)
|
||||
|
|
|
|||
|
|
@ -116,10 +116,11 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
)
|
||||
from litellm.router_utils.auto_router_model_naming import (
|
||||
AUTO_ROUTER_MODEL_PREFIX,
|
||||
GatedAutoRouterCapability,
|
||||
capability_limit_violation,
|
||||
claimed_capability,
|
||||
classify_strategy_router_model,
|
||||
count_heuristic_v2_routers,
|
||||
heuristic_v2_limit_violation,
|
||||
uses_heuristic_v2_classifier,
|
||||
count_capability_routers,
|
||||
)
|
||||
from litellm.router_utils.batch_utils import (
|
||||
_get_router_metadata_variable_name,
|
||||
|
|
@ -208,6 +209,7 @@ from litellm.types.router import (
|
|||
AlertingConfig,
|
||||
AllowedFailsPolicy,
|
||||
AssistantsTypedDict,
|
||||
AutoRouterCapabilityLimit,
|
||||
ConsumedRequestTagsStamp,
|
||||
CredentialLiteLLMParams,
|
||||
CustomRoutingStrategyBase,
|
||||
|
|
@ -215,7 +217,6 @@ from litellm.types.router import (
|
|||
DeploymentTypedDict,
|
||||
FallbackAccessCheck,
|
||||
GuardrailTypedDict,
|
||||
HeuristicV2RouterLimit,
|
||||
LiteLLM_Params,
|
||||
MockRouterTestingParams,
|
||||
ModelGroupInfo,
|
||||
|
|
@ -692,7 +693,7 @@ class Router:
|
|||
background_health_check_model_groups: Sequence[str] | None = None,
|
||||
enable_weighted_failover: bool = False,
|
||||
fallback_access_check: FallbackAccessCheck | None = None,
|
||||
heuristic_v2_router_limit: HeuristicV2RouterLimit | None = None,
|
||||
auto_router_capability_limit: AutoRouterCapabilityLimit | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
|
||||
|
|
@ -769,7 +770,7 @@ class Router:
|
|||
|
||||
self.set_verbose = set_verbose
|
||||
self.ignore_invalid_deployments = ignore_invalid_deployments
|
||||
self.heuristic_v2_router_limit = heuristic_v2_router_limit
|
||||
self.auto_router_capability_limit = auto_router_capability_limit
|
||||
self.fallback_access_check: Final = fallback_access_check
|
||||
self.debug_level = debug_level
|
||||
self.enable_pre_call_checks = enable_pre_call_checks
|
||||
|
|
@ -2596,14 +2597,20 @@ class Router:
|
|||
model_response: CustomStreamWrapper,
|
||||
messages: list[dict[str, str]],
|
||||
initial_kwargs: dict,
|
||||
deployment_slot: contextlib.AsyncExitStack | None = None,
|
||||
) -> CustomStreamWrapper:
|
||||
"""
|
||||
Helper to iterate over a streaming response.
|
||||
|
||||
Catches errors for fallbacks using the router's fallback system
|
||||
|
||||
`deployment_slot` holds the deployment's max_parallel_requests semaphore; it is
|
||||
released when the stream is exhausted, closed, or falls back to another deployment
|
||||
"""
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
|
||||
held_slot: Final = deployment_slot if deployment_slot is not None else contextlib.AsyncExitStack()
|
||||
|
||||
class FallbackStreamWrapper(CustomStreamWrapper):
|
||||
def __init__(self, async_generator: AsyncGenerator):
|
||||
# Copy attributes from the original model_response
|
||||
|
|
@ -2627,12 +2634,26 @@ class Router:
|
|||
async def __anext__(self):
|
||||
return await self._async_generator.__anext__()
|
||||
|
||||
async def close_model_response() -> None:
|
||||
if not hasattr(model_response, "aclose"):
|
||||
return
|
||||
try:
|
||||
await model_response.aclose()
|
||||
except BaseException as e:
|
||||
verbose_router_logger.debug(
|
||||
"stream_with_fallbacks: error closing model_response: %s",
|
||||
e,
|
||||
)
|
||||
|
||||
async def stream_with_fallbacks():
|
||||
fallback_response = None # Track for cleanup in finally
|
||||
try:
|
||||
async for item in model_response:
|
||||
yield item
|
||||
except MidStreamFallbackError as e:
|
||||
with anyio.CancelScope(shield=True):
|
||||
await close_model_response()
|
||||
await held_slot.aclose()
|
||||
if not e.is_pre_first_chunk and (
|
||||
e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)
|
||||
):
|
||||
|
|
@ -2706,14 +2727,8 @@ class Router:
|
|||
# (e.g. on client disconnect).
|
||||
# Shield from anyio cancellation so the awaits can complete.
|
||||
with anyio.CancelScope(shield=True):
|
||||
if hasattr(model_response, "aclose"):
|
||||
try:
|
||||
await model_response.aclose()
|
||||
except BaseException as e:
|
||||
verbose_router_logger.debug(
|
||||
"stream_with_fallbacks: error closing model_response: %s",
|
||||
e,
|
||||
)
|
||||
await close_model_response()
|
||||
await held_slot.aclose()
|
||||
if fallback_response is not None and hasattr(fallback_response, "aclose"):
|
||||
try:
|
||||
await fallback_response.aclose()
|
||||
|
|
@ -3378,61 +3393,53 @@ class Router:
|
|||
kwargs=kwargs,
|
||||
client_type="max_parallel_requests",
|
||||
)
|
||||
if rpm_semaphore is not None and isinstance(rpm_semaphore, asyncio.Semaphore):
|
||||
async with rpm_semaphore:
|
||||
"""
|
||||
- Check rpm limits before making the call
|
||||
- If allowed, increment the rpm limit (allows global value to be updated, concurrency-safe)
|
||||
"""
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment,
|
||||
logging_obj=logging_obj,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
response = await _response
|
||||
else:
|
||||
async with contextlib.AsyncExitStack() as deployment_slot:
|
||||
if isinstance(rpm_semaphore, asyncio.Semaphore):
|
||||
await deployment_slot.enter_async_context(rpm_semaphore)
|
||||
await self.async_routing_strategy_pre_call_checks(
|
||||
deployment=deployment,
|
||||
logging_obj=logging_obj,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
response = await _response
|
||||
|
||||
## CHECK CONTENT FILTER ERROR ##
|
||||
if isinstance(response, ModelResponse):
|
||||
_should_raise = self._should_raise_content_policy_error(model=model, response=response, kwargs=kwargs)
|
||||
if _should_raise:
|
||||
raise litellm.ContentPolicyViolationError(
|
||||
message="Response output was blocked.",
|
||||
model=model,
|
||||
llm_provider="",
|
||||
## CHECK CONTENT FILTER ERROR ##
|
||||
if isinstance(response, ModelResponse):
|
||||
_should_raise = self._should_raise_content_policy_error(
|
||||
model=model, response=response, kwargs=kwargs
|
||||
)
|
||||
if _should_raise:
|
||||
raise litellm.ContentPolicyViolationError(
|
||||
message="Response output was blocked.",
|
||||
model=model,
|
||||
llm_provider="",
|
||||
)
|
||||
|
||||
if (
|
||||
isinstance(response, CustomStreamWrapper)
|
||||
and response.completion_stream is None
|
||||
and response.make_call is not None
|
||||
):
|
||||
await response.fetch_stream()
|
||||
if (
|
||||
isinstance(response, CustomStreamWrapper)
|
||||
and response.completion_stream is None
|
||||
and response.make_call is not None
|
||||
):
|
||||
await response.fetch_stream()
|
||||
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name)
|
||||
# debug how often this deployment picked
|
||||
self._track_deployment_metrics(
|
||||
deployment=deployment,
|
||||
response=response,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
if isinstance(response, CustomStreamWrapper):
|
||||
return await self._acompletion_streaming_iterator(
|
||||
model_response=response,
|
||||
messages=messages,
|
||||
initial_kwargs=input_kwargs_for_streaming_fallback,
|
||||
self.success_calls[model_name] += 1
|
||||
verbose_router_logger.info("litellm.acompletion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name)
|
||||
# debug how often this deployment picked
|
||||
self._track_deployment_metrics(
|
||||
deployment=deployment,
|
||||
response=response,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
return response
|
||||
if isinstance(response, CustomStreamWrapper):
|
||||
return await self._acompletion_streaming_iterator(
|
||||
model_response=response,
|
||||
messages=messages,
|
||||
initial_kwargs=input_kwargs_for_streaming_fallback,
|
||||
deployment_slot=deployment_slot.pop_all(),
|
||||
)
|
||||
|
||||
return response
|
||||
except litellm.Timeout as e:
|
||||
deployment_request_timeout_param: Final = _timeout_debug_deployment_dict.get("litellm_params", {}).get(
|
||||
"request_timeout", None
|
||||
|
|
@ -8373,17 +8380,12 @@ class Router:
|
|||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_failure_handler(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
end_time=time.time(),
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback.format_exc()),
|
||||
).start() # log response
|
||||
_set_cooldown_deployments(
|
||||
litellm_router_instance=self,
|
||||
exception_status=e.status_code,
|
||||
|
|
@ -8396,17 +8398,12 @@ class Router:
|
|||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_failure_handler(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
end_time=time.time(),
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback.format_exc()),
|
||||
).start() # log response
|
||||
raise e
|
||||
|
||||
async def async_callback_filter_deployments(
|
||||
|
|
@ -8444,17 +8441,12 @@ class Router:
|
|||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
logging_obj.async_failure_handler(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
end_time=time.time(),
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback.format_exc()),
|
||||
).start() # log response
|
||||
raise e
|
||||
return returned_healthy_deployments
|
||||
|
||||
|
|
@ -8644,7 +8636,7 @@ class Router:
|
|||
raise ValueError(ptu_error)
|
||||
zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None
|
||||
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(
|
||||
**(
|
||||
**( # pyright: ignore[reportArgumentType] # untyped merged dict; already true for every field here
|
||||
_litellm_params
|
||||
if zeroed_pricing is None
|
||||
else MappingProxyType({**_litellm_params, **zeroed_pricing})
|
||||
|
|
@ -8811,20 +8803,21 @@ class Router:
|
|||
if not (isinstance(model_info, Mapping) and model_info.get("db_model")):
|
||||
yield deployment
|
||||
|
||||
def heuristic_v2_router_limit_violation(self) -> str | None:
|
||||
def auto_router_capability_violation(self, capability: GatedAutoRouterCapability) -> str | None:
|
||||
"""
|
||||
Why one more heuristic_v2 router cannot join this router, or None when it can.
|
||||
Why one more router claiming ``capability`` cannot join this router, or None when it can.
|
||||
|
||||
Judged against every deployment currently on the model_list; an upsert pops the row being
|
||||
edited first, so an edit of an existing heuristic_v2 router keeps its own slot. The limit is
|
||||
resolved on every call through ``heuristic_v2_router_limit``; unset means unlimited, which
|
||||
is the SDK default, and the proxy injects a resolver backed by its license.
|
||||
edited first, so an edit of an existing gated router keeps its own slot. The limit is
|
||||
resolved on every call through ``auto_router_capability_limit``; unset means unlimited,
|
||||
which is the SDK default, and the proxy injects a resolver backed by its license.
|
||||
"""
|
||||
limit: Final = self.heuristic_v2_router_limit() if self.heuristic_v2_router_limit is not None else None
|
||||
others: Final = count_heuristic_v2_routers(
|
||||
deployment for deployment in self.model_list if isinstance(deployment, Mapping)
|
||||
limit: Final = self.auto_router_capability_limit() if self.auto_router_capability_limit is not None else None
|
||||
others: Final = count_capability_routers(
|
||||
(deployment for deployment in self.model_list if isinstance(deployment, Mapping)),
|
||||
capability=capability,
|
||||
)
|
||||
return heuristic_v2_limit_violation(held=others + 1, limit=limit)
|
||||
return capability_limit_violation(capability=capability, held=others + 1, limit=limit)
|
||||
|
||||
def init_complexity_router_deployment(self, deployment: Deployment):
|
||||
"""
|
||||
|
|
@ -8843,8 +8836,9 @@ class Router:
|
|||
)
|
||||
|
||||
complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config
|
||||
if uses_heuristic_v2_classifier(complexity_router_config):
|
||||
limit_violation: Final = self.heuristic_v2_router_limit_violation()
|
||||
capability: Final = claimed_capability(complexity_router_config)
|
||||
if capability is not None:
|
||||
limit_violation: Final = self.auto_router_capability_violation(capability)
|
||||
if limit_violation is not None:
|
||||
raise ValueError(limit_violation)
|
||||
|
||||
|
|
@ -9674,13 +9668,13 @@ class Router:
|
|||
"""Put a deployment back the way it was before a failed upsert popped it.
|
||||
|
||||
A rollback re-admits state that was already serving, so it does not go through the
|
||||
heuristic_v2 ceiling a newcomer gets: with the ceiling tightened since the deployment first
|
||||
capability ceiling a newcomer gets: with the ceiling tightened since the deployment first
|
||||
registered, judging the rollback would drop a serving router over an unrelated failed edit.
|
||||
"""
|
||||
if previous_deployment is None or self.has_model_id(model_id):
|
||||
return
|
||||
limit_resolver: Final = self.heuristic_v2_router_limit
|
||||
self.heuristic_v2_router_limit = None
|
||||
limit_resolver: Final = self.auto_router_capability_limit
|
||||
self.auto_router_capability_limit = None
|
||||
try:
|
||||
self.add_deployment(deployment=previous_deployment)
|
||||
verbose_router_logger.info(
|
||||
|
|
@ -9696,7 +9690,7 @@ class Router:
|
|||
restore_error,
|
||||
)
|
||||
finally:
|
||||
self.heuristic_v2_router_limit = limit_resolver
|
||||
self.auto_router_capability_limit = limit_resolver
|
||||
|
||||
@staticmethod
|
||||
def _backend_cost_map_keys(model: str, custom_llm_provider: str | None) -> tuple[str, ...]:
|
||||
|
|
@ -12634,13 +12628,13 @@ class Router:
|
|||
logging_obj: Final = request_kwargs.get("litellm_logging_obj", None)
|
||||
|
||||
if logging_obj is not None:
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback_exception),
|
||||
).start() # log response
|
||||
# Handle any exceptions that might occur during streaming
|
||||
asyncio.create_task(logging_obj.async_failure_handler(e, traceback_exception))
|
||||
asyncio.create_task(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback_exception,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
async def async_get_available_deployment_for_pass_through(
|
||||
|
|
@ -12768,11 +12762,13 @@ class Router:
|
|||
if request_kwargs is not None:
|
||||
logging_obj: Final = request_kwargs.get("litellm_logging_obj", None)
|
||||
if logging_obj is not None:
|
||||
threading.Thread(
|
||||
target=logging_obj.failure_handler,
|
||||
args=(e, traceback_exception),
|
||||
).start()
|
||||
asyncio.create_task(logging_obj.async_failure_handler(e, traceback_exception))
|
||||
asyncio.create_task(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception=e,
|
||||
traceback_exception=traceback_exception,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
raise e
|
||||
|
||||
async def _run_routing_plugins(
|
||||
|
|
@ -13037,13 +13033,48 @@ class Router:
|
|||
)
|
||||
return None
|
||||
|
||||
pre_routing_hook_response: Final = await selected_strategy.strategy.async_pre_routing_hook(
|
||||
from litellm.proxy.guardrails.auto_router_compression import (
|
||||
messages_for_routing,
|
||||
model_hop_compression_armed,
|
||||
policy_for_model,
|
||||
team_id_from_request,
|
||||
)
|
||||
|
||||
# Same tag-aware lookup the proxy's pre-call arming used, so an alias with
|
||||
# several tag-scoped markers cannot suppress under one and route under another.
|
||||
compression_policy: Final = policy_for_model(
|
||||
llm_router=self,
|
||||
model_alias=registered_model_name,
|
||||
team_id=team_id_from_request(request_kwargs),
|
||||
request_tags=_get_tags_from_request_kwargs(request_kwargs),
|
||||
)
|
||||
# Shared compression already ran in the pre-call hook, so reuse it rather than
|
||||
# compressing twice. Conditional on arming having actually happened: only the
|
||||
# proxy arms, and on the SDK path the shortcut would skip both hops entirely.
|
||||
needs_independent_routing_compression: Final = compression_policy is not None and not (
|
||||
compression_policy.is_same and compression_policy.model is not None and model_hop_compression_armed()
|
||||
)
|
||||
routing_messages: Final = (
|
||||
await messages_for_routing(policy=compression_policy, messages=messages, request_kwargs=request_kwargs)
|
||||
if needs_independent_routing_compression
|
||||
else None
|
||||
)
|
||||
|
||||
routed: Final = await selected_strategy.strategy.async_pre_routing_hook(
|
||||
model=registered_model_name,
|
||||
request_kwargs=request_kwargs,
|
||||
messages=messages,
|
||||
messages=routing_messages if routing_messages is not None else messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
)
|
||||
# Routing-only compression must not leak into the response: the model call and
|
||||
# deployment-context filtering key off this field. Compared by value, since
|
||||
# pydantic rebuilds the list rather than keeping the object passed in.
|
||||
pre_routing_hook_response: Final = (
|
||||
routed.model_copy(update={"messages": messages}) # mutable-ok: pydantic's model_copy takes a dict
|
||||
if routed is not None and routing_messages is not None and routed.messages == routing_messages
|
||||
else routed
|
||||
)
|
||||
self._record_routing_decision(
|
||||
request_kwargs=request_kwargs,
|
||||
routing_decision=(pre_routing_hook_response.routing_decision if pre_routing_hook_response else None),
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ the router silently dropping the deployment at load time under
|
|||
``ignore_invalid_deployments``.
|
||||
"""
|
||||
|
||||
from collections.abc import Iterable, Mapping, Sequence
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
|
@ -81,6 +81,11 @@ def classify_strategy_router_model(model: str) -> StrategyRouterKind | None:
|
|||
return "semantic"
|
||||
|
||||
|
||||
def is_complexity_router_model(model: str | None) -> bool:
|
||||
"""Whether ``model`` selects the complexity-router implementation."""
|
||||
return classify_strategy_router_model(model or "") == "complexity"
|
||||
|
||||
|
||||
def _named(value: object, role: StrategyRouterDependencyRole) -> tuple[StrategyRouterDependency, ...]:
|
||||
"""One dependency from a scalar field, or none when it is absent or not a name."""
|
||||
return (StrategyRouterDependency(value, role),) if isinstance(value, str) and value else ()
|
||||
|
|
@ -168,20 +173,121 @@ def uses_heuristic_v2_classifier(complexity_router_config: object) -> bool:
|
|||
return _mapping(complexity_router_config).get("classifier_type") == "heuristic_v2"
|
||||
|
||||
|
||||
def is_heuristic_v2_router(litellm_params: Mapping[str, object]) -> bool:
|
||||
"""Whether this deployment is a complexity router that classifies with heuristic_v2."""
|
||||
return classify_strategy_router_model(str(litellm_params.get("model") or "")) == "complexity" and (
|
||||
uses_heuristic_v2_classifier(litellm_params.get("complexity_router_config"))
|
||||
def defines_custom_tiers(complexity_router_config: object) -> bool:
|
||||
"""Whether this complexity config replaces the built-in tier ladder with operator-defined tier_definitions.
|
||||
|
||||
Mirrors the SQL spelling on the capability record: only an actual array claims the capability,
|
||||
so an explicit JSON null or a malformed value does not.
|
||||
"""
|
||||
return isinstance(_mapping(complexity_router_config).get("tier_definitions"), (list, tuple))
|
||||
|
||||
|
||||
OPERATOR_CLASSIFIER_PROMPT_FIELDS: Final = ("classification_prompt", "classification_examples")
|
||||
|
||||
|
||||
def defines_custom_classifier_prompt(complexity_router_config: object) -> bool:
|
||||
"""Whether an operator wrote any part of this router's classifier prompt themselves.
|
||||
|
||||
Three spellings, all metered: a whole replacement prompt (``classifier_llm_config.system_prompt``),
|
||||
replacement opening instructions (``classification_prompt``), and replacement calibration examples
|
||||
(``classification_examples``). Choosing a shipped ``classification_rubric`` preset is not authoring.
|
||||
Scoped to the classifier types that actually call an LLM, which is also where the config validator
|
||||
accepts these fields: the heuristic scorers never read them.
|
||||
"""
|
||||
config: Final = _mapping(complexity_router_config)
|
||||
if config.get("classifier_type") not in LLM_CLASSIFIER_TYPES:
|
||||
return False
|
||||
return _mapping(config.get("classifier_llm_config")).get("system_prompt") is not None or any(
|
||||
config.get(field) is not None for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS
|
||||
)
|
||||
|
||||
|
||||
def count_heuristic_v2_routers(deployments: Iterable[Mapping[str, object]]) -> int:
|
||||
"""How many of ``deployments`` (router model_list entries or config.yaml rows) are heuristic_v2 routers."""
|
||||
return sum(1 for deployment in deployments if is_heuristic_v2_router(_mapping(deployment.get("litellm_params"))))
|
||||
def uses_custom_tier_or_classifier_prompt(complexity_router_config: object) -> bool:
|
||||
"""Whether this router replaces shipped tiers or its shipped classifier prompt."""
|
||||
return defines_custom_tiers(complexity_router_config) or defines_custom_classifier_prompt(complexity_router_config)
|
||||
|
||||
|
||||
def heuristic_v2_limit_violation(*, held: int, limit: int | None) -> str | None:
|
||||
"""Why holding ``held`` heuristic_v2 routers exceeds ``limit``, or None when it fits.
|
||||
_LLM_CLASSIFIER_TYPES_SQL: Final = ", ".join(f"'{name}'" for name in sorted(LLM_CLASSIFIER_TYPES))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GatedAutoRouterCapability:
|
||||
"""A complexity-router capability the license meters, in every spelling an enforcement point needs.
|
||||
|
||||
``uses`` and ``sql_config_predicate`` answer the same question, in process and in a DB count over
|
||||
stored ``litellm_params`` (``{config}`` is the caller's expression for the normalized
|
||||
``complexity_router_config`` jsonb, substituted as many times as the predicate needs); they live
|
||||
on one record so they cannot drift apart. ``subject`` and ``remedy`` build the shared refusal
|
||||
message. A validated config claims at most one capability, and the validator is what makes that
|
||||
true: tier_definitions rejects every heuristic classifier_type, and it also rejects the
|
||||
classifier system_prompt, which in turn only applies to the classifier types heuristic_v2 is not.
|
||||
"""
|
||||
|
||||
key: str
|
||||
subject: str
|
||||
remedy: str
|
||||
uses: Callable[[object], bool]
|
||||
sql_config_predicate: str
|
||||
|
||||
|
||||
HEURISTIC_V2_CAPABILITY: Final = GatedAutoRouterCapability(
|
||||
key="heuristic_v2",
|
||||
subject="with classifier_type 'heuristic_v2'",
|
||||
remedy="Use classifier_type 'heuristic' for this router or remove an existing heuristic_v2 router.",
|
||||
uses=uses_heuristic_v2_classifier,
|
||||
sql_config_predicate="{config} ->> 'classifier_type' = 'heuristic_v2'",
|
||||
)
|
||||
|
||||
_OPERATOR_PROMPT_FIELDS_SQL: Final = " OR ".join(
|
||||
f"{{config}} ->> '{field}' IS NOT NULL" for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS
|
||||
)
|
||||
|
||||
CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
|
||||
key="tier_or_classifier_prompt",
|
||||
subject="with operator-defined tier_definitions or an operator-written classifier prompt",
|
||||
remedy=(
|
||||
"Use the shipped tiers and classifier prompt for this router or remove an existing router "
|
||||
"with tier_definitions or its own classifier prompt."
|
||||
),
|
||||
uses=uses_custom_tier_or_classifier_prompt,
|
||||
sql_config_predicate=(
|
||||
"jsonb_typeof({config} -> 'tier_definitions') = 'array' OR "
|
||||
f"({{config}} ->> 'classifier_type' IN ({_LLM_CLASSIFIER_TYPES_SQL}) AND ("
|
||||
"{config} -> 'classifier_llm_config' ->> 'system_prompt' IS NOT NULL OR "
|
||||
f"{_OPERATOR_PROMPT_FIELDS_SQL}))"
|
||||
),
|
||||
)
|
||||
|
||||
GATED_AUTO_ROUTER_CAPABILITIES: Final = (HEURISTIC_V2_CAPABILITY, CUSTOMIZATION_CAPABILITY)
|
||||
|
||||
|
||||
def claimed_capability(complexity_router_config: object) -> GatedAutoRouterCapability | None:
|
||||
"""The licensed capability this complexity config claims, or None."""
|
||||
return next(
|
||||
(capability for capability in GATED_AUTO_ROUTER_CAPABILITIES if capability.uses(complexity_router_config)),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def gated_capability_of(litellm_params: Mapping[str, object]) -> GatedAutoRouterCapability | None:
|
||||
"""The licensed capability this deployment claims, or None unless it is a complexity router."""
|
||||
model: Final = litellm_params.get("model")
|
||||
if not is_complexity_router_model(model if isinstance(model, str) else None):
|
||||
return None
|
||||
return claimed_capability(litellm_params.get("complexity_router_config"))
|
||||
|
||||
|
||||
def count_capability_routers(
|
||||
deployments: Iterable[Mapping[str, object]], *, capability: GatedAutoRouterCapability
|
||||
) -> int:
|
||||
"""How many of ``deployments`` (router model_list entries or config.yaml rows) claim ``capability``."""
|
||||
return sum(
|
||||
1 for deployment in deployments if gated_capability_of(_mapping(deployment.get("litellm_params"))) is capability
|
||||
)
|
||||
|
||||
|
||||
def capability_limit_violation(*, capability: GatedAutoRouterCapability, held: int, limit: int | None) -> str | None:
|
||||
"""Why holding ``held`` routers claiming ``capability`` exceeds ``limit``, or None when it fits.
|
||||
|
||||
``limit`` None means unlimited. The message is shared by every enforcement point (config
|
||||
load, model writes, router registration) and stays SDK-neutral: it names the cap and what
|
||||
|
|
@ -190,8 +296,8 @@ def heuristic_v2_limit_violation(*, held: int, limit: int | None) -> str | None:
|
|||
if limit is None or held <= limit:
|
||||
return None
|
||||
return (
|
||||
f"At most {limit} auto-router(s) with classifier_type 'heuristic_v2' can be registered but this would make "
|
||||
f"{held}. Use classifier_type 'heuristic' for this router or remove an existing heuristic_v2 router."
|
||||
f"At most {limit} auto-router(s) {capability.subject} can be registered but this would make "
|
||||
f"{held}. {capability.remedy}"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -237,9 +343,7 @@ def carries_complexity_router_settings(model: str | None, present_fields: frozen
|
|||
``validate_strategy_router_model_write`` is judged on, so a router named only by its
|
||||
default model is in scope, and a field added to the table above is covered here for free.
|
||||
"""
|
||||
return classify_strategy_router_model(model or "") == "complexity" or bool(
|
||||
present_fields & _COMPLEXITY_ROUTER_FIELDS
|
||||
)
|
||||
return is_complexity_router_model(model) or bool(present_fields & _COMPLEXITY_ROUTER_FIELDS)
|
||||
|
||||
|
||||
def validate_complexity_router_config_placement(litellm_params: Mapping[str, object] | None) -> str | None:
|
||||
|
|
|
|||
|
|
@ -56,8 +56,9 @@ class PatternMatchRouter:
|
|||
This class will store a mapping for regex pattern: List[Deployments]
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, pattern_utils: type[PatternUtils] = PatternUtils):
|
||||
self.patterns: dict[str, list] = {}
|
||||
self._pattern_utils: Final = pattern_utils
|
||||
|
||||
def add_pattern(self, pattern: str, llm_deployment: dict):
|
||||
"""
|
||||
|
|
@ -69,9 +70,10 @@ class PatternMatchRouter:
|
|||
"""
|
||||
# Convert the pattern to a regex
|
||||
regex: Final = self._pattern_to_regex(pattern)
|
||||
if regex not in self.patterns:
|
||||
self.patterns[regex] = []
|
||||
self.patterns[regex].append(llm_deployment)
|
||||
if regex in self.patterns:
|
||||
self.patterns[regex].append(llm_deployment)
|
||||
return
|
||||
self.patterns = dict(self._pattern_utils.sorted_patterns({**self.patterns, regex: [llm_deployment]}))
|
||||
|
||||
def remove_deployment(self, model_id: str) -> None:
|
||||
"""
|
||||
|
|
@ -138,11 +140,12 @@ class PatternMatchRouter:
|
|||
if request is None:
|
||||
return None
|
||||
|
||||
sorted_patterns: Final = PatternUtils.sorted_patterns(self.patterns)
|
||||
regex_filtered_model_names: Final = (
|
||||
[self._pattern_to_regex(m) for m in filtered_model_names] if filtered_model_names is not None else []
|
||||
tuple(self._pattern_to_regex(m) for m in filtered_model_names)
|
||||
if filtered_model_names is not None
|
||||
else ()
|
||||
)
|
||||
for pattern, llm_deployments in sorted_patterns:
|
||||
for pattern, llm_deployments in self.patterns.items():
|
||||
if filtered_model_names is not None and pattern not in regex_filtered_model_names:
|
||||
continue
|
||||
pattern_match = re.match(pattern, request)
|
||||
|
|
|
|||
|
|
@ -361,6 +361,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
auto_router_default_model: str | None = None
|
||||
auto_router_embedding_model: str | None = None
|
||||
auto_router_max_input_chars: int | None = None
|
||||
# Compression policy for the two hops of a routed request. Both unset means the
|
||||
# request's own compression guardrails apply to both, as they always have.
|
||||
auto_router_routing_compression: str | None = None
|
||||
auto_router_model_compression: str | None = None
|
||||
|
||||
# complexity-router params
|
||||
complexity_router_config: dict | None = None
|
||||
|
|
@ -888,9 +892,9 @@ class FallbackAccessCheck(Protocol):
|
|||
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
|
||||
|
||||
|
||||
class HeuristicV2RouterLimit(Protocol):
|
||||
class AutoRouterCapabilityLimit(Protocol):
|
||||
"""
|
||||
Resolves how many heuristic_v2 complexity routers the Router may hold right now; None means unlimited.
|
||||
Resolves how many complexity routers may claim each licensed capability right now; None means unlimited.
|
||||
|
||||
The Router calls it on every registration and limit query instead of caching the answer, so the
|
||||
proxy can keep the limit on its license object (re-verified on config load) rather than hand
|
||||
|
|
|
|||
|
|
@ -3769,6 +3769,8 @@ all_litellm_params = (
|
|||
"auto_router_default_model",
|
||||
"auto_router_embedding_model",
|
||||
"auto_router_max_input_chars",
|
||||
"auto_router_routing_compression",
|
||||
"auto_router_model_compression",
|
||||
"complexity_router_config",
|
||||
"complexity_router_default_model",
|
||||
"adaptive_router_config",
|
||||
|
|
|
|||
|
|
@ -67,8 +67,8 @@ proxy = [
|
|||
"azure-identity>=1.25.2,<2.0",
|
||||
"azure-storage-blob>=12.28.0,<13.0",
|
||||
"mcp>=1.28.1,<2.0",
|
||||
"litellm-proxy-extras==0.4.93",
|
||||
"litellm-enterprise==0.1.64",
|
||||
"litellm-proxy-extras==0.4.94",
|
||||
"litellm-enterprise==0.1.65",
|
||||
"RestrictedPython>=8.5,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
"InquirerPy>=0.3.4,<1.0",
|
||||
|
|
|
|||
|
|
@ -246,7 +246,7 @@
|
|||
"limit": 109
|
||||
},
|
||||
"TRY300": {
|
||||
"limit": 849
|
||||
"limit": 848
|
||||
},
|
||||
"UP028": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -392,6 +392,63 @@ with its own provider config (one `examples/default`-style root per project),
|
|||
or fork the module to add `configuration_aliases` and pass per-instance
|
||||
`providers = { ... }`.
|
||||
|
||||
## Dependencies only (run LiteLLM on GKE)
|
||||
|
||||
Set `create_runtime = false` to provision Cloud SQL, Memorystore, GCS,
|
||||
Secret Manager, and the runtime service account without Cloud Run or the
|
||||
load balancer. For a Shared VPC, set the full host-project network ID and
|
||||
skip PSA creation after the host project has configured it:
|
||||
|
||||
```hcl
|
||||
create_runtime = false
|
||||
network_id = "projects/<host>/global/networks/<vpc>"
|
||||
create_psa_connection = false
|
||||
```
|
||||
|
||||
The host project must already have Private Services Access configured on
|
||||
that network and the Service Networking API enabled; the module cannot set
|
||||
PSA up from a service project. GKE nodes must sit on the same Shared VPC so
|
||||
the Cloud SQL and Memorystore private IPs are routable from the pods. Run
|
||||
the root with its provider pointed at the project that should own the
|
||||
dependencies. `create_runtime = true` with `network_id` set is also allowed,
|
||||
but the Serverless VPC Access connector has to live in the same project as
|
||||
the network, so that combination only works when the VPC is in the
|
||||
deployment project
|
||||
|
||||
Map the outputs into the Helm values as follows:
|
||||
|
||||
```yaml
|
||||
database:
|
||||
writer:
|
||||
host: <cloudsql_writer_ip>
|
||||
dbname: <db_name>
|
||||
passwordSecret:
|
||||
name: <kubernetes-secret-with-db-credentials>
|
||||
reader:
|
||||
host: <cloudsql_reader_ip>
|
||||
dbname: <db_name>
|
||||
passwordSecret:
|
||||
name: <kubernetes-secret-with-db-credentials>
|
||||
redis:
|
||||
host: <redis_host>
|
||||
port: <redis_port>
|
||||
masterKey:
|
||||
secretName: <kubernetes-secret-with-master-key>
|
||||
```
|
||||
|
||||
Create the database Secret with keys `username` (the `db_username` output)
|
||||
and `password` (read it with `gcloud secrets versions access latest
|
||||
--secret=<db_password_secret_id>`), and the master key Secret from
|
||||
`master_key_secret_id` the same way. Memorystore only accepts TLS by
|
||||
default, so store the `redis_server_ca_pem` output in a third Secret,
|
||||
mount it into the gateway and backend pods via `volumes` / `volumeMounts`,
|
||||
and add `REDIS_SSL=true` and `REDIS_SSL_CA_CERTS=<mount path>` to each
|
||||
component's `extraEnv`. Setting `redis_transit_encryption = false` removes
|
||||
the CA plumbing at the cost of plaintext Redis traffic inside the VPC
|
||||
|
||||
The chart's pre-install/pre-upgrade migration hook runs the Prisma
|
||||
migration, so nothing replaces the Cloud Run migrations Job in this mode
|
||||
|
||||
## Storage and database retention
|
||||
|
||||
Two opt-in tripwires guard against accidental data loss on
|
||||
|
|
@ -409,14 +466,15 @@ Flip `cloudsql_deletion_protection` to `false` or `gcs_force_destroy` to
|
|||
|
||||
## Redis encryption
|
||||
|
||||
Memorystore runs with `transit_encryption_mode = "SERVER_AUTHENTICATION"`,
|
||||
so the proxy connects via `rediss://`. The instance's self-signed CA cert
|
||||
(`server_ca_certs[0].cert`) is shipped to gateway + backend as
|
||||
`REDIS_CA_PEM_B64`; their entrypoint shell decodes it to `/tmp/redis-ca.pem`
|
||||
before uvicorn starts and points `REDIS_SSL_CA_CERTS` at that path. No
|
||||
extra config needed — but if you ever swap Memorystore for an external
|
||||
Redis, override `REDIS_HOST`/`REDIS_PORT` and either drop these env vars
|
||||
or point them at your own CA.
|
||||
By default, Memorystore runs with
|
||||
`transit_encryption_mode = "SERVER_AUTHENTICATION"`, so Cloud Run connects
|
||||
via `rediss://`. The instance's self-signed CA cert
|
||||
(`server_ca_certs[0].cert`) is shipped to gateway and backend as
|
||||
`REDIS_CA_PEM_B64`; their entrypoint shell decodes it to
|
||||
`/tmp/redis-ca.pem` before uvicorn starts and points `REDIS_SSL_CA_CERTS` at
|
||||
that path. Set `redis_transit_encryption = false` to use plaintext Redis.
|
||||
For GKE, use `redis_server_ca_pem` as described in the dependencies-only
|
||||
section, or accept the security tradeoff of disabling transit encryption
|
||||
|
||||
## Files
|
||||
|
||||
|
|
@ -434,4 +492,5 @@ or point them at your own CA.
|
|||
| `iam.tf` | Runtime SA + Cloud SQL client + Secret Manager accessor |
|
||||
| `cloudrun.tf` | 3 Cloud Run services + Cloud Run Job for migrations |
|
||||
| `load_balancer.tf`| External HTTPS LB, serverless NEGs, URL map for path routing |
|
||||
| `outputs.tf` | LB IP, service URLs, secret IDs, migration `execute` command |
|
||||
| `outputs.tf` | LB IP, service URLs, dependency endpoints, secret IDs, migration command |
|
||||
| `tests/` | Plan-only mock-provider coverage for deployment modes and Redis encryption |
|
||||
|
|
|
|||
|
|
@ -15,15 +15,17 @@
|
|||
# enough to invoke Cloud Run admin APIs (`gcloud auth login`).
|
||||
|
||||
resource "terraform_data" "migration" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
triggers_replace = {
|
||||
job_id = google_cloud_run_v2_job.migrations.id
|
||||
job_id = google_cloud_run_v2_job.migrations[0].id
|
||||
job_image = local.migrations_image
|
||||
}
|
||||
|
||||
provisioner "local-exec" {
|
||||
interpreter = ["bash", "-c"]
|
||||
environment = {
|
||||
JOB = google_cloud_run_v2_job.migrations.name
|
||||
JOB = google_cloud_run_v2_job.migrations[0].name
|
||||
REGION = var.region
|
||||
PROJECT = var.project_id
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,25 +6,28 @@ locals {
|
|||
# Memorystore exposes a self-signed CA cert per instance; we ship it as
|
||||
# a base64 env var and decode it to a file at container startup so the
|
||||
# rediss:// connection can validate. Public cert, not sensitive.
|
||||
redis_ca_pem_b64 = base64encode(google_redis_instance.this.server_ca_certs[0].cert)
|
||||
redis_ca_pem_b64 = var.redis_transit_encryption ? base64encode(google_redis_instance.this.server_ca_certs[0].cert) : ""
|
||||
|
||||
shared_env_kv = [
|
||||
{ name = "DATABASE_HOST", value = google_sql_database_instance.writer.private_ip_address },
|
||||
{ name = "DATABASE_PORT", value = "5432" },
|
||||
{ name = "DATABASE_USER", value = var.db_username },
|
||||
{ name = "DATABASE_NAME", value = var.db_name },
|
||||
{ name = "DATABASE_HOST_READ_REPLICA", value = google_sql_database_instance.reader.private_ip_address },
|
||||
{ name = "DATABASE_PORT_READ_REPLICA", value = "5432" },
|
||||
{ name = "REDIS_HOST", value = google_redis_instance.this.host },
|
||||
{ name = "REDIS_PORT", value = tostring(google_redis_instance.this.port) },
|
||||
# _redis.get_redis_url_from_environment honors REDIS_SSL to flip the
|
||||
# scheme to rediss://; REDIS_SSL_CA_CERTS is mapped via
|
||||
# _get_redis_env_kwarg_mapping → ssl_ca_certs on the redis-py client.
|
||||
{ name = "REDIS_SSL", value = "true" },
|
||||
{ name = "REDIS_SSL_CA_CERTS", value = "/tmp/redis-ca.pem" },
|
||||
{ name = "REDIS_CA_PEM_B64", value = local.redis_ca_pem_b64 },
|
||||
{ name = "GCS_BUCKET_NAME", value = google_storage_bucket.this.name },
|
||||
]
|
||||
shared_env_kv = concat(
|
||||
[
|
||||
{ name = "DATABASE_HOST", value = google_sql_database_instance.writer.private_ip_address },
|
||||
{ name = "DATABASE_PORT", value = "5432" },
|
||||
{ name = "DATABASE_USER", value = var.db_username },
|
||||
{ name = "DATABASE_NAME", value = var.db_name },
|
||||
{ name = "DATABASE_HOST_READ_REPLICA", value = google_sql_database_instance.reader.private_ip_address },
|
||||
{ name = "DATABASE_PORT_READ_REPLICA", value = "5432" },
|
||||
{ name = "REDIS_HOST", value = google_redis_instance.this.host },
|
||||
{ name = "REDIS_PORT", value = tostring(google_redis_instance.this.port) },
|
||||
],
|
||||
var.redis_transit_encryption ? [
|
||||
{ name = "REDIS_SSL", value = "true" },
|
||||
{ name = "REDIS_SSL_CA_CERTS", value = "/tmp/redis-ca.pem" },
|
||||
{ name = "REDIS_CA_PEM_B64", value = local.redis_ca_pem_b64 },
|
||||
] : [],
|
||||
[
|
||||
{ name = "GCS_BUCKET_NAME", value = google_storage_bucket.this.name },
|
||||
],
|
||||
)
|
||||
|
||||
# OTel v2 is opt-in and gated on otel_endpoint, matching the AWS stack —
|
||||
# nothing OTel-related is added to the container env until an endpoint is
|
||||
|
|
@ -126,9 +129,9 @@ locals {
|
|||
# Decode the Memorystore CA cert (passed as REDIS_CA_PEM_B64) to the
|
||||
# path REDIS_SSL_CA_CERTS points at, so the redis-py client can validate
|
||||
# the rediss:// handshake.
|
||||
redis_ca_fragment = [
|
||||
redis_ca_fragment = var.redis_transit_encryption ? [
|
||||
"python -c \"import os, base64, pathlib; pathlib.Path(os.environ['REDIS_SSL_CA_CERTS']).write_bytes(base64.b64decode(os.environ['REDIS_CA_PEM_B64']))\""
|
||||
]
|
||||
] : []
|
||||
|
||||
database_url_fragment = [
|
||||
"export DATABASE_URL=\"postgresql://$${DATABASE_USER}:$${DATABASE_PASSWORD}@$${DATABASE_HOST}:$${DATABASE_PORT}/$${DATABASE_NAME}\"",
|
||||
|
|
@ -171,29 +174,7 @@ locals {
|
|||
|
||||
# ---------- Gateway ----------
|
||||
resource "google_cloud_run_v2_service" "gateway" {
|
||||
# Metering needs a client certificate AND its key. Each secret is created only
|
||||
# when its own PEM is supplied, so an endpoint set with a missing key would
|
||||
# otherwise apply cleanly and leave the proxy logging "missing config" and
|
||||
# never exporting. ca_cert_pem stays optional: empty means fall back to the
|
||||
# system trust store.
|
||||
#
|
||||
# The guard lives here, on an unconditional resource, rather than on the cert
|
||||
# secret: that secret is count-gated on the cert itself, so it has zero
|
||||
# instances in exactly the case this must catch. Adding count or for_each to
|
||||
# this resource would silently stop the guard from evaluating.
|
||||
#
|
||||
# endpoint cert key -> result
|
||||
# "" any any -> metering off, no secrets created
|
||||
# set set set -> metering on
|
||||
# set any-missing -> plan fails here
|
||||
lifecycle {
|
||||
precondition {
|
||||
condition = var.billing_metrics_endpoint == "" || (
|
||||
var.billing_metrics_client_cert_pem != "" && var.billing_metrics_client_key_pem != ""
|
||||
)
|
||||
error_message = "billing_metrics_client_cert_pem and billing_metrics_client_key_pem are both required when billing_metrics_endpoint is set."
|
||||
}
|
||||
}
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-gateway"
|
||||
location = var.region
|
||||
|
|
@ -206,7 +187,7 @@ resource "google_cloud_run_v2_service" "gateway" {
|
|||
max_instance_request_concurrency = var.gateway_max_instance_request_concurrency
|
||||
|
||||
vpc_access {
|
||||
connector = google_vpc_access_connector.this.id
|
||||
connector = google_vpc_access_connector.this[0].id
|
||||
egress = "PRIVATE_RANGES_ONLY"
|
||||
}
|
||||
|
||||
|
|
@ -312,17 +293,7 @@ resource "google_cloud_run_v2_service" "gateway" {
|
|||
|
||||
# ---------- Backend ----------
|
||||
resource "google_cloud_run_v2_service" "backend" {
|
||||
# Same guard as the gateway: the backend meters too (it serves the named-server
|
||||
# MCP transport), and a targeted apply of just this resource must not slip a
|
||||
# billing endpoint through without the credentials to use it.
|
||||
lifecycle {
|
||||
precondition {
|
||||
condition = var.billing_metrics_endpoint == "" || (
|
||||
var.billing_metrics_client_cert_pem != "" && var.billing_metrics_client_key_pem != ""
|
||||
)
|
||||
error_message = "billing_metrics_client_cert_pem and billing_metrics_client_key_pem are both required when billing_metrics_endpoint is set."
|
||||
}
|
||||
}
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-backend"
|
||||
location = var.region
|
||||
|
|
@ -335,7 +306,7 @@ resource "google_cloud_run_v2_service" "backend" {
|
|||
max_instance_request_concurrency = var.backend_max_instance_request_concurrency
|
||||
|
||||
vpc_access {
|
||||
connector = google_vpc_access_connector.this.id
|
||||
connector = google_vpc_access_connector.this[0].id
|
||||
egress = "PRIVATE_RANGES_ONLY"
|
||||
}
|
||||
|
||||
|
|
@ -443,6 +414,8 @@ resource "google_cloud_run_v2_service" "backend" {
|
|||
# with zero IAM bindings, so a compromised UI container can't pivot to
|
||||
# Secret Manager / Cloud SQL via the metadata service.
|
||||
resource "google_cloud_run_v2_service" "ui" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-ui"
|
||||
location = var.region
|
||||
ingress = "INGRESS_TRAFFIC_INTERNAL_LOAD_BALANCER"
|
||||
|
|
@ -450,7 +423,7 @@ resource "google_cloud_run_v2_service" "ui" {
|
|||
deletion_protection = false
|
||||
|
||||
template {
|
||||
service_account = google_service_account.ui_runtime.email
|
||||
service_account = google_service_account.ui_runtime[0].email
|
||||
max_instance_request_concurrency = var.ui_max_instance_request_concurrency
|
||||
|
||||
scaling {
|
||||
|
|
@ -491,25 +464,31 @@ resource "google_cloud_run_v2_service" "ui" {
|
|||
# (LITELLM_MASTER_KEY); these IAM bindings just open up Cloud Run's invoker
|
||||
# gate so the LB request makes it to the container.
|
||||
resource "google_cloud_run_v2_service_iam_member" "gateway_allusers" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
project = var.project_id
|
||||
location = google_cloud_run_v2_service.gateway.location
|
||||
name = google_cloud_run_v2_service.gateway.name
|
||||
location = google_cloud_run_v2_service.gateway[0].location
|
||||
name = google_cloud_run_v2_service.gateway[0].name
|
||||
role = "roles/run.invoker"
|
||||
member = "allUsers"
|
||||
}
|
||||
|
||||
resource "google_cloud_run_v2_service_iam_member" "backend_allusers" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
project = var.project_id
|
||||
location = google_cloud_run_v2_service.backend.location
|
||||
name = google_cloud_run_v2_service.backend.name
|
||||
location = google_cloud_run_v2_service.backend[0].location
|
||||
name = google_cloud_run_v2_service.backend[0].name
|
||||
role = "roles/run.invoker"
|
||||
member = "allUsers"
|
||||
}
|
||||
|
||||
resource "google_cloud_run_v2_service_iam_member" "ui_allusers" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
project = var.project_id
|
||||
location = google_cloud_run_v2_service.ui.location
|
||||
name = google_cloud_run_v2_service.ui.name
|
||||
location = google_cloud_run_v2_service.ui[0].location
|
||||
name = google_cloud_run_v2_service.ui[0].name
|
||||
role = "roles/run.invoker"
|
||||
member = "allUsers"
|
||||
}
|
||||
|
|
@ -519,6 +498,8 @@ resource "google_cloud_run_v2_service_iam_member" "ui_allusers" {
|
|||
# assembles DATABASE_URL from the DATABASE_* env vars and runs `prisma
|
||||
# migrate deploy`. No proxy_config, no master key, no shell wrapper.
|
||||
resource "google_cloud_run_v2_job" "migrations" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-migrations"
|
||||
location = var.region
|
||||
labels = local.labels
|
||||
|
|
@ -529,7 +510,7 @@ resource "google_cloud_run_v2_job" "migrations" {
|
|||
service_account = google_service_account.runtime.email
|
||||
|
||||
vpc_access {
|
||||
connector = google_vpc_access_connector.this.id
|
||||
connector = google_vpc_access_connector.this[0].id
|
||||
egress = "PRIVATE_RANGES_ONLY"
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ resource "google_sql_database_instance" "writer" {
|
|||
|
||||
ip_configuration {
|
||||
ipv4_enabled = false
|
||||
private_network = google_compute_network.this.id
|
||||
private_network = local.network_id
|
||||
}
|
||||
|
||||
insights_config {
|
||||
|
|
@ -55,6 +55,11 @@ resource "google_sql_database_instance" "writer" {
|
|||
# (full data loss). Set the initial size only; let Cloud SQL own it
|
||||
# thereafter.
|
||||
ignore_changes = [settings[0].disk_size]
|
||||
|
||||
precondition {
|
||||
condition = var.create_psa_connection || var.network_id != ""
|
||||
error_message = "create_psa_connection must be true unless network_id references an existing VPC with Private Services Access configured."
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -76,7 +81,7 @@ resource "google_sql_database_instance" "reader" {
|
|||
|
||||
ip_configuration {
|
||||
ipv4_enabled = false
|
||||
private_network = google_compute_network.this.id
|
||||
private_network = local.network_id
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -31,6 +31,11 @@ module "litellm" {
|
|||
tenant = var.tenant
|
||||
env = var.env
|
||||
|
||||
create_runtime = var.create_runtime
|
||||
network_id = var.network_id
|
||||
create_psa_connection = var.create_psa_connection
|
||||
redis_transit_encryption = var.redis_transit_encryption
|
||||
|
||||
litellm_master_key = var.litellm_master_key
|
||||
litellm_license = var.litellm_license
|
||||
ui_password = var.ui_password
|
||||
|
|
|
|||
|
|
@ -38,6 +38,31 @@ output "redis_endpoint" {
|
|||
value = module.litellm.redis_endpoint
|
||||
}
|
||||
|
||||
output "redis_host" {
|
||||
description = "Memorystore Redis host."
|
||||
value = module.litellm.redis_host
|
||||
}
|
||||
|
||||
output "redis_port" {
|
||||
description = "Memorystore Redis port."
|
||||
value = module.litellm.redis_port
|
||||
}
|
||||
|
||||
output "redis_server_ca_pem" {
|
||||
description = "Memorystore server CA PEM."
|
||||
value = module.litellm.redis_server_ca_pem
|
||||
}
|
||||
|
||||
output "db_username" {
|
||||
description = "Cloud SQL application username."
|
||||
value = module.litellm.db_username
|
||||
}
|
||||
|
||||
output "db_name" {
|
||||
description = "Cloud SQL database name."
|
||||
value = module.litellm.db_name
|
||||
}
|
||||
|
||||
output "gcs_bucket" {
|
||||
description = "GCS bucket name."
|
||||
value = module.litellm.gcs_bucket
|
||||
|
|
@ -53,6 +78,11 @@ output "db_password_secret_id" {
|
|||
value = module.litellm.db_password_secret_id
|
||||
}
|
||||
|
||||
output "runtime_service_account_email" {
|
||||
description = "Runtime service account email."
|
||||
value = module.litellm.runtime_service_account_email
|
||||
}
|
||||
|
||||
output "migration_run_command" {
|
||||
description = "Break-glass command to re-run the one-off migration job."
|
||||
value = module.litellm.migration_run_command
|
||||
|
|
|
|||
|
|
@ -8,6 +8,14 @@ region = "us-central1"
|
|||
tenant = "acme"
|
||||
env = "stage"
|
||||
|
||||
# Deployment mode. For dependencies only on a Shared VPC, set
|
||||
# create_runtime = false, network_id to the full host-project network ID, and
|
||||
# create_psa_connection = false after configuring PSA on that network.
|
||||
# create_runtime = true
|
||||
# network_id = ""
|
||||
# create_psa_connection = true
|
||||
# redis_transit_encryption = true
|
||||
|
||||
# Tenant-supplied secrets. Prefer TF_VAR_litellm_master_key /
|
||||
# TF_VAR_litellm_license / TF_VAR_ui_password env vars so the values don't
|
||||
# end up in a committed tfvars file. All three are optional — when
|
||||
|
|
|
|||
|
|
@ -26,6 +26,30 @@ variable "env" {
|
|||
type = string
|
||||
}
|
||||
|
||||
variable "create_runtime" {
|
||||
description = "Create Cloud Run and load balancer resources."
|
||||
type = bool
|
||||
default = true
|
||||
}
|
||||
|
||||
variable "network_id" {
|
||||
description = "Existing VPC network resource ID. Empty creates a VPC."
|
||||
type = string
|
||||
default = ""
|
||||
}
|
||||
|
||||
variable "create_psa_connection" {
|
||||
description = "Create Private Services Access resources."
|
||||
type = bool
|
||||
default = true
|
||||
}
|
||||
|
||||
variable "redis_transit_encryption" {
|
||||
description = "Enable Memorystore transit encryption."
|
||||
type = bool
|
||||
default = true
|
||||
}
|
||||
|
||||
# Sensitive — prefer TF_VAR_litellm_master_key / TF_VAR_litellm_license /
|
||||
# TF_VAR_ui_password so values stay out of any committed tfvars file.
|
||||
variable "litellm_master_key" {
|
||||
|
|
|
|||
|
|
@ -6,6 +6,15 @@
|
|||
resource "google_service_account" "runtime" {
|
||||
account_id = "${local.name}-runtime"
|
||||
display_name = "LiteLLM Cloud Run runtime"
|
||||
|
||||
lifecycle {
|
||||
precondition {
|
||||
condition = !var.create_runtime || var.billing_metrics_endpoint == "" || (
|
||||
var.billing_metrics_client_cert_pem != "" && var.billing_metrics_client_key_pem != ""
|
||||
)
|
||||
error_message = "billing_metrics_client_cert_pem and billing_metrics_client_key_pem are both required when billing_metrics_endpoint is set and create_runtime is true."
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
# UI runtime SA — no role bindings. The UI is static nginx with no DB,
|
||||
|
|
@ -14,6 +23,8 @@ resource "google_service_account" "runtime" {
|
|||
# project's serverless service agent (not this SA), so it doesn't need
|
||||
# artifactregistry.reader either.
|
||||
resource "google_service_account" "ui_runtime" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
account_id = "${local.name}-ui-runtime"
|
||||
display_name = "LiteLLM Cloud Run UI runtime (no data-plane access)"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -14,77 +14,93 @@ locals {
|
|||
}
|
||||
|
||||
resource "google_compute_global_address" "lb" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-lb-ip"
|
||||
labels = local.labels
|
||||
}
|
||||
|
||||
# Serverless NEGs — one per Cloud Run service.
|
||||
resource "google_compute_region_network_endpoint_group" "gateway" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-gateway-neg"
|
||||
region = var.region
|
||||
network_endpoint_type = "SERVERLESS"
|
||||
|
||||
cloud_run {
|
||||
service = google_cloud_run_v2_service.gateway.name
|
||||
service = google_cloud_run_v2_service.gateway[0].name
|
||||
}
|
||||
}
|
||||
|
||||
resource "google_compute_region_network_endpoint_group" "backend" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-backend-neg"
|
||||
region = var.region
|
||||
network_endpoint_type = "SERVERLESS"
|
||||
|
||||
cloud_run {
|
||||
service = google_cloud_run_v2_service.backend.name
|
||||
service = google_cloud_run_v2_service.backend[0].name
|
||||
}
|
||||
}
|
||||
|
||||
resource "google_compute_region_network_endpoint_group" "ui" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-ui-neg"
|
||||
region = var.region
|
||||
network_endpoint_type = "SERVERLESS"
|
||||
|
||||
cloud_run {
|
||||
service = google_cloud_run_v2_service.ui.name
|
||||
service = google_cloud_run_v2_service.ui[0].name
|
||||
}
|
||||
}
|
||||
|
||||
# Backend services wrap each NEG.
|
||||
resource "google_compute_backend_service" "gateway" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-gateway-bs"
|
||||
protocol = "HTTP"
|
||||
load_balancing_scheme = "EXTERNAL_MANAGED"
|
||||
|
||||
backend {
|
||||
group = google_compute_region_network_endpoint_group.gateway.id
|
||||
group = google_compute_region_network_endpoint_group.gateway[0].id
|
||||
}
|
||||
}
|
||||
|
||||
resource "google_compute_backend_service" "backend" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-backend-bs"
|
||||
protocol = "HTTP"
|
||||
load_balancing_scheme = "EXTERNAL_MANAGED"
|
||||
|
||||
backend {
|
||||
group = google_compute_region_network_endpoint_group.backend.id
|
||||
group = google_compute_region_network_endpoint_group.backend[0].id
|
||||
}
|
||||
}
|
||||
|
||||
resource "google_compute_backend_service" "ui" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-ui-bs"
|
||||
protocol = "HTTP"
|
||||
load_balancing_scheme = "EXTERNAL_MANAGED"
|
||||
|
||||
backend {
|
||||
group = google_compute_region_network_endpoint_group.ui.id
|
||||
group = google_compute_region_network_endpoint_group.ui[0].id
|
||||
}
|
||||
}
|
||||
|
||||
# URL map. Default → backend (management API). Path matchers route the
|
||||
# gateway and UI prefixes elsewhere.
|
||||
resource "google_compute_url_map" "this" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = local.name
|
||||
default_service = google_compute_backend_service.backend.id
|
||||
default_service = google_compute_backend_service.backend[0].id
|
||||
|
||||
host_rule {
|
||||
hosts = ["*"]
|
||||
|
|
@ -93,13 +109,13 @@ resource "google_compute_url_map" "this" {
|
|||
|
||||
path_matcher {
|
||||
name = "main"
|
||||
default_service = google_compute_backend_service.backend.id
|
||||
default_service = google_compute_backend_service.backend[0].id
|
||||
|
||||
# UI paths (catch them before any /v1/* gateway rules so /favicon.ico
|
||||
# and / take precedence).
|
||||
path_rule {
|
||||
paths = local.ui_path_prefixes
|
||||
service = google_compute_backend_service.ui.id
|
||||
service = google_compute_backend_service.ui[0].id
|
||||
}
|
||||
|
||||
# Gateway path prefixes. GCP URL maps cap a path_rule at 10 path globs,
|
||||
|
|
@ -108,7 +124,7 @@ resource "google_compute_url_map" "this" {
|
|||
for_each = { for idx, chunk in chunklist(local.gateway_path_prefixes, 10) : idx => chunk }
|
||||
content {
|
||||
paths = path_rule.value
|
||||
service = google_compute_backend_service.gateway.id
|
||||
service = google_compute_backend_service.gateway[0].id
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -118,7 +134,7 @@ resource "google_compute_url_map" "this" {
|
|||
# target proxy when TLS is enabled; otherwise the regular path-routing
|
||||
# URL map is attached to the HTTP proxy and everything stays plaintext.
|
||||
resource "google_compute_url_map" "https_redirect" {
|
||||
count = local.tls_enabled ? 1 : 0
|
||||
count = var.create_runtime && local.tls_enabled ? 1 : 0
|
||||
name = "${local.name}-redirect"
|
||||
|
||||
default_url_redirect {
|
||||
|
|
@ -129,8 +145,10 @@ resource "google_compute_url_map" "https_redirect" {
|
|||
}
|
||||
|
||||
resource "google_compute_target_http_proxy" "this" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-http"
|
||||
url_map = local.tls_enabled ? google_compute_url_map.https_redirect[0].id : google_compute_url_map.this.id
|
||||
url_map = local.tls_enabled ? google_compute_url_map.https_redirect[0].id : google_compute_url_map.this[0].id
|
||||
|
||||
# Default-deny on the HTTP-only path: TLS is the supported posture.
|
||||
# Operators must either supply DNS names or explicitly opt in.
|
||||
|
|
@ -143,12 +161,14 @@ resource "google_compute_target_http_proxy" "this" {
|
|||
}
|
||||
|
||||
resource "google_compute_global_forwarding_rule" "http" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-http"
|
||||
ip_protocol = "TCP"
|
||||
port_range = "80"
|
||||
load_balancing_scheme = "EXTERNAL_MANAGED"
|
||||
ip_address = google_compute_global_address.lb.address
|
||||
target = google_compute_target_http_proxy.this.id
|
||||
ip_address = google_compute_global_address.lb[0].address
|
||||
target = google_compute_target_http_proxy.this[0].id
|
||||
labels = local.labels
|
||||
}
|
||||
|
||||
|
|
@ -161,7 +181,7 @@ resource "google_compute_global_forwarding_rule" "http" {
|
|||
# transitions to ACTIVE.
|
||||
|
||||
resource "google_compute_managed_ssl_certificate" "this" {
|
||||
count = local.tls_enabled ? 1 : 0
|
||||
count = var.create_runtime && local.tls_enabled ? 1 : 0
|
||||
|
||||
# A managed cert's `domains` is immutable, so changing var.lb_domains
|
||||
# forces replacement, and the cert is referenced by the HTTPS target
|
||||
|
|
@ -181,19 +201,19 @@ resource "google_compute_managed_ssl_certificate" "this" {
|
|||
}
|
||||
|
||||
resource "google_compute_target_https_proxy" "this" {
|
||||
count = local.tls_enabled ? 1 : 0
|
||||
count = var.create_runtime && local.tls_enabled ? 1 : 0
|
||||
name = "${local.name}-https"
|
||||
url_map = google_compute_url_map.this.id
|
||||
url_map = google_compute_url_map.this[0].id
|
||||
ssl_certificates = [google_compute_managed_ssl_certificate.this[0].id]
|
||||
}
|
||||
|
||||
resource "google_compute_global_forwarding_rule" "https" {
|
||||
count = local.tls_enabled ? 1 : 0
|
||||
count = var.create_runtime && local.tls_enabled ? 1 : 0
|
||||
name = "${local.name}-https"
|
||||
ip_protocol = "TCP"
|
||||
port_range = "443"
|
||||
load_balancing_scheme = "EXTERNAL_MANAGED"
|
||||
ip_address = google_compute_global_address.lb.address
|
||||
ip_address = google_compute_global_address.lb[0].address
|
||||
target = google_compute_target_https_proxy.this[0].id
|
||||
labels = local.labels
|
||||
}
|
||||
|
|
|
|||
|
|
@ -21,6 +21,9 @@ locals {
|
|||
var.labels,
|
||||
)
|
||||
|
||||
create_network = var.network_id == ""
|
||||
network_id = local.create_network ? google_compute_network.this[0].id : var.network_id
|
||||
|
||||
gateway_path_prefixes = [
|
||||
"/v1/chat/*", "/chat/*",
|
||||
"/v1/completions*", "/completions*",
|
||||
|
|
@ -74,7 +77,7 @@ locals {
|
|||
"/ui/*",
|
||||
]
|
||||
|
||||
proxy_config_enabled = length(keys(var.proxy_config)) > 0
|
||||
proxy_config_enabled = var.create_runtime && length(keys(var.proxy_config)) > 0
|
||||
proxy_config_yaml = local.proxy_config_enabled ? yamlencode(var.proxy_config) : ""
|
||||
|
||||
proxy_config_mount_path = "/etc/litellm"
|
||||
|
|
|
|||
|
|
@ -1,13 +1,17 @@
|
|||
resource "google_compute_network" "this" {
|
||||
count = local.create_network ? 1 : 0
|
||||
|
||||
name = local.name
|
||||
auto_create_subnetworks = false
|
||||
routing_mode = "REGIONAL"
|
||||
}
|
||||
|
||||
resource "google_compute_subnetwork" "this" {
|
||||
count = local.create_network ? 1 : 0
|
||||
|
||||
name = "${local.name}-${var.region}"
|
||||
region = var.region
|
||||
network = google_compute_network.this.id
|
||||
network = google_compute_network.this[0].id
|
||||
ip_cidr_range = var.subnet_cidr
|
||||
private_ip_google_access = true
|
||||
}
|
||||
|
|
@ -16,17 +20,21 @@ resource "google_compute_subnetwork" "this" {
|
|||
# managed services peer with the VPC over the connection below using
|
||||
# addresses from this range.
|
||||
resource "google_compute_global_address" "psa" {
|
||||
count = var.create_psa_connection ? 1 : 0
|
||||
|
||||
name = "${local.name}-psa"
|
||||
purpose = "VPC_PEERING"
|
||||
address_type = "INTERNAL"
|
||||
prefix_length = 16
|
||||
network = google_compute_network.this.id
|
||||
network = local.network_id
|
||||
}
|
||||
|
||||
resource "google_service_networking_connection" "psa" {
|
||||
network = google_compute_network.this.id
|
||||
count = var.create_psa_connection ? 1 : 0
|
||||
|
||||
network = local.network_id
|
||||
service = "servicenetworking.googleapis.com"
|
||||
reserved_peering_ranges = [google_compute_global_address.psa.name]
|
||||
reserved_peering_ranges = [google_compute_global_address.psa[0].name]
|
||||
}
|
||||
|
||||
# Serverless VPC Access connector — required so Cloud Run can reach
|
||||
|
|
@ -37,9 +45,11 @@ resource "google_service_networking_connection" "psa" {
|
|||
# for low-to-moderate Cloud Run egress; bump max if your services push
|
||||
# heavy private-network traffic.
|
||||
resource "google_vpc_access_connector" "this" {
|
||||
count = var.create_runtime ? 1 : 0
|
||||
|
||||
name = "${local.name}-conn"
|
||||
region = var.region
|
||||
network = google_compute_network.this.name
|
||||
network = local.network_id
|
||||
ip_cidr_range = var.vpc_connector_cidr
|
||||
min_instances = 2
|
||||
max_instances = 3
|
||||
|
|
|
|||
|
|
@ -1,26 +1,26 @@
|
|||
output "lb_ip" {
|
||||
description = "Global anycast IP of the external HTTPS load balancer."
|
||||
value = google_compute_global_address.lb.address
|
||||
description = "Global anycast IP of the external HTTPS load balancer. Null when create_runtime is false."
|
||||
value = var.create_runtime ? one(google_compute_global_address.lb[*].address) : null
|
||||
}
|
||||
|
||||
output "lb_url" {
|
||||
description = "Proxy URL. Switches scheme based on whether lb_domains is set; when TLS is enabled the URL points at the first listed domain (since managed certs are tied to the hostname, not the anycast IP). The dashboard is served at /, the API at /v1/*."
|
||||
value = local.tls_enabled ? "https://${var.lb_domains[0]}" : "http://${google_compute_global_address.lb.address}"
|
||||
description = "Proxy URL, or null when create_runtime is false. Switches scheme based on whether lb_domains is set."
|
||||
value = var.create_runtime ? (local.tls_enabled ? "https://${var.lb_domains[0]}" : "http://${one(google_compute_global_address.lb[*].address)}") : null
|
||||
}
|
||||
|
||||
output "gateway_service_url" {
|
||||
description = "Default Cloud Run URL for the gateway (bypasses the LB)."
|
||||
value = google_cloud_run_v2_service.gateway.uri
|
||||
description = "Default Cloud Run URL for the gateway, or null when create_runtime is false."
|
||||
value = var.create_runtime ? one(google_cloud_run_v2_service.gateway[*].uri) : null
|
||||
}
|
||||
|
||||
output "backend_service_url" {
|
||||
description = "Default Cloud Run URL for the backend (bypasses the LB)."
|
||||
value = google_cloud_run_v2_service.backend.uri
|
||||
description = "Default Cloud Run URL for the backend, or null when create_runtime is false."
|
||||
value = var.create_runtime ? one(google_cloud_run_v2_service.backend[*].uri) : null
|
||||
}
|
||||
|
||||
output "ui_service_url" {
|
||||
description = "Default Cloud Run URL for the UI (bypasses the LB)."
|
||||
value = google_cloud_run_v2_service.ui.uri
|
||||
description = "Default Cloud Run URL for the UI, or null when create_runtime is false."
|
||||
value = var.create_runtime ? one(google_cloud_run_v2_service.ui[*].uri) : null
|
||||
}
|
||||
|
||||
output "cloudsql_writer_ip" {
|
||||
|
|
@ -38,6 +38,36 @@ output "redis_endpoint" {
|
|||
value = "${google_redis_instance.this.host}:${google_redis_instance.this.port}"
|
||||
}
|
||||
|
||||
output "runtime_service_account_email" {
|
||||
description = "Runtime service account email for Cloud Run or GKE Workload Identity."
|
||||
value = google_service_account.runtime.email
|
||||
}
|
||||
|
||||
output "redis_host" {
|
||||
description = "Memorystore Redis host."
|
||||
value = google_redis_instance.this.host
|
||||
}
|
||||
|
||||
output "redis_port" {
|
||||
description = "Memorystore Redis port."
|
||||
value = google_redis_instance.this.port
|
||||
}
|
||||
|
||||
output "redis_server_ca_pem" {
|
||||
description = "Memorystore server CA PEM. Mount it in the pod and set REDIS_SSL=true and REDIS_SSL_CA_CERTS=<path> via extraEnv when transit encryption is enabled."
|
||||
value = var.redis_transit_encryption ? google_redis_instance.this.server_ca_certs[0].cert : null
|
||||
}
|
||||
|
||||
output "db_username" {
|
||||
description = "Cloud SQL application username."
|
||||
value = var.db_username
|
||||
}
|
||||
|
||||
output "db_name" {
|
||||
description = "Cloud SQL database name."
|
||||
value = var.db_name
|
||||
}
|
||||
|
||||
output "gcs_bucket" {
|
||||
description = "GCS bucket name. Exposed to gateway + backend as GCS_BUCKET_NAME. Reference from proxy_config via `os.environ/GCS_BUCKET_NAME`."
|
||||
value = google_storage_bucket.this.name
|
||||
|
|
@ -54,11 +84,11 @@ output "db_password_secret_id" {
|
|||
}
|
||||
|
||||
output "migration_run_command" {
|
||||
description = "Shell command that executes the one-off migration job against Cloud SQL. Run this once after the first apply."
|
||||
value = format(
|
||||
description = "Shell command that executes the one-off migration job against Cloud SQL, or null when create_runtime is false."
|
||||
value = var.create_runtime ? format(
|
||||
"gcloud run jobs execute %s --region %s --project %s --wait",
|
||||
google_cloud_run_v2_job.migrations.name,
|
||||
one(google_cloud_run_v2_job.migrations[*].name),
|
||||
var.region,
|
||||
var.project_id,
|
||||
)
|
||||
) : null
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ resource "google_redis_instance" "this" {
|
|||
memory_size_gb = var.redis_memory_size_gb
|
||||
region = var.region
|
||||
|
||||
authorized_network = google_compute_network.this.id
|
||||
authorized_network = local.network_id
|
||||
connect_mode = "PRIVATE_SERVICE_ACCESS"
|
||||
|
||||
redis_version = "REDIS_7_0"
|
||||
|
|
@ -16,7 +16,7 @@ resource "google_redis_instance" "this" {
|
|||
# and passed to the proxy as REDIS_CA_PEM_B64); the proxy decodes it to
|
||||
# /tmp/redis-ca.pem at startup and uses it to validate the rediss://
|
||||
# handshake. Mirrors `transit_encryption_enabled = true` on AWS.
|
||||
transit_encryption_mode = "SERVER_AUTHENTICATION"
|
||||
transit_encryption_mode = var.redis_transit_encryption ? "SERVER_AUTHENTICATION" : "DISABLED"
|
||||
|
||||
depends_on = [google_service_networking_connection.psa]
|
||||
}
|
||||
|
|
|
|||
175
terraform/litellm/gcp/tests/deps_only.tftest.hcl
Normal file
175
terraform/litellm/gcp/tests/deps_only.tftest.hcl
Normal file
|
|
@ -0,0 +1,175 @@
|
|||
mock_provider "google" {
|
||||
mock_resource "google_redis_instance" {
|
||||
defaults = {
|
||||
host = "10.0.0.4"
|
||||
port = 6379
|
||||
server_ca_certs = [{
|
||||
cert = "-----BEGIN CERTIFICATE-----\nmock\n-----END CERTIFICATE-----"
|
||||
}]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mock_provider "google-beta" {}
|
||||
mock_provider "random" {}
|
||||
|
||||
variables {
|
||||
project_id = "test-project"
|
||||
tenant = "tenant"
|
||||
env = "test"
|
||||
allow_plaintext_lb = true
|
||||
image_registry = "us-central1-docker.pkg.dev/test-project/litellm"
|
||||
}
|
||||
|
||||
run "default_creates_everything" {
|
||||
command = plan
|
||||
|
||||
assert {
|
||||
condition = alltrue([
|
||||
length(google_compute_network.this) == 1,
|
||||
length(google_compute_subnetwork.this) == 1,
|
||||
length(google_compute_global_address.psa) == 1,
|
||||
length(google_service_networking_connection.psa) == 1,
|
||||
length(google_vpc_access_connector.this) == 1,
|
||||
length(google_cloud_run_v2_service.gateway) == 1,
|
||||
length(google_cloud_run_v2_service.backend) == 1,
|
||||
length(google_cloud_run_v2_service.ui) == 1,
|
||||
length(google_cloud_run_v2_job.migrations) == 1,
|
||||
length(google_compute_global_address.lb) == 1,
|
||||
length(terraform_data.migration) == 1,
|
||||
])
|
||||
error_message = "The default mode must create networking, runtime services, the load balancer, and migrations."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = google_redis_instance.this.transit_encryption_mode == "SERVER_AUTHENTICATION"
|
||||
error_message = "Redis transit encryption must remain enabled by default."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = length(local.shared_env_kv) == 12
|
||||
error_message = "The default runtime environment must include GCS and the three Redis TLS entries."
|
||||
}
|
||||
}
|
||||
|
||||
run "deps_only_creates_no_runtime" {
|
||||
command = plan
|
||||
|
||||
variables {
|
||||
create_runtime = false
|
||||
proxy_config = {
|
||||
model_list = []
|
||||
}
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = alltrue([
|
||||
length(google_cloud_run_v2_service.gateway) == 0,
|
||||
length(google_cloud_run_v2_service.backend) == 0,
|
||||
length(google_cloud_run_v2_service.ui) == 0,
|
||||
length(google_cloud_run_v2_job.migrations) == 0,
|
||||
length(google_cloud_run_v2_service_iam_member.gateway_allusers) == 0,
|
||||
length(google_cloud_run_v2_service_iam_member.backend_allusers) == 0,
|
||||
length(google_cloud_run_v2_service_iam_member.ui_allusers) == 0,
|
||||
length(google_compute_global_address.lb) == 0,
|
||||
length(google_compute_region_network_endpoint_group.gateway) == 0,
|
||||
length(google_compute_region_network_endpoint_group.backend) == 0,
|
||||
length(google_compute_region_network_endpoint_group.ui) == 0,
|
||||
length(google_compute_backend_service.gateway) == 0,
|
||||
length(google_compute_backend_service.backend) == 0,
|
||||
length(google_compute_backend_service.ui) == 0,
|
||||
length(google_compute_url_map.this) == 0,
|
||||
length(google_compute_url_map.https_redirect) == 0,
|
||||
length(google_compute_target_http_proxy.this) == 0,
|
||||
length(google_compute_global_forwarding_rule.http) == 0,
|
||||
length(google_compute_managed_ssl_certificate.this) == 0,
|
||||
length(google_compute_target_https_proxy.this) == 0,
|
||||
length(google_compute_global_forwarding_rule.https) == 0,
|
||||
length(terraform_data.migration) == 0,
|
||||
length(google_vpc_access_connector.this) == 0,
|
||||
length(google_service_account.ui_runtime) == 0,
|
||||
length(google_storage_bucket.proxy_config) == 0,
|
||||
])
|
||||
error_message = "Dependencies-only mode must omit all runtime, load balancer, connector, UI identity, and proxy config resources."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = alltrue([
|
||||
google_sql_database_instance.writer.name == "tenant-litellm-test",
|
||||
google_sql_database_instance.reader.name == "tenant-litellm-test-reader",
|
||||
google_redis_instance.this.name == "tenant-litellm-test",
|
||||
google_storage_bucket.this.force_destroy == false,
|
||||
google_secret_manager_secret.master_key.secret_id == "tenant-litellm-test-master-key",
|
||||
google_secret_manager_secret.db_password.secret_id == "tenant-litellm-test-db-password",
|
||||
google_service_account.runtime.account_id == "tenant-litellm-test-runtime",
|
||||
])
|
||||
error_message = "Dependencies-only mode must retain data stores, secrets, and the runtime service account."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = output.lb_url == null && output.migration_run_command == null
|
||||
error_message = "Runtime outputs must be null while dependency outputs remain available."
|
||||
}
|
||||
}
|
||||
|
||||
run "existing_network_attaches_data_stores" {
|
||||
command = plan
|
||||
|
||||
variables {
|
||||
network_id = "projects/host-proj/global/networks/shared"
|
||||
create_psa_connection = false
|
||||
create_runtime = false
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = alltrue([
|
||||
length(google_compute_network.this) == 0,
|
||||
length(google_compute_subnetwork.this) == 0,
|
||||
length(google_compute_global_address.psa) == 0,
|
||||
length(google_service_networking_connection.psa) == 0,
|
||||
google_sql_database_instance.writer.settings[0].ip_configuration[0].private_network == var.network_id,
|
||||
google_redis_instance.this.authorized_network == var.network_id,
|
||||
])
|
||||
error_message = "An existing VPC must receive the Cloud SQL and Memorystore private-network attachments."
|
||||
}
|
||||
}
|
||||
|
||||
run "psa_required_without_existing_network" {
|
||||
command = plan
|
||||
|
||||
variables {
|
||||
create_psa_connection = false
|
||||
}
|
||||
|
||||
expect_failures = [
|
||||
google_sql_database_instance.writer,
|
||||
]
|
||||
}
|
||||
|
||||
run "redis_plaintext_drops_tls_env" {
|
||||
command = plan
|
||||
|
||||
variables {
|
||||
redis_transit_encryption = false
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = google_redis_instance.this.transit_encryption_mode == "DISABLED"
|
||||
error_message = "Redis transit encryption must be disabled when requested."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = length(local.shared_env_kv) == 9
|
||||
error_message = "Plaintext Redis mode must include GCS and omit the three Redis TLS entries."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = length([for env in local.shared_env_kv : env if env.name == "REDIS_SSL"]) == 0
|
||||
error_message = "Plaintext Redis mode must not set REDIS_SSL."
|
||||
}
|
||||
|
||||
assert {
|
||||
condition = length(local.redis_ca_fragment) == 0
|
||||
error_message = "Plaintext Redis mode must not decode a Redis CA at startup."
|
||||
}
|
||||
}
|
||||
|
|
@ -79,16 +79,42 @@ variable "ui_password" {
|
|||
sensitive = true
|
||||
}
|
||||
|
||||
# ---------- Deployment mode ----------
|
||||
|
||||
variable "create_runtime" {
|
||||
description = "Create Cloud Run, load balancer, VPC connector, runtime support resources, and the migration job. Set false for GKE or another external runtime."
|
||||
type = bool
|
||||
default = true
|
||||
}
|
||||
|
||||
variable "network_id" {
|
||||
description = "Existing VPC network resource ID (`projects/<host-project>/global/networks/<name>`). When set, no VPC or subnet is created. A VPC connector requires this network to be in the deployment project when create_runtime is true."
|
||||
type = string
|
||||
default = ""
|
||||
}
|
||||
|
||||
variable "create_psa_connection" {
|
||||
description = "Create the Private Services Access range and connection for Cloud SQL and Memorystore. Set false when the existing network already has PSA configured."
|
||||
type = bool
|
||||
default = true
|
||||
}
|
||||
|
||||
variable "redis_transit_encryption" {
|
||||
description = "Enable Memorystore transit encryption and inject Redis TLS settings into Cloud Run. Set false to use plaintext Redis."
|
||||
type = bool
|
||||
default = true
|
||||
}
|
||||
|
||||
# ---------- Networking ----------
|
||||
|
||||
variable "subnet_cidr" {
|
||||
description = "Primary CIDR block for the LiteLLM subnet."
|
||||
description = "Primary CIDR block for the LiteLLM subnet. Unused when network_id is set."
|
||||
type = string
|
||||
default = "10.40.0.0/16"
|
||||
}
|
||||
|
||||
variable "vpc_connector_cidr" {
|
||||
description = "CIDR for the Serverless VPC Access connector. /28 required."
|
||||
description = "CIDR for the Serverless VPC Access connector. /28 required. Unused when create_runtime is false."
|
||||
type = string
|
||||
default = "10.41.0.0/28"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 10975
|
||||
"limit": 10972
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ from access_control_client import (
|
|||
from e2e_config import unique_marker
|
||||
from e2e_http import Success, UnauthorizedError, UnknownApiError, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody
|
||||
from models import ChatBody, ChatMessage, ChatResponse, EmbedBody, LiteLLMParamsBody
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
|
@ -32,6 +32,7 @@ pytestmark = pytest.mark.e2e
|
|||
ALLOWED_MODEL = "gemini-2.5-flash"
|
||||
DISALLOWED_MODEL = "gpt-5.5"
|
||||
VIRTUAL_KEY_BACKEND = "anthropic/claude-haiku-4-5-20251001"
|
||||
EMBEDDING_MODEL = "openai-text-embedding-3-small"
|
||||
|
||||
|
||||
class TestAccessControl:
|
||||
|
|
@ -71,6 +72,31 @@ class TestAccessControl:
|
|||
f"403 body must be a model-access denial, got: {result.body[:300]}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("other.auth.virtual_key.route_group_allowed")
|
||||
def test_llm_api_routes_group_grants_every_llm_endpoint(
|
||||
self, client: AccessControlClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = client.llm_only_key()
|
||||
resources.defer(lambda: client.delete_key(key))
|
||||
|
||||
chat = client.chat_status(key, ALLOWED_MODEL, f"capital of France? {unique_marker()}")
|
||||
assert chat.status_code == 200, (
|
||||
f"llm_api_routes key must reach /chat/completions, got {chat.status_code}: {chat.body[:300]}"
|
||||
)
|
||||
assert ChatResponse.model_validate_json(chat.body).choices, (
|
||||
f"200 must carry a real completion, not an error envelope: {chat.body[:300]}"
|
||||
)
|
||||
|
||||
embedding = unwrap(
|
||||
client.proxy.embed(key, EmbedBody(model=EMBEDDING_MODEL, input=f"route group {unique_marker()}"))
|
||||
)
|
||||
assert embedding.model, f"llm_api_routes key reached /embeddings but got no model back: {embedding}"
|
||||
|
||||
denied = client.create_model_status(key, f"e2e-route-group-{unique_marker()}")
|
||||
assert denied.status_code == 403 and ROUTE_NOT_ALLOWED_MARKER in denied.body, (
|
||||
f"the same key must still be shut out of /model/new, got {denied.status_code}: {denied.body[:300]}"
|
||||
)
|
||||
|
||||
def test_llm_only_key_forbidden_from_management_route_403(
|
||||
self, client: AccessControlClient, resources: ResourceManager
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -10,6 +10,11 @@ Each case asserts the feature actually happened, not just a 200. Coverage matrix
|
|||
intentionally not covered here.
|
||||
- Vertex (gemini-2.5-flash): prompt caching via ``cache_control`` context
|
||||
caching; the second identical call must report cached prompt tokens > 0.
|
||||
- Anthropic (claude-haiku-4-5, direct): the same ``cache_control`` prefix over
|
||||
the OpenAI-compatible route; the second call must report cache-read tokens > 0.
|
||||
- OpenAI (gpt-5.6): automatic prompt caching needs no request marker, so the
|
||||
cacheable prefix goes out as a plain system string with a ``prompt_cache_key``
|
||||
and the second call must report ``prompt_tokens_details.cached_tokens`` > 0.
|
||||
|
||||
service_tier lives in test_provider_features_e2e.py.
|
||||
|
||||
|
|
@ -21,6 +26,7 @@ built from the typed content blocks shared in ``endpoints_client.py``.
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -29,7 +35,7 @@ from e2e_config import unique_marker
|
|||
from e2e_http import Result, unwrap
|
||||
from endpoints_client import CacheControl, RichMessage, TextBlock
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatResponse, LiteLLMParamsBody, Usage
|
||||
from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody, Usage
|
||||
from passthrough_client import PassthroughClient
|
||||
import os
|
||||
|
||||
|
|
@ -37,6 +43,8 @@ pytestmark = pytest.mark.e2e
|
|||
|
||||
BEDROCK_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
VERTEX_MODEL = "vertex_ai/gemini-2.5-flash"
|
||||
ANTHROPIC_MODEL = "anthropic/claude-haiku-4-5-20251001"
|
||||
OPENAI_MODEL = "openai/gpt-5.6"
|
||||
|
||||
|
||||
class CacheChatBody(BaseModel):
|
||||
|
|
@ -89,17 +97,36 @@ def _cache_chat(
|
|||
)
|
||||
|
||||
|
||||
def _plain_cache_chat(
|
||||
client: PassthroughClient, key: str, model: str, prefix: str, cache_key: str
|
||||
) -> Result[ChatResponse]:
|
||||
"""The same cacheable prefix as a plain system string, for providers that cache
|
||||
automatically and take no per-block marker (OpenAI)."""
|
||||
return client.proxy.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[
|
||||
ChatMessage(role="system", content=prefix),
|
||||
ChatMessage(role="user", content="Reply with one word."),
|
||||
],
|
||||
max_tokens=64,
|
||||
prompt_cache_key=cache_key,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _assert_cache_read_on_second_call(
|
||||
client: PassthroughClient, key: str, model: str
|
||||
model: str, send: Callable[[str], Result[ChatResponse]]
|
||||
) -> None:
|
||||
prefix = _cacheable_prefix()
|
||||
|
||||
first = unwrap(_cache_chat(client, key, model, prefix))
|
||||
first = unwrap(send(prefix))
|
||||
assert first.choices, f"{model}: first cache-priming call returned no choices: {first}"
|
||||
|
||||
deadline = time.monotonic() + 30.0
|
||||
while True:
|
||||
second = unwrap(_cache_chat(client, key, model, prefix))
|
||||
second = unwrap(send(prefix))
|
||||
read_tokens = _cached_read_tokens(second.usage)
|
||||
if read_tokens > 0 or time.monotonic() >= deadline:
|
||||
break
|
||||
|
|
@ -125,7 +152,8 @@ class TestCacheControl:
|
|||
LiteLLMParamsBody(model=BEDROCK_MODEL, aws_region_name="us-east-1"),
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
_assert_cache_read_on_second_call(client, resources.key(), model)
|
||||
key = resources.key()
|
||||
_assert_cache_read_on_second_call(model, lambda prefix: _cache_chat(client, key, model, prefix))
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.chat_completions.vertex.prompt_cache_5m.nonstream.works",
|
||||
|
|
@ -145,4 +173,40 @@ class TestCacheControl:
|
|||
),
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
_assert_cache_read_on_second_call(client, resources.key(), model)
|
||||
key = resources.key()
|
||||
_assert_cache_read_on_second_call(model, lambda prefix: _cache_chat(client, key, model, prefix))
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.chat_completions.anthropic.prompt_cache_5m.nonstream.works",
|
||||
exercised_on=[],
|
||||
)
|
||||
def test_anthropic_prompt_caching_reads_cache(
|
||||
self, client: PassthroughClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = f"e2e-anthropic-cache-{unique_marker()}"
|
||||
model_id = client.proxy.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(model=ANTHROPIC_MODEL, api_key="os.environ/ANTHROPIC_API_KEY"),
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
key = resources.key()
|
||||
_assert_cache_read_on_second_call(model, lambda prefix: _cache_chat(client, key, model, prefix))
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.chat_completions.openai.prompt_cache_5m.nonstream.works",
|
||||
exercised_on=[],
|
||||
)
|
||||
def test_openai_prompt_caching_reads_cache(
|
||||
self, client: PassthroughClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = f"e2e-openai-cache-{unique_marker()}"
|
||||
model_id = client.proxy.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(model=OPENAI_MODEL, api_key="os.environ/OPENAI_API_KEY"),
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
key = resources.key()
|
||||
cache_key = f"e2e-openai-cache-{unique_marker()}"
|
||||
_assert_cache_read_on_second_call(
|
||||
model, lambda prefix: _plain_cache_chat(client, key, model, prefix, cache_key)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ fails that provider's row here.
|
|||
|
||||
The per-provider classes below cover the OpenAI-compatible /chat/completions
|
||||
translation for providers customers reach by registering their own deployment
|
||||
via /model/new (Cohere, Gemini, hosted_vllm), each deleted on teardown.
|
||||
via /model/new (Cohere, Gemini, hosted_vllm, Anthropic), each deleted on teardown.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -46,6 +46,7 @@ pytestmark = pytest.mark.e2e
|
|||
COHERE_BACKEND = "cohere/command-r-08-2024"
|
||||
GEMINI_BACKEND = "gemini/gemini-2.5-flash"
|
||||
OPENAI_BACKEND = "openai/gpt-5.6"
|
||||
ANTHROPIC_BACKEND = "anthropic/claude-haiku-4-5-20251001"
|
||||
BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
|
||||
|
||||
|
|
@ -746,3 +747,98 @@ class TestBedrockConverseChatCompletions:
|
|||
|
||||
response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_vision_messages(), max_tokens=32)))
|
||||
_assert_describes_cat(response)
|
||||
|
||||
|
||||
class TestAnthropicChatCompletions:
|
||||
"""Anthropic via the OpenAI-compatible /chat/completions path, the translation
|
||||
customers on the OpenAI SDK rely on when they route to Claude. The streamed call
|
||||
must deliver real content deltas, and a tool-forced call must come back as a
|
||||
well-formed tool_call on both the non-streamed and streamed paths.
|
||||
"""
|
||||
|
||||
def _register(self, client: PassthroughClient, resources: ResourceManager, prefix: str) -> str:
|
||||
model = f"{prefix}-{unique_marker()}"
|
||||
model_id = client.proxy.create_model(
|
||||
model, LiteLLMParamsBody(model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY")
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
return model
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.chat_completions.anthropic.basic.stream.works",
|
||||
exercised_on=["chat_completions"],
|
||||
)
|
||||
def test_anthropic_chat_streams_real_content(
|
||||
self, client: PassthroughClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = self._register(client, resources, "e2e-anthropic-stream")
|
||||
key = resources.key()
|
||||
|
||||
result = client.proxy.chat_stream(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[
|
||||
ChatMessage(role="user", content=f"Count from 1 to 5, one number per line. {unique_marker()}")
|
||||
],
|
||||
max_tokens=64,
|
||||
stream=True,
|
||||
),
|
||||
)
|
||||
_assert_streamed_completion(result)
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.chat_completions.anthropic.tool_use.nonstream.works",
|
||||
exercised_on=["chat_completions"],
|
||||
)
|
||||
def test_anthropic_chat_returns_tool_call(
|
||||
self, client: PassthroughClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = self._register(client, resources, "e2e-anthropic-tool")
|
||||
key = resources.key()
|
||||
|
||||
response = unwrap(
|
||||
client.proxy.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[
|
||||
ChatMessage(role="user", content="What is the weather in San Francisco? Use the get_weather tool.")
|
||||
],
|
||||
tools=[_WEATHER_TOOL],
|
||||
tool_choice="required",
|
||||
max_tokens=128,
|
||||
),
|
||||
)
|
||||
)
|
||||
_assert_weather_tool_call(response)
|
||||
|
||||
@pytest.mark.covers(
|
||||
"llm.chat_completions.anthropic.tool_use.stream.works",
|
||||
exercised_on=["chat_completions"],
|
||||
)
|
||||
def test_anthropic_chat_streams_tool_call(
|
||||
self, client: PassthroughClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = self._register(client, resources, "e2e-anthropic-tool-stream")
|
||||
key = resources.key()
|
||||
|
||||
result = client.proxy.chat_stream(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[
|
||||
ChatMessage(role="user", content="What is the weather in San Francisco? Use the get_weather tool.")
|
||||
],
|
||||
tools=[_WEATHER_TOOL],
|
||||
tool_choice="required",
|
||||
max_tokens=128,
|
||||
stream=True,
|
||||
),
|
||||
)
|
||||
assert result.ok and result.is_streaming, f"tool stream was not established: {result}"
|
||||
assert result.stream_error is None, f"tool stream carried an error event: {result.stream_error}"
|
||||
name, arguments = _streamed_tool_call(result.stream_events)
|
||||
assert name == "get_weather", f"streamed tool call named {name!r}: {result.stream_events[:5]}"
|
||||
args = _WeatherArgs.model_validate_json(arguments)
|
||||
assert args.location.strip(), f"streamed tool call arguments missing location: {arguments!r}"
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Live e2e: POST /embeddings returns a real vector across OpenAI, Bedrock, Vertex.
|
||||
"""Live e2e: POST /embeddings returns a real vector across OpenAI, Bedrock, Vertex, Cohere.
|
||||
|
||||
Each test registers the deployment it needs at runtime (deleted on teardown) and
|
||||
asserts a non-empty, non-zero vector came back. The LIT-3167 guard in
|
||||
|
|
@ -86,6 +86,26 @@ class TestEmbeddingsEndpoint:
|
|||
f"embedding vector is all zeros: {result.body[:300]}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.embeddings.cohere.basic.nonstream.works")
|
||||
def test_cohere_embeddings_returns_vector(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = f"e2e-embeddings-cohere-{unique_marker()}"
|
||||
model_id = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(model="cohere/embed-v4.0", api_key="os.environ/COHERE_API_KEY"),
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
||||
result = endpoints_client.embeddings(key, model, "Say this is a test!")
|
||||
require_successful_call(result)
|
||||
parsed = EmbeddingsResult.model_validate_json(result.body)
|
||||
assert parsed.first_vector, f"/embeddings returned no vector: {result.body[:300]}"
|
||||
assert any(component != 0.0 for component in parsed.first_vector), (
|
||||
f"embedding vector is all zeros: {result.body[:300]}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.embeddings.vertex.basic.nonstream.works")
|
||||
def test_vertex_embeddings_returns_vector(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ import pytest
|
|||
from e2e_config import CHEAP_OPENAI_MODEL, unique_marker
|
||||
from e2e_http import require_successful_call, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import KeyGenerateBody, SpendLogRow
|
||||
from models import ChatResponse, KeyGenerateBody, SpendLogRow
|
||||
from passthrough_client import (
|
||||
AnthropicTool,
|
||||
GeminiFunctionDeclaration,
|
||||
|
|
@ -344,6 +344,44 @@ class TestOpenAIPassthroughSpend:
|
|||
)
|
||||
|
||||
|
||||
class TestOpenAIProviderPrefixChat:
|
||||
"""OpenAI-format chat through the raw `/openai/{endpoint}` passthrough (LIT-4752).
|
||||
|
||||
The body goes to OpenAI untranslated with the proxy's own OPENAI_API_KEY swapped
|
||||
in, so the customer gets OpenAI's real completion back, and the gateway must
|
||||
still write a costed pass_through_endpoint row for it.
|
||||
"""
|
||||
|
||||
@pytest.mark.covers("llm.chat_completions.openai.passthrough.nonstream.cost_logged")
|
||||
def test_openai_prefix_chat_returns_completion_and_logs_its_cost(
|
||||
self, client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
result = client.openai_chat(scoped_key, CHEAP_OPENAI_MODEL, f"Say hi in one word. {unique_marker()}")
|
||||
require_successful_call(result)
|
||||
|
||||
completion = ChatResponse.model_validate_json(result.body)
|
||||
assert completion.id, f"/openai/v1/chat/completions relayed no completion id: {result.body[:300]}"
|
||||
content = (
|
||||
completion.choices[0].message.content
|
||||
if completion.choices and completion.choices[0].message
|
||||
else None
|
||||
)
|
||||
assert content and content.strip(), (
|
||||
f"/openai/v1/chat/completions relayed an empty completion: {result.body[:300]}"
|
||||
)
|
||||
assert completion.usage is not None, f"the completion carried no usage to price from: {completion}"
|
||||
|
||||
row = _fetch_cost_breakdown(client, completion.id)
|
||||
assert row.prompt_tokens == completion.usage.prompt_tokens, (
|
||||
f"logged {row.prompt_tokens} prompt tokens, the completion the customer read "
|
||||
f"reported {completion.usage.prompt_tokens}"
|
||||
)
|
||||
assert row.completion_tokens == completion.usage.completion_tokens, (
|
||||
f"logged {row.completion_tokens} completion tokens, the completion the customer read "
|
||||
f"reported {completion.usage.completion_tokens}"
|
||||
)
|
||||
|
||||
|
||||
class TestOpenAIPassthroughWebsocket:
|
||||
"""The OpenAI passthrough prefixes must answer a websocket upgrade, not only a POST.
|
||||
|
||||
|
|
|
|||
|
|
@ -40,6 +40,8 @@ from models import (
|
|||
KeyListParams,
|
||||
KeyListResponse,
|
||||
KeyRegenerateBody,
|
||||
KeyResetSpendBody,
|
||||
KeyResetSpendResponse,
|
||||
KeyUpdateBody,
|
||||
ModelDeleteBody,
|
||||
OrgDeleteBody,
|
||||
|
|
@ -191,16 +193,26 @@ class ManagementClient:
|
|||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
def regenerate_key(self, key: str) -> str:
|
||||
def regenerate_key(self, key: str, *, grace_period: str | None = None) -> str:
|
||||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/key/regenerate",
|
||||
headers=self.proxy.transport.master,
|
||||
json=KeyRegenerateBody(key=key),
|
||||
json=KeyRegenerateBody(key=key, grace_period=grace_period),
|
||||
response_type=KeyGenerateResponse,
|
||||
)
|
||||
).key
|
||||
|
||||
def reset_key_spend(self, key: str, reset_to: float) -> KeyResetSpendResponse:
|
||||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
f"/key/{key}/reset_spend",
|
||||
headers=self.proxy.transport.master,
|
||||
json=KeyResetSpendBody(reset_to=reset_to),
|
||||
response_type=KeyResetSpendResponse,
|
||||
)
|
||||
)
|
||||
|
||||
def key_list(self, key_alias: str, *, caller_key: str | None = None) -> Result[KeyListResponse]:
|
||||
"""GET /key/list, the Virtual Keys page's own inventory call. `caller_key` is
|
||||
who is asking: the master key by default, or a virtual key."""
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from typing import Literal
|
|||
import pytest
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import NoBody, unwrap
|
||||
from e2e_http import NoBody, StreamingResponse, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from management_client import ManagementClient
|
||||
from models import KeyDeleteBody, KeyGenerateBody, KeyUpdateBody
|
||||
|
|
@ -26,6 +26,9 @@ from pydantic import BaseModel
|
|||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
TINY_BUDGET = 3e-6
|
||||
SPEND_MODEL = "claude-haiku-4-5"
|
||||
|
||||
|
||||
class KeyToggleBlockBody(BaseModel):
|
||||
key: str
|
||||
|
|
@ -82,6 +85,30 @@ def _generate_key(client: ManagementClient, resources: ResourceManager, body: Ke
|
|||
return key
|
||||
|
||||
|
||||
def _is_budget_block(outcome: StreamingResponse) -> bool:
|
||||
return not outcome.ok and "budget_exceeded" in outcome.body
|
||||
|
||||
|
||||
def _spend_until_budget_blocks(client: ManagementClient, key: str) -> None:
|
||||
for _ in range(40):
|
||||
outcome = client.chat_status(key, SPEND_MODEL, f"spend {unique_marker()}")
|
||||
if _is_budget_block(outcome):
|
||||
assert outcome.status_code == 429, (
|
||||
f"budget refusal must be 429, got {outcome.status_code}: {outcome.body[:200]}"
|
||||
)
|
||||
return
|
||||
assert outcome.ok, f"paid call failed before the budget tripped ({outcome.status_code}): {outcome.body[:300]}"
|
||||
time.sleep(2)
|
||||
pytest.fail(f"max_budget={TINY_BUDGET} never blocked a call on the key")
|
||||
|
||||
|
||||
def _settled_spend(client: ManagementClient, key: str) -> float | None:
|
||||
first = client.proxy.key_info(key).spend or 0.0
|
||||
time.sleep(client.proxy.poll_interval)
|
||||
second = client.proxy.key_info(key).spend or 0.0
|
||||
return second if first > 0 and first == second else None
|
||||
|
||||
|
||||
def _block(client: ManagementClient, key: str) -> None:
|
||||
_ = unwrap(
|
||||
client.proxy.transport.post(
|
||||
|
|
@ -197,6 +224,32 @@ class TestKeyManagementRoutes:
|
|||
"/key/info never reported max_budget 42.0 after /key/bulk_update before the deadline",
|
||||
)
|
||||
|
||||
@pytest.mark.covers("other.key_mgmt.spend_reset.resets_to_value")
|
||||
def test_reset_spend_zeroes_recorded_spend_and_lifts_the_budget_block(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
key = _generate_key(client, resources, KeyGenerateBody(models=[SPEND_MODEL], max_budget=TINY_BUDGET))
|
||||
_spend_until_budget_blocks(client, key)
|
||||
recorded = _poll(
|
||||
client, lambda: _settled_spend(client, key), "key spend never landed in /key/info before the deadline"
|
||||
)
|
||||
|
||||
reset = client.reset_key_spend(key, reset_to=0.0)
|
||||
assert reset.previous_spend == recorded, (
|
||||
f"reset_spend reported previous_spend {reset.previous_spend}, /key/info had recorded {recorded}"
|
||||
)
|
||||
assert reset.spend == 0.0, f"reset_spend to 0 reported spend {reset.spend}"
|
||||
assert client.proxy.key_info(key).spend == 0.0, "/key/info still reports spend after the reset to 0"
|
||||
|
||||
def call_allowed_again() -> bool | None:
|
||||
outcome = client.chat_status(key, SPEND_MODEL, f"after reset {unique_marker()}")
|
||||
if _is_budget_block(outcome):
|
||||
return None
|
||||
assert outcome.ok, f"post-reset call failed ({outcome.status_code}): {outcome.body[:300]}"
|
||||
return True
|
||||
|
||||
_ = _poll(client, call_allowed_again, "the key stayed budget-blocked after its spend was reset to 0")
|
||||
|
||||
@pytest.mark.covers("mgmt.key.generate.admin_only")
|
||||
def test_generate_forbidden_for_non_admin_key(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from __future__ import annotations
|
|||
import math
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -42,6 +43,10 @@ from models import (
|
|||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
REGENERATE_GRACE_PERIOD = "15s"
|
||||
REGENERATE_GRACE_SECONDS = 15.0
|
||||
|
||||
|
||||
def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T:
|
||||
deadline = time.monotonic() + client.proxy.poll_timeout
|
||||
while time.monotonic() < deadline:
|
||||
|
|
@ -365,6 +370,36 @@ class TestKeyRegeneration:
|
|||
client, old_rejected, "old key was still accepted after regeneration (never rejected 401) at the deadline"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("other.key_mgmt.regenerate.grace_period_honored")
|
||||
def test_regenerate_with_grace_period_keeps_old_key_until_revoked(
|
||||
self, client: ManagementClient, resources: ResourceManager
|
||||
) -> None:
|
||||
old_key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"]))
|
||||
|
||||
new_key = client.regenerate_key(old_key, grace_period=REGENERATE_GRACE_PERIOD)
|
||||
resources.defer(lambda: client.proxy.delete_key(new_key))
|
||||
revoke_at: Final = time.monotonic() + REGENERATE_GRACE_SECONDS
|
||||
assert new_key != old_key, "regenerate returned the same key string, so no rotation happened"
|
||||
|
||||
def old_accepted() -> bool | None:
|
||||
outcome = client.chat_status(old_key, "gpt-5.5", f"say hi {unique_marker()}")
|
||||
return True if outcome.ok else None
|
||||
|
||||
_ = _poll(client, old_accepted, "old key was rejected 401 inside its grace period at the deadline")
|
||||
assert time.monotonic() < revoke_at, (
|
||||
f"old key was only accepted after its {REGENERATE_GRACE_PERIOD} grace period had elapsed"
|
||||
)
|
||||
|
||||
def old_rejected() -> bool | None:
|
||||
outcome = client.chat_status(old_key, "gpt-5.5", f"say hi {unique_marker()}")
|
||||
return True if outcome.status_code == 401 else None
|
||||
|
||||
_ = _poll(
|
||||
client,
|
||||
old_rejected,
|
||||
f"old key was still accepted past its {REGENERATE_GRACE_PERIOD} grace period (never 401) at the deadline",
|
||||
)
|
||||
|
||||
|
||||
class TestTeamRoutes:
|
||||
@pytest.mark.covers("mgmt.team.new.persists")
|
||||
|
|
|
|||
|
|
@ -83,6 +83,16 @@ class KeyGenerateResponse(BaseModel):
|
|||
|
||||
class KeyRegenerateBody(BaseModel):
|
||||
key: str
|
||||
grace_period: str | None = None
|
||||
|
||||
|
||||
class KeyResetSpendBody(BaseModel):
|
||||
reset_to: float
|
||||
|
||||
|
||||
class KeyResetSpendResponse(BaseModel):
|
||||
spend: float
|
||||
previous_spend: float
|
||||
|
||||
|
||||
class KeyDeleteBody(BaseModel):
|
||||
|
|
@ -851,6 +861,15 @@ class ModelListEntry(BaseModel):
|
|||
id: str
|
||||
|
||||
|
||||
class ModelsListParams(BaseModel):
|
||||
"""Query for GET /v1/models. A wildcard route such as ``openai/gpt-5.4*`` is
|
||||
listed only under ``return_wildcard_routes``; without it the route is dropped
|
||||
and only its expansions remain, so a readiness poll for the pattern itself
|
||||
never resolves."""
|
||||
|
||||
return_wildcard_routes: bool = True
|
||||
|
||||
|
||||
class ModelsListResponse(BaseModel):
|
||||
"""GET /v1/models on the data plane: the deployments the gateway can actually
|
||||
serve right now. Used to confirm a freshly created model has propagated from
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ from models import (
|
|||
ModelMode,
|
||||
ModelNewBody,
|
||||
ModelNewResponse,
|
||||
ModelsListParams,
|
||||
ModelsListResponse,
|
||||
ModelUpdateBody,
|
||||
OcrBody,
|
||||
|
|
@ -336,7 +337,7 @@ class ProxyClient:
|
|||
lambda poll_timeout: self.transport.get(
|
||||
"/v1/models",
|
||||
headers=headers,
|
||||
params=NoBody(),
|
||||
params=ModelsListParams(),
|
||||
response_type=ModelsListResponse,
|
||||
timeout=poll_timeout,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -597,9 +597,20 @@ class TestSemanticAutoRouterResponses:
|
|||
)
|
||||
)
|
||||
assert answer.id, "/v1/responses through the semantic auto-router returned no response id"
|
||||
rows: Final = proxy.poll_logs_for_key(key, min_rows=1)
|
||||
rows: Final = proxy.poll_logs_for_key(
|
||||
key,
|
||||
min_rows=2,
|
||||
predicate=lambda logged: any(row.model == EMBEDDING_MODEL for row in logged),
|
||||
)
|
||||
embedding_rows: Final = tuple(row for row in rows if row.model == EMBEDDING_MODEL)
|
||||
assert embedding_rows, (
|
||||
"the routing embedding was not billed to the caller's key; "
|
||||
f"spend logs show {tuple(row.model for row in rows)}"
|
||||
)
|
||||
_assert_served_only_by(
|
||||
rows, CHEAP_SERVED | {semantic_auto_router.target}, "semantic auto-router /v1/responses string input"
|
||||
[row for row in rows if row.model != EMBEDDING_MODEL],
|
||||
CHEAP_SERVED | {semantic_auto_router.target},
|
||||
"semantic auto-router /v1/responses string input",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -184,6 +184,7 @@ test.describe("Admin tables scroll inside the page", () => {
|
|||
);
|
||||
try {
|
||||
await navigateToPage(page, Page.TagManagement);
|
||||
await setRowsPerPage(page, "50");
|
||||
await expectRowsAtLeast(page, SEED_ROWS);
|
||||
expect(await rowsPaintingPastAnAncestor(page)).toEqual([]);
|
||||
} finally {
|
||||
|
|
@ -207,6 +208,7 @@ test.describe("Admin tables scroll inside the page", () => {
|
|||
);
|
||||
try {
|
||||
await navigateToPage(page, Page.ModelHubTable);
|
||||
await setRowsPerPage(page, "50");
|
||||
await expectRowsAtLeast(page, SEED_ROWS);
|
||||
expect(await rowsPaintingPastAnAncestor(page)).toEqual([]);
|
||||
} finally {
|
||||
|
|
|
|||
|
|
@ -241,7 +241,7 @@ class TestRouterIndexManagement:
|
|||
"_get_deployment_by_litellm_model": "lookup by litellm_params.model, which is not indexed",
|
||||
"_finalize_adaptive_router_if_configured": 'init-time prefix scan for "auto_router/adaptive_router"; no index for prefix match',
|
||||
"config_deployments": "filters the whole list on model_info.db_model; admin path only (model add/upsert)",
|
||||
"heuristic_v2_router_limit_violation": "counts heuristic_v2 routers across the whole list; admin path only (auto-router init/upsert)",
|
||||
"auto_router_capability_violation": "counts gated auto-routers across the whole list; admin path only (auto-router init/upsert)",
|
||||
}
|
||||
|
||||
# Get path to router.py
|
||||
|
|
|
|||
|
|
@ -10,6 +10,8 @@ Covers the three defects from the ticket:
|
|||
handling live only on the native path).
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm_enterprise.enterprise_callbacks.secret_detection import (
|
||||
|
|
@ -19,12 +21,16 @@ from litellm.caching.caching import DualCache
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
AWS_KEY = "AKIAIOSFODNN7EXAMPLE"
|
||||
OPENAI_KEY = "sk-test-abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGH"
|
||||
SHORT_OPENAI_KEY = "sk-12345"
|
||||
UNICODE_DIGIT_SUFFIX = "sk-notification٣"
|
||||
STRIPE_LIVE_KEY = f"sk_live_{'1234567890' * 3}"
|
||||
URL_ENCODED_KEY = "Bearer%20sk-Ab3dEf6Gh7Ij8Kl9Mn0Pq2Rs3Tu4Vw5X"
|
||||
AWS_KEYS = [f"AKIAIOSFODNN7EXAMPL{suffix}" for suffix in "FEDCBA"]
|
||||
|
||||
|
||||
def _guardrail() -> _ENTERPRISE_SecretDetection:
|
||||
return _ENTERPRISE_SecretDetection(
|
||||
guardrail_name="hide-secrets", event_hook="pre_call", default_on=True
|
||||
)
|
||||
return _ENTERPRISE_SecretDetection(guardrail_name="hide-secrets", event_hook="pre_call", default_on=True)
|
||||
|
||||
|
||||
def _recorded(request_data: dict) -> dict:
|
||||
|
|
@ -33,6 +39,91 @@ def _recorded(request_data: dict) -> dict:
|
|||
return entries[0]
|
||||
|
||||
|
||||
def test_scan_message_preserves_benign_identifiers_and_xml_tags():
|
||||
guardrail = _guardrail()
|
||||
content = "<task-notification> model: claude-sonnet-4-5-20250929 </task-notification>"
|
||||
|
||||
assert guardrail.scan_message_for_secrets(content) == []
|
||||
assert guardrail.redact_text(content) == content
|
||||
assert guardrail.redact_text("result = compute(x) </task-notification>") == (
|
||||
"result = compute(x) </task-notification>"
|
||||
)
|
||||
|
||||
|
||||
def test_scan_message_preserves_quoted_benign_identifiers():
|
||||
guardrail = _guardrail()
|
||||
content = '{"content-type": "application/json", "model": "claude-sonnet-4-5-20250929"}'
|
||||
|
||||
assert guardrail.scan_message_for_secrets(content) == []
|
||||
assert guardrail.redact_text(content) == content
|
||||
|
||||
|
||||
def test_scan_message_redacts_every_openai_key_occurrence():
|
||||
guardrail = _guardrail()
|
||||
content = f"first {OPENAI_KEY}, second {OPENAI_KEY}"
|
||||
|
||||
assert guardrail.redact_text(content) == "first [REDACTED], second [REDACTED]"
|
||||
|
||||
|
||||
def test_scan_message_redacts_short_numeric_openai_like_values():
|
||||
guardrail = _guardrail()
|
||||
|
||||
assert guardrail.redact_text(f"value {SHORT_OPENAI_KEY}") == "value [REDACTED]"
|
||||
|
||||
|
||||
def test_scan_message_requires_ascii_digits_for_openai_like_values():
|
||||
guardrail = _guardrail()
|
||||
|
||||
assert guardrail.scan_message_for_secrets(UNICODE_DIGIT_SUFFIX) == []
|
||||
assert guardrail.redact_text(UNICODE_DIGIT_SUFFIX) == UNICODE_DIGIT_SUFFIX
|
||||
|
||||
|
||||
def test_scan_message_redacts_openai_key_after_separator():
|
||||
guardrail = _guardrail()
|
||||
|
||||
assert guardrail.redact_text(f"openai_{OPENAI_KEY} key-{OPENAI_KEY}") == (
|
||||
"openai_[REDACTED] key-[REDACTED]"
|
||||
)
|
||||
assert guardrail.redact_text(URL_ENCODED_KEY) == "Bearer%20[REDACTED]"
|
||||
|
||||
|
||||
def test_scan_message_does_not_stop_openai_key_at_token_characters():
|
||||
guardrail = _guardrail()
|
||||
|
||||
assert guardrail.redact_text("key sk-proj-abcde12345/extra") == "key [REDACTED]/extra"
|
||||
|
||||
|
||||
def test_scan_message_stays_linear_on_repeated_sk_separators():
|
||||
guardrail = _guardrail()
|
||||
content = "-sk-" * 25_000
|
||||
|
||||
started = time.perf_counter()
|
||||
assert guardrail.scan_message_for_secrets(content) == []
|
||||
assert time.perf_counter() - started < 2.0
|
||||
|
||||
|
||||
def test_scan_message_redacts_whole_stripe_live_key():
|
||||
guardrail = _guardrail()
|
||||
|
||||
assert guardrail.redact_text(f"stripe {STRIPE_LIVE_KEY} end") == "stripe [REDACTED] end"
|
||||
|
||||
|
||||
def test_scan_message_returns_matches_in_stable_order():
|
||||
guardrail = _guardrail()
|
||||
detected = guardrail.scan_message_for_secrets(" ".join(AWS_KEYS))
|
||||
|
||||
assert [secret["value"] for secret in detected] == sorted(AWS_KEYS)
|
||||
|
||||
|
||||
def test_scan_message_replaces_longest_overlapping_match_first():
|
||||
guardrail = _guardrail()
|
||||
content = f'token = "{OPENAI_KEY}/extra"'
|
||||
|
||||
detected = guardrail.scan_message_for_secrets(content)
|
||||
assert [secret["value"] for secret in detected] == [f"{OPENAI_KEY}/extra", OPENAI_KEY]
|
||||
assert guardrail.redact_text(content) == 'token = "[REDACTED]"'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_redacts_secrets():
|
||||
"""Playground path: the returned texts must carry [REDACTED], not the secret."""
|
||||
|
|
@ -199,9 +290,7 @@ async def test_apply_guardrail_without_texts_records_nothing():
|
|||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": "https://x/y.png"}}
|
||||
],
|
||||
"content": [{"type": "image_url", "image_url": {"url": "https://x/y.png"}}],
|
||||
}
|
||||
],
|
||||
"metadata": {},
|
||||
|
|
|
|||
|
|
@ -69,6 +69,30 @@ class TestCloudZeroStreamer:
|
|||
assert "2025-01-19" in result
|
||||
assert len(result["2025-01-19"]) == 1
|
||||
|
||||
def test_group_by_date_infers_schema_from_every_row(self):
|
||||
"""Test daily batches retain optional string columns that are null for thousands of leading rows."""
|
||||
streamer = CloudZeroStreamer("test-key", "test-connection")
|
||||
leading_nulls = 10_000
|
||||
rows = [
|
||||
{"time/usage_start": "2025-01-19T10:30:00Z", "resource/tag:team_alias": None}
|
||||
for _ in range(leading_nulls)
|
||||
]
|
||||
rows.append(
|
||||
{"time/usage_start": "2025-01-19T10:30:00Z", "resource/tag:team_alias": "team-alias"}
|
||||
)
|
||||
data = pl.DataFrame(
|
||||
rows,
|
||||
schema={"time/usage_start": pl.String, "resource/tag:team_alias": pl.String},
|
||||
)
|
||||
|
||||
result = streamer._group_by_date(data)
|
||||
|
||||
batch = result["2025-01-19"]
|
||||
assert len(batch) == leading_nulls + 1
|
||||
assert batch.schema["resource/tag:team_alias"] == pl.String
|
||||
assert batch["resource/tag:team_alias"].null_count() == leading_nulls
|
||||
assert batch.tail(1).item(0, "resource/tag:team_alias") == "team-alias"
|
||||
|
||||
def test_parse_and_convert_timestamp_utc(self):
|
||||
"""Test _parse_and_convert_timestamp method with UTC timestamp."""
|
||||
streamer = CloudZeroStreamer("test-key", "test-connection")
|
||||
|
|
|
|||
|
|
@ -86,6 +86,33 @@ class TestCBFTransformer:
|
|||
|
||||
assert result.is_empty()
|
||||
|
||||
def test_transform_keeps_tags_first_seen_after_row_100(self):
|
||||
transformer = CBFTransformer()
|
||||
teamless_rows = 101
|
||||
team_rows = 2
|
||||
total_rows = teamless_rows + team_rows
|
||||
data = pl.DataFrame(
|
||||
{
|
||||
"date": ["2025-01-19"] * total_rows,
|
||||
"successful_requests": [1] * total_rows,
|
||||
"spend": [0.5] * total_rows,
|
||||
"prompt_tokens": [10] * total_rows,
|
||||
"completion_tokens": [5] * total_rows,
|
||||
"model": ["gpt-4"] * total_rows,
|
||||
"custom_llm_provider": ["openai"] * total_rows,
|
||||
"api_key": ["sk-late-team"] * total_rows,
|
||||
"team_id": pl.Series([None] * teamless_rows + ["team-late"] * team_rows, dtype=pl.String),
|
||||
"team_alias": pl.Series([None] * teamless_rows + ["Late Team"] * team_rows, dtype=pl.String),
|
||||
}
|
||||
)
|
||||
|
||||
result = transformer.transform(data)
|
||||
|
||||
assert len(result) == total_rows
|
||||
assert "resource/tag:team_alias" in result.columns
|
||||
assert result["resource/tag:team_alias"].to_list() == [None] * teamless_rows + ["Late Team"] * team_rows
|
||||
assert result["resource/tag:entity_id"].to_list() == [None] * teamless_rows + ["Late Team"] * team_rows
|
||||
|
||||
def test_create_cbf_record(self):
|
||||
"""Test _create_cbf_record method with valid row data."""
|
||||
transformer = CBFTransformer()
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
from typing import TYPE_CHECKING, Literal, Optional
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -11,9 +12,11 @@ from litellm.integrations.custom_guardrail import (
|
|||
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailTracingDetail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class TestCustomGuardrailDeploymentHook:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_pre_call_deployment_hook_no_guardrails(self):
|
||||
"""Test that method returns kwargs unchanged when no guardrails are present"""
|
||||
|
|
@ -26,18 +29,14 @@ class TestCustomGuardrailDeploymentHook:
|
|||
"guardrails": None,
|
||||
}
|
||||
|
||||
result = await custom_guardrail.async_pre_call_deployment_hook(
|
||||
kwargs=kwargs, call_type=CallTypes.completion
|
||||
)
|
||||
result = await custom_guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.completion)
|
||||
|
||||
assert result == kwargs
|
||||
|
||||
# Test with guardrails as non-list
|
||||
kwargs["guardrails"] = "not_a_list"
|
||||
|
||||
result = await custom_guardrail.async_pre_call_deployment_hook(
|
||||
kwargs=kwargs, call_type=CallTypes.completion
|
||||
)
|
||||
result = await custom_guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.completion)
|
||||
|
||||
assert result == kwargs
|
||||
|
||||
|
|
@ -64,9 +63,7 @@ class TestCustomGuardrailDeploymentHook:
|
|||
"user_api_key_request_route": "test_route",
|
||||
}
|
||||
|
||||
result = await custom_guardrail.async_pre_call_deployment_hook(
|
||||
kwargs=kwargs, call_type=CallTypes.completion
|
||||
)
|
||||
result = await custom_guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.completion)
|
||||
|
||||
# Verify async_pre_call_hook was called with correct parameters
|
||||
custom_guardrail.async_pre_call_hook.assert_called_once()
|
||||
|
|
@ -99,9 +96,7 @@ class TestCustomGuardrailDeploymentHook:
|
|||
super().__init__(guardrail_name="g1", default_on=True)
|
||||
self.pre_call_count = 0
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict, cache, data, call_type
|
||||
):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.pre_call_count += 1
|
||||
return data
|
||||
|
||||
|
|
@ -114,9 +109,7 @@ class TestCustomGuardrailDeploymentHook:
|
|||
}
|
||||
|
||||
guardrail.mark_pre_call_hook_ran(kwargs)
|
||||
await guardrail.async_pre_call_deployment_hook(
|
||||
kwargs=kwargs, call_type=CallTypes.completion
|
||||
)
|
||||
await guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.completion)
|
||||
|
||||
assert guardrail.pre_call_count == 0
|
||||
|
||||
|
|
@ -130,9 +123,7 @@ class TestCustomGuardrailDeploymentHook:
|
|||
super().__init__(guardrail_name="g1", default_on=True)
|
||||
self.pre_call_count = 0
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict, cache, data, call_type
|
||||
):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.pre_call_count += 1
|
||||
return data
|
||||
|
||||
|
|
@ -144,9 +135,7 @@ class TestCustomGuardrailDeploymentHook:
|
|||
"metadata": {},
|
||||
}
|
||||
|
||||
await guardrail.async_pre_call_deployment_hook(
|
||||
kwargs=kwargs, call_type=CallTypes.completion
|
||||
)
|
||||
await guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.completion)
|
||||
|
||||
assert guardrail.pre_call_count == 1
|
||||
|
||||
|
|
@ -175,9 +164,7 @@ class TestCustomGuardrailDeploymentHook:
|
|||
super().__init__(guardrail_name="g1", default_on=True)
|
||||
self.pre_call_count = 0
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict, cache, data, call_type
|
||||
):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.pre_call_count += 1
|
||||
return data
|
||||
|
||||
|
|
@ -189,15 +176,12 @@ class TestCustomGuardrailDeploymentHook:
|
|||
"metadata": {PRE_CALL_EXECUTED_GUARDRAILS_KEY: ["g1"]},
|
||||
}
|
||||
|
||||
await guardrail.async_pre_call_deployment_hook(
|
||||
kwargs=kwargs, call_type=CallTypes.completion
|
||||
)
|
||||
await guardrail.async_pre_call_deployment_hook(kwargs=kwargs, call_type=CallTypes.completion)
|
||||
|
||||
assert guardrail.pre_call_count == 1
|
||||
|
||||
|
||||
class TestCustomGuardrailShouldRunGuardrail:
|
||||
|
||||
def test_should_run_guardrail_with_litellm_metadata(self):
|
||||
"""Test that should_run_guardrail works with litellm_metadata pattern"""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
|
@ -214,9 +198,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"litellm_metadata": {"guardrails": ["test_guardrail"]},
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
result = custom_guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
|
@ -236,9 +218,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"metadata": {"guardrails": ["test_guardrail"]},
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
result = custom_guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
|
@ -255,9 +235,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
# Test with guardrails at root level
|
||||
data = {"model": "gpt-3.5-turbo", "guardrails": ["test_guardrail"]}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
result = custom_guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
|
@ -277,9 +255,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"litellm_metadata": {"guardrails": ["different_guardrail"]},
|
||||
}
|
||||
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
result = custom_guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call)
|
||||
|
||||
assert result is False
|
||||
|
||||
|
|
@ -298,9 +274,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
result = custom_guardrail.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call)
|
||||
assert result is True, "Global guardrail should run when default_on=True"
|
||||
|
||||
# Test 2: User-injected disable at root level is IGNORED
|
||||
|
|
@ -312,9 +286,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_with_disable_root, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is True
|
||||
), "User-injected disable_global_guardrails should be ignored"
|
||||
assert result is True, "User-injected disable_global_guardrails should be ignored"
|
||||
|
||||
# Test 3: User-injected disable in metadata is IGNORED
|
||||
data_with_disable_metadata = {
|
||||
|
|
@ -345,12 +317,8 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}},
|
||||
"litellm_metadata": {"request_tags": ["user-supplied"]},
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_cross_key, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), "Admin config in metadata must not be shadowed by user-supplied litellm_metadata"
|
||||
result = custom_guardrail.should_run_guardrail(data=data_cross_key, event_type=GuardrailEventHooks.pre_call)
|
||||
assert result is False, "Admin config in metadata must not be shadowed by user-supplied litellm_metadata"
|
||||
|
||||
# Test 6: After the pre-call strip runs, user-injected
|
||||
# user_api_key_metadata in the non-authoritative metadata key is gone.
|
||||
|
|
@ -361,12 +329,8 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}},
|
||||
"litellm_metadata": {}, # post-strip: attacker payload removed
|
||||
}
|
||||
result = custom_guardrail.should_run_guardrail(
|
||||
data=data_post_strip, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), "Admin config in metadata must be respected when other metadata key is empty"
|
||||
result = custom_guardrail.should_run_guardrail(data=data_post_strip, event_type=GuardrailEventHooks.pre_call)
|
||||
assert result is False, "Admin config in metadata must be respected when other metadata key is empty"
|
||||
|
||||
def test_should_run_guardrail_key_disable_global_not_overruled_by_team_guardrail_list(
|
||||
self,
|
||||
|
|
@ -432,12 +396,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"messages": [{"role": "user", "content": "test"}],
|
||||
"opted_out_global_guardrails": ["global_guardrail"],
|
||||
}
|
||||
assert (
|
||||
custom_guardrail.should_run_guardrail(
|
||||
data=data_root, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert custom_guardrail.should_run_guardrail(data=data_root, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
# Test 2: User-injected opt-out in metadata is IGNORED
|
||||
data_metadata = {
|
||||
|
|
@ -446,10 +405,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"metadata": {"opted_out_global_guardrails": ["global_guardrail"]},
|
||||
}
|
||||
assert (
|
||||
custom_guardrail.should_run_guardrail(
|
||||
data=data_metadata, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is True
|
||||
custom_guardrail.should_run_guardrail(data=data_metadata, event_type=GuardrailEventHooks.pre_call) is True
|
||||
)
|
||||
|
||||
# Test 4: a different guardrail in the opt-out list → still runs
|
||||
|
|
@ -458,12 +414,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"messages": [{"role": "user", "content": "test"}],
|
||||
"metadata": {"opted_out_global_guardrails": ["some_other_guardrail"]},
|
||||
}
|
||||
assert (
|
||||
custom_guardrail.should_run_guardrail(
|
||||
data=data_other, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert custom_guardrail.should_run_guardrail(data=data_other, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
# Test 5: empty opt-out list → still runs
|
||||
data_empty = {
|
||||
|
|
@ -471,12 +422,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"messages": [{"role": "user", "content": "test"}],
|
||||
"metadata": {"opted_out_global_guardrails": []},
|
||||
}
|
||||
assert (
|
||||
custom_guardrail.should_run_guardrail(
|
||||
data=data_empty, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert custom_guardrail.should_run_guardrail(data=data_empty, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
# Test 6: malformed value (bool instead of list) → safely ignored, guardrail runs
|
||||
data_malformed = {
|
||||
|
|
@ -485,10 +431,7 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"metadata": {"opted_out_global_guardrails": True},
|
||||
}
|
||||
assert (
|
||||
custom_guardrail.should_run_guardrail(
|
||||
data=data_malformed, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is True
|
||||
custom_guardrail.should_run_guardrail(data=data_malformed, event_type=GuardrailEventHooks.pre_call) is True
|
||||
)
|
||||
|
||||
def test_should_run_guardrail_opt_out_does_not_affect_non_global(self):
|
||||
|
|
@ -511,12 +454,69 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
"guardrails": ["opt_in_guardrail"],
|
||||
},
|
||||
}
|
||||
assert (
|
||||
non_global.should_run_guardrail(
|
||||
data=data, event_type=GuardrailEventHooks.pre_call
|
||||
)
|
||||
is True
|
||||
assert non_global.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
def test_should_run_guardrail_suppressed_by_auto_router_compression(self):
|
||||
"""An auto router's own compression policy can suppress an otherwise-eligible
|
||||
guardrail, even one that is default_on and explicitly requested."""
|
||||
from litellm.proxy.guardrails import auto_router_compression
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
always_on = CustomGuardrail(
|
||||
guardrail_name="headroom-default",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
token = auto_router_compression._suppressed_compression_guardrails.set(frozenset({"headroom-default"}))
|
||||
try:
|
||||
assert (
|
||||
always_on.should_run_guardrail(data={"model": "smart-router"}, event_type=GuardrailEventHooks.pre_call)
|
||||
is False
|
||||
)
|
||||
finally:
|
||||
auto_router_compression._suppressed_compression_guardrails.reset(token)
|
||||
|
||||
def test_should_run_guardrail_suppression_does_not_affect_other_names(self):
|
||||
from litellm.proxy.guardrails import auto_router_compression
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
always_on = CustomGuardrail(
|
||||
guardrail_name="headroom-default",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
token = auto_router_compression._suppressed_compression_guardrails.set(frozenset({"some-other-guardrail"}))
|
||||
try:
|
||||
assert (
|
||||
always_on.should_run_guardrail(data={"model": "smart-router"}, event_type=GuardrailEventHooks.pre_call)
|
||||
is True
|
||||
)
|
||||
finally:
|
||||
auto_router_compression._suppressed_compression_guardrails.reset(token)
|
||||
|
||||
def test_request_metadata_can_never_suppress_a_guardrail(self):
|
||||
"""Regression (security): suppression state is request-scoped and server-set,
|
||||
never read from metadata. Metadata reaches spend logs the caller can read, so
|
||||
anything honored from there is something a later request could replay to switch
|
||||
off a PII or content-filter guardrail for itself."""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
always_on = CustomGuardrail(
|
||||
guardrail_name="headroom-default",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
forged = {
|
||||
"model": "smart-router",
|
||||
"metadata": {
|
||||
"_auto_router_suppressed_compression_guardrails": [
|
||||
"headroom-default",
|
||||
"any-token:headroom-default",
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
assert always_on.should_run_guardrail(data=forged, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
|
||||
class TestApplyGuardrailCheck:
|
||||
|
|
@ -555,35 +555,33 @@ class TestApplyGuardrailCheck:
|
|||
child_with_override = ChildGuardrailWithOverride()
|
||||
|
||||
# Test: CustomGuardrail itself has apply_guardrail in its __dict__
|
||||
assert (
|
||||
"apply_guardrail" in type(CustomGuardrail()).__dict__
|
||||
), "CustomGuardrail should have apply_guardrail in its own __dict__"
|
||||
assert "apply_guardrail" in type(CustomGuardrail()).__dict__, (
|
||||
"CustomGuardrail should have apply_guardrail in its own __dict__"
|
||||
)
|
||||
|
||||
# Test: ParentGuardrail inherits but doesn't override, so it should NOT be in __dict__
|
||||
assert (
|
||||
"apply_guardrail" not in type(parent_instance).__dict__
|
||||
), "ParentGuardrail should NOT have apply_guardrail in its own __dict__ (only inherited)"
|
||||
assert "apply_guardrail" not in type(parent_instance).__dict__, (
|
||||
"ParentGuardrail should NOT have apply_guardrail in its own __dict__ (only inherited)"
|
||||
)
|
||||
|
||||
# Test: ChildGuardrailWithoutOverride only inherits, should NOT be in __dict__
|
||||
assert (
|
||||
"apply_guardrail" not in type(child_without_override).__dict__
|
||||
), "ChildGuardrailWithoutOverride should NOT have apply_guardrail in its own __dict__ (only inherited)"
|
||||
assert "apply_guardrail" not in type(child_without_override).__dict__, (
|
||||
"ChildGuardrailWithoutOverride should NOT have apply_guardrail in its own __dict__ (only inherited)"
|
||||
)
|
||||
|
||||
# Test: ChildGuardrailWithOverride overrides the method, SHOULD be in __dict__
|
||||
assert (
|
||||
"apply_guardrail" in type(child_with_override).__dict__
|
||||
), "ChildGuardrailWithOverride SHOULD have apply_guardrail in its own __dict__ (overridden)"
|
||||
assert "apply_guardrail" in type(child_with_override).__dict__, (
|
||||
"ChildGuardrailWithOverride SHOULD have apply_guardrail in its own __dict__ (overridden)"
|
||||
)
|
||||
|
||||
# Verify that all instances still have the method via inheritance (hasattr)
|
||||
assert hasattr(
|
||||
parent_instance, "apply_guardrail"
|
||||
), "All instances should have apply_guardrail via inheritance"
|
||||
assert hasattr(
|
||||
child_without_override, "apply_guardrail"
|
||||
), "All instances should have apply_guardrail via inheritance"
|
||||
assert hasattr(
|
||||
child_with_override, "apply_guardrail"
|
||||
), "All instances should have apply_guardrail via inheritance"
|
||||
assert hasattr(parent_instance, "apply_guardrail"), "All instances should have apply_guardrail via inheritance"
|
||||
assert hasattr(child_without_override, "apply_guardrail"), (
|
||||
"All instances should have apply_guardrail via inheritance"
|
||||
)
|
||||
assert hasattr(child_with_override, "apply_guardrail"), (
|
||||
"All instances should have apply_guardrail via inheritance"
|
||||
)
|
||||
|
||||
|
||||
class TestGuardrailLoggingAggregation:
|
||||
|
|
@ -610,11 +608,7 @@ class TestGuardrailLoggingAggregation:
|
|||
|
||||
def test_appends_to_existing_metadata_list(self):
|
||||
request_data = {
|
||||
"metadata": {
|
||||
"standard_logging_guardrail_information": [
|
||||
{"guardrail_name": "existing_guardrail"}
|
||||
]
|
||||
}
|
||||
"metadata": {"standard_logging_guardrail_information": [{"guardrail_name": "existing_guardrail"}]}
|
||||
}
|
||||
|
||||
self._invoke_add_log(request_data)
|
||||
|
|
@ -626,11 +620,7 @@ class TestGuardrailLoggingAggregation:
|
|||
assert info[1]["guardrail_name"] == "test_guardrail"
|
||||
|
||||
def test_converts_existing_metadata_dict_to_list(self):
|
||||
request_data = {
|
||||
"metadata": {
|
||||
"standard_logging_guardrail_information": {"guardrail_name": "legacy"}
|
||||
}
|
||||
}
|
||||
request_data = {"metadata": {"standard_logging_guardrail_information": {"guardrail_name": "legacy"}}}
|
||||
|
||||
self._invoke_add_log(request_data)
|
||||
|
||||
|
|
@ -642,18 +632,12 @@ class TestGuardrailLoggingAggregation:
|
|||
|
||||
def test_appends_to_litellm_metadata(self):
|
||||
request_data = {
|
||||
"litellm_metadata": {
|
||||
"standard_logging_guardrail_information": [
|
||||
{"guardrail_name": "litellm_existing"}
|
||||
]
|
||||
}
|
||||
"litellm_metadata": {"standard_logging_guardrail_information": [{"guardrail_name": "litellm_existing"}]}
|
||||
}
|
||||
|
||||
self._invoke_add_log(request_data)
|
||||
|
||||
info = request_data["litellm_metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
]
|
||||
info = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
|
||||
assert isinstance(info, list)
|
||||
assert len(info) == 2
|
||||
assert info[1]["guardrail_name"] == "test_guardrail"
|
||||
|
|
@ -670,12 +654,10 @@ class TestGuardrailLoggingAggregation:
|
|||
|
||||
self._invoke_add_log(request_data)
|
||||
|
||||
assert (
|
||||
"standard_logging_guardrail_information" not in request_data["metadata"]
|
||||
), "entry landed in the caller's metadata, where the spend log does not read it"
|
||||
info = request_data["litellm_metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
]
|
||||
assert "standard_logging_guardrail_information" not in request_data["metadata"], (
|
||||
"entry landed in the caller's metadata, where the spend log does not read it"
|
||||
)
|
||||
info = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(info) == 1
|
||||
assert info[0]["guardrail_name"] == "test_guardrail"
|
||||
|
||||
|
|
@ -693,9 +675,7 @@ class TestGuardrailLoggingAggregation:
|
|||
}
|
||||
|
||||
self._invoke_add_log(request_data)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name="test_guardrail"
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name="test_guardrail")
|
||||
|
||||
buckets = {
|
||||
key
|
||||
|
|
@ -741,9 +721,7 @@ class TestGuardrailOtelSpanEmission:
|
|||
|
||||
assert len(captured) == 1
|
||||
emitted = captured[0]
|
||||
recorded = request_data["metadata"]["standard_logging_guardrail_information"][
|
||||
-1
|
||||
]
|
||||
recorded = request_data["metadata"]["standard_logging_guardrail_information"][-1]
|
||||
assert emitted is recorded
|
||||
assert emitted["guardrail_name"] == "emit_guard"
|
||||
assert emitted["start_time"] == 1.0
|
||||
|
|
@ -753,9 +731,7 @@ class TestGuardrailOtelSpanEmission:
|
|||
def _boom(_entry):
|
||||
raise RuntimeError("otel exporter down")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.integrations.otel.logger.emit_guardrail_span", _boom
|
||||
)
|
||||
monkeypatch.setattr("litellm.integrations.otel.logger.emit_guardrail_span", _boom)
|
||||
|
||||
request_data = {"metadata": {}}
|
||||
self._record(self._make_guardrail(), request_data)
|
||||
|
|
@ -852,9 +828,7 @@ class TestGuardrailSensitiveFieldStripping:
|
|||
duration=1.0,
|
||||
)
|
||||
|
||||
logged_response = request_data["metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
][0]["guardrail_response"]
|
||||
logged_response = request_data["metadata"]["standard_logging_guardrail_information"][0]["guardrail_response"]
|
||||
assert "secret_fields" not in logged_response
|
||||
assert "sk-live-SHOULD-NOT-APPEAR" not in json.dumps(logged_response)
|
||||
|
||||
|
|
@ -867,9 +841,7 @@ class TestGuardrailSensitiveFieldStripping:
|
|||
guardrail_json_response=[
|
||||
{
|
||||
"result": "ok",
|
||||
"secret_fields": {
|
||||
"raw_headers": {"authorization": "Bearer sk-secret"}
|
||||
},
|
||||
"secret_fields": {"raw_headers": {"authorization": "Bearer sk-secret"}},
|
||||
},
|
||||
{"result": "also_ok"},
|
||||
],
|
||||
|
|
@ -923,9 +895,7 @@ class TestGuardrailResponseCredentialMasking:
|
|||
duration=1.0,
|
||||
)
|
||||
|
||||
logged = request_data["metadata"]["standard_logging_guardrail_information"][0][
|
||||
"guardrail_response"
|
||||
]
|
||||
logged = request_data["metadata"]["standard_logging_guardrail_information"][0]["guardrail_response"]
|
||||
|
||||
masked_key = logged["metadata_snapshot"]["callback_vars"]["langsmith_api_key"]
|
||||
assert masked_key != plaintext_key
|
||||
|
|
@ -934,10 +904,7 @@ class TestGuardrailResponseCredentialMasking:
|
|||
|
||||
assert logged["model"] == "gpt-4o-mini"
|
||||
assert logged["messages"] == [{"role": "user", "content": "hi"}]
|
||||
assert (
|
||||
logged["metadata_snapshot"]["callback_vars"]["langsmith_project"]
|
||||
== "proj-name"
|
||||
)
|
||||
assert logged["metadata_snapshot"]["callback_vars"]["langsmith_project"] == "proj-name"
|
||||
|
||||
def test_nested_user_api_key_auth_metadata_is_masked(self):
|
||||
import json
|
||||
|
|
@ -996,9 +963,7 @@ class TestGuardrailResponseCredentialMasking:
|
|||
request_data: dict = {"metadata": {}}
|
||||
|
||||
guardrail.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response={
|
||||
"filters": [{"regex": r"\d{3}-\d{2}-\d{4}", "action": "BLOCKED"}]
|
||||
},
|
||||
guardrail_json_response={"filters": [{"regex": r"\d{3}-\d{2}-\d{4}", "action": "BLOCKED"}]},
|
||||
request_data=request_data,
|
||||
guardrail_status="success",
|
||||
)
|
||||
|
|
@ -1021,9 +986,7 @@ class TestGuardrailResponseCredentialMasking:
|
|||
guardrail_status="success",
|
||||
)
|
||||
|
||||
logged = request_data["metadata"]["standard_logging_guardrail_information"][0][
|
||||
"guardrail_response"
|
||||
]
|
||||
logged = request_data["metadata"]["standard_logging_guardrail_information"][0]["guardrail_response"]
|
||||
assert logged["flagged"] is True
|
||||
assert logged["score"] == 0.94
|
||||
assert logged["tokens_used"] == 42
|
||||
|
|
@ -1035,18 +998,14 @@ class TestGuardrailResponseCredentialMasking:
|
|||
plaintext = "lsv2_pt_abcdef1234567890"
|
||||
|
||||
guardrail.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response={
|
||||
"metadata_snapshot": {
|
||||
"callback_vars": {"langsmith_api_key": plaintext}
|
||||
}
|
||||
},
|
||||
guardrail_json_response={"metadata_snapshot": {"callback_vars": {"langsmith_api_key": plaintext}}},
|
||||
request_data=request_data,
|
||||
guardrail_status="success",
|
||||
)
|
||||
|
||||
masked = request_data["metadata"]["standard_logging_guardrail_information"][0][
|
||||
"guardrail_response"
|
||||
]["metadata_snapshot"]["callback_vars"]["langsmith_api_key"]
|
||||
masked = request_data["metadata"]["standard_logging_guardrail_information"][0]["guardrail_response"][
|
||||
"metadata_snapshot"
|
||||
]["callback_vars"]["langsmith_api_key"]
|
||||
assert masked != plaintext
|
||||
assert masked.startswith(plaintext[:4])
|
||||
assert masked.endswith(plaintext[-4:])
|
||||
|
|
@ -1540,9 +1499,7 @@ class TestEventTypeLogging:
|
|||
guardrail = TestGuardrail()
|
||||
request_data = {"metadata": {}}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["x"]}, request_data=request_data
|
||||
)
|
||||
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data)
|
||||
|
||||
logged_info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(logged_info) == 1, (
|
||||
|
|
@ -1584,9 +1541,7 @@ class TestEventTypeLogging:
|
|||
request_data = {"metadata": {}}
|
||||
|
||||
with pytest.raises(ValueError, match="blocked"):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["x"]}, request_data=request_data
|
||||
)
|
||||
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data)
|
||||
|
||||
logged_info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(logged_info) == 1
|
||||
|
|
@ -1715,9 +1670,7 @@ class TestTracingFieldsPopulation:
|
|||
guardrail_json_response="blocked",
|
||||
request_data=request_data,
|
||||
guardrail_status="guardrail_intervened",
|
||||
tracing_detail=GuardrailTracingDetail(
|
||||
policy_template="EU AI Act Article 5"
|
||||
),
|
||||
tracing_detail=GuardrailTracingDetail(policy_template="EU AI Act Article 5"),
|
||||
)
|
||||
|
||||
slg_list = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
|
|
@ -1759,13 +1712,7 @@ class TestCustomGuardrailSpendLogMatchRedaction:
|
|||
cg = CustomGuardrail(guardrail_name="test-rail")
|
||||
raw = {
|
||||
"assessments": [
|
||||
{
|
||||
"sensitiveInformationPolicy": {
|
||||
"piiEntities": [
|
||||
{"type": "NAME", "match": "GG", "action": "BLOCKED"}
|
||||
]
|
||||
}
|
||||
}
|
||||
{"sensitiveInformationPolicy": {"piiEntities": [{"type": "NAME", "match": "GG", "action": "BLOCKED"}]}}
|
||||
]
|
||||
}
|
||||
request_data: dict = {"metadata": {}}
|
||||
|
|
@ -1776,17 +1723,10 @@ class TestCustomGuardrailSpendLogMatchRedaction:
|
|||
)
|
||||
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert (
|
||||
slg["guardrail_response"]["assessments"][0]["sensitiveInformationPolicy"][
|
||||
"piiEntities"
|
||||
][0]["match"]
|
||||
slg["guardrail_response"]["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0]["match"]
|
||||
== "[REDACTED]"
|
||||
)
|
||||
assert (
|
||||
raw["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0][
|
||||
"match"
|
||||
]
|
||||
== "GG"
|
||||
)
|
||||
assert raw["assessments"][0]["sensitiveInformationPolicy"]["piiEntities"][0]["match"] == "GG"
|
||||
|
||||
def test_add_standard_logging_redacts_regex_field(self):
|
||||
cg = CustomGuardrail(guardrail_name="test-rail")
|
||||
|
|
@ -2239,6 +2179,170 @@ class TestRecordsOwnGuardrailInformation:
|
|||
assert _guardrail_entries(request_data) == []
|
||||
|
||||
|
||||
class _UndecoratedGuardrail(CustomGuardrail):
|
||||
"""apply_guardrail written like the docs example: no @log_guardrail_information."""
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
||||
if any("forbidden" in text for text in inputs.get("texts") or []):
|
||||
raise GuardrailRaisedException(guardrail_name=self.guardrail_name, message="Content blocked")
|
||||
return inputs
|
||||
|
||||
|
||||
class _UndecoratedSelfRecordingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response={"custom": True},
|
||||
request_data=request_data,
|
||||
guardrail_status="success",
|
||||
start_time=0.0,
|
||||
end_time=0.0,
|
||||
duration=0.0,
|
||||
)
|
||||
return inputs
|
||||
|
||||
|
||||
class _InheritedApplyGuardrail(_UndecoratedGuardrail):
|
||||
pass
|
||||
|
||||
|
||||
class TestUndecoratedApplyGuardrailIsLogged:
|
||||
"""LIT-5983 regression: a custom guardrail that overrides apply_guardrail without the
|
||||
@log_guardrail_information decorator must still record guardrail information, and the
|
||||
auto-wrap must not double-record decorated or self-recording implementations."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_undecorated_success_is_recorded(self):
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
guardrail = _UndecoratedGuardrail(guardrail_name="docs-style", event_hook=GuardrailEventHooks.pre_call)
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["hello"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
entries = _guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == "docs-style"
|
||||
assert entries[0]["guardrail_mode"] == "pre_call"
|
||||
assert entries[0]["guardrail_status"] == "success"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_undecorated_block_is_recorded_and_reraised(self):
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
||||
guardrail = _UndecoratedGuardrail(guardrail_name="docs-style")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["forbidden"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
entries = _guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == "docs-style"
|
||||
assert entries[0]["guardrail_status"] == "guardrail_intervened"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_undecorated_bare_exception_is_recorded_as_failed_to_respond(self):
|
||||
class _BareExceptionGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
raise Exception("Content blocked: policy violation")
|
||||
|
||||
guardrail = _BareExceptionGuardrail(guardrail_name="docs-style")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
with pytest.raises(Exception, match="Content blocked"):
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["x"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
entries = _guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inherited_apply_guardrail_is_recorded_once(self):
|
||||
guardrail = _InheritedApplyGuardrail(guardrail_name="child")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["hello"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert len(_guardrail_entries(request_data)) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_undecorated_self_recording_apply_guardrail_is_recorded_once(self):
|
||||
guardrail = _UndecoratedSelfRecordingGuardrail(guardrail_name="self-recording")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["hello"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
entries = _guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_response"] == {"custom": True}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_base_apply_guardrail_is_not_recorded(self):
|
||||
guardrail = CustomGuardrail(guardrail_name="base")
|
||||
request_data: dict = {"model": "gpt-4o"}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["hello"]),
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert _guardrail_entries(request_data) == []
|
||||
|
||||
def test_subclass_keywords_reach_cooperative_init_subclass(self):
|
||||
class _LabelMixin:
|
||||
seen_label: str = ""
|
||||
|
||||
def __init_subclass__(cls, label: str = "", **kwargs: object) -> None:
|
||||
super().__init_subclass__(**kwargs)
|
||||
cls.seen_label = label
|
||||
|
||||
class _Labelled(CustomGuardrail, _LabelMixin, label="docs-style"):
|
||||
pass
|
||||
|
||||
assert _Labelled.seen_label == "docs-style"
|
||||
|
||||
|
||||
class _ApplyOnlyObserver(CustomGuardrail):
|
||||
"""Overrides only apply_guardrail, like panw_prisma_airs; inherits async_logging_hook."""
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ the detached pipeline's single attempt-row write, and the cache-first job lookup
|
|||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -145,6 +146,48 @@ def _shadow_reply_router(message, finish_reason="stop", routed_model="cheap-mode
|
|||
return router
|
||||
|
||||
|
||||
def _reasoning_judge_router(
|
||||
reasoning_tokens: int, verdict: str = '{"preference": "A", "confidence": 0.9}'
|
||||
) -> MagicMock:
|
||||
"""A router whose judge arm reasons before it answers, the way a deployment carrying an
|
||||
elevated reasoning_effort does: reasoning bills against the caller's own max_tokens and
|
||||
the reply is cut off at that cap. One character stands in for one token."""
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == SHADOW_EVAL_ROUTER_CALL_ORIGIN:
|
||||
kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": "cheap-model"}
|
||||
return {"choices": [{"message": {"content": "shadow answer"}}]}
|
||||
budget_for_the_answer: Final = kwargs["max_tokens"] - reasoning_tokens
|
||||
return {"choices": [{"message": {"content": verdict[: max(0, budget_for_the_answer)]}}]}
|
||||
|
||||
router.acompletion = MagicMock(side_effect=acompletion)
|
||||
return router
|
||||
|
||||
|
||||
def _judge_reply_router(content: str | None, finish_reason: str = "stop", served_model: str = "judge-pick") -> MagicMock:
|
||||
"""A router whose judge arm returns a caller-shaped reply, so the shapes that all land
|
||||
on the same parser error can be posed apart: no content at all, versus JSON cut off
|
||||
mid-object."""
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == SHADOW_EVAL_ROUTER_CALL_ORIGIN:
|
||||
kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": "cheap-model"}
|
||||
return {"choices": [{"message": {"content": "shadow answer"}}]}
|
||||
return ModelResponse(
|
||||
model=served_model,
|
||||
choices=[{"index": 0, "finish_reason": finish_reason, "message": {"role": "assistant", "content": content}}],
|
||||
)
|
||||
|
||||
router.acompletion = MagicMock(side_effect=acompletion)
|
||||
return router
|
||||
|
||||
|
||||
TOOL_CALL_MESSAGE = {
|
||||
"content": None,
|
||||
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}],
|
||||
|
|
@ -1258,6 +1301,79 @@ class TestShadowPipeline:
|
|||
assert row["judge_cost"] == expected_cost
|
||||
assert row["shadow_cost"] == expected_shadow_cost
|
||||
|
||||
async def _judge_error(self, router: MagicMock, monkeypatch: pytest.MonkeyPatch) -> str:
|
||||
import litellm as litellm_module
|
||||
|
||||
monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.007)
|
||||
prisma = _prisma()
|
||||
await _logger(router=router, prisma=prisma)._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
real_text="real answer",
|
||||
real_model="claude-opus",
|
||||
real_cost=0.0,
|
||||
real_classifier_cost=0.0,
|
||||
real_cache_hit=False,
|
||||
control_tier=None,
|
||||
shadow_params={},
|
||||
parent_metadata={},
|
||||
)
|
||||
return prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]["error"]
|
||||
|
||||
async def test_a_judge_that_answered_nothing_is_told_apart_from_one_cut_off(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""Both land on the same parser message, and they want opposite fixes: a judge
|
||||
returning no content points at the reply never being text, while one cut off
|
||||
mid-object points at the output cap. The row has to say which."""
|
||||
truncated = '{"preference": "A", "confidence": 0.9, "reasoning": "'
|
||||
answered_nothing = await self._judge_error(_judge_reply_router(None), monkeypatch)
|
||||
cut_off = await self._judge_error(
|
||||
_judge_reply_router(truncated, finish_reason="length"), monkeypatch
|
||||
)
|
||||
|
||||
assert "content=no content" in answered_nothing
|
||||
assert "finish_reason=stop" in answered_nothing
|
||||
assert f"content={len(truncated)} chars" in cut_off
|
||||
assert "finish_reason=length" in cut_off
|
||||
|
||||
async def test_an_unparseable_verdict_names_the_model_that_served_it(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A judge_model that fans out over deployments hides which one truncates: without
|
||||
the served model the operator cannot tell a bad deployment from a bad cap."""
|
||||
error = await self._judge_error(_judge_reply_router(None, served_model="claude-sonnet-5"), monkeypatch)
|
||||
|
||||
assert "model=claude-sonnet-5" in error
|
||||
|
||||
async def test_a_diagnosed_verdict_error_stays_groupable(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""The customer groups attempt rows by error text. Every varying part has to sit
|
||||
after the first semicolon or each row becomes its own group."""
|
||||
first = await self._judge_error(_judge_reply_router(None, served_model="model-a"), monkeypatch)
|
||||
second = await self._judge_error(_judge_reply_router(None, served_model="model-b"), monkeypatch)
|
||||
|
||||
assert first != second
|
||||
assert first.split(";")[0] == second.split(";")[0]
|
||||
|
||||
async def test_a_judge_reply_that_cannot_be_read_still_records_an_error(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""The shape reader runs inside the failure path: it must never raise a second time
|
||||
and cost the row entirely."""
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == SHADOW_EVAL_ROUTER_CALL_ORIGIN:
|
||||
kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": "cheap-model"}
|
||||
return {"choices": [{"message": {"content": "shadow answer"}}]}
|
||||
return {"choices": []}
|
||||
|
||||
router.acompletion = MagicMock(side_effect=acompletion)
|
||||
|
||||
error = await self._judge_error(router, monkeypatch)
|
||||
|
||||
assert "unparseable judge verdict" in error
|
||||
assert "unreadable judge reply" in error
|
||||
|
||||
async def test_an_empty_shadow_reply_still_bills_its_cost(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A shadow call that returns no extractable text has still billed; pricing it at
|
||||
zero would keep the dollar gate open while shadow calls keep charging the key."""
|
||||
|
|
@ -1287,6 +1403,34 @@ class TestShadowPipeline:
|
|||
assert row["shadow_cost"] == 0.007
|
||||
assert logger._test_counter["spend:shadow_eval:job-1"] == 0.007
|
||||
|
||||
async def test_the_judge_output_cap_leaves_room_for_a_reasoning_judge(self):
|
||||
"""The output cap covers reasoning tokens as well as the answer, and a judge_model
|
||||
deployment carrying an elevated reasoning_effort spends that budget before it writes
|
||||
anything. A cap sized for the verdict JSON alone goes entirely to reasoning and the
|
||||
reply arrives empty, which the attempt records as an unparseable verdict rather than
|
||||
a result. The judge here burns a reasoning budget a live claude-sonnet-5 call was
|
||||
measured at, so the cap has to clear it for the verdict to survive."""
|
||||
reasoning_tokens = 2000
|
||||
logger = _logger(router=_reasoning_judge_router(reasoning_tokens), prisma=(prisma := _prisma()))
|
||||
|
||||
await logger._run_shadow_eval(
|
||||
job=_job(),
|
||||
request_id="req-1",
|
||||
messages=({"role": "user", "content": "hi"},),
|
||||
real_text="real answer",
|
||||
real_model="claude-opus",
|
||||
real_cost=0.0,
|
||||
real_classifier_cost=0.0,
|
||||
real_cache_hit=False,
|
||||
control_tier=None,
|
||||
shadow_params={},
|
||||
parent_metadata={},
|
||||
)
|
||||
|
||||
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
|
||||
assert row["outcome"] in ("real", "shadow", "tie"), row["error"]
|
||||
assert row["error"] is None
|
||||
|
||||
async def _no_text_error(self, router) -> str:
|
||||
prisma = _prisma()
|
||||
await _logger(router=router, prisma=prisma)._run_shadow_eval(
|
||||
|
|
|
|||
|
|
@ -995,6 +995,35 @@ async def test_anthropic_messages_marks_litellm_params_async():
|
|||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_marks_litellm_params_async(monkeypatch):
|
||||
"""LIT-6973: ``_arealtime`` must plant ``_arealtime`` in ``litellm_params`` so
|
||||
``_is_sync_litellm_request`` classifies the session async and a failed session
|
||||
reaches a CustomLogger's failure hook once, through the async path only, even
|
||||
though the sync ``failure_handler`` still runs ahead of the async one."""
|
||||
captured = {}
|
||||
async_logged = asyncio.Event()
|
||||
|
||||
class CaptureLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
captured["litellm_params"] = kwargs.get("litellm_params", {})
|
||||
async_logged.set()
|
||||
|
||||
logger = CaptureLogger()
|
||||
logger.log_failure_event = MagicMock()
|
||||
monkeypatch.setattr(litellm, "callbacks", [logger])
|
||||
monkeypatch.setattr(litellm, "failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_failure_callback", [])
|
||||
monkeypatch.setattr(litellm, "success_callback", [])
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [])
|
||||
with pytest.raises(ValueError, match="Unsupported model"):
|
||||
await litellm._arealtime(model="anthropic/claude-x", websocket=MagicMock())
|
||||
await asyncio.wait_for(async_logged.wait(), timeout=10)
|
||||
logger.log_failure_event.assert_not_called()
|
||||
assert captured["litellm_params"].get("_arealtime") is True
|
||||
assert LitellmLogging._is_sync_litellm_request(captured["litellm_params"]) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agenerate_content_marks_litellm_params_async():
|
||||
"""LIT-4475: the async ``agenerate_content`` entrypoint must plant
|
||||
|
|
@ -1085,6 +1114,56 @@ async def test_logging_non_streaming_request():
|
|||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_success_handler_truncates_large_base64_off_the_event_loop(monkeypatch):
|
||||
"""The standard logging payload's base64 scan of a large multimodal request must not run on the loop thread."""
|
||||
import threading
|
||||
|
||||
from litellm.litellm_core_utils import logging_utils
|
||||
|
||||
loop_thread = threading.get_ident()
|
||||
scan_threads: list[int] = []
|
||||
original_scan = logging_utils._truncate_base64_in_string
|
||||
|
||||
def recording_scan(value: str) -> str:
|
||||
scan_threads.append(threading.get_ident())
|
||||
return original_scan(value)
|
||||
|
||||
monkeypatch.setattr(logging_utils, "_truncate_base64_in_string", recording_scan)
|
||||
monkeypatch.setattr(logging_utils, "BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS", 1_000)
|
||||
|
||||
logged = asyncio.Event()
|
||||
captured: dict = {}
|
||||
|
||||
class CaptureLogger(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
captured["standard_logging_object"] = kwargs["standard_logging_object"]
|
||||
logged.set()
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [CaptureLogger()])
|
||||
payload = "L" * 20_000
|
||||
await litellm.acompletion(
|
||||
model="openai/gpt-5.6",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "describe"},
|
||||
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{payload}"}},
|
||||
],
|
||||
}
|
||||
],
|
||||
mock_response="ok",
|
||||
)
|
||||
await asyncio.wait_for(logged.wait(), timeout=10)
|
||||
|
||||
logged_url = captured["standard_logging_object"]["messages"][0]["content"][1]["image_url"]["url"]
|
||||
assert "base64_data truncated" in logged_url
|
||||
assert payload not in logged_url
|
||||
assert scan_threads
|
||||
assert loop_thread not in scan_threads
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"async_flag",
|
||||
[
|
||||
|
|
@ -1180,6 +1259,7 @@ def test_is_sync_litellm_request():
|
|||
assert LitellmLogging._is_sync_litellm_request({}) is True
|
||||
assert LitellmLogging._is_sync_litellm_request({"acompletion": True}) is False
|
||||
assert LitellmLogging._is_sync_litellm_request({"allm_passthrough_route": True}) is False
|
||||
assert LitellmLogging._is_sync_litellm_request({"_arealtime": True}) is False
|
||||
assert LitellmLogging._is_sync_litellm_request({"aanthropic_messages": True}) is False
|
||||
assert LitellmLogging._is_sync_litellm_request({"agenerate_content": True}) is False
|
||||
assert LitellmLogging._is_sync_litellm_request({"agenerate_content_stream": True}) is False
|
||||
|
|
@ -1466,6 +1546,62 @@ async def test_dispatch_failure_handlers_async_completes_before_sync_submit(
|
|||
assert events == ["async_start", "async_end", "sync_submit"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_failure_handlers_submits_sync_handler_when_task_is_cancelled(
|
||||
logging_obj,
|
||||
):
|
||||
"""Cancelling the dispatch task mid-await still submits the sync failure_handler.
|
||||
|
||||
Router failure paths fire the dispatcher with ``asyncio.create_task`` and raise
|
||||
right away. When the event loop is torn down before the task finishes (a short
|
||||
``asyncio.run`` in the SDK), the cancelled task must still hand the sync callbacks
|
||||
to the executor, as the old raw-thread path did, and only once the async handler
|
||||
has stopped.
|
||||
"""
|
||||
exception = ValueError("boom")
|
||||
traceback_exception = "traceback"
|
||||
events: list[str] = []
|
||||
async_started = asyncio.Event()
|
||||
|
||||
async def _async_failure(exc, tb, **kwargs):
|
||||
events.append("async_start")
|
||||
async_started.set()
|
||||
await asyncio.sleep(10)
|
||||
events.append("async_end")
|
||||
|
||||
def _submit(*args, **kwargs):
|
||||
events.append("sync_submit")
|
||||
|
||||
logging_obj.model_call_details["litellm_params"] = {}
|
||||
|
||||
with (
|
||||
patch.object(logging_obj, "async_failure_handler", side_effect=_async_failure),
|
||||
patch.object(logging_obj, "failure_handler", new_callable=MagicMock),
|
||||
patch.object(
|
||||
logging_obj,
|
||||
"_should_run_sync_failure_callbacks_for_async_calls",
|
||||
return_value=True,
|
||||
),
|
||||
patch( # test-quality-ok: the executor submit is the observable
|
||||
"litellm.litellm_core_utils.litellm_logging.executor.submit",
|
||||
side_effect=_submit,
|
||||
),
|
||||
):
|
||||
task = asyncio.create_task(
|
||||
logging_obj.dispatch_failure_handlers(
|
||||
exception,
|
||||
traceback_exception,
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
await async_started.wait()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert events == ["async_start", "sync_submit"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_failure_handlers_submits_sync_handler_for_failure_only_callbacks(
|
||||
logging_obj,
|
||||
|
|
@ -5997,6 +6133,34 @@ def test_failure_handler_helper_fn_builds_payload_once_per_exception():
|
|||
assert obj.model_call_details["standard_logging_object"] is not first_payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_failure_handler_reuses_payload_after_callable_async_callback():
|
||||
"""Regression for LIT-6886: the proxy runs async_failure_handler, then the threaded
|
||||
failure_handler, for every rejected request. A plain-function async callback (the
|
||||
Router registers one) is dispatched through CustomLogger.async_log_event, which
|
||||
restamps log_event_type on the shared model_call_details; the sync handler then
|
||||
rebuilt the standardized payload, doubling the redaction and payload cost of a 403."""
|
||||
router_style_callback = AsyncMock()
|
||||
obj = LitellmLogging(
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hey"}],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="lit-6886-1",
|
||||
function_id="f",
|
||||
dynamic_async_failure_callbacks=[router_style_callback],
|
||||
)
|
||||
exc = _raise_and_catch(_ClientError(status_code=403, message="key not allowed to access model"))
|
||||
await obj.async_failure_handler(exception=exc, traceback_exception="")
|
||||
first_payload = obj.model_call_details["standard_logging_object"]
|
||||
assert first_payload is not None
|
||||
assert router_style_callback.await_count == 1
|
||||
|
||||
obj.failure_handler(exc, "")
|
||||
assert obj.model_call_details["standard_logging_object"] is first_payload
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_hook_injection_marker_recorded_for_every_surface(logging_obj):
|
||||
"""The savings gate reads litellm_gateway_injected_cache from the request's
|
||||
|
|
|
|||
|
|
@ -2,12 +2,16 @@
|
|||
Tests for litellm.litellm_core_utils.logging_utils — base64 truncation helpers.
|
||||
"""
|
||||
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils import logging_utils
|
||||
from litellm.litellm_core_utils.logging_utils import (
|
||||
_format_base64_size,
|
||||
_truncate_base64_in_string,
|
||||
truncate_base64_in_messages,
|
||||
truncate_base64_in_messages_async,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -157,3 +161,70 @@ class TestTruncateBase64InMessages:
|
|||
result[0]["content"][0]["image_url"]["url"]
|
||||
== f"data:image/png;base64,{short}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# truncate_base64_in_messages_async
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _image_messages(payload: str) -> list:
|
||||
return [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "describe"},
|
||||
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{payload}"}},
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def scan_threads(monkeypatch):
|
||||
"""Record the thread that runs every base64 regex scan."""
|
||||
threads: list[int] = []
|
||||
original = logging_utils._truncate_base64_in_string
|
||||
|
||||
def recording_scan(value: str) -> str:
|
||||
threads.append(threading.get_ident())
|
||||
return original(value)
|
||||
|
||||
monkeypatch.setattr(logging_utils, "_truncate_base64_in_string", recording_scan)
|
||||
return threads
|
||||
|
||||
|
||||
class TestTruncateBase64InMessagesAsync:
|
||||
@pytest.mark.asyncio
|
||||
async def test_large_payload_is_scanned_off_the_event_loop(self, monkeypatch, scan_threads):
|
||||
monkeypatch.setattr(logging_utils, "BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS", 1_000)
|
||||
payload = "I" * 20_000
|
||||
messages = _image_messages(payload)
|
||||
|
||||
result = await truncate_base64_in_messages_async(messages)
|
||||
offload_threads = tuple(scan_threads)
|
||||
|
||||
assert result == truncate_base64_in_messages(messages)
|
||||
assert payload not in result[0]["content"][1]["image_url"]["url"]
|
||||
assert payload in messages[0]["content"][1]["image_url"]["url"]
|
||||
assert offload_threads
|
||||
assert threading.get_ident() not in offload_threads
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_small_payload_stays_on_the_calling_thread(self, monkeypatch, scan_threads):
|
||||
monkeypatch.setattr(logging_utils, "BASE64_TRUNCATION_OFFLOAD_THRESHOLD_CHARS", 1_000)
|
||||
messages = _image_messages("J" * 200)
|
||||
|
||||
result = await truncate_base64_in_messages_async(messages)
|
||||
|
||||
assert result == truncate_base64_in_messages(messages)
|
||||
assert scan_threads
|
||||
assert set(scan_threads) == {threading.get_ident()}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_and_disabled_truncation_short_circuit(self, monkeypatch, scan_threads):
|
||||
assert await truncate_base64_in_messages_async(None) is None
|
||||
monkeypatch.setattr(logging_utils, "MAX_BASE64_LENGTH_FOR_LOGGING", 0)
|
||||
messages = _image_messages("K" * 20_000)
|
||||
assert await truncate_base64_in_messages_async(messages) is messages
|
||||
assert scan_threads == []
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
|
||||
client_close_code,
|
||||
realtime_error_event,
|
||||
websocket_close_reason,
|
||||
)
|
||||
|
|
@ -42,3 +44,11 @@ def test_websocket_close_reason_truncates_multibyte_message_by_bytes():
|
|||
assert len(reason.encode("utf-8")) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES
|
||||
assert reason == "あ" * (WEBSOCKET_CLOSE_REASON_MAX_BYTES // 3)
|
||||
assert "<EFBFBD>" not in reason
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("upstream_code", "expected"),
|
||||
[(1000, 1000), (1008, 1008), (1011, 1011), (4001, 4001), (1005, 1011), (1006, 1011), (1015, 1011), (2999, 1011)],
|
||||
)
|
||||
def test_client_close_code_only_forwards_codes_a_server_may_send(upstream_code, expected):
|
||||
assert client_close_code(upstream_code) == expected
|
||||
|
|
|
|||
|
|
@ -1,14 +1,20 @@
|
|||
import asyncio
|
||||
import json
|
||||
from collections.abc import Coroutine
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from websockets.frames import Close
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.realtime_streaming import (
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
RealTimeStreaming,
|
||||
client_sent_openai_beta_realtime_header,
|
||||
)
|
||||
|
|
@ -2941,13 +2947,11 @@ async def test_log_messages_routes_async_logging_through_bounded_worker():
|
|||
realtime turn leaves a suspended task pinning its response in memory -> an
|
||||
unbounded leak. Regression for that fix."""
|
||||
logging_obj = MagicMock()
|
||||
streaming = RealTimeStreaming(MagicMock(), MagicMock(), logging_obj)
|
||||
mock_worker = MagicMock()
|
||||
streaming = RealTimeStreaming(MagicMock(), MagicMock(), logging_obj, logging_worker=mock_worker)
|
||||
streaming.messages = [{"type": "session.created"}]
|
||||
|
||||
with (
|
||||
patch("litellm.litellm_core_utils.realtime_streaming.GLOBAL_LOGGING_WORKER") as mock_worker,
|
||||
patch("litellm.litellm_core_utils.realtime_streaming.asyncio.create_task") as mock_create_task,
|
||||
):
|
||||
with patch("litellm.litellm_core_utils.realtime_streaming.asyncio.create_task") as mock_create_task:
|
||||
await streaming.log_messages()
|
||||
|
||||
mock_worker.ensure_initialized_and_enqueue.assert_called_once()
|
||||
|
|
@ -3028,12 +3032,12 @@ async def test_session_close_flushes_unbilled_transcription_usage():
|
|||
messages before log_messages runs, and never forwarded to the client."""
|
||||
from typing import Final
|
||||
|
||||
from litellm.types.realtime import RealtimeInputAudioTranscriptionUsage
|
||||
from litellm.types.realtime import RealtimeInputAudioTranscriptionUsage, RealtimeResponseTypedDict
|
||||
|
||||
client_ws: Final = MagicMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
backend_ws: Final = MagicMock()
|
||||
backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None))
|
||||
backend_ws.recv = AsyncMock(side_effect=[b'{"serverContent": {}}', ConnectionClosed(None, None)])
|
||||
logging_obj: Final = MagicMock()
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.success_handler = MagicMock()
|
||||
|
|
@ -3045,7 +3049,24 @@ async def test_session_close_flushes_unbilled_transcription_usage():
|
|||
"total_tokens": 171,
|
||||
"input_token_details": {"text_tokens": 0, "audio_tokens": 153},
|
||||
}
|
||||
transcript_frame: Final[RealtimeResponseTypedDict] = {
|
||||
"response": {
|
||||
"type": "conversation.item.input_audio_transcription.completed",
|
||||
"event_id": "event_1",
|
||||
"transcript": "ahoy",
|
||||
"item_id": "item_1",
|
||||
"content_index": 0,
|
||||
},
|
||||
"current_output_item_id": None,
|
||||
"current_response_id": None,
|
||||
"current_delta_chunks": None,
|
||||
"current_conversation_id": None,
|
||||
"current_item_chunks": None,
|
||||
"current_delta_type": None,
|
||||
"session_configuration_request": None,
|
||||
}
|
||||
provider_config: Final = MagicMock()
|
||||
provider_config.transform_realtime_response = MagicMock(return_value=transcript_frame)
|
||||
provider_config.unbilled_usage_on_session_close = MagicMock(return_value=usage)
|
||||
|
||||
streaming: Final = RealTimeStreaming(
|
||||
|
|
@ -3077,7 +3098,9 @@ async def test_session_close_flushes_unbilled_transcription_usage():
|
|||
)
|
||||
assert len(flushed) == 1
|
||||
assert flushed[0] in logged_snapshots[0]
|
||||
assert not client_ws.send_text.called
|
||||
forwarded: Final = tuple(json.loads(call.args[0]) for call in client_ws.send_text.await_args_list)
|
||||
assert [event.get("transcript") for event in forwarded] == ["ahoy"]
|
||||
assert all("usage" not in event for event in forwarded)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3111,3 +3134,281 @@ async def test_session_close_flush_noop_without_unbilled_usage():
|
|||
isinstance(message, dict) and message.get("type") == "conversation.item.input_audio_transcription.completed"
|
||||
for message in streaming.messages
|
||||
)
|
||||
|
||||
|
||||
|
||||
_UPSTREAM_REFUSAL: Final = "Publisher model `publishers/google/models/gemini-live-2.5-flash` was not found"
|
||||
|
||||
|
||||
class _InlineLoggingWorker:
|
||||
def __init__(self) -> None:
|
||||
self.enqueued: tuple[Coroutine[object, object, None], ...] = ()
|
||||
|
||||
def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None:
|
||||
self.enqueued = (*self.enqueued, async_coroutine)
|
||||
|
||||
async def drain(self) -> None:
|
||||
for coroutine in self.enqueued:
|
||||
await coroutine
|
||||
|
||||
|
||||
class _RecordingLogging:
|
||||
def __init__(self) -> None:
|
||||
self.model_call_details: dict[str, object] = {}
|
||||
self.logged_sessions: tuple[tuple[dict, ...], ...] = ()
|
||||
self.logged_failures: tuple[Exception, ...] = ()
|
||||
|
||||
def pre_call(self, input: str | dict, api_key: str) -> None:
|
||||
return None
|
||||
|
||||
async def dispatch_success_handlers(self, result: list[dict], prefer_async_handlers: bool = False) -> None:
|
||||
self.logged_sessions = (*self.logged_sessions, tuple(result))
|
||||
|
||||
async def dispatch_failure_handlers(
|
||||
self, exception: Exception, traceback_exception: str, prefer_async_handlers: bool = False
|
||||
) -> None:
|
||||
self.logged_failures = (*self.logged_failures, exception)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RelaySession:
|
||||
streaming: RealTimeStreaming
|
||||
logging: _RecordingLogging
|
||||
worker: _InlineLoggingWorker
|
||||
|
||||
async def run(self) -> None:
|
||||
await asyncio.wait_for(self.streaming.bidirectional_forward(), timeout=2)
|
||||
await self.worker.drain()
|
||||
|
||||
|
||||
async def _wait_forever() -> str:
|
||||
await asyncio.Event().wait()
|
||||
raise AssertionError("unreachable")
|
||||
|
||||
|
||||
def _client_ws_that_never_sends() -> MagicMock:
|
||||
client_ws: Final = MagicMock()
|
||||
client_ws.headers = {}
|
||||
client_ws.receive_text = AsyncMock(side_effect=_wait_forever)
|
||||
client_ws.send_text = AsyncMock()
|
||||
client_ws.close = AsyncMock()
|
||||
return client_ws
|
||||
|
||||
|
||||
def _backend_ws_closing_with(*frames: bytes | Exception) -> MagicMock:
|
||||
backend_ws: Final = MagicMock()
|
||||
backend_ws.recv = AsyncMock(side_effect=list(frames))
|
||||
return backend_ws
|
||||
|
||||
|
||||
def _relay_session(client_ws: MagicMock, backend_ws: MagicMock) -> _RelaySession:
|
||||
logging: Final = _RecordingLogging()
|
||||
worker: Final = _InlineLoggingWorker()
|
||||
streaming: Final = RealTimeStreaming(
|
||||
client_ws, backend_ws, logging, model="gpt-realtime", logging_worker=worker
|
||||
)
|
||||
return _RelaySession(streaming=streaming, logging=logging, worker=worker)
|
||||
|
||||
|
||||
def _error_events_sent_to(client_ws: MagicMock) -> list[dict]:
|
||||
events: Final = (json.loads(call.args[0]) for call in client_ws.send_text.await_args_list)
|
||||
return [event for event in events if event.get("type") == "error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidirectional_forward_relays_upstream_policy_close_to_client():
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
session: Final = _relay_session(client_ws, _backend_ws_closing_with(upstream_close))
|
||||
|
||||
await session.run()
|
||||
|
||||
(error_event,) = _error_events_sent_to(client_ws)
|
||||
assert error_event["error"]["type"] == "server_error"
|
||||
assert "1008" in error_event["error"]["message"]
|
||||
assert _UPSTREAM_REFUSAL in error_event["error"]["message"]
|
||||
client_ws.close.assert_awaited_once_with(code=1008, reason=_UPSTREAM_REFUSAL)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"leaked_detail",
|
||||
(
|
||||
pytest.param("sk-live-abcdef0123456789abcdef0123", id="credential"),
|
||||
pytest.param("vertex-int.svc.cluster.local", id="internal-hostname"),
|
||||
pytest.param("/etc/litellm/service-account.json", id="filesystem-path"),
|
||||
),
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_close_details_are_scrubbed_before_reaching_the_client(leaked_detail: str):
|
||||
"""LIT-6973: the relayed close goes through the proxy's client-facing redaction, so an upstream
|
||||
error echoing a credential, an internal host, or a server path never reaches the client verbatim."""
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, f"upstream rejected: {leaked_detail}"), None)
|
||||
session: Final = _relay_session(client_ws, _backend_ws_closing_with(upstream_close))
|
||||
|
||||
await session.run()
|
||||
|
||||
(error_event,) = _error_events_sent_to(client_ws)
|
||||
assert leaked_detail not in error_event["error"]["message"]
|
||||
relayed_reason: Final = client_ws.close.await_args.kwargs["reason"]
|
||||
assert leaked_detail not in relayed_reason
|
||||
assert "REDACTED" in relayed_reason
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidirectional_forward_maps_abnormal_upstream_close_to_internal_error():
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
session: Final = _relay_session(client_ws, _backend_ws_closing_with(ConnectionClosed(None, None)))
|
||||
|
||||
await session.run()
|
||||
|
||||
(error_event,) = _error_events_sent_to(client_ws)
|
||||
assert "1006" in error_event["error"]["message"]
|
||||
client_ws.close.assert_awaited_once()
|
||||
assert client_ws.close.await_args.kwargs["code"] == 1011
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bidirectional_forward_relays_normal_upstream_close_without_error_event():
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
session: Final = _relay_session(client_ws, _backend_ws_closing_with(ConnectionClosed(Close(1000, ""), None)))
|
||||
|
||||
await session.run()
|
||||
|
||||
assert _error_events_sent_to(client_ws) == []
|
||||
client_ws.close.assert_awaited_once()
|
||||
assert client_ws.close.await_args.kwargs["code"] == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_refusal_before_any_frame_logs_a_failure_not_a_success():
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
session: Final = _relay_session(_client_ws_that_never_sends(), _backend_ws_closing_with(upstream_close))
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.logging.logged_failures == (upstream_close,)
|
||||
assert session.logging.logged_sessions == ()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_refusal_after_a_synthetic_session_created_still_logs_a_failure():
|
||||
"""LIT-6973: deferred Gemini Live setup stores a synthetic ``session.created`` before
|
||||
the relay starts. It is not an upstream frame, so a refusal after it is still a refusal."""
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
session: Final = _relay_session(_client_ws_that_never_sends(), _backend_ws_closing_with(upstream_close))
|
||||
session.streaming.store_message(json.dumps({"type": "session.created", "session": {"id": "sess_synthetic"}}))
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.logging.logged_failures == (upstream_close,)
|
||||
assert session.logging.logged_sessions == ()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_close_after_relayed_events_still_logs_the_session_as_success():
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
session_created: Final = json.dumps({"type": "session.created", "session": {"id": "sess_1"}}).encode()
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
session: Final = _relay_session(client_ws, _backend_ws_closing_with(session_created, upstream_close))
|
||||
|
||||
await session.run()
|
||||
|
||||
(logged_session,) = session.logging.logged_sessions
|
||||
assert [event["type"] for event in logged_session] == ["session.created"]
|
||||
assert session.logging.logged_failures == ()
|
||||
client_ws.close.assert_awaited_once_with(code=1008, reason=_UPSTREAM_REFUSAL)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_closing_while_a_client_message_is_forwarded_still_reaches_the_client():
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
backend_closed: Final = asyncio.Event()
|
||||
client_messages: Final = iter((json.dumps({"type": "response.create"}),))
|
||||
|
||||
async def receive_text() -> str:
|
||||
message = next(client_messages, None)
|
||||
return message if message is not None else await _wait_forever()
|
||||
|
||||
async def send_to_backend(_message: str) -> None:
|
||||
backend_closed.set()
|
||||
raise upstream_close
|
||||
|
||||
async def recv_from_backend() -> bytes:
|
||||
await backend_closed.wait()
|
||||
raise upstream_close
|
||||
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
client_ws.receive_text = receive_text
|
||||
backend_ws: Final = MagicMock()
|
||||
backend_ws.send = send_to_backend
|
||||
backend_ws.recv = recv_from_backend
|
||||
session: Final = _relay_session(client_ws, backend_ws)
|
||||
|
||||
await session.run()
|
||||
|
||||
(error_event,) = _error_events_sent_to(client_ws)
|
||||
assert _UPSTREAM_REFUSAL in error_event["error"]["message"]
|
||||
client_ws.close.assert_awaited_once_with(code=1008, reason=_UPSTREAM_REFUSAL)
|
||||
assert session.logging.logged_failures == (upstream_close,)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_hanging_up_first_ends_the_session_without_a_relayed_close():
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
client_ws.receive_text = AsyncMock(side_effect=RuntimeError("client went away"))
|
||||
backend_ws: Final = MagicMock()
|
||||
backend_ws.recv = AsyncMock(side_effect=_wait_forever)
|
||||
session: Final = _relay_session(client_ws, backend_ws)
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.logging.logged_sessions == ((),)
|
||||
assert session.logging.logged_failures == ()
|
||||
client_ws.close.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_hanging_up_with_a_websockets_close_is_not_mistaken_for_the_backend_closing():
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
client_ws.receive_text = AsyncMock(side_effect=ConnectionClosed(None, None))
|
||||
backend_ws: Final = MagicMock()
|
||||
backend_ws.recv = AsyncMock(side_effect=_wait_forever)
|
||||
session: Final = _relay_session(client_ws, backend_ws)
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.logging.logged_sessions == ((),)
|
||||
assert session.logging.logged_failures == ()
|
||||
client_ws.close.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_logging_stamps_the_reservation_ownership_marker():
|
||||
"""LIT-6973: only the success path enqueues the cost callback that settles the
|
||||
session's budget reservation, so it stamps REALTIME_SESSION_SUCCESS_LOGGED_KEY on
|
||||
the shared logging object. The proxy endpoint reads that stamp to decide whether to
|
||||
release the reservation itself, so a logged-as-success session must carry it."""
|
||||
client_ws: Final = _client_ws_that_never_sends()
|
||||
session_created: Final = json.dumps({"type": "session.created", "session": {"id": "sess_1"}}).encode()
|
||||
upstream_close: Final = ConnectionClosed(Close(1000, ""), None)
|
||||
session: Final = _relay_session(client_ws, _backend_ws_closing_with(session_created, upstream_close))
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.logging.logged_sessions != ()
|
||||
assert session.logging.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refused_session_does_not_stamp_the_reservation_ownership_marker():
|
||||
"""A refused session logs a failure, not a success, so it must not stamp
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY. If it did, the proxy endpoint would skip its
|
||||
own reservation release and the refused session's reservation would stay pinned."""
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
session: Final = _relay_session(_client_ws_that_never_sends(), _backend_ws_closing_with(upstream_close))
|
||||
|
||||
await session.run()
|
||||
|
||||
assert session.logging.logged_failures == (upstream_close,)
|
||||
assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details
|
||||
|
|
|
|||
|
|
@ -1075,6 +1075,37 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
|||
assert result == responses_so_far
|
||||
|
||||
|
||||
class TestUndecoratedGuardrailIsRecorded:
|
||||
"""LIT-5983 regression: the handler calls apply_guardrail bare, so a custom guardrail
|
||||
without @log_guardrail_information must still end up in the request's guardrail
|
||||
information on both the request and response paths."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_path_records_undecorated_guardrail(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="docs-style")
|
||||
data = {"messages": [{"role": "user", "content": "hello"}], "metadata": {}}
|
||||
|
||||
await handler.process_input_messages(data, guardrail)
|
||||
|
||||
entries = data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert [(e["guardrail_name"], e["guardrail_status"]) for e in entries] == [("docs-style", "success")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_path_records_undecorated_guardrail(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="docs-style")
|
||||
response = ModelResponse(
|
||||
choices=[Choices(finish_reason="stop", index=0, message=Message(content="hi", role="assistant"))]
|
||||
)
|
||||
request_data: dict = {"metadata": {}}
|
||||
|
||||
await handler.process_output_response(response, guardrail, request_data=request_data)
|
||||
|
||||
entries = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert [(e["guardrail_name"], e["guardrail_status"]) for e in entries] == [("docs-style", "success")]
|
||||
|
||||
|
||||
class TestGetStructuredMessages:
|
||||
"""Test the get_structured_messages method."""
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from fastapi import HTTPException
|
|||
from pydantic import ValidationError
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import MCPServerURLCredentialsError
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
oauth_protected_resource_path,
|
||||
raise_public,
|
||||
|
|
@ -442,6 +443,17 @@ def test_raise_public_maps_each_error_to_its_status(error, status):
|
|||
assert exc_info.value.status_code == status
|
||||
|
||||
|
||||
def test_raise_public_marks_only_url_credentials_error_as_safe_for_preview():
|
||||
with pytest.raises(HTTPException) as generic_exc_info:
|
||||
raise_public(CredError.of_misconfigured("private operator detail"))
|
||||
assert not isinstance(generic_exc_info.value, MCPServerURLCredentialsError)
|
||||
|
||||
error = CredError.of_url_credentials_not_allowed()
|
||||
with pytest.raises(MCPServerURLCredentialsError) as url_exc_info:
|
||||
raise_public(error)
|
||||
assert url_exc_info.value.detail == error.summary
|
||||
|
||||
|
||||
def test_raise_public_emits_unauthorized_challenge():
|
||||
body = {"error": "byok_auth_required", "server_id": "s1"}
|
||||
error = CredError.of_unauthorized("needs key", www_authenticate='Bearer resource_metadata="/x"', body=body)
|
||||
|
|
|
|||
|
|
@ -9,9 +9,11 @@ returning the stub.
|
|||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import httpx
|
||||
import jwt as pyjwt
|
||||
import pytest
|
||||
from pydantic import SecretStr
|
||||
|
||||
|
|
@ -42,6 +44,11 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_sto
|
|||
OAuthToken,
|
||||
TokenStoreUnavailable,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_refresher import (
|
||||
RefreshingSSOAssertionStore,
|
||||
SSOAssertionRefresher,
|
||||
SSOClientConfig,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
AssertionStoreUnavailable,
|
||||
SSOIdentityAssertion,
|
||||
|
|
@ -116,6 +123,35 @@ async def test_none_mode_yields_a_no_op_auth():
|
|||
assert isinstance(result.ok, NoOpAuth)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_mode_rejects_url_userinfo():
|
||||
spec = ServerSpec(
|
||||
server_id="s",
|
||||
resource="https://lit-user:s3cr3t@upstream.example.com/mcp",
|
||||
config=NoneConfig(),
|
||||
)
|
||||
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(_SUBJECT, spec)
|
||||
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "url_credentials_not_allowed"
|
||||
assert "Basic Auth" in result.error.summary
|
||||
assert "auth_type: basic" in result.error.summary
|
||||
assert "auth_value: username:password" in result.error.summary
|
||||
assert "lit-user" not in result.error.summary
|
||||
assert "s3cr3t" not in result.error.summary
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_mode_does_not_validate_non_credential_resource():
|
||||
spec = ServerSpec(server_id="s", resource="https://[::1", config=NoneConfig())
|
||||
|
||||
result = await UpstreamCredentialProvider().resolve_credentials(_SUBJECT, spec)
|
||||
|
||||
assert isinstance(result, Ok)
|
||||
assert isinstance(result.ok, NoOpAuth)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_key_shared_emits_the_configured_header():
|
||||
config = ApiKeyConfig(
|
||||
|
|
@ -560,6 +596,73 @@ async def test_id_jag_refuses_an_expired_stored_assertion_without_calling_the_id
|
|||
assert endpoint.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_id_jag_renews_an_expired_stored_assertion_instead_of_challenging():
|
||||
"""The unattended-agent case end to end: the user last signed in more than an id_token lifetime
|
||||
ago, so without renewal this is the 412 above. With the renewing store wired the arm resolves,
|
||||
and leg 1 asserts the renewed token rather than the one that ran out."""
|
||||
renewed_id_token = pyjwt.encode(
|
||||
{"iss": "https://idp.example.com", "sub": "alice", "exp": int(time.time()) + 3600},
|
||||
"test-idp-signing-key-32-bytes-long-xxxx",
|
||||
algorithm="HS256",
|
||||
)
|
||||
expired = SSOIdentityAssertion(
|
||||
id_token=SecretStr("stale-id-token"),
|
||||
refresh_token=SecretStr("rt_1"),
|
||||
expires_at=datetime.now(timezone.utc) - timedelta(seconds=1),
|
||||
)
|
||||
rows = {"alice": expired}
|
||||
|
||||
async def _read(user_id: str) -> SSOIdentityAssertion | None:
|
||||
return rows.get(user_id)
|
||||
|
||||
async def _write(user_id: str, assertion: SSOIdentityAssertion) -> None:
|
||||
rows[user_id] = assertion
|
||||
|
||||
class _Inner:
|
||||
async def fetch(self, user_id: str) -> SSOIdentityAssertion | None:
|
||||
return await _read(user_id)
|
||||
|
||||
class _Transport:
|
||||
async def post(self, url, form, headers):
|
||||
return Ok({"access_token": "at", "id_token": renewed_id_token})
|
||||
|
||||
refresher = SSOAssertionRefresher(
|
||||
_Transport(),
|
||||
client_config=lambda: SSOClientConfig(
|
||||
token_endpoint="https://idp.example.com/token",
|
||||
client_id="litellm",
|
||||
client_secret=SecretStr("s"),
|
||||
auth_method="client_secret_basic",
|
||||
),
|
||||
read=_read,
|
||||
write=_write,
|
||||
)
|
||||
endpoint = _FakeTokenEndpoint(_two_leg_ok("final-access"))
|
||||
provider = UpstreamCredentialProvider(
|
||||
token_endpoint=endpoint,
|
||||
sso_assertion_store=RefreshingSSOAssertionStore(
|
||||
_Inner(), refresher, fresh_read=_read, coordinator_factory=lambda: None
|
||||
),
|
||||
)
|
||||
|
||||
result = await provider.resolve_credentials(
|
||||
Subject(tenant_id="", subject_id="alice"), _spec(_id_jag_config())
|
||||
)
|
||||
|
||||
assert isinstance(result, Ok)
|
||||
_, _, leg1_params = endpoint.calls[0]
|
||||
assert leg1_params["subject_token"] == renewed_id_token
|
||||
|
||||
|
||||
def test_the_resolver_defaults_to_the_renewing_assertion_store():
|
||||
"""A resolver built without collaborators is what production gets, so the default has to renew;
|
||||
the plain database reader would strand every agent an id_token lifetime after its user's login."""
|
||||
provider = UpstreamCredentialProvider()
|
||||
|
||||
assert isinstance(provider._sso_assertion_store, RefreshingSSOAssertionStore) # noqa: SLF001 # the wiring is the assertion
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_id_jag_accepts_a_stored_assertion_that_declares_no_expiry():
|
||||
endpoint = _FakeTokenEndpoint(_two_leg_ok("final-access"))
|
||||
|
|
|
|||
|
|
@ -0,0 +1,794 @@
|
|||
"""Tests for renewing the stored SSO identity assertion behind the ID-JAG arm.
|
||||
|
||||
Pins the contract an unattended agent depends on: an assertion that has run out is renewed from the
|
||||
refresh token captured beside it instead of stranding the agent until its user signs in again, the
|
||||
IdP sees one redemption per user no matter how many tool calls arrive at once, a rotation is written
|
||||
back without overwriting a sign-in that landed mid-renewal, and the two failure kinds stay
|
||||
distinguishable - a dead refresh token still challenges the user, an unreachable IdP does not.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import itertools
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import httpx
|
||||
import jwt as pyjwt
|
||||
import pytest
|
||||
from pydantic import SecretStr
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
|
||||
Error,
|
||||
Ok,
|
||||
Result,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_refresher import (
|
||||
HttpxTokenEndpointTransport,
|
||||
RefreshFailure,
|
||||
RefreshingSSOAssertionStore,
|
||||
SSOAssertionRefresher,
|
||||
SSOClientConfig,
|
||||
default_sso_assertion_store,
|
||||
sso_client_config,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
AssertionStoreUnavailable,
|
||||
SSOIdentityAssertion,
|
||||
)
|
||||
|
||||
SIGNING_KEY = "test-idp-signing-key-32-bytes-long-xxxx"
|
||||
ISSUER = "https://idp.example.com"
|
||||
TOKEN_ENDPOINT = "https://idp.example.com/token"
|
||||
|
||||
_CLIENT = SSOClientConfig(
|
||||
token_endpoint=TOKEN_ENDPOINT,
|
||||
client_id="litellm",
|
||||
client_secret=SecretStr("s3cret"),
|
||||
auth_method="client_secret_basic",
|
||||
)
|
||||
_POST_CLIENT = SSOClientConfig(
|
||||
token_endpoint=TOKEN_ENDPOINT,
|
||||
client_id="litellm",
|
||||
client_secret=SecretStr("s3cret"),
|
||||
auth_method="client_secret_post",
|
||||
)
|
||||
|
||||
|
||||
_MINTED = itertools.count()
|
||||
|
||||
|
||||
def _id_token(subject: str = "u1", exp_offset: int = 3600) -> str:
|
||||
"""A distinct token per call. Two mints with the same claims in the same second would encode
|
||||
identically, which would let a test that means "the renewed token replaced the old one" pass
|
||||
while comparing a value to itself."""
|
||||
return pyjwt.encode(
|
||||
{"iss": ISSUER, "sub": subject, "exp": int(time.time()) + exp_offset, "jti": f"t{next(_MINTED)}"},
|
||||
SIGNING_KEY,
|
||||
algorithm="HS256",
|
||||
)
|
||||
|
||||
|
||||
def _stored(id_token: str, *, expires_in: int, refresh_token: str | None = "rt_1") -> SSOIdentityAssertion:
|
||||
"""A row as the SSO callback wrote it: ``expires_in`` seconds from now, mirroring the id_token."""
|
||||
return SSOIdentityAssertion(
|
||||
id_token=SecretStr(id_token),
|
||||
refresh_token=SecretStr(refresh_token) if refresh_token else None,
|
||||
issuer=ISSUER,
|
||||
expires_at=datetime.now(timezone.utc) + timedelta(seconds=expires_in),
|
||||
)
|
||||
|
||||
|
||||
class _FakeRows:
|
||||
"""The one assertion row per user: the inner read seam and the refresher's read/write pair."""
|
||||
|
||||
def __init__(self, rows: dict[str, SSOIdentityAssertion] | None = None) -> None:
|
||||
self.rows: dict[str, SSOIdentityAssertion] = dict(rows or {})
|
||||
self.cached_rows: dict[str, SSOIdentityAssertion] = {}
|
||||
self.reads: list[str] = []
|
||||
self.writes: list[tuple[str, SSOIdentityAssertion]] = []
|
||||
|
||||
async def fetch(self, user_id: str) -> SSOIdentityAssertion | None:
|
||||
self.reads.append(user_id)
|
||||
# A real suspension point, so concurrent callers interleave here instead of running to
|
||||
# completion one at a time and never actually racing.
|
||||
await asyncio.sleep(0)
|
||||
return self.cached_rows.get(user_id, self.rows.get(user_id))
|
||||
|
||||
async def fetch_fresh(self, user_id: str) -> SSOIdentityAssertion | None:
|
||||
self.reads.append(user_id)
|
||||
await asyncio.sleep(0)
|
||||
return self.rows.get(user_id)
|
||||
|
||||
async def write(self, user_id: str, assertion: SSOIdentityAssertion) -> None:
|
||||
self.writes.append((user_id, assertion))
|
||||
self.rows[user_id] = assertion
|
||||
|
||||
|
||||
class _FakeTransport:
|
||||
"""Answers every refresh with the same canned result, optionally holding until ``gate`` opens."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
response: Result[Mapping[str, object], RefreshFailure],
|
||||
*,
|
||||
gate: asyncio.Event | None = None,
|
||||
on_call: Callable[[], None] | None = None,
|
||||
) -> None:
|
||||
self._response = response
|
||||
self._gate = gate
|
||||
self._on_call = on_call
|
||||
self.calls: list[tuple[str, dict[str, str]]] = []
|
||||
self.headers: list[dict[str, str]] = []
|
||||
|
||||
async def post(
|
||||
self, url: str, form: Mapping[str, str], headers: Mapping[str, str]
|
||||
) -> Result[Mapping[str, object], RefreshFailure]:
|
||||
self.calls.append((url, dict(form)))
|
||||
self.headers.append(dict(headers))
|
||||
if self._on_call is not None:
|
||||
self._on_call()
|
||||
if self._gate is not None:
|
||||
await self._gate.wait()
|
||||
return self._response
|
||||
|
||||
|
||||
def _renewal(id_token: str, refresh_token: str | None = None) -> Result[Mapping[str, object], RefreshFailure]:
|
||||
body: dict[str, object] = {"access_token": "at", "id_token": id_token, "token_type": "Bearer"}
|
||||
return Ok({**body, "refresh_token": refresh_token} if refresh_token else body)
|
||||
|
||||
|
||||
def _store(
|
||||
rows: _FakeRows,
|
||||
transport: _FakeTransport,
|
||||
*,
|
||||
client_config: Callable[[], SSOClientConfig | None] = lambda: _CLIENT,
|
||||
coordinator_factory: Callable[[], object] = lambda: None,
|
||||
) -> RefreshingSSOAssertionStore:
|
||||
refresher = SSOAssertionRefresher(transport, client_config=client_config, read=rows.fetch, write=rows.write)
|
||||
return RefreshingSSOAssertionStore(
|
||||
rows,
|
||||
refresher,
|
||||
fresh_read=rows.fetch_fresh,
|
||||
coordinator_factory=coordinator_factory, # pyright: ignore[reportArgumentType] # test doubles stand in for the runtime factory
|
||||
)
|
||||
|
||||
|
||||
async def _until(predicate: Callable[[], bool]) -> None:
|
||||
for _ in range(2000):
|
||||
if predicate():
|
||||
return
|
||||
await asyncio.sleep(0)
|
||||
raise AssertionError("condition never became true")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_expiring_assertion_is_renewed_and_the_renewal_is_what_the_reader_gets():
|
||||
"""The whole point: an agent calling after its user's id_token ran out keeps working."""
|
||||
stale, fresh = _id_token(exp_offset=-1), _id_token()
|
||||
rows = _FakeRows({"alice": _stored(stale, expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(fresh))
|
||||
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == fresh
|
||||
assert len(transport.calls) == 1
|
||||
url, form = transport.calls[0]
|
||||
assert url == TOKEN_ENDPOINT
|
||||
assert form["grant_type"] == "refresh_token"
|
||||
assert form["refresh_token"] == "rt_1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_basic_auth_login_gets_a_basic_auth_refresh():
|
||||
"""The non-PKCE login always sends HTTP Basic, so the renewal must too; credentials in the body
|
||||
would 401 against an IdP application registered for Basic."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
await _store(rows, transport).fetch("alice")
|
||||
|
||||
expected = base64.b64encode(b"litellm:s3cret").decode()
|
||||
assert transport.headers[0]["Authorization"] == f"Basic {expected}"
|
||||
_url, form = transport.calls[0]
|
||||
assert "client_secret" not in form
|
||||
assert "client_id" not in form
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_body_credential_login_gets_a_body_credential_refresh():
|
||||
"""The mirror case. A PKCE deployment with GENERIC_INCLUDE_CLIENT_ID set signs in with the
|
||||
credentials in the body, so Basic here would 401 against an application registered for post; the
|
||||
renewal has to follow the login rather than a constant."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
await _store(rows, transport, client_config=lambda: _POST_CLIENT).fetch("alice")
|
||||
|
||||
assert "Authorization" not in transport.headers[0]
|
||||
_url, form = transport.calls[0]
|
||||
assert form["client_id"] == "litellm"
|
||||
assert form["client_secret"] == "s3cret"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("include_client_id", "expected"),
|
||||
[
|
||||
(None, "client_secret_basic"),
|
||||
("false", "client_secret_basic"),
|
||||
("TRUE", "client_secret_post"),
|
||||
("true", "client_secret_post"),
|
||||
],
|
||||
)
|
||||
def test_the_auth_method_follows_the_flag_the_login_reads(include_client_id, expected):
|
||||
"""``GENERIC_INCLUDE_CLIENT_ID`` is what the PKCE login branches on, parsed the same way it
|
||||
parses it, so the renewal cannot pick a method the sign-in did not use."""
|
||||
env = {
|
||||
"GENERIC_TOKEN_ENDPOINT": TOKEN_ENDPOINT,
|
||||
"GENERIC_CLIENT_ID": "litellm",
|
||||
"GENERIC_CLIENT_SECRET": "s3cret",
|
||||
**({"GENERIC_INCLUDE_CLIENT_ID": include_client_id} if include_client_id is not None else {}),
|
||||
}
|
||||
|
||||
config = sso_client_config(env)
|
||||
|
||||
assert config is not None
|
||||
assert config.auth_method == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_assertion_well_inside_its_lifetime_never_reaches_the_idp():
|
||||
"""The common path must cost exactly what it did before this store existed."""
|
||||
current = _id_token()
|
||||
rows = _FakeRows({"alice": _stored(current, expires_in=1800)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == current
|
||||
assert transport.calls == []
|
||||
assert rows.writes == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_renewal_starts_inside_the_skew_rather_than_after_expiry():
|
||||
"""A token that would die between resolution and the second exchange leg is replaced first."""
|
||||
about_to_expire, fresh = _id_token(), _id_token()
|
||||
assert about_to_expire != fresh
|
||||
rows = _FakeRows({"alice": _stored(about_to_expire, expires_in=30)})
|
||||
transport = _FakeTransport(_renewal(fresh))
|
||||
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == fresh
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_user_with_no_stored_assertion_is_still_absent():
|
||||
rows = _FakeRows()
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
assert await _store(rows, transport).fetch("nobody") is None
|
||||
assert transport.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_refused_refresh_leaves_the_expired_assertion_for_the_reader_to_reject():
|
||||
"""A dead refresh token is the user's problem, and the reader's expiry guard is what tells them;
|
||||
swapping in a renewed-looking value or hiding the row would break that challenge."""
|
||||
stale = _id_token(exp_offset=-1)
|
||||
rows = _FakeRows({"alice": _stored(stale, expires_in=-1)})
|
||||
transport = _FakeTransport(Error(RefreshFailure.of_rejected("the IdP refused the refresh with status 400")))
|
||||
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == stale
|
||||
assert rows.writes == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_unreachable_idp_is_a_store_outage_not_a_sign_in_again_challenge():
|
||||
"""503, not 412: the user has nothing to fix by signing in again while the IdP is down."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(Error(RefreshFailure.of_unavailable("the IdP token endpoint is unreachable")))
|
||||
|
||||
with pytest.raises(AssertionStoreUnavailable):
|
||||
await _store(rows, transport).fetch("alice")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_missing_refresh_token_names_the_scope_the_operator_has_to_set(caplog):
|
||||
"""Nothing to redeem is the default state of a deployment, so the log has to say what to change
|
||||
or the feature stays silently inert."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1, refresh_token=None)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert transport.calls == []
|
||||
assert "GENERIC_SCOPE" in caplog.text
|
||||
assert "offline_access" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_unconfigured_sso_client_never_calls_the_idp(caplog):
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
served = await _store(rows, transport, client_config=lambda: None).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert transport.calls == []
|
||||
assert "GENERIC_TOKEN_ENDPOINT" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_refresh_response_carrying_no_id_token_is_refused(caplog):
|
||||
"""An access token is not an identity assertion, so there is nothing to assert upstream."""
|
||||
stale = _id_token(exp_offset=-1)
|
||||
rows = _FakeRows({"alice": _stored(stale, expires_in=-1)})
|
||||
transport = _FakeTransport(Ok({"access_token": "at", "token_type": "Bearer"}))
|
||||
|
||||
with caplog.at_level(logging.WARNING):
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == stale
|
||||
assert rows.writes == []
|
||||
assert "openid" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_rotated_refresh_token_replaces_the_stored_one():
|
||||
"""An IdP that rotates invalidates the old token, so keeping it would cost a sign-in next time."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token(), refresh_token="rt_2"))
|
||||
|
||||
await _store(rows, transport).fetch("alice")
|
||||
|
||||
stored = rows.rows["alice"]
|
||||
assert stored.refresh_token is not None
|
||||
assert stored.refresh_token.get_secret_value() == "rt_2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_omitted_refresh_token_carries_the_previous_one_forward():
|
||||
"""An IdP that does not rotate expects the original to keep working; dropping it would strand
|
||||
the user after exactly one renewal."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
|
||||
await _store(rows, transport).fetch("alice")
|
||||
|
||||
stored = rows.rows["alice"]
|
||||
assert stored.refresh_token is not None
|
||||
assert stored.refresh_token.get_secret_value() == "rt_1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_renewed_expiry_moves_forward_so_the_next_read_does_not_refresh_again():
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token(exp_offset=3600)))
|
||||
store = _store(rows, transport)
|
||||
|
||||
await store.fetch("alice")
|
||||
await store.fetch("alice")
|
||||
|
||||
assert len(transport.calls) == 1
|
||||
|
||||
|
||||
async def _explode(user_id: str, assertion: SSOIdentityAssertion) -> None:
|
||||
raise RuntimeError("write failed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_renewal_that_cannot_be_recorded_is_reported_as_transient():
|
||||
"""The store is what every caller reads, so a renewal nobody can see is not a success. Calling it
|
||||
one would hand back a token the gateway failed to record."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
refresher = SSOAssertionRefresher(
|
||||
_FakeTransport(_renewal(_id_token())), client_config=lambda: _CLIENT, read=rows.fetch, write=_explode
|
||||
)
|
||||
|
||||
outcome = await refresher.refresh("alice", rows.rows["alice"])
|
||||
|
||||
assert isinstance(outcome, Error)
|
||||
assert outcome.error.kind == "unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failed_write_does_not_tell_the_user_to_sign_in_again():
|
||||
"""A database that cannot take the write is not something signing in again fixes, so the reader
|
||||
has to see an outage rather than the stale row's expiry."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
refresher = SSOAssertionRefresher(transport, client_config=lambda: _CLIENT, read=rows.fetch, write=_explode)
|
||||
store = RefreshingSSOAssertionStore(
|
||||
rows,
|
||||
refresher,
|
||||
fresh_read=rows.fetch_fresh,
|
||||
coordinator_factory=lambda: None, # pyright: ignore[reportArgumentType] # test double stands in for the runtime factory
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionStoreUnavailable):
|
||||
await store.fetch("alice")
|
||||
|
||||
assert len(transport.calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_reads_for_one_user_redeem_the_refresh_token_once():
|
||||
"""A burst of tool calls must not replay one refresh token N times: an IdP that rotates reads
|
||||
that as reuse and can revoke the whole grant chain."""
|
||||
gate = asyncio.Event()
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
fresh = _id_token()
|
||||
transport = _FakeTransport(_renewal(fresh), gate=gate)
|
||||
store = _store(rows, transport)
|
||||
|
||||
callers = [asyncio.create_task(store.fetch("alice")) for _ in range(8)]
|
||||
await _until(lambda: len(transport.calls) >= 1 and len(rows.reads) >= 8)
|
||||
# Guards against a vacuous pass: every caller must have read the expired row and entered the
|
||||
# renewal branch while the winner is still blocked, otherwise they never raced at all.
|
||||
assert len(rows.reads) >= 8
|
||||
assert not any(task.done() for task in callers)
|
||||
|
||||
gate.set()
|
||||
served = await asyncio.gather(*callers)
|
||||
|
||||
assert len(transport.calls) == 1
|
||||
assert {assertion.id_token.get_secret_value() for assertion in served if assertion is not None} == {fresh}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_reads_for_different_users_each_get_their_own_refresh():
|
||||
"""Single-flight is per user; collapsing across users would leave everyone but one stranded."""
|
||||
gate = asyncio.Event()
|
||||
rows = _FakeRows(
|
||||
{
|
||||
"alice": _stored(_id_token("alice", exp_offset=-1), expires_in=-1),
|
||||
"bob": _stored(_id_token("bob", exp_offset=-1), expires_in=-1),
|
||||
}
|
||||
)
|
||||
transport = _FakeTransport(_renewal(_id_token()), gate=gate)
|
||||
store = _store(rows, transport)
|
||||
|
||||
callers = [asyncio.create_task(store.fetch(user)) for user in ("alice", "bob")]
|
||||
await _until(lambda: len(transport.calls) >= 2)
|
||||
gate.set()
|
||||
await asyncio.gather(*callers)
|
||||
|
||||
assert len(transport.calls) == 2
|
||||
assert {form["refresh_token"] for _url, form in transport.calls} == {"rt_1"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_renewal_writes_back_when_the_row_did_not_move():
|
||||
"""The refresh-then-sign-in ordering: nothing displaced the row, so the rotation must land."""
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
fresh = _id_token()
|
||||
transport = _FakeTransport(_renewal(fresh, refresh_token="rt_2"))
|
||||
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert [user_id for user_id, _assertion in rows.writes] == ["alice"]
|
||||
assert rows.rows["alice"].id_token.get_secret_value() == fresh
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == fresh
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_sign_in_landing_mid_renewal_is_not_overwritten():
|
||||
"""The sign-in-then-refresh ordering. The login wrote a newer assertion while the IdP call was in
|
||||
flight; overwriting it would put back a refresh token the IdP has already rotated away, costing
|
||||
that user a sign-in later."""
|
||||
from_login = _stored(_id_token("alice"), expires_in=3600, refresh_token="rt_from_login")
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
|
||||
def _login_lands() -> None:
|
||||
rows.rows["alice"] = from_login
|
||||
|
||||
transport = _FakeTransport(_renewal(_id_token(), refresh_token="rt_2"), on_call=_login_lands)
|
||||
|
||||
served = await _store(rows, transport).fetch("alice")
|
||||
|
||||
assert rows.writes == []
|
||||
stored = rows.rows["alice"]
|
||||
assert stored.refresh_token is not None
|
||||
assert stored.refresh_token.get_secret_value() == "rt_from_login"
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == from_login.id_token.get_secret_value()
|
||||
|
||||
|
||||
class _RecordingCoordinator:
|
||||
"""Stands in for the cross-replica coordinator, running the winner's refresh inline."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.runs: list[tuple[str, str]] = []
|
||||
|
||||
async def run(
|
||||
self,
|
||||
user_id: str,
|
||||
server_id: str,
|
||||
refresh: Callable[[], Awaitable[None]],
|
||||
reread: Callable[[], Awaitable[None]],
|
||||
) -> None:
|
||||
self.runs.append((user_id, server_id))
|
||||
return await refresh()
|
||||
|
||||
|
||||
class _ReplaceThenRefreshCoordinator:
|
||||
"""Replaces the row before running the elected refresh."""
|
||||
|
||||
def __init__(self, replace: Callable[[], None]) -> None:
|
||||
self._replace = replace
|
||||
self.runs: list[tuple[str, str]] = []
|
||||
|
||||
async def run(
|
||||
self,
|
||||
user_id: str,
|
||||
server_id: str,
|
||||
refresh: Callable[[], Awaitable[None]],
|
||||
reread: Callable[[], Awaitable[None]],
|
||||
) -> None:
|
||||
self.runs.append((user_id, server_id))
|
||||
self._replace()
|
||||
return await refresh()
|
||||
|
||||
|
||||
class _HeldCoordinator:
|
||||
"""Emulates a cross-replica holder finishing before the loser re-reads."""
|
||||
|
||||
def __init__(self, before_reread: Callable[[], None] | None = None) -> None:
|
||||
self._before_reread = before_reread
|
||||
self.runs: list[tuple[str, str]] = []
|
||||
|
||||
async def run(
|
||||
self,
|
||||
user_id: str,
|
||||
server_id: str,
|
||||
refresh: Callable[[], Awaitable[None]],
|
||||
reread: Callable[[], Awaitable[None]],
|
||||
) -> None:
|
||||
self.runs.append((user_id, server_id))
|
||||
if self._before_reread is not None:
|
||||
self._before_reread()
|
||||
return await reread()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_elected_renewal_redeems_the_row_it_re_reads_not_the_one_it_entered_with():
|
||||
stale = _id_token(exp_offset=-1)
|
||||
fresh = _stored(_id_token(), expires_in=3600, refresh_token="rt_2")
|
||||
rows = _FakeRows({"alice": _stored(stale, expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
coordinator = _ReplaceThenRefreshCoordinator(lambda: rows.rows.__setitem__("alice", fresh))
|
||||
|
||||
served = await _store(rows, transport, coordinator_factory=lambda: coordinator).fetch("alice")
|
||||
|
||||
assert transport.calls == []
|
||||
assert served is not None
|
||||
assert served.id_token.get_secret_value() == fresh.id_token.get_secret_value()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_cross_replica_loser_whose_winner_renewed_reads_the_renewal_without_redeeming():
|
||||
fresh = _stored(_id_token(), expires_in=3600, refresh_token="rt_2")
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
coordinator = _HeldCoordinator(before_reread=lambda: rows.rows.__setitem__("alice", fresh))
|
||||
|
||||
served = await _store(rows, transport, coordinator_factory=lambda: coordinator).fetch("alice")
|
||||
|
||||
assert served is fresh
|
||||
assert transport.calls == []
|
||||
assert len(coordinator.runs) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_cross_replica_loser_rereads_past_a_stale_process_local_cache():
|
||||
stale = _stored(_id_token(exp_offset=-1), expires_in=-1)
|
||||
fresh = _stored(_id_token(), expires_in=3600, refresh_token="rt_2")
|
||||
rows = _FakeRows({"alice": fresh})
|
||||
rows.cached_rows["alice"] = stale
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
coordinator = _HeldCoordinator()
|
||||
|
||||
served = await _store(rows, transport, coordinator_factory=lambda: coordinator).fetch("alice")
|
||||
|
||||
assert served is fresh
|
||||
assert transport.calls == []
|
||||
assert len(coordinator.runs) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_cross_replica_loser_does_not_turn_a_write_failure_into_a_sign_in_challenge():
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token()))
|
||||
refresher = SSOAssertionRefresher(transport, client_config=lambda: _CLIENT, read=rows.fetch, write=_explode)
|
||||
coordinator = _HeldCoordinator()
|
||||
store = RefreshingSSOAssertionStore(
|
||||
rows,
|
||||
refresher,
|
||||
fresh_read=rows.fetch_fresh,
|
||||
coordinator_factory=lambda: coordinator, # pyright: ignore[reportArgumentType] # test double stands in for the runtime factory
|
||||
)
|
||||
|
||||
with pytest.raises(AssertionStoreUnavailable):
|
||||
await store.fetch("alice")
|
||||
|
||||
assert transport.calls == []
|
||||
assert len(coordinator.runs) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_cross_replica_loser_never_redeems_the_token_the_holder_may_have_rotated():
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(Error(RefreshFailure.of_rejected("dead")))
|
||||
coordinator = _HeldCoordinator()
|
||||
|
||||
with pytest.raises(AssertionStoreUnavailable):
|
||||
await _store(rows, transport, coordinator_factory=lambda: coordinator).fetch("alice")
|
||||
|
||||
assert transport.calls == []
|
||||
assert len(coordinator.runs) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_cross_replica_coordinator_is_used_and_built_once():
|
||||
"""Redis elects one refresher across the fleet; rebuilding its client per renewal would open a
|
||||
connection every time."""
|
||||
coordinator = _RecordingCoordinator()
|
||||
builds: list[int] = []
|
||||
|
||||
def _factory() -> object:
|
||||
builds.append(1)
|
||||
return coordinator
|
||||
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token(exp_offset=-1)))
|
||||
store = _store(rows, transport, coordinator_factory=_factory)
|
||||
|
||||
await store.fetch("alice")
|
||||
await store.fetch("alice")
|
||||
|
||||
assert len(builds) == 1
|
||||
assert coordinator.runs == [("alice", "sso_identity_assertion"), ("alice", "sso_identity_assertion")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_in_process_coordinator_is_retried_until_redis_appears():
|
||||
"""A proxy that gains Redis after boot must stop electing a winner per worker."""
|
||||
coordinator = _RecordingCoordinator()
|
||||
available: list[bool] = [False]
|
||||
|
||||
def _factory() -> object | None:
|
||||
return coordinator if available[0] else None
|
||||
|
||||
rows = _FakeRows({"alice": _stored(_id_token(exp_offset=-1), expires_in=-1)})
|
||||
transport = _FakeTransport(_renewal(_id_token(exp_offset=-1)))
|
||||
store = _store(rows, transport, coordinator_factory=_factory)
|
||||
|
||||
await store.fetch("alice")
|
||||
assert coordinator.runs == []
|
||||
|
||||
available[0] = True
|
||||
await store.fetch("alice")
|
||||
assert coordinator.runs == [("alice", "sso_identity_assertion")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env",
|
||||
[
|
||||
{},
|
||||
{"GENERIC_CLIENT_ID": "litellm", "GENERIC_CLIENT_SECRET": "s"},
|
||||
{"GENERIC_TOKEN_ENDPOINT": TOKEN_ENDPOINT, "GENERIC_CLIENT_SECRET": "s"},
|
||||
{"GENERIC_TOKEN_ENDPOINT": TOKEN_ENDPOINT, "GENERIC_CLIENT_ID": "litellm"},
|
||||
{"GENERIC_TOKEN_ENDPOINT": "", "GENERIC_CLIENT_ID": "litellm", "GENERIC_CLIENT_SECRET": "s"},
|
||||
],
|
||||
)
|
||||
def test_a_partial_sso_client_is_no_client(env):
|
||||
"""Redeeming against a half-configured client would post credentials nowhere useful; the arm
|
||||
treats it as "cannot renew" and falls back to the sign-in challenge."""
|
||||
assert sso_client_config(env) is None
|
||||
|
||||
|
||||
def test_the_configured_sso_client_is_the_one_the_login_used():
|
||||
config = sso_client_config(
|
||||
{
|
||||
"GENERIC_TOKEN_ENDPOINT": TOKEN_ENDPOINT,
|
||||
"GENERIC_CLIENT_ID": "litellm",
|
||||
"GENERIC_CLIENT_SECRET": "s3cret",
|
||||
}
|
||||
)
|
||||
|
||||
assert config is not None
|
||||
assert config.token_endpoint == TOKEN_ENDPOINT
|
||||
assert config.client_id == "litellm"
|
||||
assert config.client_secret.get_secret_value() == "s3cret"
|
||||
|
||||
|
||||
def test_the_live_store_renews_over_the_database_reader():
|
||||
"""The composition root has to produce a renewing store, or none of this runs in production."""
|
||||
assert isinstance(default_sso_assertion_store(), RefreshingSSOAssertionStore)
|
||||
|
||||
|
||||
def _responding(response: httpx.Response | None) -> HttpxTokenEndpointTransport:
|
||||
async def _post(url: str, form: Mapping[str, str], headers: Mapping[str, str]) -> httpx.Response | None:
|
||||
return response
|
||||
|
||||
return HttpxTokenEndpointTransport(_post)
|
||||
|
||||
|
||||
def _json_response(status: int, payload: dict[str, object]) -> httpx.Response:
|
||||
return httpx.Response(status, json=payload, request=httpx.Request("POST", TOKEN_ENDPOINT))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", [400, 401, 403])
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_idp_declining_the_grant_is_a_refusal_the_user_must_act_on(status):
|
||||
"""A 4xx means this refresh token is finished; calling that an outage would sit the user behind a
|
||||
503 forever instead of telling them to sign in."""
|
||||
outcome = await _responding(_json_response(status, {"error": "invalid_grant"})).post(TOKEN_ENDPOINT, {}, {})
|
||||
|
||||
assert isinstance(outcome, Error)
|
||||
assert outcome.error.kind == "rejected"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", [500, 502, 503])
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failing_idp_is_an_outage_not_a_refusal(status):
|
||||
"""The refresh token is probably fine; telling the user to sign in again would blame them for
|
||||
someone else's outage, and would burn their session for nothing."""
|
||||
outcome = await _responding(_json_response(status, {})).post(TOKEN_ENDPOINT, {}, {})
|
||||
|
||||
assert isinstance(outcome, Error)
|
||||
assert outcome.error.kind == "unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_unreachable_endpoint_is_an_outage():
|
||||
async def _post(url: str, form: Mapping[str, str], headers: Mapping[str, str]) -> httpx.Response | None:
|
||||
raise httpx.ConnectError("connection refused")
|
||||
|
||||
outcome = await HttpxTokenEndpointTransport(_post).post(TOKEN_ENDPOINT, {}, {})
|
||||
|
||||
assert isinstance(outcome, Error)
|
||||
assert outcome.error.kind == "unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_non_json_body_is_an_outage():
|
||||
response = httpx.Response(200, text="<html>maintenance</html>", request=httpx.Request("POST", TOKEN_ENDPOINT))
|
||||
|
||||
outcome = await _responding(response).post(TOKEN_ENDPOINT, {}, {})
|
||||
|
||||
assert isinstance(outcome, Error)
|
||||
assert outcome.error.kind == "unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_missing_response_is_an_outage():
|
||||
outcome = await _responding(None).post(TOKEN_ENDPOINT, {}, {})
|
||||
|
||||
assert isinstance(outcome, Error)
|
||||
assert outcome.error.kind == "unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_successful_grant_is_handed_back_as_the_parsed_body():
|
||||
outcome = await _responding(_json_response(200, {"access_token": "at", "id_token": "idt"})).post(
|
||||
TOKEN_ENDPOINT, {"grant_type": "refresh_token"}, {}
|
||||
)
|
||||
|
||||
assert isinstance(outcome, Ok)
|
||||
assert outcome.ok["id_token"] == "idt"
|
||||
|
|
@ -75,6 +75,15 @@ def test_crederror_factory_sets_the_matching_tag(factory, expected_tag):
|
|||
assert "detail text" in err.summary
|
||||
|
||||
|
||||
def test_url_credentials_error_has_a_fixed_actionable_summary():
|
||||
err = CredError.of_url_credentials_not_allowed()
|
||||
|
||||
assert err.tag == "url_credentials_not_allowed"
|
||||
assert "Basic Auth" in err.summary
|
||||
assert "auth_type: basic" in err.summary
|
||||
assert "auth_value: username:password" in err.summary
|
||||
|
||||
|
||||
def test_apikeyconfig_requires_a_key_source():
|
||||
with pytest.raises(ValidationError):
|
||||
ApiKeyConfig() # type: ignore[call-arg]
|
||||
|
|
|
|||
|
|
@ -461,6 +461,30 @@ class TestMCPServerManager:
|
|||
base.update(overrides)
|
||||
return {"m2mserver": base}
|
||||
|
||||
def _id_jag_config(self):
|
||||
return {
|
||||
"idjag_server": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"auth_type": MCPAuth.oauth2_id_jag,
|
||||
"client_id": "cid",
|
||||
"client_secret": "csec",
|
||||
"token_exchange_endpoint": "https://idp.example.com/token",
|
||||
"id_jag_resource_token_endpoint": "https://resource.example.com/token",
|
||||
"id_jag_resource": "https://resource.example.com",
|
||||
}
|
||||
}
|
||||
|
||||
def _clear_sso_env(self, monkeypatch):
|
||||
for env_var in (
|
||||
"GOOGLE_CLIENT_ID",
|
||||
"MICROSOFT_CLIENT_ID",
|
||||
"GENERIC_CLIENT_ID",
|
||||
"SAML_IDP_METADATA_URL",
|
||||
"SAML_IDP_METADATA_XML",
|
||||
):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
@pytest.mark.parametrize("value", ["1", "true", "TRUE", "yes", "on"])
|
||||
def test_mcp_oauth_discovery_on_startup_true_values(self, value):
|
||||
with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": value}):
|
||||
|
|
@ -1130,6 +1154,72 @@ class TestMCPServerManager:
|
|||
server = next(iter(manager.config_mcp_servers.values()))
|
||||
assert server.oauth2_flow is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_warns_for_id_jag_with_google_sso(self, monkeypatch, caplog):
|
||||
self._clear_sso_env(monkeypatch)
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
|
||||
manager = MCPServerManager()
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
|
||||
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
|
||||
caplog.at_level(logging.WARNING, logger="LiteLLM"),
|
||||
):
|
||||
await manager.load_servers_from_config(self._id_jag_config())
|
||||
|
||||
warnings = [message for message in caplog.messages if "oauth2_id_jag" in message]
|
||||
assert len(warnings) == 1
|
||||
assert "idjag_server" in warnings[0]
|
||||
assert "GENERIC_CLIENT_ID" in warnings[0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_does_not_warn_for_id_jag_without_sso(self, monkeypatch, caplog):
|
||||
self._clear_sso_env(monkeypatch)
|
||||
manager = MCPServerManager()
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
|
||||
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
|
||||
caplog.at_level(logging.WARNING, logger="LiteLLM"),
|
||||
):
|
||||
await manager.load_servers_from_config(self._id_jag_config())
|
||||
|
||||
assert not any("oauth2_id_jag" in message for message in caplog.messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, monkeypatch, caplog):
|
||||
self._clear_sso_env(monkeypatch)
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
|
||||
manager = MCPServerManager()
|
||||
config = {
|
||||
"api_key_server": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"auth_type": MCPAuth.api_key,
|
||||
"auth_value": "upstream-secret",
|
||||
}
|
||||
}
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
|
||||
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
|
||||
caplog.at_level(logging.WARNING, logger="LiteLLM"),
|
||||
):
|
||||
await manager.load_servers_from_config(config)
|
||||
|
||||
assert not any("oauth2_id_jag" in message for message in caplog.messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_does_not_warn_for_id_jag_with_generic_sso(self, monkeypatch, caplog):
|
||||
self._clear_sso_env(monkeypatch)
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
manager = MCPServerManager()
|
||||
with (
|
||||
patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)),
|
||||
patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()),
|
||||
caplog.at_level(logging.WARNING, logger="LiteLLM"),
|
||||
):
|
||||
await manager.load_servers_from_config(self._id_jag_config())
|
||||
|
||||
assert not any("oauth2_id_jag" in message for message in caplog.messages)
|
||||
|
||||
def _client_forwarded_config(self, auth_type, **overrides):
|
||||
base = {
|
||||
"url": "https://example.com/mcp",
|
||||
|
|
@ -8966,6 +9056,24 @@ class TestCreateMcpClientV2Graft:
|
|||
assert isinstance(client._resolved_auth, NoOpAuth)
|
||||
assert client._mcp_auth_value is None
|
||||
|
||||
@pytest.mark.parametrize("auth_type", [None, MCPAuth.none])
|
||||
async def test_none_mode_rejects_url_userinfo(self, auth_type):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await MCPServerManager()._create_mcp_client(
|
||||
self._http_server(
|
||||
auth_type=auth_type,
|
||||
url="https://lit-user:s3cr3t@upstream.example.com/mcp",
|
||||
)
|
||||
)
|
||||
|
||||
detail = str(exc_info.value.detail)
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "Basic Auth" in detail
|
||||
assert "auth_type: basic" in detail
|
||||
assert "auth_value: username:password" in detail
|
||||
assert "lit-user" not in detail
|
||||
assert "s3cr3t" not in detail
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, token, expected_name, expected_value",
|
||||
[
|
||||
|
|
@ -11369,6 +11477,30 @@ class TestResolveOpenapiToolAuth:
|
|||
|
||||
assert "Authorization" not in (forwarded or {})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_mode_without_url_keeps_spec_path_server_unauthenticated(self):
|
||||
server = MCPServer(
|
||||
server_id="openapi-only",
|
||||
name="report_api",
|
||||
server_name="report_api",
|
||||
url=None,
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
spec_path="https://api.example.com/openapi.json",
|
||||
)
|
||||
|
||||
resolved, forwarded = await MCPServerManager().resolve_openapi_upstream_auth(
|
||||
mcp_server=server,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
mcp_auth_header=None,
|
||||
user_api_key_auth=None,
|
||||
forwarded_headers={"X-Trace": "trace-id"},
|
||||
)
|
||||
|
||||
assert resolved is None
|
||||
assert forwarded == {"X-Trace": "trace-id"}
|
||||
|
||||
|
||||
class TestOpenApiHandlerRelaysUpstreamAuth:
|
||||
"""`_call_openapi_tool_handler` must not flatten a re-auth signal into a generic message.
|
||||
|
|
|
|||
|
|
@ -177,6 +177,36 @@ class TestExecuteWithMcpClient:
|
|||
assert "https://api.example.com/mcp/" in message
|
||||
assert "30s" in message
|
||||
|
||||
def test_connection_error_message_hides_arbitrary_http_exception_detail(self):
|
||||
message = rest_endpoints._connection_error_message(
|
||||
HTTPException(status_code=500, detail="secret upstream detail"),
|
||||
"https://api.example.com/mcp/",
|
||||
30.0,
|
||||
)
|
||||
|
||||
assert "secret upstream detail" not in message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_mode_url_credentials_returns_actionable_redacted_error(self):
|
||||
async def unreached_operation(client):
|
||||
raise AssertionError("operation must not run for an invalid server configuration")
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://lit-user:s3cr3t@upstream.example.com/mcp",
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(payload, unreached_operation)
|
||||
|
||||
message = str(result["message"])
|
||||
assert result["error"] is True
|
||||
assert "Basic Auth" in message
|
||||
assert "auth_type: basic" in message
|
||||
assert "auth_value: username:password" in message
|
||||
assert "lit-user" not in message
|
||||
assert "s3cr3t" not in message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_static_headers(self, monkeypatch):
|
||||
"""Ensure static_headers are forwarded to the MCP client during test calls.
|
||||
|
|
|
|||
|
|
@ -34,27 +34,27 @@ def test_is_over_limit():
|
|||
assert license_check.is_over_limit(99) is False
|
||||
|
||||
|
||||
def test_heuristic_v2_router_limit() -> None:
|
||||
def test_auto_router_capability_limit() -> None:
|
||||
"""Only the signed license's auto_router feature lifts the one-router limit; an API-verified
|
||||
license (no airgapped data) and an airgapped license without the feature keep it."""
|
||||
license_check = LicenseCheck()
|
||||
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["auto_router"]}
|
||||
assert license_check.heuristic_v2_router_limit() is None
|
||||
assert license_check.auto_router_capability_limit() is None
|
||||
|
||||
license_check.airgapped_license_data = {
|
||||
"expiration_date": "2999-01-01",
|
||||
"allowed_features": ["sso", "auto_router", "audit_logs"],
|
||||
}
|
||||
assert license_check.heuristic_v2_router_limit() is None
|
||||
assert license_check.auto_router_capability_limit() is None
|
||||
|
||||
license_check.airgapped_license_data = {"expiration_date": "2999-01-01", "allowed_features": ["sso"]}
|
||||
assert license_check.heuristic_v2_router_limit() == 1
|
||||
assert license_check.auto_router_capability_limit() == 1
|
||||
|
||||
license_check.airgapped_license_data = {"expiration_date": "2999-01-01"}
|
||||
assert license_check.heuristic_v2_router_limit() == 1
|
||||
assert license_check.auto_router_capability_limit() == 1
|
||||
|
||||
license_check.airgapped_license_data = None
|
||||
assert license_check.heuristic_v2_router_limit() == 1
|
||||
assert license_check.auto_router_capability_limit() == 1
|
||||
|
||||
|
||||
def _signed_license(expiration_date: str) -> tuple[RSAPublicKey, str]:
|
||||
|
|
@ -81,12 +81,12 @@ def test_expired_or_unreadable_license_grants_no_features() -> None:
|
|||
license_check = LicenseCheck()
|
||||
public_key, valid_key = _signed_license("2999-01-01")
|
||||
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=valid_key) is True
|
||||
assert license_check.heuristic_v2_router_limit() is None
|
||||
assert license_check.auto_router_capability_limit() is None
|
||||
|
||||
_, expired_key = _signed_license("2000-01-01")
|
||||
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=expired_key) is not True
|
||||
assert license_check.airgapped_license_data is None
|
||||
assert license_check.heuristic_v2_router_limit() == 1
|
||||
assert license_check.auto_router_capability_limit() == 1
|
||||
|
||||
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=valid_key) is True
|
||||
assert license_check.verify_license_without_api_request(public_key=public_key, license_key="not-a-license") is not True
|
||||
|
|
@ -98,4 +98,4 @@ def test_valid_signed_license_with_auto_router_lifts_the_limit() -> None:
|
|||
public_key, license_key = _signed_license("2999-01-01")
|
||||
|
||||
assert license_check.verify_license_without_api_request(public_key=public_key, license_key=license_key) is True
|
||||
assert license_check.heuristic_v2_router_limit() is None
|
||||
assert license_check.auto_router_capability_limit() is None
|
||||
|
|
|
|||
|
|
@ -1137,7 +1137,11 @@ async def test_bedrock_apply_guardrail_response_uses_OUTPUT_source():
|
|||
mock_api.assert_called_once()
|
||||
kwargs = mock_api.call_args.kwargs
|
||||
assert kwargs["source"] == "OUTPUT"
|
||||
assert kwargs["request_data"] == {"model": "gpt-4o"}
|
||||
assert kwargs["request_data"]["model"] == "gpt-4o"
|
||||
recorded = kwargs["request_data"]["metadata"]["standard_logging_guardrail_information"]
|
||||
assert [(e["guardrail_name"], e["guardrail_status"]) for e in recorded] == [
|
||||
(guardrail.guardrail_name, "success")
|
||||
]
|
||||
synthetic = kwargs["response"]
|
||||
assert isinstance(synthetic, ModelResponse)
|
||||
assert len(synthetic.choices) == 2
|
||||
|
|
|
|||
|
|
@ -0,0 +1,381 @@
|
|||
"""Unit tests for litellm.proxy.guardrails.auto_router_compression."""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy.guardrails import auto_router_compression
|
||||
from litellm.proxy.guardrails.auto_router_compression import (
|
||||
AutoRouterCompressionPolicy,
|
||||
arm_pre_call,
|
||||
messages_for_routing,
|
||||
policy_for_model,
|
||||
policy_from_litellm_params,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
class TestPolicyFromLitellmParams:
|
||||
def test_neither_key_set_is_no_policy(self):
|
||||
assert policy_from_litellm_params({}) is None
|
||||
|
||||
def test_routing_only(self):
|
||||
policy = policy_from_litellm_params({"auto_router_routing_compression": "headroom-a"})
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-a", model=None)
|
||||
|
||||
def test_none_sentinel_normalizes_to_no_compression(self):
|
||||
policy = policy_from_litellm_params(
|
||||
{"auto_router_routing_compression": "headroom-a", "auto_router_model_compression": "none"}
|
||||
)
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-a", model=None)
|
||||
|
||||
def test_none_sentinel_is_case_insensitive(self):
|
||||
policy = policy_from_litellm_params({"auto_router_routing_compression": "NONE"})
|
||||
assert policy == AutoRouterCompressionPolicy(routing=None, model=None)
|
||||
|
||||
def test_is_same_true_for_matching_names(self):
|
||||
policy = policy_from_litellm_params(
|
||||
{"auto_router_routing_compression": "x", "auto_router_model_compression": "x"}
|
||||
)
|
||||
assert policy.is_same is True
|
||||
|
||||
def test_is_same_false_for_different_names(self):
|
||||
policy = policy_from_litellm_params(
|
||||
{"auto_router_routing_compression": "x", "auto_router_model_compression": "y"}
|
||||
)
|
||||
assert policy.is_same is False
|
||||
|
||||
def test_is_same_true_when_both_no_compression(self):
|
||||
policy = policy_from_litellm_params(
|
||||
{"auto_router_routing_compression": "none", "auto_router_model_compression": "none"}
|
||||
)
|
||||
assert policy.is_same is True
|
||||
|
||||
|
||||
class _FakeRouter:
|
||||
"""Minimal stand-in for litellm.Router.get_model_list, for policy_for_model."""
|
||||
|
||||
def __init__(self, deployments: list[dict[str, Any]]):
|
||||
self._deployments = deployments
|
||||
|
||||
def get_model_list(self, model_name, team_id=None):
|
||||
return [d for d in self._deployments if d.get("model_name") == model_name]
|
||||
|
||||
|
||||
def _marker(compression: dict[str, str], tags: list[str] | None = None) -> dict[str, Any]:
|
||||
return {
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
**compression,
|
||||
**({"tags": tags} if tags is not None else {}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class TestPolicyForModel:
|
||||
def test_no_router_returns_none(self):
|
||||
assert policy_for_model(llm_router=None, model_alias="smart-router", team_id=None, request_tags=()) is None
|
||||
|
||||
def test_no_marker_deployment_returns_none(self):
|
||||
router = _FakeRouter([{"model_name": "smart-router", "litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
assert policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=()) is None
|
||||
|
||||
def test_marker_deployment_without_policy_returns_none(self):
|
||||
router = _FakeRouter(
|
||||
[{"model_name": "smart-router", "litellm_params": {"model": "auto_router/complexity_router"}}]
|
||||
)
|
||||
assert policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=()) is None
|
||||
|
||||
def test_marker_deployment_with_policy_is_found(self):
|
||||
router = _FakeRouter(
|
||||
[_marker({"auto_router_routing_compression": "headroom-a", "auto_router_model_compression": "none"})]
|
||||
)
|
||||
policy = policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=())
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-a", model=None)
|
||||
|
||||
def test_picks_the_marker_whose_tags_the_request_carries(self):
|
||||
"""Regression: an alias with several tag-scoped markers must not suppress one
|
||||
marker's guardrail and then route under a different marker's policy."""
|
||||
router = _FakeRouter(
|
||||
[
|
||||
_marker({"auto_router_routing_compression": "headroom-eu"}, tags=["eu"]),
|
||||
_marker({"auto_router_routing_compression": "headroom-us"}, tags=["us"]),
|
||||
]
|
||||
)
|
||||
|
||||
eu = policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=("eu",))
|
||||
us = policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=("us",))
|
||||
|
||||
assert eu == AutoRouterCompressionPolicy(routing="headroom-eu", model=None)
|
||||
assert us == AutoRouterCompressionPolicy(routing="headroom-us", model=None)
|
||||
|
||||
def test_untagged_marker_matches_any_request(self):
|
||||
router = _FakeRouter([_marker({"auto_router_routing_compression": "headroom-a"})])
|
||||
policy = policy_for_model(
|
||||
llm_router=router, model_alias="smart-router", team_id=None, request_tags=("anything",)
|
||||
)
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-a", model=None)
|
||||
|
||||
def test_a_marker_scoped_to_other_tags_is_never_the_fallback(self):
|
||||
"""Regression: a "us" request must not fall back to an "eu" marker's policy."""
|
||||
router = _FakeRouter(
|
||||
[
|
||||
_marker({"auto_router_routing_compression": "headroom-eu"}, tags=["eu"]),
|
||||
_marker({"auto_router_routing_compression": "headroom-default"}),
|
||||
]
|
||||
)
|
||||
policy = policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=("us",))
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-default", model=None)
|
||||
|
||||
def test_no_untagged_fallback_means_no_policy(self):
|
||||
"""No matching marker means no policy, not an unrelated slice's compression."""
|
||||
router = _FakeRouter([_marker({"auto_router_routing_compression": "headroom-eu"}, tags=["eu"])])
|
||||
assert (
|
||||
policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=("us",)) is None
|
||||
)
|
||||
|
||||
def test_tag_scoped_marker_takes_precedence_over_untagged(self):
|
||||
"""Regression: when multiple markers exist, the tag-scoped one the request
|
||||
actually matches should be used, not the first untagged one."""
|
||||
router = _FakeRouter(
|
||||
[
|
||||
_marker({"auto_router_routing_compression": "headroom-untagged"}),
|
||||
_marker({"auto_router_routing_compression": "headroom-eu"}, tags=["eu"]),
|
||||
]
|
||||
)
|
||||
policy = policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=("eu",))
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-eu", model=None)
|
||||
|
||||
|
||||
class _RecordingCompressionGuardrail(CustomGuardrail):
|
||||
"""A guardrail whose apply_guardrail marks every text message as compressed."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(guardrail_name=guardrail_name)
|
||||
self.request_data_seen: list[dict] = []
|
||||
|
||||
async def apply_guardrail(
|
||||
self, inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: str, logging_obj=None
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.request_data_seen.append(request_data)
|
||||
structured_messages = inputs.get("structured_messages") or []
|
||||
compressed = [{**m, "content": f"[COMPRESSED] {m.get('content')}"} for m in structured_messages]
|
||||
return {**inputs, "structured_messages": compressed}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registered_guardrail(monkeypatch):
|
||||
import litellm
|
||||
from litellm.proxy.guardrails import guardrail_registry
|
||||
|
||||
# Registered under a compression provider name: both hops refuse a name that does
|
||||
# not resolve to one, so a bare callback would (correctly) never be used.
|
||||
monkeypatch.setitem(guardrail_registry.guardrail_class_registry, "headroom", _RecordingCompressionGuardrail)
|
||||
guardrail = _RecordingCompressionGuardrail(guardrail_name="fake-compress")
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
yield guardrail
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
|
||||
|
||||
|
||||
class _NonCompressionGuardrail(CustomGuardrail):
|
||||
"""A guardrail that is not a compression provider, e.g. a PII or content filter."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(guardrail_name=guardrail_name)
|
||||
self.called = False
|
||||
|
||||
async def apply_guardrail(
|
||||
self, inputs: GenericGuardrailAPIInputs, request_data: dict, input_type: str, logging_obj=None
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.called = True
|
||||
return inputs
|
||||
|
||||
|
||||
class TestArmPreCall:
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_router_is_noop(self):
|
||||
data = {"model": "smart-router", "messages": [{"role": "user", "content": "hi"}]}
|
||||
await arm_pre_call(data=data, llm_router=None)
|
||||
assert "metadata" not in data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_policy_does_not_create_metadata_bucket(self):
|
||||
router = _FakeRouter([{"model_name": "smart-router", "litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
data = {"model": "smart-router", "messages": [{"role": "user", "content": "hi"}]}
|
||||
await arm_pre_call(data=data, llm_router=router)
|
||||
assert "metadata" not in data
|
||||
assert "litellm_metadata" not in data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_policy_suppresses_active_compression_guardrails(self, monkeypatch):
|
||||
from litellm.proxy.guardrails import guardrail_registry
|
||||
|
||||
monkeypatch.setitem(
|
||||
guardrail_registry.guardrail_class_registry, "fake-provider", _RecordingCompressionGuardrail
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.guardrails.auto_router_compression.COMPRESSION_GUARDRAIL_PROVIDERS",
|
||||
frozenset({"fake-provider"}),
|
||||
)
|
||||
import litellm
|
||||
|
||||
always_on = _RecordingCompressionGuardrail(guardrail_name="always-on-compression")
|
||||
litellm.logging_callback_manager.add_litellm_callback(always_on)
|
||||
try:
|
||||
router = _FakeRouter(
|
||||
[
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"auto_router_routing_compression": "headroom-a",
|
||||
"auto_router_model_compression": "none",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
data = {"model": "smart-router", "messages": [{"role": "user", "content": "hi"}]}
|
||||
await arm_pre_call(data=data, llm_router=router)
|
||||
assert auto_router_compression.suppressed_compression_guardrails() == frozenset({"always-on-compression"})
|
||||
# Suppression state must never ride along in metadata: that reaches spend
|
||||
# logs the caller can read, and anything there is replayable.
|
||||
assert "always-on-compression" not in json.dumps(data.get("metadata", {}))
|
||||
assert always_on.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is False
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(always_on)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_suppression_state_never_enters_request_metadata(self):
|
||||
"""Regression (security): metadata reaches spend logs, so a suppression list
|
||||
there is one a caller could read back and replay to disable a guardrail."""
|
||||
guardrail = _RecordingCompressionGuardrail(guardrail_name="always-on-compression")
|
||||
import litellm
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(guardrail)
|
||||
try:
|
||||
router = _FakeRouter([_marker({"auto_router_routing_compression": "headroom-a"})])
|
||||
data = {"model": "smart-router", "messages": [{"role": "user", "content": "hi"}]}
|
||||
await arm_pre_call(data=data, llm_router=router)
|
||||
assert "suppress" not in json.dumps(data).lower()
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(guardrail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_side_guardrail_is_requested_even_when_not_default_on(self, monkeypatch):
|
||||
import litellm
|
||||
from litellm.proxy.guardrails import guardrail_registry
|
||||
|
||||
monkeypatch.setitem(guardrail_registry.guardrail_class_registry, "headroom", _RecordingCompressionGuardrail)
|
||||
active = _RecordingCompressionGuardrail(guardrail_name="headroom-b")
|
||||
litellm.logging_callback_manager.add_litellm_callback(active)
|
||||
router = _FakeRouter(
|
||||
[
|
||||
{
|
||||
"model_name": "smart-router",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"auto_router_routing_compression": "none",
|
||||
"auto_router_model_compression": "headroom-b",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
data = {"model": "smart-router", "messages": [{"role": "user", "content": "hi"}]}
|
||||
try:
|
||||
await arm_pre_call(data=data, llm_router=router)
|
||||
assert data["metadata"]["guardrails"] == ["headroom-b"]
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(active)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arm_pre_call_keeps_no_copy_of_the_prompt(self):
|
||||
"""Regression (security): arm_pre_call runs before the guardrails, so any copy it
|
||||
kept would be pre-masking text that routing then POSTs to an external service."""
|
||||
router = _FakeRouter([_marker({"auto_router_routing_compression": "headroom-a"})])
|
||||
data = {"model": "smart-router", "messages": [{"role": "user", "content": "my ssn is 123-45-6789"}]}
|
||||
|
||||
await arm_pre_call(data=data, llm_router=router)
|
||||
|
||||
assert "123-45-6789" not in json.dumps(data.get("metadata", {}))
|
||||
assert not hasattr(auto_router_compression, "_routing_messages_snapshot")
|
||||
|
||||
|
||||
class TestMessagesForRouting:
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_policy_returns_none(self):
|
||||
assert await messages_for_routing(policy=None, messages=[], request_kwargs={}) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_none_with_no_model_compression_returns_none(self):
|
||||
"""Nothing compressed either hop, so the caller's own messages are already right."""
|
||||
policy = AutoRouterCompressionPolicy(routing=None, model=None)
|
||||
assert await messages_for_routing(policy=policy, messages=[], request_kwargs={}) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_none_never_reaches_for_a_pre_guardrail_copy(self):
|
||||
"""No uncompressed copy survives the model hop, and keeping one would mean
|
||||
retaining the pre-masking text. Routing reads what it has."""
|
||||
policy = AutoRouterCompressionPolicy(routing=None, model="headroom-a")
|
||||
model_compressed = [{"role": "user", "content": "[COMPRESSED] the full original conversation"}]
|
||||
|
||||
assert await messages_for_routing(policy=policy, messages=model_compressed, request_kwargs={}) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_guardrail_name_routes_on_the_uncompressed_messages(self):
|
||||
policy = AutoRouterCompressionPolicy(routing="does-not-exist", model=None)
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
result = await messages_for_routing(policy=policy, messages=messages, request_kwargs={})
|
||||
assert result == messages
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_compresses_via_the_named_guardrail(self, registered_guardrail):
|
||||
policy = AutoRouterCompressionPolicy(routing="fake-compress", model=None)
|
||||
messages = [{"role": "user", "content": "hello world"}]
|
||||
result = await messages_for_routing(policy=policy, messages=messages, request_kwargs={})
|
||||
assert result == [{"role": "user", "content": "[COMPRESSED] hello world"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_compresses_what_the_other_guardrails_left_behind(self, registered_guardrail):
|
||||
"""Regression (security): routing POSTs its input out, so it must read what the
|
||||
earlier guardrails left behind, not a pre-masking copy."""
|
||||
policy = AutoRouterCompressionPolicy(routing="fake-compress", model="headroom-b")
|
||||
masked = [{"role": "user", "content": "my ssn is [REDACTED]"}]
|
||||
|
||||
result = await messages_for_routing(policy=policy, messages=masked, request_kwargs={})
|
||||
|
||||
assert result == [{"role": "user", "content": "[COMPRESSED] my ssn is [REDACTED]"}]
|
||||
assert registered_guardrail.request_data_seen[0]["messages"] == masked
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_non_compression_guardrail_is_never_invoked_for_routing(self, monkeypatch):
|
||||
"""Regression (security): naming an ordinary guardrail must not turn the routing
|
||||
hop into a way to ship prompts to whatever service backs it."""
|
||||
import litellm
|
||||
|
||||
other = _NonCompressionGuardrail(guardrail_name="pii-filter")
|
||||
litellm.logging_callback_manager.add_litellm_callback(other)
|
||||
try:
|
||||
policy = AutoRouterCompressionPolicy(routing="pii-filter", model=None)
|
||||
messages = [{"role": "user", "content": "my ssn is 123-45-6789"}]
|
||||
|
||||
result = await messages_for_routing(policy=policy, messages=messages, request_kwargs={})
|
||||
|
||||
assert other.called is False
|
||||
assert result == messages
|
||||
finally:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(other)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_receives_a_throwaway_request_data_not_the_real_request_kwargs(self, registered_guardrail):
|
||||
"""Regression: a guardrail writes stats onto the request_data it is given, so
|
||||
passing the caller's own would double-count into extract_compression_saved_tokens."""
|
||||
policy = AutoRouterCompressionPolicy(routing="fake-compress", model=None)
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
request_kwargs = {"metadata": {}}
|
||||
await messages_for_routing(policy=policy, messages=messages, request_kwargs=request_kwargs)
|
||||
assert registered_guardrail.request_data_seen[0] is not request_kwargs
|
||||
assert request_kwargs == {"metadata": {}}
|
||||
|
|
@ -0,0 +1,117 @@
|
|||
import pytest
|
||||
|
||||
from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import (
|
||||
ActiveSSOProvider,
|
||||
active_sso_provider,
|
||||
id_jag_assertion_capture_gap,
|
||||
id_jag_assertion_capture_gap_at_startup,
|
||||
)
|
||||
|
||||
_SSO_ENV_VARS = (
|
||||
"GOOGLE_CLIENT_ID",
|
||||
"MICROSOFT_CLIENT_ID",
|
||||
"GENERIC_CLIENT_ID",
|
||||
"SAML_IDP_METADATA_URL",
|
||||
"SAML_IDP_METADATA_XML",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolated_sso_env(monkeypatch):
|
||||
"""Every SSO selector is read from the process environment, so a value left behind by
|
||||
another test would silently decide this one's answer."""
|
||||
for name in _SSO_ENV_VARS:
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
|
||||
|
||||
class TestActiveSSOProviderMirrorsTheCallback:
|
||||
"""The gap warning is only as good as its agreement with the branch the login callback
|
||||
actually takes, so provider selection is asserted branch by branch, including the
|
||||
precedence that makes a co-configured generic client unreachable."""
|
||||
|
||||
def test_google_client_id_selects_google(self, monkeypatch):
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
|
||||
assert active_sso_provider() is ActiveSSOProvider.google
|
||||
|
||||
def test_microsoft_client_id_selects_microsoft(self, monkeypatch):
|
||||
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-cid")
|
||||
assert active_sso_provider() is ActiveSSOProvider.microsoft
|
||||
|
||||
def test_generic_client_id_selects_generic(self, monkeypatch):
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
assert active_sso_provider() is ActiveSSOProvider.generic
|
||||
|
||||
def test_saml_metadata_selects_saml(self, monkeypatch):
|
||||
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata")
|
||||
assert active_sso_provider() is ActiveSSOProvider.saml
|
||||
|
||||
def test_nothing_configured_selects_none(self):
|
||||
assert active_sso_provider() is ActiveSSOProvider.none
|
||||
|
||||
def test_google_outranks_a_co_configured_generic_client(self, monkeypatch):
|
||||
"""The callback tests GOOGLE_CLIENT_ID first, so the generic arm never runs here and
|
||||
no assertion is captured; reporting generic would clear a gap that is still open."""
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
assert active_sso_provider() is ActiveSSOProvider.google
|
||||
|
||||
def test_microsoft_outranks_a_co_configured_generic_client(self, monkeypatch):
|
||||
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "ms-cid")
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
assert active_sso_provider() is ActiveSSOProvider.microsoft
|
||||
|
||||
def test_generic_outranks_saml(self, monkeypatch):
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata")
|
||||
assert active_sso_provider() is ActiveSSOProvider.generic
|
||||
|
||||
|
||||
class TestIdJagAssertionCaptureGap:
|
||||
def test_generic_oidc_has_no_gap(self, monkeypatch):
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
assert id_jag_assertion_capture_gap() is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env_var, provider_label",
|
||||
[
|
||||
("GOOGLE_CLIENT_ID", "google"),
|
||||
("MICROSOFT_CLIENT_ID", "microsoft"),
|
||||
("SAML_IDP_METADATA_URL", "saml"),
|
||||
],
|
||||
)
|
||||
def test_non_capturing_provider_is_named_with_the_remedy(self, monkeypatch, env_var, provider_label):
|
||||
monkeypatch.setenv(env_var, "configured")
|
||||
gap = id_jag_assertion_capture_gap()
|
||||
assert gap is not None
|
||||
assert provider_label in gap
|
||||
assert "GENERIC_CLIENT_ID" in gap
|
||||
|
||||
def test_no_sso_configured_reports_a_gap(self):
|
||||
gap = id_jag_assertion_capture_gap()
|
||||
assert gap is not None
|
||||
assert "no SSO provider is configured" in gap
|
||||
|
||||
def test_google_beside_generic_still_reports_a_gap(self, monkeypatch):
|
||||
"""The precedence trap in operator terms: adding a generic client id without removing
|
||||
GOOGLE_CLIENT_ID does not fix the deployment, so the gap must not clear."""
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
gap = id_jag_assertion_capture_gap()
|
||||
assert gap is not None
|
||||
assert "google" in gap
|
||||
|
||||
|
||||
class TestIdJagAssertionCaptureGapAtStartup:
|
||||
def test_no_provider_at_startup_is_not_yet_a_gap(self):
|
||||
assert id_jag_assertion_capture_gap_at_startup() is None
|
||||
|
||||
def test_google_provider_at_startup_reports_the_capture_gap(self, monkeypatch):
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid")
|
||||
startup_gap = id_jag_assertion_capture_gap_at_startup()
|
||||
callback_gap = id_jag_assertion_capture_gap()
|
||||
assert startup_gap is not None
|
||||
assert startup_gap == callback_gap
|
||||
|
||||
def test_generic_provider_at_startup_has_no_gap(self, monkeypatch):
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "generic-cid")
|
||||
assert id_jag_assertion_capture_gap_at_startup() is None
|
||||
|
|
@ -2,6 +2,7 @@ import os
|
|||
import sys
|
||||
import types
|
||||
import json
|
||||
import logging
|
||||
from contextlib import ExitStack
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
|
|
@ -3840,6 +3841,155 @@ class TestAddMCPServerAtomicity:
|
|||
mock_manager.reload_servers_from_database.assert_not_awaited()
|
||||
|
||||
|
||||
class TestIdJagRegistrationWarnsAboutTheSSOGap:
|
||||
"""An `oauth2_id_jag` server only ever works when the login path captures an IdP identity
|
||||
assertion, and only the generic OIDC arm does. Registering one under Google or Microsoft
|
||||
succeeds and then fails for every user on every call, so the mismatch has to be said at
|
||||
registration time, while the admin is still looking at the configuration."""
|
||||
|
||||
@staticmethod
|
||||
def _clear_sso_env(monkeypatch):
|
||||
for name in (
|
||||
"GOOGLE_CLIENT_ID",
|
||||
"MICROSOFT_CLIENT_ID",
|
||||
"GENERIC_CLIENT_ID",
|
||||
"SAML_IDP_METADATA_URL",
|
||||
"SAML_IDP_METADATA_XML",
|
||||
):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
|
||||
@staticmethod
|
||||
def _id_jag_warnings(caplog) -> list[str]:
|
||||
return [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if record.levelno == logging.WARNING and "oauth2_id_jag" in record.getMessage()
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _server_record(auth_type) -> LiteLLM_MCPServerTable:
|
||||
record = generate_mock_mcp_server_db_record(server_id="ema-1", alias="ema")
|
||||
record.auth_type = auth_type
|
||||
return record
|
||||
|
||||
async def _run_create(self, monkeypatch, provider_env, auth_type, caplog):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
add_mcp_server,
|
||||
)
|
||||
|
||||
self._clear_sso_env(monkeypatch)
|
||||
for name, value in provider_env.items():
|
||||
monkeypatch.setenv(name, value)
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.add_server = AsyncMock()
|
||||
mock_manager.reload_servers_from_database = AsyncMock()
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs MCP server creation
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_mcp_server",
|
||||
AsyncMock(return_value=self._server_record(auth_type)),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint reads the global MCP manager
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await add_mcp_server(
|
||||
payload=NewMCPServerRequest(
|
||||
alias="ema",
|
||||
url="https://ema.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
),
|
||||
user_api_key_dict=generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"
|
||||
),
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"provider_env, expected_fragment",
|
||||
[
|
||||
({"GOOGLE_CLIENT_ID": "cid"}, "google"),
|
||||
({"MICROSOFT_CLIENT_ID": "cid"}, "microsoft"),
|
||||
({"SAML_IDP_METADATA_URL": "https://idp.example.com/metadata"}, "saml"),
|
||||
({}, "no SSO provider is configured"),
|
||||
],
|
||||
)
|
||||
async def test_create_warns_under_a_provider_that_captures_nothing(
|
||||
self, monkeypatch, caplog, provider_env, expected_fragment
|
||||
):
|
||||
await self._run_create(monkeypatch, provider_env, MCPAuth.oauth2_id_jag, caplog)
|
||||
warnings = self._id_jag_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert expected_fragment in str(warnings[0])
|
||||
assert "ema-1" in str(warnings[0])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_is_silent_under_generic_oidc(self, monkeypatch, caplog):
|
||||
await self._run_create(monkeypatch, {"GENERIC_CLIENT_ID": "cid"}, MCPAuth.oauth2_id_jag, caplog)
|
||||
assert self._id_jag_warnings(caplog) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_is_silent_for_other_auth_types(self, monkeypatch, caplog):
|
||||
"""Nothing but the id_jag arm sources credentials from a stored SSO assertion, so no
|
||||
other server registered under Google has anything to warn about."""
|
||||
await self._run_create(monkeypatch, {"GOOGLE_CLIENT_ID": "cid"}, MCPAuth.api_key, caplog)
|
||||
assert self._id_jag_warnings(caplog) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_to_id_jag_warns(self, monkeypatch, caplog):
|
||||
"""Switching an existing server onto id_jag opens the same gap a create does."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
edit_mcp_server,
|
||||
)
|
||||
|
||||
self._clear_sso_env(monkeypatch)
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.update_server = AsyncMock()
|
||||
mock_manager.reload_servers_from_database = AsyncMock()
|
||||
with (
|
||||
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs the MCP server lookup
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(return_value=self._server_record(MCPAuth.api_key)),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs MCP server updates
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
|
||||
AsyncMock(return_value=self._server_record(MCPAuth.oauth2_id_jag)),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs credential cleanup
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.purge_user_oauth_credentials_for_server",
|
||||
AsyncMock(return_value=0),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint reads the global MCP manager
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await edit_mcp_server(
|
||||
payload=UpdateMCPServerRequest(server_id="ema-1", auth_type=MCPAuth.oauth2_id_jag),
|
||||
user_api_key_dict=generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user"
|
||||
),
|
||||
)
|
||||
|
||||
warnings = self._id_jag_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert "google" in str(warnings[0])
|
||||
|
||||
|
||||
class TestHealthCheckServers:
|
||||
"""Test suite for health check servers endpoint"""
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import inspect
|
|||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Dict, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -4048,6 +4049,72 @@ class TestStrategyRouterWriteValidation:
|
|||
)
|
||||
assert _strategy_router_write_violation(incoming_params=None, existing_params=None) is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config",
|
||||
[
|
||||
{"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}},
|
||||
{
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tier_definitions": [
|
||||
{"name": "routine", "description": "routine drafting"},
|
||||
{"name": "hard", "description": "hard reasoning"},
|
||||
],
|
||||
"tiers": {"routine": "gpt-4o-mini", "hard": "gpt-4o"},
|
||||
"fallback_tier": "routine",
|
||||
},
|
||||
],
|
||||
)
|
||||
def test_model_less_patch_cannot_attach_router_config_to_a_regular_model(self, config: dict[str, object]) -> None:
|
||||
"""The license gate applies only to complexity routers, so a partial PATCH cannot poison a regular
|
||||
model with a capability-shaped config and make it occupy a slot."""
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
violation = _strategy_router_write_violation(
|
||||
incoming_params=updateLiteLLMParams(complexity_router_config=config),
|
||||
existing_params=LiteLLM_Params(model="openai/gpt-4o-mini"),
|
||||
)
|
||||
|
||||
assert violation is not None
|
||||
assert "does not start with 'auto_router/'" in violation
|
||||
assert "complexity_router_config" in violation
|
||||
|
||||
def test_effective_params_decrypts_a_stored_complexity_router_model(self, monkeypatch) -> None:
|
||||
"""A database row encrypts model, so the model-aware gate must not accidentally rely on plaintext mocks."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_effective_complexity_router_params,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt")
|
||||
encrypted_model = encrypt_value_helper("auto_router/complexity_router")
|
||||
effective_params = _effective_complexity_router_params(
|
||||
updateLiteLLMParams(complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}),
|
||||
LiteLLM_Params(model=encrypted_model),
|
||||
)
|
||||
|
||||
assert effective_params["model"] == "auto_router/complexity_router"
|
||||
|
||||
def test_model_less_patch_keeps_a_complexity_router_in_scope(self) -> None:
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
)
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
assert (
|
||||
_strategy_router_write_violation(
|
||||
incoming_params=updateLiteLLMParams(
|
||||
complexity_router_config={"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}}
|
||||
),
|
||||
existing_params=self._stored_complexity_params(),
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_restore_of_corrupted_row_is_allowed(self):
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_strategy_router_write_violation,
|
||||
|
|
@ -4354,33 +4421,33 @@ class TestStrategyRouterWriteValidation:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _live_router_holding_one_heuristic_v2(limit: int | None) -> Router:
|
||||
def _live_router_holding_one_capability(limit: int | None, config: Mapping[str, object]) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "k"}},
|
||||
{
|
||||
"model_name": "held-v2",
|
||||
"model_name": "held",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}},
|
||||
"complexity_router_config": config,
|
||||
},
|
||||
"model_info": {"id": "held-id"},
|
||||
},
|
||||
],
|
||||
heuristic_v2_router_limit=lambda: limit,
|
||||
auto_router_capability_limit=lambda: limit,
|
||||
)
|
||||
|
||||
class _FakeTx:
|
||||
"""Stands in for a prisma transaction: records the raw statements and exposes the model table."""
|
||||
"""Stands in for a prisma transaction: records raw statements and returns encrypted-model candidates."""
|
||||
|
||||
def __init__(self, db_held: int) -> None:
|
||||
self.db_held = db_held
|
||||
def __init__(self, db_models: list[str]) -> None:
|
||||
self.db_models = db_models
|
||||
self.raw_calls: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.litellm_proxymodeltable = MagicMock(create=AsyncMock(), update=AsyncMock())
|
||||
|
||||
async def query_raw(self, sql: str, *args: object) -> list[dict[str, object]]:
|
||||
self.raw_calls.append((sql, args))
|
||||
return [{"held": self.db_held}] if "count(*)" in sql else []
|
||||
return [{"model": model} for model in self.db_models] if "AS model" in sql else []
|
||||
|
||||
async def __aenter__(self) -> "TestStrategyRouterWriteValidation._FakeTx":
|
||||
return self
|
||||
|
|
@ -4391,9 +4458,9 @@ class TestStrategyRouterWriteValidation:
|
|||
class _FakeDb:
|
||||
"""Stands in for prisma_client: the plain client and the transaction it opens are told apart by identity."""
|
||||
|
||||
def __init__(self, db_held: int, existing_row: object = None) -> None:
|
||||
def __init__(self, db_models: list[str], existing_row: object = None) -> None:
|
||||
self.db = self
|
||||
self.tx_obj = TestStrategyRouterWriteValidation._FakeTx(db_held)
|
||||
self.tx_obj = TestStrategyRouterWriteValidation._FakeTx(db_models)
|
||||
self.litellm_proxymodeltable = MagicMock(
|
||||
create=AsyncMock(), update=AsyncMock(), find_unique=AsyncMock(return_value=existing_row)
|
||||
)
|
||||
|
|
@ -4403,6 +4470,43 @@ class TestStrategyRouterWriteValidation:
|
|||
|
||||
_V2 = {"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}}
|
||||
_V1 = {"classifier_type": "heuristic", "tiers": {"SIMPLE": "gpt-4o-mini"}}
|
||||
_CUSTOM_TIERS = {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tier_definitions": [
|
||||
{"name": "routine", "description": "routine drafting"},
|
||||
{"name": "hard", "description": "hard reasoning"},
|
||||
],
|
||||
"tiers": {"routine": "gpt-4o-mini", "hard": "gpt-4o"},
|
||||
"fallback_tier": "routine",
|
||||
}
|
||||
_TIER_LABELS_ONLY = {
|
||||
"classifier_type": "heuristic",
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
"tier_labels": {"SIMPLE": "Cheap"},
|
||||
}
|
||||
_CUSTOM_PROMPT = {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini", "system_prompt": "judge it my way"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
}
|
||||
_OPERATOR_EXAMPLES = {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
"classification_examples": '- "reset my password" -> SIMPLE',
|
||||
}
|
||||
_OPERATOR_OPENING_PROMPT = {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
"classification_prompt": "Grade by data sensitivity",
|
||||
}
|
||||
_SHIPPED_RUBRIC = {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini", "classification_rubric": "agentic"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"incoming,existing,expected",
|
||||
|
|
@ -4431,41 +4535,55 @@ class TestStrategyRouterWriteValidation:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"limit,effective_config,db_held,config_holds_one,model_id,expected",
|
||||
"limit,effective_params,db_models,config_config,model_id,expected",
|
||||
[
|
||||
(1, _V2, 1, False, None, "refused"),
|
||||
(1, _V2, 0, True, None, "refused"),
|
||||
(1, _V2, 0, False, None, "reserved"),
|
||||
(1, _V2, 0, False, "held-id", "reserved"),
|
||||
(2, _V2, 1, False, None, "reserved"),
|
||||
(1, _V1, 5, True, None, "plain"),
|
||||
(1, None, 5, True, None, "plain"),
|
||||
(None, _V2, 5, True, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, ["auto_router/complexity_router"], None, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, [], _V2, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, [], None, None, "reserved"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, [], None, "held-id", "reserved"),
|
||||
(2, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, ["auto_router/complexity_router"], None, None, "reserved"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CUSTOM_TIERS}, ["openai/gpt-4o"], None, None, "reserved"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CUSTOM_TIERS}, [], _CUSTOM_PROMPT, None, "refused"),
|
||||
(1, {"model": "openai/gpt-4o", "complexity_router_config": _CUSTOM_TIERS}, ["auto_router/complexity_router"], None, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V1}, ["auto_router/complexity_router"], _V2, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": None}, ["auto_router/complexity_router"], _V2, None, "plain"),
|
||||
(None, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, ["auto_router/complexity_router"], _V2, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _TIER_LABELS_ONLY}, ["auto_router/complexity_router"], _CUSTOM_TIERS, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CUSTOM_PROMPT}, ["auto_router/complexity_router"], None, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _OPERATOR_EXAMPLES}, [], _CUSTOM_TIERS, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _OPERATOR_OPENING_PROMPT}, [], _CUSTOM_PROMPT, None, "refused"),
|
||||
(None, {"model": "auto_router/complexity_router", "complexity_router_config": _OPERATOR_EXAMPLES}, ["auto_router/complexity_router"], _CUSTOM_TIERS, None, "plain"),
|
||||
],
|
||||
)
|
||||
async def test_heuristic_v2_slot_matrix(
|
||||
async def test_auto_router_capability_slot_matrix(
|
||||
self,
|
||||
limit: int | None,
|
||||
effective_config: object,
|
||||
db_held: int,
|
||||
config_holds_one: bool,
|
||||
effective_params: Mapping[str, object],
|
||||
db_models: list[str],
|
||||
config_config: Mapping[str, object] | None,
|
||||
model_id: str | None,
|
||||
expected: str,
|
||||
) -> None:
|
||||
"""The slot is claimed inside a locked transaction only for a heuristic_v2 write under a limit; the DB rows
|
||||
(other pods included) plus config.yaml routers decide, the row being edited is excluded through the SQL
|
||||
parameter, and every other write runs on the plain client with no lock."""
|
||||
"""The slot is claimed inside a locked transaction only for a write that claims a licensed capability
|
||||
under a limit; the DB rows (other pods included) plus config.yaml routers decide, the row being edited
|
||||
is excluded through the SQL parameter, and every other write runs on the plain client with no lock.
|
||||
|
||||
heuristic_v2 has its own slot, while custom tier definitions and custom prompts count into one shared
|
||||
customization slot. Renaming built-in tiers through tier_labels claims nothing at all."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
HEURISTIC_V2_SLOT_LOCK_KEY,
|
||||
_heuristic_v2_slot,
|
||||
AUTO_ROUTER_CAPABILITY_SLOT_LOCK_KEY,
|
||||
_auto_router_capability_slot,
|
||||
)
|
||||
from litellm.router_utils.auto_router_model_naming import gated_capability_of
|
||||
|
||||
fake = self._FakeDb(db_held)
|
||||
live_router = self._live_router_holding_one_heuristic_v2(limit) if config_holds_one else None
|
||||
capability = gated_capability_of(effective_params)
|
||||
|
||||
fake = self._FakeDb(db_models)
|
||||
live_router = self._live_router_holding_one_capability(limit, config_config) if config_config is not None else None
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: limit), # test-quality-ok: the guard reads the proxy license singleton with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.auto_router_capability_limit", lambda: limit), # test-quality-ok: the guard reads the proxy license singleton with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", live_router), # test-quality-ok: the guard reads the proxy router global with no injection seam
|
||||
patch( # test-quality-ok: the cross-pod publish is the side effect under test; redis is not configured here
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.publish_config_change",
|
||||
|
|
@ -4474,13 +4592,15 @@ class TestStrategyRouterWriteValidation:
|
|||
):
|
||||
if expected == "refused":
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
async with _heuristic_v2_slot(fake, effective_config=effective_config, model_id=model_id):
|
||||
async with _auto_router_capability_slot(fake, effective_params=effective_params, model_id=model_id):
|
||||
pass
|
||||
assert exc_info.value.status_code == 403
|
||||
assert capability is not None
|
||||
assert "At most 1 auto-router" in str(exc_info.value.detail)
|
||||
assert capability.subject in str(exc_info.value.detail)
|
||||
assert "'auto_router' feature lifts the limit" in str(exc_info.value.detail)
|
||||
return
|
||||
async with _heuristic_v2_slot(fake, effective_config=effective_config, model_id=model_id) as tables:
|
||||
async with _auto_router_capability_slot(fake, effective_params=effective_params, model_id=model_id) as tables:
|
||||
handle = tables
|
||||
if expected == "plain":
|
||||
await handle.create(data={})
|
||||
|
|
@ -4489,10 +4609,13 @@ class TestStrategyRouterWriteValidation:
|
|||
return
|
||||
assert handle is fake.tx_obj.litellm_proxymodeltable
|
||||
published.assert_awaited_once_with(redis_cache=None, object_type="litellm_proxymodeltable")
|
||||
(lock_sql, lock_params), (_count_sql, count_params) = fake.tx_obj.raw_calls
|
||||
(lock_sql, lock_params), (count_sql, count_params) = fake.tx_obj.raw_calls
|
||||
assert "pg_advisory_xact_lock($1)" in lock_sql and "count" not in lock_sql
|
||||
assert lock_params == (HEURISTIC_V2_SLOT_LOCK_KEY,)
|
||||
assert lock_params == (AUTO_ROUTER_CAPABILITY_SLOT_LOCK_KEY,)
|
||||
assert count_params == (model_id or "",)
|
||||
assert "AS model" in count_sql
|
||||
assert capability is not None
|
||||
assert capability.sql_config_predicate.split("{config}")[-1].strip() in count_sql
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_bookkeeping_runs_after_the_slot_is_released(self) -> None:
|
||||
|
|
@ -4549,14 +4672,14 @@ class TestStrategyRouterWriteValidation:
|
|||
)
|
||||
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
fake = self._FakeDb(db_held=1)
|
||||
fake = self._FakeDb(["auto_router/complexity_router"])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.auto_router_capability_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch( # test-quality-ok: prior auth check needs a live DB; only the license limit is under test
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
|
|
@ -4579,6 +4702,93 @@ class TestStrategyRouterWriteValidation:
|
|||
fake.tx_obj.litellm_proxymodeltable.create.assert_not_awaited()
|
||||
fake.litellm_proxymodeltable.create.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_less_patch_rejects_router_config_on_a_regular_model(self) -> None:
|
||||
"""PATCH rejects the poison before its row write or the capability slot."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
|
||||
from litellm.types.router import updateLiteLLMParams
|
||||
|
||||
model_id = "regular-model"
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
regular = Deployment(
|
||||
model_name="regular-model",
|
||||
litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"),
|
||||
model_info={"id": model_id},
|
||||
)
|
||||
fake = self._FakeDb([])
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: config-based lookup must be absent to drive the stored-row branch
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reaches its DB-write branch only with this process setting
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: authorization branch reads the proxy-wide premium flag
|
||||
patch( # test-quality-ok: inject stored regular row without a database
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.get_db_model",
|
||||
new=AsyncMock(return_value=regular),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint must reject before database authorization needs a live store
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await patch_model(
|
||||
model_id=model_id,
|
||||
patch_data=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=self._CUSTOM_TIERS)
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert "does not start with 'auto_router/'" in str(exc_info.value.message)
|
||||
assert fake.tx_obj.raw_calls == []
|
||||
assert fake.tx_obj.litellm_proxymodeltable.update.await_count == 0
|
||||
assert fake.litellm_proxymodeltable.update.await_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_less_legacy_update_rejects_router_config_on_a_regular_model(self) -> None:
|
||||
"""The legacy update endpoint enforces the same boundary before its row write or slot."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import update_model
|
||||
from litellm.types.router import ModelInfo, updateLiteLLMParams
|
||||
|
||||
model_id = "regular-model"
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
regular = Deployment(
|
||||
model_name="regular-model",
|
||||
litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"),
|
||||
model_info={"id": model_id},
|
||||
)
|
||||
existing_row = MagicMock()
|
||||
existing_row.model_dump.return_value = regular.model_dump()
|
||||
existing_row.litellm_params = regular.litellm_params.model_dump()
|
||||
fake = self._FakeDb([], existing_row=existing_row)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: config-based lookup must be absent to drive the stored-row branch
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reaches its DB-write branch only with this process setting
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: authorization branch reads the proxy-wide premium flag
|
||||
patch( # test-quality-ok: endpoint must reject before database authorization needs a live store
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await update_model(
|
||||
model_params=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=self._CUSTOM_TIERS),
|
||||
model_info=ModelInfo(id=model_id),
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert "does not start with 'auto_router/'" in str(exc_info.value.message)
|
||||
assert fake.tx_obj.raw_calls == []
|
||||
assert fake.tx_obj.litellm_proxymodeltable.update.await_count == 0
|
||||
assert fake.litellm_proxymodeltable.update.await_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_model_refuses_switching_another_router_to_heuristic_v2(self) -> None:
|
||||
"""patch_model relays HTTPException as-is, so the license refusal reaches the client as a plain 403."""
|
||||
|
|
@ -4591,14 +4801,14 @@ class TestStrategyRouterWriteValidation:
|
|||
|
||||
model_id = "other-id"
|
||||
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
fake = self._FakeDb(db_held=1)
|
||||
fake = self._FakeDb(["auto_router/complexity_router"])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.auto_router_capability_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch( # test-quality-ok: the write must be refused before this DB step runs
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.get_db_model",
|
||||
new=AsyncMock(return_value=self._db_complexity_router(model_id)),
|
||||
|
|
@ -4643,14 +4853,14 @@ class TestStrategyRouterWriteValidation:
|
|||
"model_info": {"id": model_id},
|
||||
}
|
||||
existing_row.litellm_params = existing_row.model_dump.return_value["litellm_params"]
|
||||
fake = self._FakeDb(db_held=1, existing_row=existing_row)
|
||||
fake = self._FakeDb(["auto_router/complexity_router"], existing_row=existing_row)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", fake), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", None), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch("litellm.proxy.proxy_server._license_check.auto_router_capability_limit", lambda: 1), # test-quality-ok: endpoint reads proxy server globals with no injection seam
|
||||
patch( # test-quality-ok: prior auth check needs a live DB; only the license limit is under test
|
||||
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
|
||||
new=AsyncMock(return_value=None),
|
||||
|
|
|
|||
|
|
@ -1181,3 +1181,17 @@ async def test_find_member_if_email_missing_row_raises_documented_400():
|
|||
"non-existent user_email in LiteLLM_UserTable. Use 'user_id' instead."
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
def test_v2_update_organization_is_in_openapi_schema():
|
||||
"""PATCH /v2/organization/{organization_id} is documented in the generated OpenAPI spec."""
|
||||
from fastapi import FastAPI
|
||||
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import router
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
|
||||
v2_path = app.openapi()["paths"]["/v2/organization/{organization_id}"]
|
||||
assert v2_path["patch"]["tags"] == ["organization management"]
|
||||
assert "OrganizationUpdateRequestV2" in json.dumps(v2_path["patch"]["requestBody"])
|
||||
|
|
|
|||
|
|
@ -1,17 +1,16 @@
|
|||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from contextlib import ExitStack, asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm._uuid import uuid
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import LiteLLM_UserTable, NewUserResponse
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
|
||||
|
|
@ -1615,8 +1614,8 @@ async def test_get_generic_sso_response_with_empty_headers():
|
|||
async def test_get_generic_sso_response_includes_token_claims_when_enabled(monkeypatch):
|
||||
import jwt as pyjwt
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_jwt_handler = MagicMock(spec=JWTHandler)
|
||||
|
|
@ -2321,10 +2320,10 @@ class TestCustomUISSO:
|
|||
async def test_handle_custom_ui_sso_sign_in_success(self):
|
||||
"""Test successful custom UI SSO sign-in with valid headers"""
|
||||
from fastapi_sso.sso.base import OpenID
|
||||
|
||||
from litellm_enterprise.proxy.auth.custom_sso_handler import (
|
||||
EnterpriseCustomSSOHandler,
|
||||
)
|
||||
|
||||
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
|
||||
|
||||
# Mock request with custom headers
|
||||
|
|
@ -2400,6 +2399,7 @@ class TestCustomUISSO:
|
|||
from litellm_enterprise.proxy.auth.custom_sso_handler import (
|
||||
EnterpriseCustomSSOHandler,
|
||||
)
|
||||
|
||||
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -2436,10 +2436,10 @@ class TestCustomUISSO:
|
|||
and its methods are called with the correct parameters
|
||||
"""
|
||||
from fastapi_sso.sso.base import OpenID
|
||||
|
||||
from litellm_enterprise.proxy.auth.custom_sso_handler import (
|
||||
EnterpriseCustomSSOHandler,
|
||||
)
|
||||
|
||||
from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler
|
||||
|
||||
# Create a real custom handler class instance
|
||||
|
|
@ -8167,6 +8167,128 @@ async def test_debug_sso_callback_handles_missing_raw_response():
|
|||
assert "user@example.com" in body
|
||||
|
||||
|
||||
# ── The debug page is where an operator lands when ID-JAG is failing ──────────
|
||||
|
||||
_GOOGLE_DEBUG_CLIENT_ID = "debug-google-client-id"
|
||||
_GENERIC_DEBUG_CLIENT_ID = "debug-generic-client-id"
|
||||
|
||||
|
||||
async def _render_debug_page(provider_env, id_jag_registered, force_inert=False):
|
||||
"""Drive /sso/debug/callback and return the raw response body."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import GoogleSSOHandler, debug_sso_callback
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://proxy.example.com/"
|
||||
mock_request.cookies = {}
|
||||
mock_request.query_params = {}
|
||||
|
||||
parsed = {"sub": "user_123", "email": "u@example.com"}
|
||||
|
||||
async def fake_generic(**kwargs):
|
||||
return parsed, {"sub": "user_123"}, {"scope": "openid"}, None
|
||||
|
||||
async def fake_google(**kwargs):
|
||||
return parsed
|
||||
|
||||
stack = [
|
||||
patch.dict(os.environ, provider_env, clear=False),
|
||||
patch( # test-quality-ok: endpoint test stubs the upstream generic IdP boundary
|
||||
"litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response", side_effect=fake_generic
|
||||
),
|
||||
patch.object( # test-quality-ok: endpoint test stubs the upstream Google IdP boundary
|
||||
GoogleSSOHandler, "get_google_callback_response", side_effect=fake_google
|
||||
),
|
||||
patch( # test-quality-ok: debug endpoint reads this module global without an injection seam
|
||||
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
|
||||
AsyncMock(return_value=id_jag_registered),
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: debug endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: debug endpoint reads proxy DB
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: debug endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.jwt_handler", MagicMock(spec=JWTHandler)), # test-quality-ok: debug endpoint reads proxy globals
|
||||
]
|
||||
if force_inert:
|
||||
stack.append(
|
||||
patch( # test-quality-ok: force-inert reference isolates the endpoint's pre-change response
|
||||
"litellm.proxy.management_endpoints.ui_sso.warn_if_id_jag_capture_gap",
|
||||
AsyncMock(return_value=None),
|
||||
)
|
||||
)
|
||||
|
||||
with ExitStack() as es:
|
||||
for ctx in stack:
|
||||
es.enter_context(ctx)
|
||||
for var in ("MICROSOFT_CLIENT_ID", "GOOGLE_CLIENT_ID", "GENERIC_CLIENT_ID", "SAML_IDP_METADATA_URL"):
|
||||
if var not in provider_env:
|
||||
os.environ.pop(var, None)
|
||||
response = await debug_sso_callback(mock_request)
|
||||
|
||||
return response.body.decode()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_debug_page_logs_the_capture_gap_but_never_renders_it(caplog):
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
body = await _render_debug_page({"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID}, id_jag_registered=True)
|
||||
|
||||
warnings = _id_jag_gap_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert "google" in warnings[0]
|
||||
assert "GENERIC_CLIENT_ID" in warnings[0]
|
||||
assert "id_jag" not in body
|
||||
assert "GENERIC_CLIENT_ID" not in body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_debug_page_is_byte_identical_when_the_provider_captures():
|
||||
"""A deployment with no gap must get the page it got before this change, to the byte. The
|
||||
comparison is against the endpoint with the diagnostic forced inert, not against a guess."""
|
||||
with_feature = await _render_debug_page(
|
||||
{"GENERIC_CLIENT_ID": _GENERIC_DEBUG_CLIENT_ID}, id_jag_registered=True
|
||||
)
|
||||
pre_change = await _render_debug_page(
|
||||
{"GENERIC_CLIENT_ID": _GENERIC_DEBUG_CLIENT_ID},
|
||||
id_jag_registered=True,
|
||||
force_inert=True,
|
||||
)
|
||||
|
||||
assert with_feature == pre_change
|
||||
assert "id_jag" not in with_feature
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_debug_page_is_byte_identical_when_no_id_jag_server_is_registered():
|
||||
"""Most deployments run Google SSO and no id_jag server at all; their debug page must not
|
||||
grow an ID-JAG section about a feature they do not use."""
|
||||
with_feature = await _render_debug_page(
|
||||
{"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID}, id_jag_registered=False
|
||||
)
|
||||
pre_change = await _render_debug_page(
|
||||
{"GOOGLE_CLIENT_ID": _GOOGLE_DEBUG_CLIENT_ID},
|
||||
id_jag_registered=False,
|
||||
force_inert=True,
|
||||
)
|
||||
|
||||
assert with_feature == pre_change
|
||||
assert "id_jag" not in with_feature
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_debug_page_survives_a_store_outage(monkeypatch, caplog):
|
||||
"""The page's job is to render claims; an unreachable MCP table must cost it the annotation,
|
||||
not the page."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import warn_if_id_jag_capture_gap
|
||||
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", _GOOGLE_DEBUG_CLIENT_ID)
|
||||
retention_check = AsyncMock(side_effect=Exception("db down"))
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
assert await warn_if_id_jag_capture_gap(retention_enabled=retention_check) is None
|
||||
|
||||
retention_check.assert_awaited_once()
|
||||
|
||||
assert _id_jag_gap_warnings(caplog) == []
|
||||
|
||||
|
||||
async def _render_legacy_login_page(env_overrides, general_settings):
|
||||
from litellm.proxy.management_endpoints.ui_sso import google_login
|
||||
|
||||
|
|
@ -8261,8 +8383,8 @@ async def test_saml_callback_enforces_free_sso_user_limit_after_validation():
|
|||
that /sso/key/generate enforces; the ACS re-checks it after validating the assertion,
|
||||
so the entitlement DB query never runs on unvalidated input."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import saml_callback
|
||||
from litellm.proxy.management_endpoints.types import CustomOpenID
|
||||
from litellm.proxy.management_endpoints.ui_sso import saml_callback
|
||||
|
||||
call_order: list[str] = []
|
||||
|
||||
|
|
@ -8681,6 +8803,248 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
|
|||
assert response.status_code == 200
|
||||
|
||||
|
||||
def _id_jag_gap_warnings(caplog) -> list[str]:
|
||||
return [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if record.levelno == logging.WARNING and "oauth2_id_jag" in record.getMessage()
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"provider_env, expected_fragment",
|
||||
[
|
||||
({"GOOGLE_CLIENT_ID": "cid"}, "google"),
|
||||
({"MICROSOFT_CLIENT_ID": "cid", "MICROSOFT_TENANT": "t"}, "microsoft"),
|
||||
({}, "no SSO provider is configured"),
|
||||
],
|
||||
)
|
||||
async def test_uncaptured_assertion_warns_when_an_id_jag_server_is_registered(
|
||||
monkeypatch, caplog, provider_env, expected_fragment
|
||||
):
|
||||
"""A provider with no capture path leaves ID-JAG permanently broken, and the only place
|
||||
that is knowable is the login itself; without this line the operator sees nothing at all."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
warn_if_id_jag_assertion_uncaptured,
|
||||
)
|
||||
|
||||
for name in ("GOOGLE_CLIENT_ID", "MICROSOFT_CLIENT_ID", "GENERIC_CLIENT_ID", "SAML_IDP_METADATA_URL"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
for name, value in provider_env.items():
|
||||
monkeypatch.setenv(name, value)
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await warn_if_id_jag_assertion_uncaptured(None, retention_enabled=AsyncMock(return_value=True))
|
||||
|
||||
warnings = _id_jag_gap_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert expected_fragment in str(warnings[0])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generic_provider_that_returned_no_id_token_still_warns(monkeypatch, caplog):
|
||||
"""Generic OIDC has a capture path, so there is no configuration gap to report; the login
|
||||
still handed the id_jag arm nothing, and that must not pass silently."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
warn_if_id_jag_assertion_uncaptured,
|
||||
)
|
||||
|
||||
for name in ("GOOGLE_CLIENT_ID", "MICROSOFT_CLIENT_ID", "SAML_IDP_METADATA_URL"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
monkeypatch.setenv("GENERIC_CLIENT_ID", "cid")
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await warn_if_id_jag_assertion_uncaptured(None, retention_enabled=AsyncMock(return_value=True))
|
||||
|
||||
warnings = _id_jag_gap_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert "no usable id_token" in str(warnings[0])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_warning_when_the_assertion_was_captured(monkeypatch, caplog):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import (
|
||||
assertion_from_sso_login,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
warn_if_id_jag_assertion_uncaptured,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
|
||||
assertion = assertion_from_sso_login(_ema_id_token(), None)
|
||||
assert assertion is not None
|
||||
|
||||
retention_mock = AsyncMock(return_value=True)
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await warn_if_id_jag_assertion_uncaptured(assertion, retention_enabled=retention_mock)
|
||||
|
||||
assert _id_jag_gap_warnings(caplog) == []
|
||||
retention_mock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_warning_when_no_id_jag_server_is_registered(monkeypatch, caplog):
|
||||
"""Most deployments never register one; a warning about ID-JAG on every login there would
|
||||
be pure noise and would train operators to ignore it."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
warn_if_id_jag_assertion_uncaptured,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await warn_if_id_jag_assertion_uncaptured(None, retention_enabled=AsyncMock(return_value=False))
|
||||
|
||||
assert _id_jag_gap_warnings(caplog) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_outage_does_not_break_the_login(monkeypatch, caplog):
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
warn_if_id_jag_assertion_uncaptured,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
assert (
|
||||
await warn_if_id_jag_assertion_uncaptured(
|
||||
None, retention_enabled=AsyncMock(side_effect=Exception("db down"))
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
assert _id_jag_gap_warnings(caplog) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_browser_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog):
|
||||
"""Wiring: the browser login path must reach the diagnostic, not just define it."""
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
mock_request.cookies = {}
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: endpoint test stubs the Prisma client lookup
|
||||
"litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}), # test-quality-ok: endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.premium_user", False), # test-quality-ok: endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.user_custom_sso", None), # test-quality-ok: endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), # test-quality-ok: endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.redis_usage_cache", None), # test-quality-ok: endpoint reads proxy globals
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: endpoint reads proxy globals
|
||||
patch( # test-quality-ok: endpoint test stubs key generation at its module boundary
|
||||
"litellm.proxy.proxy_server.generate_key_helper_fn",
|
||||
AsyncMock(return_value={"token": "sk-ui-key", "user_id": "canonical-user-id"}),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs the user database lookup
|
||||
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs the admin database lookup
|
||||
"litellm.proxy.management_endpoints.ui_sso.check_and_update_if_proxy_admin_id",
|
||||
AsyncMock(return_value="internal_user"),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs assertion persistence
|
||||
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
|
||||
AsyncMock(),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint reads this module global without an injection seam
|
||||
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
):
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await SSOAuthenticationHandler.get_redirect_response_from_openid(
|
||||
result=CustomOpenID(
|
||||
id="raw-idp-subject",
|
||||
email="u@example.com",
|
||||
first_name="U",
|
||||
last_name="Ser",
|
||||
display_name="U Ser",
|
||||
provider="google",
|
||||
team_ids=[],
|
||||
user_role=None,
|
||||
),
|
||||
request=mock_request,
|
||||
received_response=None,
|
||||
generic_client_id=None,
|
||||
ui_access_mode=None,
|
||||
access_token_payload=None,
|
||||
jwt_handler=None,
|
||||
sso_assertion=None,
|
||||
)
|
||||
|
||||
warnings = _id_jag_gap_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert "google" in str(warnings[0])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog):
|
||||
"""Wiring: the CLI login path shares the gap, so it must share the diagnostic."""
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_complete_cli_sso_callback_session,
|
||||
)
|
||||
|
||||
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "cid")
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
|
||||
user_info = MagicMock()
|
||||
user_info.user_id = "cli-user-id"
|
||||
user_info.user_role = "internal_user"
|
||||
user_info.models = []
|
||||
user_info.teams = []
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: endpoint test stubs the user database lookup
|
||||
"litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db",
|
||||
AsyncMock(return_value=user_info),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs CLI team lookup
|
||||
"litellm.proxy.management_endpoints.ui_sso.fetch_cli_sso_team_details",
|
||||
AsyncMock(return_value=[]),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs attribution metadata
|
||||
"litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata",
|
||||
return_value={},
|
||||
),
|
||||
patch( # test-quality-ok: endpoint test stubs assertion persistence
|
||||
"litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema",
|
||||
AsyncMock(),
|
||||
),
|
||||
patch( # test-quality-ok: endpoint reads this module global without an injection seam
|
||||
"litellm.proxy.management_endpoints.ui_sso.ema_assertion_retention_enabled",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
):
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await _complete_cli_sso_callback_session(
|
||||
request=mock_request,
|
||||
key="cli-login-id",
|
||||
flow={},
|
||||
result={"sub": "raw-idp-subject"},
|
||||
parsed_openid_result={
|
||||
"user_id": "raw-idp-subject",
|
||||
"user_email": "u@example.com",
|
||||
"user_role": None,
|
||||
},
|
||||
user_defined_values=None,
|
||||
prisma_client=MagicMock(),
|
||||
user_api_key_cache=MagicMock(),
|
||||
cli_sso_session_cache=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
sso_assertion=None,
|
||||
)
|
||||
|
||||
warnings = _id_jag_gap_warnings(caplog)
|
||||
assert len(warnings) == 1
|
||||
assert "microsoft" in str(warnings[0])
|
||||
|
||||
|
||||
def _cli_callback_kwargs(flow):
|
||||
return {
|
||||
"request": _cli_callback_request(),
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from litellm.proxy.proxy_server import (
|
|||
resolve_routing_plugins,
|
||||
validate_deployment_complexity_router_placement,
|
||||
validate_deployment_max_agentic_loops,
|
||||
validate_heuristic_v2_router_limit,
|
||||
validate_auto_router_capability_limits,
|
||||
)
|
||||
|
||||
from .conftest import normalize
|
||||
|
|
@ -204,13 +204,71 @@ def _heuristic_v2_row(model_name: str, classifier_type: str = "heuristic_v2") ->
|
|||
}
|
||||
|
||||
|
||||
def test_validate_heuristic_v2_router_limit_refuses_to_start_over_the_limit() -> None:
|
||||
def _custom_tier_row(model_name: str) -> dict[str, object]:
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
"classifier_type": "llm",
|
||||
"tier_definitions": [
|
||||
{"name": "routine", "description": "routine drafting"},
|
||||
{"name": "hard", "description": "hard reasoning"},
|
||||
],
|
||||
"tiers": {"routine": "gpt-4o-mini", "hard": "gpt-4o"},
|
||||
"fallback_tier": "routine",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _operator_examples_row(model_name: str) -> dict[str, object]:
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
"classification_examples": '- "reset my password" -> SIMPLE',
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _custom_prompt_row(model_name: str) -> dict[str, object]:
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini", "system_prompt": "judge it my way"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"over_limit_rows,subject",
|
||||
[
|
||||
([_heuristic_v2_row("a"), _heuristic_v2_row("b"), _heuristic_v2_row("c", "heuristic")], "heuristic_v2"),
|
||||
([_custom_tier_row("a"), _custom_tier_row("b"), _heuristic_v2_row("c", "heuristic")], "tier_definitions"),
|
||||
([_custom_prompt_row("a"), _custom_prompt_row("b"), _heuristic_v2_row("c", "heuristic")], "operator-written classifier prompt"),
|
||||
([_custom_tier_row("a"), _custom_prompt_row("b"), _heuristic_v2_row("c", "heuristic")], "operator-written classifier prompt"),
|
||||
([_operator_examples_row("a"), _custom_tier_row("b"), _heuristic_v2_row("c", "heuristic")], "operator-written classifier prompt"),
|
||||
],
|
||||
)
|
||||
def test_validate_auto_router_capability_limits_refuses_to_start_over_the_limit(
|
||||
over_limit_rows: list[dict[str, object]], subject: str
|
||||
) -> None:
|
||||
"""Same reason as the two validators above: the proxy router swallows registration errors, so
|
||||
an over-limit config.yaml must fail here instead of booting with a silently missing router."""
|
||||
with pytest.raises(ValueError, match=re.escape("At most 1 auto-router")) as exc_info:
|
||||
validate_heuristic_v2_router_limit(
|
||||
[_heuristic_v2_row("a"), _heuristic_v2_row("b"), _heuristic_v2_row("c", "heuristic")], limit=1
|
||||
)
|
||||
validate_auto_router_capability_limits(over_limit_rows, limit=1)
|
||||
assert subject in str(exc_info.value)
|
||||
assert "'auto_router' feature lifts the limit" in str(exc_info.value)
|
||||
|
||||
|
||||
|
|
@ -220,12 +278,15 @@ def test_validate_heuristic_v2_router_limit_refuses_to_start_over_the_limit() ->
|
|||
([_heuristic_v2_row("a"), _heuristic_v2_row("b")], None),
|
||||
([_heuristic_v2_row("a"), _heuristic_v2_row("c", "heuristic")], 1),
|
||||
([{"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}], 1),
|
||||
([_custom_tier_row("a"), _custom_tier_row("b")], None),
|
||||
([_custom_tier_row("a"), _heuristic_v2_row("b")], 1),
|
||||
],
|
||||
)
|
||||
def test_validate_heuristic_v2_router_limit_leaves_configs_within_the_limit_alone(
|
||||
def test_validate_auto_router_capability_limits_leaves_configs_within_the_limit_alone(
|
||||
model_list: list[dict[str, object]], limit: int | None
|
||||
) -> None:
|
||||
assert validate_heuristic_v2_router_limit(model_list, limit=limit) is None
|
||||
"""The last case is the separate-ceiling invariant: one router of each capability fits under a limit of one."""
|
||||
assert validate_auto_router_capability_limits(model_list, limit=limit) is None
|
||||
|
||||
|
||||
_TWO_HEURISTIC_V2_ROUTERS_YAML = (
|
||||
|
|
@ -247,7 +308,7 @@ _TWO_HEURISTIC_V2_ROUTERS_YAML = (
|
|||
" classifier_type: heuristic_v2\n"
|
||||
" tiers: {SIMPLE: gpt-4o-mini}\n"
|
||||
"router_settings:\n"
|
||||
" heuristic_v2_router_limit: 99\n"
|
||||
" auto_router_capability_limit: 99\n"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -256,7 +317,7 @@ _TWO_HEURISTIC_V2_ROUTERS_YAML = (
|
|||
async def test_ProxyConfig_load_config_takes_the_heuristic_v2_limit_from_the_license_only(
|
||||
tmp_path, monkeypatch, license_limit: int | None
|
||||
) -> None:
|
||||
"""`router_settings.heuristic_v2_router_limit` is managed outside config.yaml: an operator
|
||||
"""`router_settings.auto_router_capability_limit` is managed outside config.yaml: an operator
|
||||
cannot grant the entitlement by editing the config, and a licensed proxy boots both routers."""
|
||||
f = tmp_path / "c.yaml"
|
||||
f.write_text(_TWO_HEURISTIC_V2_ROUTERS_YAML)
|
||||
|
|
@ -264,15 +325,15 @@ async def test_ProxyConfig_load_config_takes_the_heuristic_v2_limit_from_the_lic
|
|||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: license_limit
|
||||
"litellm.proxy.proxy_server._license_check.auto_router_capability_limit", lambda: license_limit
|
||||
)
|
||||
|
||||
if license_limit is None:
|
||||
router, _model_list, _general_settings = await ProxyConfig().load_config(
|
||||
router=None, config_file_path=str(f)
|
||||
)
|
||||
assert router.heuristic_v2_router_limit is not None
|
||||
assert router.heuristic_v2_router_limit() is None
|
||||
assert router.auto_router_capability_limit is not None
|
||||
assert router.auto_router_capability_limit() is None
|
||||
assert sorted(router.complexity_routers) == ["v2-a", "v2-b"]
|
||||
return
|
||||
|
||||
|
|
@ -296,12 +357,12 @@ async def test_ProxyConfig_load_config_router_refuses_a_db_heuristic_v2_router_b
|
|||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server._license_check.heuristic_v2_router_limit", lambda: 1)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server._license_check.auto_router_capability_limit", lambda: 1)
|
||||
|
||||
router, _model_list, _general_settings = await ProxyConfig().load_config(router=None, config_file_path=str(f))
|
||||
|
||||
assert router.heuristic_v2_router_limit is not None
|
||||
assert router.heuristic_v2_router_limit() == 1
|
||||
assert router.auto_router_capability_limit is not None
|
||||
assert router.auto_router_capability_limit() == 1
|
||||
assert sorted(router.complexity_routers) == ["v1-b", "v2-a"]
|
||||
db_row = Deployment(**_heuristic_v2_row("v2-from-db"), model_info={"id": "db-id"})
|
||||
assert router.upsert_deployment(db_row) is None
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue