Merge remote-tracking branch 'origin/main' into litellm_langfuse_otel_metadata_keys

This commit is contained in:
yucheng 2026-09-23 20:47:54 +00:00
commit d0dfafbf89
27 changed files with 4103 additions and 107 deletions

View file

@ -11,7 +11,7 @@ run_full() {
[ -n "${CIRCLE_PULL_REQUEST:-}" ] || run_full "not a pull request"
candidate_bases="main"
candidate_bases="${PATH_FILTER_BASE_BRANCH:-main}"
merge_base=""
for base in $candidate_bases; do
git fetch --quiet origin "$base" 2>/dev/null || continue

292
.circleci/tests.yml Normal file
View file

@ -0,0 +1,292 @@
version: 2.1
commands:
wait_for_service:
parameters:
url:
type: string
timeout:
type: string
default: "60"
steps:
- run:
name: "Wait for << parameters.url >>"
command: |
TIMEOUT=<< parameters.timeout >>
URL="<< parameters.url >>"
ELAPSED=0
echo "Waiting up to ${TIMEOUT}s for ${URL} ..."
if echo "$URL" | grep -q '^tcp://'; then
HOST=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f1)
PORT=$(echo "$URL" | sed 's|tcp://||' | cut -d: -f2)
while ! bash -c "echo > /dev/tcp/$HOST/$PORT" 2>/dev/null; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
else
while ! curl -sf --max-time 5 "$URL" > /dev/null 2>&1; do
sleep 2; ELAPSED=$((ELAPSED+2))
if [ "$ELAPSED" -ge "$TIMEOUT" ]; then echo "Timed out"; exit 1; fi
done
fi
echo "Service ready after ${ELAPSED}s"
install_uv:
steps:
- run:
name: Install uv (pinned 0.10.9)
command: |
curl -LsSf -o /tmp/uv-install.sh https://astral.sh/uv/0.10.9/install.sh
echo "7fc46e39cb97290b57169c0c813a17970585ac519139f19006453c99b5f2f45f /tmp/uv-install.sh" | sha256sum -c -
env UV_NO_MODIFY_PATH=1 sh /tmp/uv-install.sh
rm -f /tmp/uv-install.sh
echo 'export PATH="$HOME/.local/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.local/bin:$PATH"
install_rust:
steps:
- run:
name: Install Rust (rustup 1.28.2, toolchain 1.98.0)
command: |
case "$(uname -m)" in
x86_64)
RUSTUP_TRIPLE=x86_64-unknown-linux-gnu
RUSTUP_SHA256=20a06e644b0d9bd2fbdbfd52d42540bdde820ea7df86e92e533c073da0cdd43c
;;
aarch64)
RUSTUP_TRIPLE=aarch64-unknown-linux-gnu
RUSTUP_SHA256=e3853c5a252fca15252d07cb23a1bdd9377a8c6f3efa01531109281ae47f841c
;;
*)
echo "install_rust: unsupported architecture $(uname -m)" >&2
exit 1
;;
esac
curl -sSLf -o /tmp/rustup-init \
"https://static.rust-lang.org/rustup/archive/1.28.2/${RUSTUP_TRIPLE}/rustup-init"
echo "${RUSTUP_SHA256} /tmp/rustup-init" | sha256sum -c -
chmod +x /tmp/rustup-init
/tmp/rustup-init -y --no-modify-path --profile minimal --default-toolchain 1.98.0
rm -f /tmp/rustup-init
echo 'export PATH="$HOME/.cargo/bin:$PATH"' >> "$BASH_ENV"
export PATH="$HOME/.cargo/bin:$PATH"
rustc --version
cargo --version
install_codecov_cli:
steps:
- run:
name: Install Codecov CLI (pinned v11.3.1)
command: |
curl -sSLf -o /tmp/codecov https://cli.codecov.io/v11.3.1/linux/codecov
curl -sSLf -o /tmp/codecov.SHA256SUM https://cli.codecov.io/v11.3.1/linux/codecov.SHA256SUM
[ "$(cat /tmp/codecov.SHA256SUM)" = "ca1d64196d2d34771084afe76ea657d581bf628e31d993ff8e52ea09cc88a56d codecov" ]
(cd /tmp && sha256sum -c codecov.SHA256SUM)
chmod +x /tmp/codecov
mkdir -p "$HOME/.local/bin"
mv /tmp/codecov "$HOME/.local/bin/codecov"
setup_litellm_enterprise_pip:
steps:
- run:
name: "Install local version of litellm-enterprise"
command: |
uv run --no-sync python -c "import litellm_enterprise; print('litellm-enterprise OK:', litellm_enterprise.__file__)"
setup_test_deps:
steps:
- checkout
- install_uv
- install_rust
- restore_cache:
keys:
- v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Install Dependencies
command: |
uv sync --frozen --all-groups --all-extras --python 3.12
- setup_litellm_enterprise_pip
- save_cache:
paths:
- ~/.cache/uv
key: v3-integration-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
skip_unless_relevant:
parameters:
category:
type: string
default: backend
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
steps:
- run:
name: "Skip job when no << parameters.category >>-relevant files changed"
command: |
export CIRCLE_PULL_REQUEST="${CIRCLE_PULL_REQUEST:-<< parameters.pull_request_url >>}"
export PATH_FILTER_BASE_BRANCH="<< parameters.base_ref >>"
[ -n "$PATH_FILTER_BASE_BRANCH" ] || unset PATH_FILTER_BASE_BRANCH
bash .circleci/scripts/path_filter.sh << parameters.category >>
start_postgres:
parameters:
db_name:
type: string
default: circle_test
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
steps:
- run:
name: Start PostgreSQL
command: |
docker run -d \
--name postgres-db \
-e POSTGRES_USER=postgres \
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
<< parameters.image >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
start_redis:
steps:
- run:
name: Start Redis
command: |
docker run -d \
--name redis-cache \
-p 6379:6379 \
redis:7-alpine@sha256:7aec734b2bb298a1d769fd8729f13b8514a41bf90fcdd1f38ec52267fbaa8ee6
- wait_for_service:
url: tcp://localhost:6379
timeout: "60"
jobs:
unit:
parameters:
tests_path:
type: string
default: tests/unit
flag:
type: string
default: unit
shards:
type: integer
default: 6
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
parallelism: << parameters.shards >>
environment:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
steps:
- setup_test_deps
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- run:
name: "Run << parameters.tests_path >> shard"
no_output_timeout: 20m
command: |
mkdir -p test-results/<< parameters.flag >>
mapfile -t files < <(find << parameters.tests_path >> -name 'test_*.py' | sort | circleci tests split --split-by=timings --timings-type=filename)
if [ "${#files[@]}" -eq 0 ]; then echo "shard ${CIRCLE_NODE_INDEX} received no << parameters.tests_path >> files; nothing to run"; exit 0; fi
set +e
uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --durations=20 -o junit_family=xunit1 --junitxml=test-results/<< parameters.flag >>/junit.xml --cov=./litellm --cov=./enterprise/litellm_enterprise --cov-report=xml:coverage.xml --cov-config=pyproject.toml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from the shard; passing"; exit 0; fi
exit "$status"
- install_codecov_cli
- run:
name: Upload coverage
when: always
command: |
[ -f coverage.xml ] || { echo "no coverage.xml produced; skipping upload"; exit 0; }
codecov upload-process --disable-search -f coverage.xml -F << parameters.flag >> -C "$CIRCLE_SHA1" -n "<< parameters.flag >>-${CIRCLE_NODE_INDEX}-${CIRCLE_BUILD_NUM}" --git-service github
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
- store_artifacts:
path: coverage.xml
documentation:
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- run:
name: Checkout litellm-docs
command: rm -rf docs/my-website && git clone --depth 1 https://github.com/BerriAI/litellm-docs.git docs/my-website
- run:
name: Run documentation validation
command: |
uv run --no-sync python ./tests/documentation_tests/test_env_keys.py
uv run --no-sync python ./tests/documentation_tests/test_router_settings.py
uv run --no-sync python ./tests/documentation_tests/test_api_docs.py
uv run --no-sync python ./tests/documentation_tests/test_circular_imports.py
integration:
parameters:
suite:
type: string
base_ref:
type: string
default: ""
pull_request_url:
type: string
default: ""
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_test_deps
- skip_unless_relevant:
base_ref: << parameters.base_ref >>
pull_request_url: << parameters.pull_request_url >>
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- start_redis
- run:
name: Run owned integration contracts
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
workflows:
tests:
when: (pipeline.event.name == "push" and pipeline.git.branch == "main") or pipeline.event.name == "api" or (pipeline.event.name == "pull_request" and (pipeline.event.github.pull_request.base.ref == "main" or pipeline.event.github.pull_request.base.ref starts-with "litellm_"))
jobs:
- unit:
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
- documentation
- integration:
name: integration-<< matrix.suite >>
matrix:
parameters:
suite: [sdk]
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>

View file

@ -50,14 +50,14 @@ def _build_tool_result_message(tool_results: Sequence[Mapping[str, object]]) ->
"""Turn executed tool results into the user message Anthropic expects."""
return AnthropicMessagesUserMessageParam(
role="user",
content=tuple(
content=[
AnthropicMessagesToolResultParam(
type="tool_result",
tool_use_id=str(result.get("tool_call_id") or ""),
content=str(result.get("result") or ""),
)
for result in tool_results
),
],
)

View file

@ -56322,6 +56322,38 @@
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"openai.gpt-6-sol": {
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_none_reasoning_effort": false,
"supports_tool_choice": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_xhigh_reasoning_effort": true,
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"global.openai.gpt-6-sol": {
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
@ -56354,6 +56386,38 @@
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"openai.gpt-6-luna": {
"input_cost_per_token": 1e-07,
"input_cost_per_token_above_272k_tokens": 2e-07,
"cache_creation_input_token_cost": 1.25e-07,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-07,
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_above_272k_tokens": 2e-08,
"output_cost_per_token": 5e-07,
"output_cost_per_token_above_272k_tokens": 7.5e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_none_reasoning_effort": false,
"supports_tool_choice": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_xhigh_reasoning_effort": true,
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"global.openai.gpt-6-luna": {
"input_cost_per_token": 1e-07,
"input_cost_per_token_above_272k_tokens": 2e-07,
@ -64072,7 +64136,7 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"source": "https://www.baseten.co/library/glm-53-fast/",
"source": "https://inference.baseten.co/v1/models",
"supported_modalities": [
"text",
"image"
@ -64358,6 +64422,24 @@
"supports_prompt_caching": true,
"supports_web_search": false
},
"openrouter/qwen/qwen3.8-max-prime": {
"input_cost_per_token": 4e-06,
"output_cost_per_token": 1.2e-05,
"cache_read_input_token_cost": 5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_video_input": true,
"supports_prompt_caching": true
},
"openrouter/deepseek/deepseek-v4-flash-0731": {
"input_cost_per_token": 4e-08,
"output_cost_per_token": 6.4e-07,
@ -72219,14 +72301,15 @@
"supports_web_search": false
},
"openrouter/stealth/space-bunny-alpha": {
"input_cost_per_token": 0,
"deprecation_date": "2098-12-31",
"input_cost_per_token": 0.0,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 0,
"source": "https://openrouter.ai/stealth/space-bunny-alpha",
"output_cost_per_token": 0.0,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
@ -72729,6 +72812,7 @@
},
"openrouter/z-ai/glm-5.3-flashx": {
"cache_read_input_token_cost": 7.5e-08,
"deprecation_date": "2098-12-31",
"input_cost_per_token": 3.7e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,

View file

@ -1,4 +1,6 @@
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, Literal
from pydantic import BaseModel
import litellm
from litellm.types.guardrails import SupportedGuardrailIntegrations
@ -8,6 +10,14 @@ from .straiker import StraikerGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
class _V3Routing(BaseModel):
api_version: Literal["v1", "v3"] | None = None
agent_ref: str | None = None
client: str | None = None
format_hint: Literal["anthropic.messages", "openai.chat"] | None = None
_OPTIONAL_INIT_FIELDS: Final = (
"timeout",
"max_retries",
@ -48,6 +58,12 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
for value in [_get_config_value(litellm_params, optional_params, field)]
if value is not None
}
routing: Final = _V3Routing.model_validate(
{
field: _get_config_value(litellm_params, optional_params, field)
for field in ("api_version", "agent_ref", "client", "format_hint")
}
)
_callback: Final = StraikerGuardrail(
api_key=api_key,
api_base=api_base if isinstance(api_base, str) else "https://api.prod.straiker.ai",
@ -55,6 +71,10 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", "straiker"),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
api_version=routing.api_version,
agent_ref=routing.agent_ref,
client=routing.client,
format_hint=routing.format_hint,
**kwargs,
)

View file

@ -1,9 +1,12 @@
from __future__ import annotations
import asyncio
import hashlib
import json
import random
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn
from urllib.parse import urlsplit
@ -12,6 +15,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.exceptions import (
BadRequestError,
GuardrailRaisedException,
@ -29,6 +33,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import SpecialProxyStrings
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
STRAIKER_WEBHOOK_SCHEMA_VERSION,
@ -43,7 +48,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
StraikerWebhookStream,
StraikerWebhookUsage,
)
from litellm.types.utils import GenericGuardrailAPIInputs
from litellm.types.utils import CallTypes, GenericGuardrailAPIInputs, ModelResponse, TextCompletionResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -54,6 +59,93 @@ DEFAULT_BLOCK_MESSAGE: Final = "Content violates policy"
DEFAULT_API_BASE: Final = "https://api.prod.straiker.ai"
DEFAULT_MAX_PAYLOAD_BYTES: Final = 524288
WEBHOOK_PATH: Final = "/api/v1/detect/webhook"
V3_DETECT_PATH: Final = "/api/v3/detect"
V3_KEY_PREFIX: Final = "sk_agt_"
V3_SESSION_HEADER: Final = "x-claude-code-session-id"
V3_CLIENT_HEADER: Final = "x-s6r-client"
V3_FORMAT_HEADER: Final = "x-s6r-format"
# (User-Agent prefix, Straiker client value, display name). Straiker recognises a coding agent
# from the system prompt of its main turns only; Claude Code's title and topic sidecars carry
# other prompts and would split the session across two agents. The User-Agent is on every call.
_V3_CLIENT_BY_USER_AGENT: Final = (("claude-cli/", "claude", "Claude"),)
V3_GATEWAY_NAME: Final = "LiteLLM"
V3_DERIVED_SESSION_PREFIX: Final = "litellm-"
V3_AGENT_HEADER: Final = "x-s6r-agent"
V3_RESPONSE_PHASE: Final = "response-sync"
V3_BLOCK_DECISIONS: Final = frozenset({"block", "deny"})
V3_BLOCKED_TURN_MEMORY: Final = 10_000
V3_BLOCKED_TURN_TTL_SECONDS: Final = 24 * 60 * 60
# An allowlist: the hook's request dict merges the client body with proxy state (`deployment`
# carries the resolved credential), so only fields named here are relayed.
_V3_PROVIDER_BODY_KEYS: Final = frozenset(
{
"model",
"messages",
"tools",
"tool_choice",
"functions",
"function_call",
"temperature",
"top_p",
"n",
"stream",
"stream_options",
"stop",
"max_tokens",
"max_completion_tokens",
"presence_penalty",
"frequency_penalty",
"logit_bias",
"user",
"response_format",
"seed",
"logprobs",
"top_logprobs",
"parallel_tool_calls",
"reasoning_effort",
"modalities",
"audio",
"prediction",
"store",
"service_tier",
"web_search_options",
"prompt",
"suffix",
"echo",
"best_of",
"system",
"stop_sequences",
"top_k",
"thinking",
"container",
"mcp_servers",
"context_management",
"output_format",
"input",
"instructions",
"previous_response_id",
"truncation",
"text",
"include",
"reasoning",
"max_output_tokens",
"background",
"conversation",
"session_id",
}
)
# The scrub of these is one level deep on purpose: a function schema that defines a `token` or
# `headers` property lives under `function.parameters` and must be relayed as sent.
_V3_CREDENTIAL_FIELDS: Final = frozenset({"authorization_token", "authorization", "headers"})
_V3_REDACTED_VALUE: Final = "[redacted]"
_V3_REDACTED_KEYS: Final = frozenset({"tools", "mcp_servers"})
_V3_IDENTITY_METADATA_KEYS: Final = (
"user_api_key_end_user_id",
"user_api_key_user_email",
"user_api_key_user_id",
"user_api_key_alias",
"user_api_key_team_id",
)
RETRY_STATUS: Final = frozenset({408, 429, 500, 502, 503, 504})
UNREACHABLE_STATUS: Final = frozenset({502, 503, 504})
_APPLICATION_METADATA_KEYS: Final = frozenset({"agent_id", "app_name"})
@ -65,13 +157,29 @@ _JSON_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
class _WebhookFailure:
message: str
is_unreachable: bool
retryable: bool = False
def _status_failure(status: int, text: str) -> _WebhookFailure:
return _WebhookFailure(
f"HTTP {status}: {text[:200]}",
is_unreachable=status in UNREACHABLE_STATUS,
retryable=status in RETRY_STATUS,
)
def _error_response_text(response: httpx.Response) -> str:
try:
return response.text
except Exception: # noqa: BLE001 # a masked response may carry no body
return ""
def _as_dict(value: object) -> dict:
return value if isinstance(value, dict) else {}
def _merged_metadata(request_data: dict) -> dict:
def _merged_metadata(request_data: Mapping[str, object]) -> dict:
return {
**_as_dict(request_data.get("metadata")),
**_as_dict(request_data.get("litellm_metadata")),
@ -268,6 +376,478 @@ def _is_streamed_request(request_data: dict) -> bool:
return body.get("stream") is True
# What the proxy stamps on a master-key call in place of a person. Sent onward, either
# would be recorded as an identity and every master-key turn filed under it.
_PLACEHOLDER_IDENTITIES: Final = frozenset({SpecialProxyStrings.default_user_id.value, "litellm_proxy_master_key"})
def _real_identity(value: object) -> str | None:
"""LiteLLM's proxy-admin placeholders are not a person."""
identity: Final = _as_optional_str(value)
return None if identity in _PLACEHOLDER_IDENTITIES else identity
def _request_header(request_data: Mapping[str, object], name: str | None) -> str | None:
"""A header from the inbound request, when LiteLLM kept it on the request data."""
if not name:
return None
proxy_request: Final = request_data.get("proxy_server_request")
headers: Final = proxy_request.get("headers") if isinstance(proxy_request, Mapping) else None
if not isinstance(headers, Mapping):
return None
wanted: Final = name.lower()
for key, value in headers.items():
if str(key).lower() == wanted and isinstance(value, str) and value.strip():
return value.strip()
return None
def _frozen(pairs: Iterable[tuple[str, object]]) -> Mapping[str, object]:
return MappingProxyType(dict(pairs))
def _json_default(value: object) -> object:
if isinstance(value, Mapping):
return dict(value) # mutable-ok: the JSON encoder needs a dict view of a frozen mapping
return str(value)
def _v3_identity_metadata(request_data: Mapping[str, object]) -> Mapping[str, str]:
"""The proxy-resolved identity fields, and only those, for the relayed body."""
merged: Final = _merged_metadata(request_data)
return MappingProxyType(
{key: value for key in _V3_IDENTITY_METADATA_KEYS if (value := _real_identity(merged.get(key)))}
)
def _v3_request_body(request_data: Mapping[str, object]) -> Mapping[str, object]:
"""The provider body LiteLLM received, stripped of everything the proxy added.
The hook sees the client's request merged with proxy bookkeeping: logging objects,
the resolved key, the inbound headers. Only the provider body is Straiker's to read,
and the client's Authorization header must not travel. Identity survives as the
metadata subset the Straiker LiteLLM adapter reads.
"""
identity: Final = _v3_identity_metadata(request_data)
turns: Final = (
_v3_prompt_as_messages(request_data.get("prompt"))
if _v3_text_completion_route(request_data) and "messages" not in request_data
else None
)
provider: Final = (
(key, _v3_without_credentials(value) if key in _V3_REDACTED_KEYS else value)
for key, value in request_data.items()
if key in _V3_PROVIDER_BODY_KEYS and not (turns is not None and key == "prompt")
)
prompt_turns: Final = (("messages", turns),) if turns is not None else ()
return _frozen((*provider, *prompt_turns, *((("metadata", identity),) if identity else ())))
def _v3_without_credentials(entries: object) -> object:
if not isinstance(entries, (list, tuple)):
return entries
return tuple(
_frozen(
(str(key), _V3_REDACTED_VALUE if str(key).lower() in _V3_CREDENTIAL_FIELDS else item)
for key, item in entry.items()
)
if isinstance(entry, Mapping)
else entry
for entry in entries
)
def _v3_route_is(request_data: Mapping[str, object], call_type: CallTypes) -> bool:
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
route: Final = _merged_metadata(request_data).get("user_api_key_request_route")
if not isinstance(route, str) or not route:
return False
return call_type in (get_call_types_for_route(route) or ())
def _v3_anthropic_messages_route(request_data: Mapping[str, object]) -> bool:
return _v3_route_is(request_data, CallTypes.anthropic_messages)
def _v3_text_completion_route(request_data: Mapping[str, object]) -> bool:
return _v3_route_is(request_data, CallTypes.text_completion)
def _v3_is_token_list(value: object) -> bool:
return (
isinstance(value, (list, tuple))
and bool(value)
and all(isinstance(token, int) and not isinstance(token, bool) for token in value)
)
def _v3_decode_tokens(tokens: Iterable[object]) -> str | None:
ids: Final = [token for token in tokens if isinstance(token, int)] # mutable-ok: tiktoken decodes a list
try:
import tiktoken
return tiktoken.encoding_for_model("text-davinci-003").decode(ids)
except Exception: # noqa: BLE001 # no tokenizer available: the raw prompt is relayed instead
return None
def _v3_prompt_texts(prompt: object) -> tuple[str, ...] | None:
"""The text the model receives for a completions `prompt`, in the proxy's own terms.
LiteLLM accepts a string, a list of strings, a list of token ids, or a list of token-id
lists, and decodes token ids with the text-davinci-003 tokenizer before calling the model.
The same decoding here means Straiker screens what the model gets. None when the prompt
is a shape this cannot render, so the caller relays it untouched rather than screening
something else.
"""
if isinstance(prompt, str):
return (prompt,)
if not isinstance(prompt, (list, tuple)) or not prompt:
return None
if all(isinstance(item, str) for item in prompt):
return tuple(str(item) for item in prompt)
if _v3_is_token_list(prompt):
decoded: Final = _v3_decode_tokens(prompt)
return (decoded,) if decoded is not None else None
if all(_v3_is_token_list(item) for item in prompt):
decoded_each: Final = tuple(_v3_decode_tokens(item) for item in prompt)
return None if any(text is None for text in decoded_each) else tuple(text or "" for text in decoded_each)
return None
def _v3_prompt_as_messages(prompt: object) -> tuple[Mapping[str, object], ...] | None:
texts: Final = _v3_prompt_texts(prompt)
if texts is None:
return None
return tuple(_frozen((("role", "user"), ("content", text))) for text in texts)
def _v3_answer(request_data: Mapping[str, object], model: str | None) -> Mapping[str, object] | None:
"""The answer in the API shape the client spoke, which is what a relay forwards.
On a streamed Messages call the proxy rebuilds the answer as a chat completion before
the hook runs. Straiker's coding-agent reader parses a Messages answer, so a Claude Code
turn sent as a chat completion scores nothing; the proxy's own adapter turns it back.
"""
response: Final = request_data.get("response")
if isinstance(response, TextCompletionResponse):
return _v3_text_completion_as_chat(response)
if not isinstance(response, ModelResponse) or not _v3_anthropic_messages_route(request_data):
return _jsonable_dict(response)
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
)
translated: Final = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(response=response)
re_keyed: Final = dict(translated, model=response.model or model) # mutable-ok: adapter TypedDict re-keyed
return _jsonable_dict(re_keyed)
def _v3_text_completion_as_chat(response: TextCompletionResponse) -> Mapping[str, object]:
"""A legacy completion answer in the chat shape the platform scores.
Straiker has no reader for a `text_completion` answer on a gateway: the request phase
of a /v1/completions call is scored, the response phase is refused. A completion is one
user turn and one assistant turn, so both phases are presented as that exchange.
"""
choices: Final = tuple(
_frozen(
(
("index", index),
("finish_reason", getattr(choice, "finish_reason", None)),
("message", _frozen((("role", "assistant"), ("content", getattr(choice, "text", "") or "")))),
)
)
for index, choice in enumerate(response.choices)
)
usage: Final = _jsonable_dict(getattr(response, "usage", None))
return _frozen(
(
("id", response.id),
("object", "chat.completion"),
("created", response.created),
("model", response.model),
("choices", choices),
*((("usage", usage),) if usage else ()),
)
)
def _v3_answer_json(
inputs: GenericGuardrailAPIInputs, request_data: Mapping[str, object], model: str | None
) -> str | None:
"""The model's answer as the raw response body Straiker parses on the response phase.
The real response object carries tool calls, which a coding-agent turn is scored on,
so it is preferred. A streamed answer reaches the hook already assembled into texts,
and those become a minimal chat completion so the answer is still scored.
"""
response: Final = _v3_answer(request_data, model)
if response:
return json.dumps(response, default=_json_default)
texts: Final = tuple(t for t in (inputs.get("texts") or []) if t)
if not texts:
return None
message: Final = _frozen((("role", "assistant"), ("content", "\n".join(texts))))
choice: Final = _frozen((("index", 0), ("finish_reason", "stop"), ("message", message)))
return json.dumps(_frozen((("object", "chat.completion"), ("choices", (choice,)))), default=_json_default)
def _v3_payload(
envelope: StraikerWebhookRequest,
inputs: GenericGuardrailAPIInputs,
request_data: Mapping[str, object],
input_type: Literal["request", "response"],
) -> Mapping[str, object]:
"""The /api/v3/detect body for one phase of a turn, the unified Kong plugin's contract.
Request phase: the provider body itself. Response phase: the answer beside the request
it answers, `{straiker_phase, sse, model, request}`, which is how Straiker classifies a
tool call the model just made. Straiker parses either and derives prompt, answer, agent
and archetype from the traffic; nothing is pre-digested here. Identity and session ride
on both phases the way Kong sends them.
"""
context: Final = envelope.context
request_body: Final = _v3_request_body(request_data)
answer_json: Final = _v3_answer_json(inputs, request_data, context.model) if input_type == "response" else None
phase: Final = (
tuple(request_body.items())
if input_type == "request"
else (
("straiker_phase", V3_RESPONSE_PHASE),
("model", context.model),
("request", request_body),
*((("sse", answer_json),) if answer_json is not None else ()),
)
)
session: Final = _v3_session_id(envelope, request_data, request_body)
user: Final = _v3_user(envelope)
return _frozen(
(
*phase,
*((("session_id", session),) if session else ()),
*(
(("original", _frozen((("processed", _frozen((("Meta", _frozen((("user", user),))),))),))),)
if user
else ()
),
)
)
def _v3_conversation_prefixes(request_body: Mapping[str, object]) -> tuple[str, ...]:
"""A fingerprint of the conversation after each of its messages, first to last.
The last one names the conversation as sent; the earlier ones let a request that
carries a blocked exchange as its history be recognised, not only an exact resend.
A `prompt` or a string `input` has one fingerprint.
"""
messages: Final = _v3_messages(request_body)
if messages:
digest: Final = hashlib.sha256()
def after(message: object) -> str:
digest.update(json.dumps(message, sort_keys=True, default=str).encode("utf-8"))
digest.update(b"\x1e")
return digest.copy().hexdigest()
return tuple(after(message) for message in messages)
plain: Final = request_body.get("input") if "input" in request_body else request_body.get("prompt")
if plain is None:
return ()
return (hashlib.sha256(json.dumps(plain, sort_keys=True, default=str).encode("utf-8")).hexdigest(),)
def _v3_session_id(
envelope: StraikerWebhookRequest,
request_data: Mapping[str, object],
request_body: Mapping[str, object],
) -> str | None:
"""A stable id for the conversation, in Kong's order of precedence.
Claude Code names its session on the wire and that wins. Then the session LiteLLM
resolved from its own metadata. Then, for a conversation that states none, a hash of
the principal, the system prompt and the first message: a chat client replays the
whole conversation on every turn, so that triple is constant for its lifetime and
groups the turns. A fresh synthetic id per request would group nothing.
The principal is in the hash because Straiker skips turns it has already scored for a
session. Two users who open with the same words are two conversations; hashed on the
words alone they shared one session, and the second user's copy of an attack came
back as a replay, unscored and allowed (measured 2026-09-20).
"""
supplied: Final = _request_header(request_data, V3_SESSION_HEADER)
if supplied:
return supplied
if envelope.context.session_id:
return envelope.context.session_id
conversation: Final = f"{_v3_system_text(request_body) or ''}\0{_v3_first_message_text(request_body)}"
if conversation == "\0":
return None
seed: Final = f"{_v3_user(envelope) or ''}\0{conversation}"
return V3_DERIVED_SESSION_PREFIX + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:32]
_V3_PREAMBLE_ROLES: Final = frozenset({"system", "developer"})
def _v3_message_text(message: object) -> str:
"""Every text block of a message, so a turn that opens with an image or a document still
seeds on what the user wrote."""
content: Final = message.get("content") if isinstance(message, Mapping) else None
if isinstance(content, str):
return content
if isinstance(content, (list, tuple)):
return "\n".join(
str(block["text"]) for block in content if isinstance(block, Mapping) and isinstance(block.get("text"), str)
)
return ""
def _v3_messages(request_body: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
messages: Final = request_body.get("messages") or request_body.get("input")
if isinstance(messages, (list, tuple)):
return tuple(message for message in messages if isinstance(message, Mapping))
return ()
def _v3_system_text(request_body: Mapping[str, object]) -> str | None:
"""The preamble, wherever the API puts it: Anthropic's `system`, the Responses API's
`instructions`, or the leading system or developer message of an OpenAI chat body."""
system: Final = request_body.get("system")
if isinstance(system, str):
return system
if system is not None:
return json.dumps(system, default=str)
instructions: Final = request_body.get("instructions")
if isinstance(instructions, str):
return instructions
preamble: Final = next((m for m in _v3_messages(request_body) if m.get("role") in _V3_PREAMBLE_ROLES), None)
return _v3_message_text(preamble) if preamble is not None else None
def _v3_first_message_text(request_body: Mapping[str, object]) -> str:
"""What the user first said: the first `user` message, never the system prompt that an
OpenAI chat body carries as `messages[0]`, else a Responses `input` string, else `prompt`."""
first_user: Final = next((m for m in _v3_messages(request_body) if m.get("role") == "user"), None)
if first_user is not None:
return _v3_message_text(first_user)
plain: Final = (
request_body.get("input") if isinstance(request_body.get("input"), str) else request_body.get("prompt")
)
return plain if isinstance(plain, str) else ""
def _v3_user(envelope: StraikerWebhookRequest) -> str | None:
"""Who is asking: the key's own user first, then the end user the request named.
The key is the authenticated principal, the way a Kong consumer is, so a per-user key
names the person even when the client packs something else into the body. Claude Code
packs a hashed account-and-session token into `metadata.user_id`, which is what the end
user resolves to when nothing better is set; it is a session, not a person, and only
surfaces when the key names nobody. A master-key call resolves to LiteLLM's
`default_user_id`; sent as an identity it would become one.
"""
identity: Final = envelope.identity
for candidate in (identity.litellm_user_email, identity.litellm_user_id, identity.end_user_id):
real = _real_identity(candidate)
if real:
return real
return None
def _v3_client_from_user_agent(request_data: Mapping[str, object]) -> tuple[str, str] | None:
"""`(client, agent name)` for a User-Agent this gateway recognises, else None."""
user_agent: Final = (_request_header(request_data, "user-agent") or "").lower()
return next(
(
(client, f"{display} ({V3_GATEWAY_NAME})")
for prefix, client, display in _V3_CLIENT_BY_USER_AGENT
if user_agent.startswith(prefix)
),
None,
)
def _v3_headers(
request_data: Mapping[str, object],
agent_ref: str | None = None,
client: str | None = None,
format_hint: str | None = None,
) -> Mapping[str, str]:
"""Per-call routing hints, the unified Kong plugin's set. All optional.
`x-s6r-agent` names ONE application when a gateway fronts several: the route's
`agent_ref`, else the caller's own header, else the agent this gateway names from the
User-Agent. The operator's value comes first because the header is caller-supplied, and
honouring it over a pinned route would let any key file its traffic under another
application's agent and controls. `x-s6r-client` is the route's `client` config, else
the client the User-Agent names. `x-s6r-format` comes from config alone. Claude Code's own session header is
forwarded when the client sent it, which is how a coding session groups the way the
native hook would.
"""
session: Final = _request_header(request_data, V3_SESSION_HEADER)
recognised: Final = _v3_client_from_user_agent(request_data)
agent: Final = (
agent_ref or _request_header(request_data, V3_AGENT_HEADER) or (recognised[1] if recognised else None)
)
named_client: Final = client or (recognised[0] if recognised else None)
candidates: Final = (
(V3_SESSION_HEADER, session),
(V3_AGENT_HEADER, agent),
(V3_CLIENT_HEADER, named_client),
(V3_FORMAT_HEADER, format_hint),
)
return MappingProxyType({name: value for name, value in candidates if value})
def _v3_decision(body: Mapping[str, object]) -> tuple[str | None, Mapping[str, object]]:
"""``(decision, verdict)``: the enforceable decision and the object carrying it.
Straiker answers in two envelopes. A relayed body gets the hook contract,
`hookSpecificOutput.permissionDecision`, with the flat fields nested under `straiker`;
a flat call answers `action` at the top level. Reading only one of them would silently
make block mode a no-op on the other.
"""
nested: Final = body.get("straiker")
verdict: Final = nested if isinstance(nested, Mapping) else body
hook: Final = body.get("hookSpecificOutput")
decision: Final = hook.get("permissionDecision") if isinstance(hook, Mapping) else None
if isinstance(decision, str) and decision:
return decision.lower(), verdict
action: Final = verdict.get("action")
return (action.lower() if isinstance(action, str) and action else None), verdict
def _v3_response(body: Mapping[str, object]) -> StraikerWebhookResponse:
"""Map a v3 verdict onto the action the guardrail already acts on.
A detect-mode control fires into `controls` without changing the decision, so it
correctly reads NONE. `blocked_by` is the block-mode subset and is honoured even if a
build answers it without flipping the decision.
"""
decision, verdict = _v3_decision(body)
raw_blocked_by: Final = verdict.get("blocked_by")
blocked_by: Final = tuple(sorted(str(c) for c in raw_blocked_by)) if isinstance(raw_blocked_by, list) else ()
blocked: Final = decision in V3_BLOCK_DECISIONS or bool(blocked_by)
stated: Final = (verdict.get("block_message"), verdict.get("deny_reason"), body.get("stopReason"))
reason: Final = (
next(
(text.strip() for text in stated if isinstance(text, str) and text.strip()),
f"Straiker blocked this turn: {', '.join(blocked_by) or 'policy'}",
)
if blocked
else None
)
return StraikerWebhookResponse(
action="BLOCKED" if blocked else "NONE",
blocked_reason=reason,
blocked_by=blocked_by,
turnId=_as_optional_str(verdict.get("turn_id")) or _as_optional_str(body.get("turn_id")),
)
class StraikerGuardrail(CustomGuardrail):
@staticmethod
def get_config_model() -> type[GuardrailConfigModel]:
@ -284,6 +864,10 @@ class StraikerGuardrail(CustomGuardrail):
self,
api_key: str,
api_base: str = DEFAULT_API_BASE,
api_version: Literal["v1", "v3"] | None = None,
agent_ref: str | None = None,
client: str | None = None,
format_hint: Literal["anthropic.messages", "openai.chat"] | None = None,
source: str = "LiteLLM Gateway",
timeout: float = 5.0,
max_retries: int = 2,
@ -302,9 +886,28 @@ class StraikerGuardrail(CustomGuardrail):
raise ValueError("api_key must be non-empty")
if unreachable_fallback not in ("fail_open", "fail_closed"):
raise ValueError(f"unreachable_fallback must be 'fail_open' or 'fail_closed'; got {unreachable_fallback!r}")
if api_version is None:
# The key names the platform: a v3 integration key cannot call v1 and a v1
# collection key cannot call v3, so an unset version follows the key.
api_version = "v3" if api_key.startswith(V3_KEY_PREFIX) else "v1"
if api_version not in ("v1", "v3"):
raise ValueError(f"api_version must be 'v1' or 'v3'; got {api_version!r}")
self.api_key = api_key
self.api_base = api_base.rstrip("/")
self.api_version = api_version
self.agent_ref = _as_optional_str(agent_ref)
self.client = _as_optional_str(client)
if format_hint is not None and format_hint not in ("anthropic.messages", "openai.chat"):
raise ValueError(f"format_hint must be 'anthropic.messages' or 'openai.chat'; got {format_hint!r}")
self.format_hint = format_hint
# Blocked conversations by session, so a resend or a conversation grown past a blocked
# turn is blocked again here: Straiker de-duplicates turns it has already scored per
# session and answers a replay `allow`, whatever the original verdict was (measured
# 2026-09-20). Per process; a replica that did not see the block asks Straiker.
self._v3_blocked_turns = InMemoryCache(
max_size_in_memory=V3_BLOCKED_TURN_MEMORY, default_ttl=V3_BLOCKED_TURN_TTL_SECONDS
)
self.source = source
self.timeout = float(timeout)
self.max_retries = max(0, int(max_retries))
@ -330,17 +933,18 @@ class StraikerGuardrail(CustomGuardrail):
self.configured_modes = _configured_modes(self.event_hook)
def _webhook_url(self) -> str:
return f"{self.api_base}{WEBHOOK_PATH}"
return f"{self.api_base}{V3_DETECT_PATH if self.api_version == 'v3' else WEBHOOK_PATH}"
def _headers(self) -> dict[str, str]:
reserved: Final = {"authorization", "content-type", "x-straiker-webhook-format"}
extra: Final = {k: v for k, v in self.custom_headers.items() if k.lower() not in reserved}
return {
headers: Final = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
"X-Straiker-Webhook-Format": "litellm",
**extra,
}
if self.api_version != "v3":
headers["X-Straiker-Webhook-Format"] = "litellm"
return {**headers, **extra}
def _build_application(self, request_data: dict) -> StraikerWebhookApplication:
meta: Final = _merged_metadata(request_data)
@ -417,9 +1021,11 @@ class StraikerGuardrail(CustomGuardrail):
metadata=_build_webhook_metadata(request_data, self.default_metadata),
)
async def _post_webhook(self, payload: dict) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
async def _post_webhook(
self, payload: Mapping[str, object], headers: Mapping[str, str] | None = None
) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
try:
body = json.dumps(payload).encode("utf-8")
body: Final = json.dumps(payload, default=_json_default).encode("utf-8")
except (TypeError, ValueError, OverflowError) as error:
return None, _WebhookFailure(f"request serialization failed: {error}", is_unreachable=False)
body_bytes: Final = len(body)
@ -430,7 +1036,7 @@ class StraikerGuardrail(CustomGuardrail):
)
url: Final = self._webhook_url()
headers: Final = self._headers()
merged_headers: Final = {**self._headers(), **(headers or {})}
attempts: Final = self.max_retries + 1
last_failure: _WebhookFailure | None = None
@ -443,48 +1049,58 @@ class StraikerGuardrail(CustomGuardrail):
"bytes": body_bytes,
"payload": payload,
},
default=str,
default=_json_default,
)
)
for attempt in range(attempts):
try:
resp = await self.async_handler.post(url, content=body, headers=headers, timeout=self.timeout)
if resp.status_code == 200:
try:
body = resp.json()
parsed = StraikerWebhookResponse.model_validate(body)
except (ValidationError, json.JSONDecodeError) as ve:
return None, _WebhookFailure(f"invalid response schema: {ve}", is_unreachable=False)
if self.verbose:
verbose_proxy_logger.info(
json.dumps(
{
"event": "straiker.webhook_response",
"status_code": resp.status_code,
"body": body,
},
default=str,
)
)
return parsed, None
last_failure = _WebhookFailure(
f"HTTP {resp.status_code}: {resp.text[:200]}",
is_unreachable=resp.status_code in UNREACHABLE_STATUS,
)
if resp.status_code not in RETRY_STATUS:
return None, last_failure
except (httpx.RequestError, asyncio.TimeoutError, Timeout) as e:
last_failure = _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=True)
except (json.JSONDecodeError, TypeError, ValueError) as e:
return None, _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=False)
parsed, last_failure = await self._attempt(url, body, merged_headers)
if last_failure is None or not last_failure.retryable:
return parsed, last_failure
if attempt < attempts - 1:
backoff = min(self.initial_backoff * (2**attempt), self.max_backoff)
await asyncio.sleep(random.uniform(0, backoff))
return None, last_failure or _WebhookFailure("unknown error", is_unreachable=True)
async def _attempt(
self, url: str, body: bytes, headers: dict[str, str]
) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
try:
resp: Final = await self.async_handler.post(url, content=body, headers=headers, timeout=self.timeout)
except httpx.HTTPStatusError as status_error:
return None, _status_failure(status_error.response.status_code, _error_response_text(status_error.response))
except (httpx.RequestError, asyncio.TimeoutError, Timeout) as e:
return None, _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=True, retryable=True)
except (json.JSONDecodeError, TypeError, ValueError) as e:
return None, _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=False)
if resp is None:
return None, _WebhookFailure("no response", is_unreachable=True, retryable=True)
if resp.status_code == 200:
return self._parse_verdict(resp)
return None, _status_failure(resp.status_code, resp.text)
def _parse_verdict(self, resp: httpx.Response) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
try:
body: Final = resp.json()
if not isinstance(body, Mapping):
return None, _WebhookFailure(
f"invalid response schema: expected an object, got {type(body).__name__}", is_unreachable=False
)
parsed: Final = (
_v3_response(body) if self.api_version == "v3" else StraikerWebhookResponse.model_validate(body)
)
except (ValidationError, json.JSONDecodeError) as ve:
return None, _WebhookFailure(f"invalid response schema: {ve}", is_unreachable=False)
if self.verbose:
verbose_proxy_logger.info(
json.dumps(
{"event": "straiker.webhook_response", "status_code": resp.status_code, "body": body},
default=_json_default,
)
)
return parsed, None
def _record(
self,
*,
@ -519,7 +1135,7 @@ class StraikerGuardrail(CustomGuardrail):
"error": error,
"fail_open": fail_open,
},
default=str,
default=_json_default,
)
)
if fail_open:
@ -564,6 +1180,76 @@ class StraikerGuardrail(CustomGuardrail):
return_inputs["texts"] = parsed.texts
return return_inputs
async def _apply_v3(
self,
*,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: LiteLLMLoggingObj | None,
) -> GenericGuardrailAPIInputs:
"""One phase of a turn against /api/v3/detect: relay, read the decision, enforce."""
try:
envelope: Final = self._build_envelope(
inputs=inputs,
request_data=request_data,
input_type=input_type,
logging_obj=logging_obj,
)
payload: Final = _v3_payload(envelope, inputs, request_data, input_type)
headers: Final = _v3_headers(request_data, self.agent_ref, self.client, self.format_hint)
request_body: Final = _v3_request_body(request_data)
# The memory is scoped by the session, else by the principal; a request that has
# neither is never remembered, so no two callers can share a block.
scope: Final = _v3_session_id(envelope, request_data, request_body) or _v3_user(envelope) or ""
prefixes: Final = _v3_conversation_prefixes(request_body) if scope else ()
except (ValidationError, TypeError, ValueError) as error:
return self._fail(
inputs=inputs,
request_data=request_data,
input_type=input_type,
error=str(error),
is_unreachable=False,
)
replayed: Final = self._v3_replayed_block(scope, prefixes) if input_type == "request" else None
if replayed is not None:
self._block(request_data=request_data, input_type=input_type, message=replayed, blocked_content=True)
parsed, failure = await self._post_webhook(payload, headers)
if failure is not None or parsed is None:
return self._fail(
inputs=inputs,
request_data=request_data,
input_type=input_type,
error=failure.message if failure is not None else "empty response from Straiker",
is_unreachable=failure.is_unreachable if failure is not None else False,
)
self._record(request_data=request_data, logging_obj=logging_obj, parsed=parsed)
if parsed.action == "BLOCKED":
message: Final = parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE
# Only a block that names a control is remembered. The same words are the same
# attack tomorrow, but a block that comes from state -- an engaged kill switch,
# a governance action -- is lifted by an administrator, and a remembered copy
# would keep refusing a conversation the platform now allows.
if prefixes and parsed.blocked_by:
self._v3_blocked_turns.set_cache(f"{scope}\0{prefixes[-1]}", message)
self._block(request_data=request_data, input_type=input_type, message=message, blocked_content=True)
return inputs
def _v3_replayed_block(self, scope: str, prefixes: tuple[str, ...]) -> str | None:
"""The block message a conversation already earned, when this request repeats or
extends a conversation this process blocked in the same scope (session or principal)."""
for prefix in prefixes:
message: str | None = self._v3_blocked_turns.get_cache(f"{scope}\0{prefix}")
if message is not None:
if self.verbose:
verbose_proxy_logger.info(
json.dumps({"event": "straiker.replay_blocked", "scope": scope, "prefix": prefix})
)
return message
return None
@log_guardrail_information
async def apply_guardrail(
self,
@ -572,6 +1258,10 @@ class StraikerGuardrail(CustomGuardrail):
input_type: Literal["request", "response"],
logging_obj: LiteLLMLoggingObj | None = None,
) -> GenericGuardrailAPIInputs:
if self.api_version == "v3":
return await self._apply_v3(
inputs=inputs, request_data=request_data, input_type=input_type, logging_obj=logging_obj
)
try:
envelope: Final = self._build_envelope(
inputs=inputs,

View file

@ -83,6 +83,9 @@ class StraikerWebhookResponse(BaseModel):
action: StraikerWebhookAction = "NONE"
blocked_reason: str | None = None
#: The controls that blocked this turn, when the platform names them. Empty for a block
#: that comes from state rather than content, such as an engaged kill switch.
blocked_by: tuple[str, ...] = ()
texts: list[str] | None = None
schema_version: str | None = None
turn_id: str | None = Field(default=None, alias="turnId")
@ -125,6 +128,36 @@ class StraikerGuardrailConfigModelOptionalParams(BaseModel):
gt=0,
description="Maximum serialized webhook payload size sent to Straiker.",
)
api_version: Literal["v1", "v3"] | None = Field(
default=None,
description=(
"Straiker detect API the gateway calls. 'v1' posts the structured webhook envelope "
"to /api/v1/detect/webhook (legacy Defend, UUID collection key). 'v3' relays the "
"provider request and response to /api/v3/detect, the v3 platform's only detect "
"route, which accepts only an sk_agt_ integration key. Unset: chosen from the key "
"prefix, so a v3 key needs no extra configuration."
),
)
agent_ref: str | None = Field(
default=None,
description=(
"v3 only. Names the Straiker agent this route's traffic belongs to when one gateway "
"fronts several applications, sent as x-s6r-agent. A client-supplied x-s6r-agent header "
"wins. Names ONE agent, never a kind of agent: Straiker keys per-agent state on it, so "
"sharing a value across applications merges them into one agent."
),
)
client: str | None = Field(
default=None,
description=(
"v3 only. Optional x-s6r-client routing hint. Leave unset on a shared gateway; set it on a "
"route that serves a single application."
),
)
format_hint: Literal["anthropic.messages", "openai.chat"] | None = Field(
default=None,
description="v3 only. Optional x-s6r-format hint. Only breaks the messages-array tie between formats.",
)
custom_headers: dict[str, str] | None = Field(
default=None,
description="Additional HTTP headers sent to Straiker, excluding Authorization and the webhook-format header.",

View file

@ -56322,6 +56322,38 @@
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"openai.gpt-6-sol": {
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_272k_tokens": 1.5e-05,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_none_reasoning_effort": false,
"supports_tool_choice": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_xhigh_reasoning_effort": true,
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"global.openai.gpt-6-sol": {
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_272k_tokens": 4e-06,
@ -56354,6 +56386,38 @@
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"openai.gpt-6-luna": {
"input_cost_per_token": 1e-07,
"input_cost_per_token_above_272k_tokens": 2e-07,
"cache_creation_input_token_cost": 1.25e-07,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-07,
"cache_read_input_token_cost": 1e-08,
"cache_read_input_token_cost_above_272k_tokens": 2e-08,
"output_cost_per_token": 5e-07,
"output_cost_per_token_above_272k_tokens": 7.5e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_none_reasoning_effort": false,
"supports_tool_choice": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_xhigh_reasoning_effort": true,
"supports_vision": true,
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"global.openai.gpt-6-luna": {
"input_cost_per_token": 1e-07,
"input_cost_per_token_above_272k_tokens": 2e-07,
@ -64072,7 +64136,7 @@
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 6.6e-06,
"source": "https://www.baseten.co/library/glm-53-fast/",
"source": "https://inference.baseten.co/v1/models",
"supported_modalities": [
"text",
"image"
@ -64358,6 +64422,24 @@
"supports_prompt_caching": true,
"supports_web_search": false
},
"openrouter/qwen/qwen3.8-max-prime": {
"input_cost_per_token": 4e-06,
"output_cost_per_token": 1.2e-05,
"cache_read_input_token_cost": 5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_video_input": true,
"supports_prompt_caching": true
},
"openrouter/deepseek/deepseek-v4-flash-0731": {
"input_cost_per_token": 4e-08,
"output_cost_per_token": 6.4e-07,
@ -72219,14 +72301,15 @@
"supports_web_search": false
},
"openrouter/stealth/space-bunny-alpha": {
"input_cost_per_token": 0,
"deprecation_date": "2098-12-31",
"input_cost_per_token": 0.0,
"litellm_provider": "openrouter",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 0,
"source": "https://openrouter.ai/stealth/space-bunny-alpha",
"output_cost_per_token": 0.0,
"source": "https://openrouter.ai/api/v1/models",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
@ -72729,6 +72812,7 @@
},
"openrouter/z-ai/glm-5.3-flashx": {
"cache_read_input_token_cost": 7.5e-08,
"deprecation_date": "2098-12-31",
"input_cost_per_token": 3.7e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,

View file

@ -1,3 +1,9 @@
[
"tests/e2e/ui/tests/integrationCritical/projectDetachment.spec.ts::project creation and explicit detachment preserve saved scope and restore serving"
"tests/e2e/ui/tests/integrationCritical/projectDetachment.spec.ts::project creation and explicit detachment preserve saved scope and restore serving",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::per-user MCP env var stays updatable and clearable from the card after it is set",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::cancelling the clear confirmation keeps the stored value and sends no delete",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::pressing Enter on Update opens the credentials modal instead of the server editor",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server with two per-user variables reports the remaining gap until both are saved",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server without per-user variables shows no credential row",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::clearing credentials for a server deleted underneath the modal reports the failure without losing the page"
]

View file

@ -0,0 +1,416 @@
import {
test,
expect,
APIRequestContext,
Locator,
Page as PlaywrightPage,
} from "@playwright/test";
import { randomUUID } from "node:crypto";
import { Page } from "../../fixtures/pages";
import { navigateToPage } from "../../helpers/navigation";
import { captureRequestBody } from "../../helpers/roundTrip";
const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master";
const headers = { Authorization: `Bearer ${master}` };
const TOKEN = "USER_TOKEN";
type EnvVarStatus = {
missing_count: number;
required: { name: string; is_set: boolean }[];
};
type Server = {
name: string;
id: string;
statusUrl: string;
status: () => Promise<EnvVarStatus>;
remove: () => Promise<void>;
};
async function createServer(
request: APIRequestContext,
variables: string[],
): Promise<Server> {
const name = `int_mcp_${randomUUID().replace(/-/g, "").slice(0, 12)}`;
const created = await request.post("/v1/mcp/server", {
headers,
data: {
server_name: name,
url: `${process.env.INTEGRATION_UPSTREAM_URL}/mcp`,
transport: "http",
auth_type: "none",
env_vars: variables.map((variable) => ({
name: variable,
scope: "user",
description: `Per-user ${variable}`,
})),
static_headers: Object.fromEntries(
variables.map((variable, index) => [
`X-User-${index}`,
`\${${variable}}`,
]),
),
},
});
expect(created.ok(), await created.text()).toBe(true);
const id = (await created.json()).server_id as string;
const statusUrl = `/v1/mcp/server/${id}/user-env-vars`;
return {
name,
id,
statusUrl,
status: async () => {
const response = await request.get(statusUrl, { headers });
expect(response.ok(), await response.text()).toBe(true);
return response.json() as Promise<EnvVarStatus>;
},
remove: async () => {
const removed = await request.delete(`/v1/mcp/server/${id}`, { headers });
expect(
removed.ok() || removed.status() === 404,
await removed.text(),
).toBe(true);
},
};
}
async function openMcpServers(page: PlaywrightPage): Promise<void> {
await page.goto("/ui/login");
await page.getByPlaceholder("Enter your username").fill("admin");
await page.getByPlaceholder("Enter your password").fill(master);
await page.getByRole("button", { name: "Login", exact: true }).click();
await expect(page).toHaveURL(
(url) => url.pathname.startsWith("/ui") && !url.pathname.includes("login"),
);
await navigateToPage(page, Page.McpServers);
}
function cardFor(page: PlaywrightPage, server: Server): Locator {
return page.getByRole("button").filter({ hasText: server.name }).first();
}
function credentialsDialog(page: PlaywrightPage): Locator {
return page.getByRole("dialog").filter({ hasText: "Set your credentials" });
}
async function saveValues(
page: PlaywrightPage,
server: Server,
values: Record<string, string>,
): Promise<void> {
const dialog = credentialsDialog(page);
for (const [variable, value] of Object.entries(values)) {
await dialog.getByLabel(variable).fill(value);
}
const body = await captureRequestBody(
page,
{ method: "POST", urlIncludes: server.statusUrl },
async () => {
await dialog.getByRole("button", { name: "Save Credentials" }).click();
},
);
expect(body).toEqual({ values });
await expect(dialog).toHaveCount(0);
}
test("per-user MCP env var stays updatable and clearable from the card after it is set", async ({
page,
request,
}) => {
const server = await createServer(request, [TOKEN]);
try {
await openMcpServers(page);
const card = cardFor(page, server);
const dialog = credentialsDialog(page);
await expect(
card.getByText("1 user field missing", { exact: true }),
).toBeVisible();
await card.getByRole("button", { name: "Set", exact: true }).click();
await saveValues(page, server, { [TOKEN]: "first-token" });
expect(await server.status()).toMatchObject({
missing_count: 0,
required: [{ name: TOKEN, is_set: true }],
});
await expect(
card.getByText("1 user field missing", { exact: true }),
).toHaveCount(0);
await page.reload();
const update = card.getByRole("button", { name: "Update", exact: true });
await expect(
update,
"a set per-user variable must keep an update entry point on the card",
).toBeVisible();
await update.click();
await expect(dialog.getByText("Set", { exact: true })).toBeVisible();
await saveValues(page, server, { [TOKEN]: "rotated-token" });
expect(await server.status()).toMatchObject({
missing_count: 0,
required: [{ name: TOKEN, is_set: true }],
});
await update.click();
const cleared = page.waitForResponse(
(response) =>
response.request().method() === "DELETE" &&
response.url().includes(server.statusUrl),
);
await dialog.getByRole("button", { name: "Clear", exact: true }).click();
const confirm = page.getByRole("alertdialog", {
name: "Clear saved credentials",
});
await expect(confirm).toContainText(server.name);
await confirm
.getByRole("button", { name: "Clear credentials", exact: true })
.click();
const clearResponse = await cleared;
expect(clearResponse.ok(), await clearResponse.text()).toBe(true);
await expect(dialog).toHaveCount(0);
expect(await server.status()).toMatchObject({
missing_count: 1,
required: [{ name: TOKEN, is_set: false }],
});
await expect(
card.getByText("1 user field missing", { exact: true }),
).toBeVisible();
await expect(
card.getByRole("button", { name: "Set", exact: true }),
).toBeVisible();
} finally {
await server.remove();
}
});
test("cancelling the clear confirmation keeps the stored value and sends no delete", async ({
page,
request,
}) => {
const server = await createServer(request, [TOKEN]);
try {
const stored = await request.post(server.statusUrl, {
headers,
data: { values: { [TOKEN]: "keep-me" } },
});
expect(stored.ok(), await stored.text()).toBe(true);
await openMcpServers(page);
const card = cardFor(page, server);
const dialog = credentialsDialog(page);
const deletes: string[] = [];
page.on("request", (sent) => {
if (sent.method() === "DELETE" && sent.url().includes(server.statusUrl))
deletes.push(sent.url());
});
await card.getByRole("button", { name: "Update", exact: true }).click();
await dialog.getByRole("button", { name: "Clear", exact: true }).click();
const confirm = page.getByRole("alertdialog", {
name: "Clear saved credentials",
});
await expect(confirm).toBeVisible();
await confirm.getByRole("button", { name: "Cancel", exact: true }).click();
await expect(confirm).toHaveCount(0);
await expect(
dialog,
"cancelling the confirmation must leave the credentials modal open",
).toBeVisible();
await dialog.getByRole("button", { name: "Cancel", exact: true }).click();
await expect(dialog).toHaveCount(0);
await card.getByRole("button", { name: "Update", exact: true }).click();
await expect(
confirm,
"a cancelled confirmation must not reappear on reopen",
).toHaveCount(0);
await page.keyboard.press("Escape");
await expect(dialog).toHaveCount(0);
expect(deletes).toEqual([]);
expect(await server.status()).toMatchObject({
missing_count: 0,
required: [{ name: TOKEN, is_set: true }],
});
await expect(
card.getByRole("button", { name: "Update", exact: true }),
).toBeVisible();
} finally {
await server.remove();
}
});
test("pressing Enter on Update opens the credentials modal instead of the server editor", async ({
page,
request,
}) => {
const server = await createServer(request, [TOKEN]);
try {
const stored = await request.post(server.statusUrl, {
headers,
data: { values: { [TOKEN]: "keyboard" } },
});
expect(stored.ok(), await stored.text()).toBe(true);
await openMcpServers(page);
const card = cardFor(page, server);
const update = card.getByRole("button", { name: "Update", exact: true });
await update.focus();
await page.keyboard.press("Enter");
const dialog = credentialsDialog(page);
await expect(dialog).toBeVisible();
await expect(
page.getByRole("button", { name: "Back to All Servers" }),
).toHaveCount(0);
await saveValues(page, server, { [TOKEN]: "keyboard-rotated" });
await expect(
page.getByRole("button", { name: "Back to All Servers" }),
).toHaveCount(0);
await expect(card).toBeVisible();
expect(await server.status()).toMatchObject({
missing_count: 0,
required: [{ name: TOKEN, is_set: true }],
});
await card.click();
await expect(
page.getByRole("button", { name: "Back to All Servers" }),
).toBeVisible();
} finally {
await server.remove();
}
});
test("a server with two per-user variables reports the remaining gap until both are saved", async ({
page,
request,
}) => {
const second = "WORKSPACE";
const server = await createServer(request, [TOKEN, second]);
try {
await openMcpServers(page);
const card = cardFor(page, server);
const dialog = credentialsDialog(page);
await expect(
card.getByText("2 user fields missing", { exact: true }),
).toBeVisible();
await card.getByRole("button", { name: "Set", exact: true }).click();
await dialog.getByLabel(TOKEN).fill("only-token");
const posts: string[] = [];
page.on("request", (sent) => {
if (sent.method() === "POST" && sent.url().includes(server.statusUrl))
posts.push(sent.url());
});
await dialog.getByRole("button", { name: "Save Credentials" }).click();
await expect(dialog.getByRole("alert")).toHaveText(`${second} is required`);
expect(posts, "a missing required field must block the save").toEqual([]);
await dialog.getByRole("button", { name: "Cancel", exact: true }).click();
await expect(dialog).toHaveCount(0);
const partial = await request.post(server.statusUrl, {
headers,
data: { values: { [TOKEN]: "only-token" } },
});
expect(partial.ok(), await partial.text()).toBe(true);
await page.reload();
await expect(
card.getByText("1 user field missing", { exact: true }),
).toBeVisible();
await expect(
card.getByRole("button", { name: "Update", exact: true }),
).toHaveCount(0);
await card.getByRole("button", { name: "Set", exact: true }).click();
await expect(dialog.getByText("Set", { exact: true })).toHaveCount(1);
await saveValues(page, server, {
[TOKEN]: "",
[second]: "workspace-value",
});
expect(await server.status()).toMatchObject({
missing_count: 0,
required: [
{ name: TOKEN, is_set: true },
{ name: second, is_set: true },
],
});
await expect(
card.getByRole("button", { name: "Update", exact: true }),
).toBeVisible();
await expect(card.getByText(/user fields? missing/)).toHaveCount(0);
} finally {
await server.remove();
}
});
test("a server without per-user variables shows no credential row", async ({
page,
request,
}) => {
const server = await createServer(request, []);
const withVariable = await createServer(request, [TOKEN]);
try {
await openMcpServers(page);
const plain = cardFor(page, server);
await expect(plain).toBeVisible();
await expect(
cardFor(page, withVariable).getByRole("button", {
name: "Set",
exact: true,
}),
).toBeVisible();
await expect(plain.getByText("Per-user credentials")).toHaveCount(0);
await expect(
plain.getByRole("button", { name: "Set", exact: true }),
).toHaveCount(0);
await expect(
plain.getByRole("button", { name: "Update", exact: true }),
).toHaveCount(0);
await expect(plain.getByText(/user fields? missing/)).toHaveCount(0);
} finally {
await server.remove();
await withVariable.remove();
}
});
test("clearing credentials for a server deleted underneath the modal reports the failure without losing the page", async ({
page,
request,
}) => {
const server = await createServer(request, [TOKEN]);
const survivor = await createServer(request, [TOKEN]);
try {
const stored = await request.post(server.statusUrl, {
headers,
data: { values: { [TOKEN]: "doomed" } },
});
expect(stored.ok(), await stored.text()).toBe(true);
await openMcpServers(page);
const card = cardFor(page, server);
const dialog = credentialsDialog(page);
await card.getByRole("button", { name: "Update", exact: true }).click();
await expect(dialog).toBeVisible();
await server.remove();
const cleared = page.waitForResponse(
(response) =>
response.request().method() === "DELETE" &&
response.url().includes(server.statusUrl),
);
await dialog.getByRole("button", { name: "Clear", exact: true }).click();
await page
.getByRole("alertdialog", { name: "Clear saved credentials" })
.getByRole("button", {
name: "Clear credentials",
exact: true,
})
.click();
const clearResponse = await cleared;
expect(clearResponse.status()).toBe(404);
await expect(page.getByText(/Failed to clear env vars/)).toBeVisible();
await expect(
dialog,
"a failed clear must keep the modal open for the user",
).toBeVisible();
await page.keyboard.press("Escape");
await expect(dialog).toHaveCount(0);
await page.reload();
await expect(
cardFor(page, survivor).getByRole("button", { name: "Set", exact: true }),
).toBeVisible();
await expect(page.getByText(server.name)).toHaveCount(0);
} finally {
await server.remove();
await survivor.remove();
}
});

View file

@ -0,0 +1,21 @@
import pytest
from integration._support.client import Gateway, object_value
from pydantic import JsonValue
def _assert_responses_post_is_documented(openapi: dict[str, JsonValue]) -> None:
post: dict[str, JsonValue] = object_value(object_value(object_value(openapi["paths"])["/v1/responses"])["post"])
body: dict[str, JsonValue] = object_value(post["requestBody"])
schema: dict[str, JsonValue] = object_value(
object_value(object_value(body["content"])["application/json"])["schema"]
)
properties: dict[str, JsonValue] = object_value(schema.get("properties"))
assert "model" in properties and "input" in properties, schema
ok: dict[str, JsonValue] = object_value(object_value(object_value(post)["responses"])["200"])
assert "schema" in object_value(object_value(ok["content"])["application/json"]), ok
def test_v1_responses_post_declares_a_request_body_and_response_schema(gateway: Gateway) -> None:
pytest.skip("BUG: POST /v1/responses takes a raw Request, so /openapi.json documents no body or response schema")
openapi: dict[str, JsonValue] = gateway.get("/openapi.json")
_assert_responses_post_is_documented(openapi)

View file

@ -0,0 +1,30 @@
import uuid
from typing import Final
from integration._support.client import Gateway, object_value, string_value
from pydantic import JsonValue
def test_scim_group_patch_add_member_provisions_the_missing_user(gateway: Gateway) -> None:
missing_user: Final = f"scim-pending-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
team: Final = scenario.team()
response: Final = gateway.request(
"PATCH",
f"/scim/v2/Groups/{team}",
{
"schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"],
"Operations": [
{"op": "add", "path": "members", "value": [{"value": missing_user}]}
],
},
)
scenario.cleanups.callback(scenario.delete_user, missing_user)
assert response.status_code == 200, response.text
team_info: dict[str, JsonValue] = gateway.get("/team/info", {"team_id": team})
members: Final = object_value(team_info["team_info"]).get("members_with_roles") or []
member_ids: Final = [
string_value(object_value(member)["user_id"]) for member in members if isinstance(member, dict)
]
assert missing_user in member_ids, members

View file

@ -258,16 +258,6 @@ def _peer_add_calls(peer: McpPeer) -> tuple[dict[str, object], ...]:
)
def _skip_if_bridge_drops_tool_result(
rig: Rig, requests: tuple[tuple[str, ...], ...], calls: tuple[object, ...]
) -> None:
if rig.surface == "messages_bridge" and len(calls) > 1 and len(requests) > 2:
pytest.skip(
"BUG: /v1/messages MCP tool loop over a non-Anthropic model drops the tool_result message, "
"so the tool is re-executed until the iteration cap"
)
@pytest.mark.parametrize("surface", SURFACES)
def test_auto_approved_gateway_tool_is_listed_executed_once_and_fed_back(gateway: Gateway, surface: Surface) -> None:
with _rig(gateway, surface) as rig:
@ -276,7 +266,6 @@ def test_auto_approved_gateway_tool_is_listed_executed_once_and_fed_back(gateway
assert response.status_code == 200, response.text
calls: Final = _peer_add_calls(rig.peer)
requests: Final = rig.upstream_tools()
_skip_if_bridge_drops_tool_result(rig, requests, calls)
assert [call["body"]["params"]["name"] for call in calls] == ["add"], calls
assert calls[0]["body"]["params"]["arguments"] == ADD, calls
assert len(requests) == 2, requests

View file

@ -0,0 +1,359 @@
import signal
import uuid
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from functools import partial
from pathlib import Path
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names
from integration._support.process import owned_proxy_process
from pydantic import JsonValue, TypeAdapter
TOKEN: Final = "USER_TOKEN"
WORKSPACE: Final = "WORKSPACE"
METHODS: Final = ("GET", "POST", "DELETE")
@dataclass(frozen=True, slots=True)
class UpstreamCall:
body: dict[str, JsonValue]
headers: dict[bytes, bytes]
UPSTREAM_CALLS: Final = TypeAdapter(tuple[UpstreamCall, ...])
STATUS_LISTING: Final = TypeAdapter(list[dict[str, JsonValue]])
JSON_BODY: Final = TypeAdapter(dict[str, JsonValue])
def body(response: httpx.Response) -> dict[str, JsonValue]:
return JSON_BODY.validate_json(response.content)
def register_user_var_server(scenario: Scenario, peer: McpPeer, *names: str) -> str:
return register_mcp(
scenario,
peer,
"integration" + uuid.uuid4().hex,
auth_type="none",
env_vars=[{"name": name, "scope": "user", "description": f"per-user {name}"} for name in names],
static_headers={
"Authorization": f"Bearer ${{{TOKEN}}}",
**({"X-Workspace": f"${{{WORKSPACE}}}"} if WORKSPACE in names else {}),
},
)
def grants(*identities: str) -> JsonValue:
return {"mcp_servers": list(identities)}
def user_key(scenario: Scenario, identity: str) -> str:
return scenario.key(user_id=scenario.user(), object_permission=grants(identity))
def env_status(gateway: Gateway, key: str, identity: str) -> httpx.Response:
return gateway.request("GET", f"/v1/mcp/server/{identity}/user-env-vars", key=key)
def store(gateway: Gateway, key: str, identity: str, values: Mapping[str, str]) -> httpx.Response:
return gateway.request("POST", f"/v1/mcp/server/{identity}/user-env-vars", {"values": dict(values)}, key=key)
def clear(gateway: Gateway, key: str, identity: str) -> httpx.Response:
return gateway.request("DELETE", f"/v1/mcp/server/{identity}/user-env-vars", key=key)
def set_names(response: httpx.Response) -> dict[str, bool]:
assert response.status_code == 200, response.text
status: Final = body(response)
assert isinstance(status["required"], list)
return {
string_value(object_value(spec)["name"]): object_value(spec)["is_set"] is True for spec in status["required"]
}
def tool_calls(peer: McpPeer) -> tuple[UpstreamCall, ...]:
return tuple(
call for call in UPSTREAM_CALLS.validate_python(peer.drain()) if call.body.get("method") == "tools/call"
)
def add_upstream_headers(gateway: Gateway, peer: McpPeer, key: str, identity: str, a: int = 2) -> dict[bytes, bytes]:
peer.drain()
response: Final = call_tool(gateway, key, identity, tool_names(gateway, key, identity)["add"], {"a": a, "b": 3})
assert response.status_code == 200, response.text
calls: Final = tool_calls(peer)
assert len(calls) == 1, calls
return calls[0].headers
def add_upstream_authorization(gateway: Gateway, peer: McpPeer, key: str, identity: str) -> bytes:
return add_upstream_headers(gateway, peer, key, identity)[b"authorization"]
def list_tools_status(target: Gateway, key: str, identity: str) -> int:
return target.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
).status_code
def wait_for_tools(target: Gateway, key: str, identity: str) -> dict[str, str]:
eventually(lambda: list_tools_status(target, key, identity), lambda status: status == 200, seconds=60)
return eventually(lambda: tool_names(target, key, identity), lambda names: "add" in names, seconds=60)
def assert_forwarded_eventually(target: Gateway, upstream: McpPeer, key: str, identity: str, expected: bytes) -> None:
observed: Final = eventually(
lambda: add_upstream_authorization(target, upstream, key, identity), lambda value: value == expected, seconds=75
)
assert observed == expected
def assert_precondition_failed(gateway: Gateway, key: str, identity: str, *missing: str) -> None:
response: Final = call_tool(gateway, key, identity, tool_names(gateway, key, identity)["add"], {"a": 2, "b": 3})
assert response.status_code == 412, response.text
detail: Final = object_value(body(response)["detail"])
assert detail["error"] == "missing_user_env_vars"
assert detail["server_id"] == identity
assert isinstance(detail["missing"], list)
assert sorted(string_value(name) for name in detail["missing"]) == sorted(missing)
assert string_value(detail["setup_url"]).endswith(f"fill_env_vars={identity}")
def stored_user_ids(identity: str) -> tuple[JsonValue, ...]:
return tuple(
row["user_id"]
for row in read_rows('SELECT user_id FROM "LiteLLM_MCPUserEnvVars" WHERE server_id = %s', (identity,))
)
def missing_count(response: httpx.Response) -> JsonValue:
return body(response)["missing_count"]
def test_stored_value_is_forwarded_rotated_and_cleared(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, peer, TOKEN)
key: Final = user_key(scenario, identity)
before: Final = env_status(gateway, key, identity)
assert set_names(before) == {TOKEN: False}
assert missing_count(before) == 1
assert string_value(body(before)["setup_url"]).endswith(f"fill_env_vars={identity}")
assert_precondition_failed(gateway, key, identity, TOKEN)
first: Final = store(gateway, key, identity, {TOKEN: "first-secret"})
assert set_names(first) == {TOKEN: True}
assert missing_count(first) == 0
assert add_upstream_authorization(gateway, peer, key, identity) == b"Bearer first-secret"
rotated: Final = store(gateway, key, identity, {TOKEN: "second-secret"})
assert set_names(rotated) == {TOKEN: True}
assert add_upstream_authorization(gateway, peer, key, identity) == b"Bearer second-secret"
assert len(stored_user_ids(identity)) == 1
cleared: Final = clear(gateway, key, identity)
assert set_names(cleared) == {TOKEN: False}
assert missing_count(cleared) == 1
assert stored_user_ids(identity) == ()
assert set_names(env_status(gateway, key, identity)) == {TOKEN: False}
assert_precondition_failed(gateway, key, identity, TOKEN)
assert set_names(clear(gateway, key, identity)) == {TOKEN: False}
def test_store_merges_per_variable_and_drops_undeclared_or_empty_values(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, peer, TOKEN, WORKSPACE)
key: Final = user_key(scenario, identity)
assert set_names(env_status(gateway, key, identity)) == {TOKEN: False, WORKSPACE: False}
assert_precondition_failed(gateway, key, identity, TOKEN, WORKSPACE)
partial: Final = store(gateway, key, identity, {TOKEN: "tok", "NOT_DECLARED": "x", "": "y"})
assert set_names(partial) == {TOKEN: True, WORKSPACE: False}
assert missing_count(partial) == 1
assert_precondition_failed(gateway, key, identity, WORKSPACE)
long_value: Final = "w" * 5120
complete: Final = store(gateway, key, identity, {WORKSPACE: long_value})
assert set_names(complete) == {TOKEN: True, WORKSPACE: True}
forwarded: Final = add_upstream_headers(gateway, peer, key, identity)
assert forwarded[b"authorization"] == b"Bearer tok"
assert forwarded[b"x-workspace"] == long_value.encode()
kept: Final = store(gateway, key, identity, {TOKEN: "", WORKSPACE: ""})
assert set_names(kept) == {TOKEN: True, WORKSPACE: True}
assert add_upstream_authorization(gateway, peer, key, identity) == b"Bearer tok"
assert set_names(store(gateway, key, identity, {TOKEN: "tok"})) == {TOKEN: True, WORKSPACE: True}
assert len(stored_user_ids(identity)) == 1
def test_malformed_bodies_missing_users_and_foreign_servers_are_rejected(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, peer, TOKEN)
other: Final = register_user_var_server(scenario, peer, TOKEN)
key: Final = user_key(scenario, identity)
userless: Final = scenario.key(object_permission=grants(identity))
path: Final = f"/v1/mcp/server/{identity}/user-env-vars"
payload: Final[dict[str, JsonValue]] = {"values": {TOKEN: "x"}}
malformed: Final[tuple[dict[str, JsonValue], ...]] = ({"values": {TOKEN: 7}}, {"values": ["a"]}, {})
assert [gateway.request("POST", path, body, key=key).status_code for body in malformed] == [422, 422, 422]
assert set_names(env_status(gateway, key, identity)) == {TOKEN: False}
assert [gateway.client.request(method, path, json=payload).status_code for method in METHODS] == [401, 401, 401]
no_user: Final = tuple(gateway.request(method, path, payload, key=userless) for method in METHODS)
assert [response.status_code for response in no_user] == [400, 400, 400], [r.text for r in no_user]
assert [object_value(body(r)["detail"])["error"] for r in no_user] == ["User ID not found in token"] * 3
foreign: Final = f"/v1/mcp/server/{other}/user-env-vars"
assert [gateway.request(method, foreign, payload, key=key).status_code for method in METHODS] == [403, 403, 403]
unknown: Final = f"/v1/mcp/server/{uuid.uuid4()}/user-env-vars"
assert [gateway.request(method, unknown, payload).status_code for method in METHODS] == [404, 404, 404]
assert stored_user_ids(identity) == () and stored_user_ids(other) == ()
def test_status_list_keeps_fully_set_servers_and_is_scoped_to_the_caller(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
per_user: Final = register_user_var_server(scenario, peer, TOKEN)
global_only: Final = register_mcp(
scenario,
peer,
"integration" + uuid.uuid4().hex,
env_vars=[{"name": "GLOBAL_TOKEN", "scope": "global", "description": "shared"}],
)
plain: Final = register_mcp(scenario, peer, "integration" + uuid.uuid4().hex)
first_user: Final = scenario.key(
user_id=scenario.user(), object_permission=grants(per_user, global_only, plain)
)
second_user: Final = scenario.key(
user_id=scenario.user(), object_permission=grants(per_user, global_only, plain)
)
def listing(key: str) -> dict[str, JsonValue]:
response: Final = gateway.request("GET", "/v1/mcp/user-env-vars/status", key=key)
assert response.status_code == 200, response.text
return {
string_value(entry["server_id"]): entry["missing_count"]
for entry in STATUS_LISTING.validate_json(response.content)
if entry["server_id"] in {per_user, global_only, plain}
}
assert listing(first_user) == {per_user: 1}
assert set_names(store(gateway, first_user, per_user, {TOKEN: "mine"})) == {TOKEN: True}
assert listing(first_user) == {per_user: 0}
assert listing(second_user) == {per_user: 1}
assert set_names(env_status(gateway, second_user, per_user)) == {TOKEN: False}
assert add_upstream_authorization(gateway, peer, first_user, per_user) == b"Bearer mine"
assert_precondition_failed(gateway, second_user, per_user, TOKEN)
assert set_names(clear(gateway, second_user, per_user)) == {TOKEN: False}
assert listing(first_user) == {per_user: 0}
assert add_upstream_authorization(gateway, peer, first_user, per_user) == b"Bearer mine"
def test_store_and_clear_on_one_process_are_honored_by_the_other(gateway: Gateway, peer: Gateway) -> None:
with mcp_peer() as upstream, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, upstream, TOKEN)
key: Final = user_key(scenario, identity)
assert_precondition_failed(gateway, key, identity, TOKEN)
wait_for_tools(peer, key, identity)
assert_precondition_failed(peer, key, identity, TOKEN)
assert set_names(store(gateway, key, identity, {TOKEN: "from-a"})) == {TOKEN: True}
assert set_names(env_status(peer, key, identity)) == {TOKEN: True}
assert add_upstream_authorization(peer, upstream, key, identity) == b"Bearer from-a"
assert set_names(store(peer, key, identity, {TOKEN: "from-b"})) == {TOKEN: True}
assert_forwarded_eventually(gateway, upstream, key, identity, b"Bearer from-b")
assert set_names(clear(gateway, key, identity)) == {TOKEN: False}
assert set_names(env_status(peer, key, identity)) == {TOKEN: False}
assert stored_user_ids(identity) == ()
def peer_status() -> int:
names: Final = tool_names(peer, key, identity)
return call_tool(peer, key, identity, names["add"], {"a": 1, "b": 1}).status_code
assert eventually(peer_status, lambda code: code == 412, seconds=75) == 412
assert_precondition_failed(gateway, key, identity, TOKEN)
@pytest.mark.timeout(240)
def test_concurrent_users_across_processes_never_leak_and_survive_a_killed_process(
gateway: Gateway, peer: Gateway, tmp_path: Path
) -> None:
with mcp_peer() as upstream, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, upstream, TOKEN)
users: Final = tuple(scenario.user() for _ in range(4))
keys: Final = {user: scenario.key(user_id=user, object_permission=grants(identity)) for user in users}
assert [set_names(store(gateway, keys[user], identity, {TOKEN: f"seed-{user}"})) for user in users] == [
{TOKEN: True}
] * len(users)
names: Final = wait_for_tools(gateway, keys[users[0]], identity)
wait_for_tools(peer, keys[users[0]], identity)
def operation(target: Gateway, user: str, index: int) -> httpx.Response:
if index % 4 == 1:
return env_status(target, keys[user], identity)
if index % 4 == 2:
return call_tool(target, keys[user], identity, names["add"], {"a": users.index(user), "b": 0})
return store(target, keys[user], identity, {TOKEN: f"{user}-{index}"})
def outcome(targets: tuple[Gateway, ...], job: tuple[str, int]) -> tuple[int, int]:
return job[1], operation(targets[job[1] % len(targets)], job[0], job[1]).status_code
def burst(pool: ThreadPoolExecutor, targets: tuple[Gateway, ...]) -> tuple[tuple[int, int], ...]:
jobs: Final = tuple((user, index) for user in users for index in range(6))
return tuple(pool.map(partial(outcome, targets), jobs))
def allowed_authorizations(item: UpstreamCall) -> tuple[str, frozenset[bytes]]:
arguments: Final = object_value(object_value(item.body["params"])["arguments"])
owner: Final = users[int(string_value(str(arguments["a"])))]
return owner, frozenset(
{f"Bearer seed-{owner}".encode()} | {f"Bearer {owner}-{i}".encode() for i in range(6)}
)
with owned_proxy_process(gateway, tmp_path, {}) as doomed, ThreadPoolExecutor(max_workers=8) as pool:
wait_for_tools(doomed.gateway, keys[users[0]], identity)
upstream.drain()
outcomes: Final = burst(pool, (gateway, peer, doomed.gateway))
assert all(code in {200, 412} for _, code in outcomes), outcomes
assert all(code == 200 for index, code in outcomes if index % 4 != 2), outcomes
doomed.process.send_signal(signal.SIGKILL)
doomed.process.wait(timeout=10)
after_kill: Final = burst(pool, (gateway, peer))
assert all(code in {200, 412} for _, code in after_kill), after_kill
assert all(code == 200 for index, code in after_kill if index % 4 != 2), after_kill
forwarded: Final = tool_calls(upstream)
assert forwarded
leaked: Final = tuple(
(owner, item.headers[b"authorization"])
for item in forwarded
for owner, allowed in (allowed_authorizations(item),)
if item.headers[b"authorization"] not in allowed
)
assert leaked == ()
assert [set_names(env_status(gateway, keys[user], identity)) for user in users] == [{TOKEN: True}] * len(users)
assert [set_names(env_status(peer, keys[user], identity)) for user in users] == [{TOKEN: True}] * len(users)
assert [set_names(store(gateway, keys[user], identity, {TOKEN: f"final-{user}"})) for user in users] == [
{TOKEN: True}
] * len(users)
for user in users:
assert_forwarded_eventually(peer, upstream, keys[user], identity, f"Bearer final-{user}".encode())
assert sorted(string_value(user_id) for user_id in stored_user_ids(identity)) == sorted(users)
def test_concurrent_stores_of_different_variables_do_not_lose_an_update(gateway: Gateway, peer: Gateway) -> None:
with mcp_peer() as upstream, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, upstream, TOKEN, WORKSPACE)
key: Final = user_key(scenario, identity)
wait_for_tools(peer, key, identity)
def race_once(pool: ThreadPoolExecutor) -> None:
assert set_names(clear(gateway, key, identity)) == {TOKEN: False, WORKSPACE: False}
first: Final = pool.submit(store, gateway, key, identity, {TOKEN: "racing-token"})
second: Final = pool.submit(store, peer, key, identity, {WORKSPACE: "racing-workspace"})
assert first.result().status_code == 200, first.result().text
assert second.result().status_code == 200, second.result().text
assert set_names(env_status(gateway, key, identity)) == {TOKEN: True, WORKSPACE: True}
assert set_names(env_status(peer, key, identity)) == {TOKEN: True, WORKSPACE: True}
assert len(stored_user_ids(identity)) == 1
forwarded: Final = add_upstream_headers(gateway, upstream, key, identity, a=1)
assert forwarded[b"authorization"] == b"Bearer racing-token"
assert forwarded[b"x-workspace"] == b"racing-workspace"
with ThreadPoolExecutor(max_workers=2) as pool:
for _ in range(5):
race_once(pool)

View file

@ -0,0 +1,56 @@
import json
import uuid
from typing import Final
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
_MODEL: Final = "bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0"
_TOKEN: Final = "synthetic-bedrock-bearer"
def test_bedrock_500_keeps_amzn_request_id_on_error_headers_and_failure_log(gateway: Gateway) -> None:
identity: Final = f"bedrock-request-id-{uuid.uuid4().hex}"
amzn_request_id: Final = str(uuid.uuid4())
prompt: Final = f"failure probe {identity}"
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/model/anthropic.claude-sonnet-4-5-20250929-v1%3A0/converse", request.target
return Reply(
status=500,
headers={"x-amzn-RequestId": amzn_request_id},
body=b'{"message":"synthetic bedrock failure"}',
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=_MODEL,
api_key=_TOKEN,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
num_retries=0,
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}]},
)
assert response.status_code >= 400, response.text
assert response.headers.get("llm_provider-x-amzn-requestid") == amzn_request_id, dict(response.headers)
call_id: Final = response.headers["x-litellm-call-id"]
assert len(wire.drain()) == 1
rows: Final = eventually(
lambda: read_rows(
'SELECT status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)
),
lambda values: len(values) == 1,
seconds=70,
)
row: Final = rows[0]
assert row["status"] == "failure", row
metadata: Final = row["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
error_information: Final = object_value(parsed["error_information"])
assert error_information["error_provider_request_id"] == amzn_request_id, error_information

View file

@ -0,0 +1,105 @@
import json
import uuid
from typing import Final
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "gpt-5.4-mini"
_API_KEY: Final = "synthetic-openai-key"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _claude_code_user_id(device_id: str, session_id: str) -> str:
return json.dumps({"device_id": device_id, "account_uuid": "", "session_id": session_id})
def _responses_reply(identity: str) -> bytes:
return json.dumps(
{
"id": f"resp_{identity}",
"object": "response",
"created_at": 1789788253,
"status": "completed",
"model": _BACKEND,
"output": [
{
"type": "message",
"id": f"msg_{identity}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "ok", "annotations": []}],
}
],
"usage": {"input_tokens": 10, "output_tokens": 2, "total_tokens": 12},
}
).encode()
def test_prompt_cache_key_is_derived_from_claude_code_session_id_not_device_id(gateway: Gateway) -> None:
identity: Final = f"claude-code-cache-key-{uuid.uuid4().hex}"
device_one: Final = "a" * 64
device_two: Final = "b" * 64
session_one: Final = str(uuid.uuid4())
session_two: Final = str(uuid.uuid4())
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/responses"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
return Reply(body=_responses_reply(identity))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
def send(user_id: str, probe: str) -> None:
response: Final = gateway.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 16,
"metadata": {"user_id": user_id},
"messages": [{"role": "user", "content": probe}],
},
)
assert response.status_code == 200, response.text
send(_claude_code_user_id(device_one, session_one), f"probe one {identity}")
send(_claude_code_user_id(device_one, session_two), f"probe two {identity}")
send(_claude_code_user_id(device_two, session_two), f"probe three {identity}")
keys: Final = [
_JSON_OBJECT.validate_json(request.body).get("prompt_cache_key") for request in wire.drain()
]
assert keys[0] == session_one, keys
assert keys[1] == session_two, keys
assert keys[2] == session_two, keys
assert keys[0] != keys[1] and keys[1] == keys[2]
def test_explicit_prompt_cache_key_wins_over_derived_session_key(gateway: Gateway) -> None:
identity: Final = f"claude-code-explicit-key-{uuid.uuid4().hex}"
explicit: Final = "explicit-client-cache-key"
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/responses"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["prompt_cache_key"] == explicit, body
return Reply(body=_responses_reply(identity))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 16,
"prompt_cache_key": explicit,
"metadata": {"user_id": _claude_code_user_id("c" * 64, str(uuid.uuid4()))},
"messages": [{"role": "user", "content": f"explicit key probe {identity}"}],
},
)
assert response.status_code == 200, response.text
assert len(wire.drain()) == 1

View file

@ -0,0 +1,128 @@
import json
import uuid
from typing import Final
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_MODEL: Final = "claude-sonnet-4-5-20250929"
_API_KEY: Final = "synthetic-anthropic-key"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _anthropic_reply(identity: str, text: str) -> bytes:
return json.dumps(
{
"id": identity,
"type": "message",
"role": "assistant",
"model": _MODEL,
"content": [{"type": "text", "text": text}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 3, "cache_creation_input_tokens": 12},
}
).encode()
def _assert_system_block(body: dict[str, JsonValue], policy: str) -> None:
assert body["model"] == _MODEL, body
assert body["system"] == [{"type": "text", "text": policy, "cache_control": {"type": "ephemeral"}}], body
def test_chat_completions_system_block_list_carries_cache_control_to_anthropic_system(gateway: Gateway) -> None:
identity: Final = f"anthropic-system-cc-{uuid.uuid4().hex}"
policy: Final = f"policy {identity}"
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/messages"
assert request.headers["x-api-key"] == _API_KEY
_assert_system_block(_JSON_OBJECT.validate_json(request.body), policy)
return Reply(body=_anthropic_reply(identity, "done"))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"max_tokens": 16,
"messages": [
{
"role": "system",
"content": [
{"type": "text", "text": policy, "cache_control": {"type": "ephemeral"}}
],
},
{"role": "user", "content": "hi"},
],
},
)
assert response.status_code == 200, response.text
assert len(wire.drain()) == 1
def test_chat_completions_system_string_with_message_cache_control_reaches_anthropic_system(
gateway: Gateway,
) -> None:
identity: Final = f"anthropic-system-str-{uuid.uuid4().hex}"
policy: Final = f"policy {identity}"
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/messages"
_assert_system_block(_JSON_OBJECT.validate_json(request.body), policy)
return Reply(body=_anthropic_reply(identity, "done"))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"max_tokens": 16,
"messages": [
{"role": "system", "content": policy, "cache_control": {"type": "ephemeral"}},
{"role": "user", "content": "hi"},
],
},
)
assert response.status_code == 200, response.text
assert len(wire.drain()) == 1
def test_responses_system_input_item_carries_cache_control_to_anthropic_system(gateway: Gateway) -> None:
identity: Final = f"responses-system-cc-{uuid.uuid4().hex}"
policy: Final = f"policy {identity}"
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/messages"
_assert_system_block(_JSON_OBJECT.validate_json(request.body), policy)
return Reply(body=_anthropic_reply(identity, "done"))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/responses",
{
"model": model,
"input": [
{
"role": "system",
"content": [
{"type": "input_text", "text": policy, "cache_control": {"type": "ephemeral"}}
],
},
{"role": "user", "content": "hi"},
],
},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["status"] == "completed", response.text
assert any(item.get("type") == "message" for item in payload.get("output", []) if isinstance(item, dict))
assert len(wire.drain()) == 1

View file

@ -0,0 +1,60 @@
import uuid
from typing import Final
from integration._support.client import Gateway, eventually
def test_spend_over_a_tag_max_budget_rejects_the_next_request(gateway: Gateway) -> None:
tag: Final = f"tag-budget-{uuid.uuid4().hex}"
def delete_tag() -> None:
gateway.post("/tag/delete", {"name": tag})
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.01, output_cost_per_token=0.01)
gateway.post("/tag/new", {"name": tag, "max_budget": 0.0001})
scenario.cleanups.callback(delete_tag)
first: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag spend {tag}"}],
"metadata": {"tags": [tag]},
},
)
assert first.status_code == 200, first.text
def rejection() -> int:
return gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag budget probe {tag}"}],
"metadata": {"tags": [tag]},
},
).status_code
status: Final = eventually(rejection, lambda code: code != 200, seconds=70)
assert status in (400, 422, 429), status
blocked: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag budget probe {tag}"}],
"metadata": {"tags": [tag]},
},
)
assert "budget" in blocked.text.lower(), blocked.text
control: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"untagged probe {tag}"}],
"metadata": {"tags": [f"other-{tag}"]},
},
)
assert control.status_code == 200, control.text

View file

@ -3,7 +3,9 @@ from unittest.mock import AsyncMock, patch
import pytest
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
)
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
)
@ -59,7 +61,7 @@ def test_anthropic_messages_handler_skips_the_gateway_on_recursion():
"litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
new=AsyncMock(return_value={"routed": True}),
) as routed:
with pytest.raises(ValueError, match='anthropic_messages_handler is not implemented for sync calls'):
with pytest.raises(ValueError, match="anthropic_messages_handler is not implemented for sync calls"):
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
@ -78,7 +80,7 @@ def test_anthropic_messages_handler_leaves_native_tools_alone():
"litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
new=AsyncMock(return_value={"routed": True}),
) as routed:
with pytest.raises(ValueError, match='anthropic_messages_handler is not implemented for sync calls'):
with pytest.raises(ValueError, match="anthropic_messages_handler is not implemented for sync calls"):
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
@ -115,8 +117,31 @@ def test_build_tool_result_message_uses_anthropic_tool_result_blocks():
message = _build_tool_result_message([{"tool_call_id": "toolu_1", "result": "9 sections", "name": "read_wiki"}])
assert message["role"] == "user"
assert list(message["content"]) == [
{"type": "tool_result", "tool_use_id": "toolu_1", "content": "9 sections"}
assert message["content"] == [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "9 sections"}]
def test_build_tool_result_message_survives_the_chat_completions_bridge():
"""
Regression test (LIT-8474): a non-Anthropic model behind /v1/messages must see
the executed tool result as a role="tool" message keyed by the tool_call_id.
The bridge only translates list content, so a tuple-shaped user message was
dropped and the model re-requested the tool until the iteration cap.
"""
message = _build_tool_result_message(
[
{"tool_call_id": "call_1", "result": "5", "name": "add"},
{"tool_call_id": "call_2", "result": "7", "name": "add"},
]
)
translated = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai(
[message], model="hosted_vllm/gpt-4o-mini", custom_llm_provider="hosted_vllm"
)
assert translated == [
{"role": "tool", "tool_call_id": "call_1", "content": "5"},
{"role": "tool", "tool_call_id": "call_2", "content": "7"},
]
@ -157,19 +182,23 @@ async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials(
{"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]},
]
with patch.object(MCPRequestContext, "resolve", return_value=context), patch.object(
mcp_handler.LiteLLM_Proxy_MCP_Handler
if hasattr(mcp_handler, "LiteLLM_Proxy_MCP_Handler")
else __import__(
"litellm.responses.mcp.litellm_proxy_mcp_handler", fromlist=["LiteLLM_Proxy_MCP_Handler"]
).LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
new=process,
), patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, "_execute_tool_calls",
new=execute,
), patch(
"litellm.anthropic_messages", new=AsyncMock(side_effect=responses)
with (
patch.object(MCPRequestContext, "resolve", return_value=context),
patch.object(
mcp_handler.LiteLLM_Proxy_MCP_Handler
if hasattr(mcp_handler, "LiteLLM_Proxy_MCP_Handler")
else __import__(
"litellm.responses.mcp.litellm_proxy_mcp_handler", fromlist=["LiteLLM_Proxy_MCP_Handler"]
).LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
new=process,
),
patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
new=execute,
),
patch("litellm.anthropic_messages", new=AsyncMock(side_effect=responses)),
):
await mcp_handler.anthropic_messages_with_mcp(
max_tokens=100,
@ -220,16 +249,19 @@ async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped
}
anthropic_messages_mock = AsyncMock(return_value=tool_use_response)
with patch.object(
MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")
), patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, "_process_mcp_tools_without_openai_transform",
new=AsyncMock(return_value=([], {})),
), patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler, "_execute_tool_calls",
new=AsyncMock(return_value=[]),
), patch(
"litellm.anthropic_messages", new=anthropic_messages_mock
with (
patch.object(MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")),
patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
"_process_mcp_tools_without_openai_transform",
new=AsyncMock(return_value=([], {})),
),
patch.object(
import_module("litellm.responses.mcp.litellm_proxy_mcp_handler").LiteLLM_Proxy_MCP_Handler,
"_execute_tool_calls",
new=AsyncMock(return_value=[]),
),
patch("litellm.anthropic_messages", new=anthropic_messages_mock),
):
result = await mcp_handler.anthropic_messages_with_mcp(
max_tokens=100,

View file

@ -12,6 +12,7 @@ from hypothesis import HealthCheck, given, settings
from hypothesis import strategies as st
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging
@ -420,7 +421,7 @@ async def test_native_aocr_state_stashed_before_a_blocking_hook_raises_reaches_f
class Blocked(Exception):
pass
class Block(CustomLogger):
class Block(CustomGuardrail):
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
request_data["litellm_logging_obj"].model_call_details["blocked-by"] = token
raise Blocked("blocked after the provider answered")

View file

@ -1,5 +1,5 @@
import React from "react";
import { render, screen } from "@testing-library/react";
import { fireEvent, render, screen } from "@testing-library/react";
import { describe, it, expect, vi, afterEach } from "vitest";
import MCPServerCard from "./MCPServerCard";
import type { MCPServer } from "@/components/mcp_tools/types";
@ -67,3 +67,48 @@ describe("MCPServerCard logo", () => {
expect(screen.getByText("DE")).toBeInTheDocument();
});
});
describe("MCPServerCard per-user credentials", () => {
const renderUserFields = (props: { missingUserFields?: string[]; hasUserFields?: boolean }) => {
const onOpenFillFields = vi.fn();
const onClick = vi.fn();
render(<MCPServerCard server={baseServer} onClick={onClick} onOpenFillFields={onOpenFillFields} {...props} />);
return { onOpenFillFields, onClick };
};
it("offers Set while a field is missing", () => {
const { onOpenFillFields, onClick } = renderUserFields({ missingUserFields: ["USER_TOKEN"], hasUserFields: true });
expect(screen.getByText("1 user field missing")).toBeInTheDocument();
expect(screen.queryByRole("button", { name: "Update" })).not.toBeInTheDocument();
fireEvent.click(screen.getByRole("button", { name: "Set" }));
expect(onOpenFillFields).toHaveBeenCalledTimes(1);
expect(onClick).not.toHaveBeenCalled();
});
it("keeps an Update entry point once every field is set", () => {
const { onOpenFillFields, onClick } = renderUserFields({ missingUserFields: [], hasUserFields: true });
expect(screen.getByText("Per-user credentials")).toBeInTheDocument();
expect(screen.getByText("Set")).toBeInTheDocument();
expect(screen.queryByRole("button", { name: "Set" })).not.toBeInTheDocument();
expect(screen.queryByText(/user field/)).not.toBeInTheDocument();
fireEvent.click(screen.getByRole("button", { name: "Update" }));
expect(onOpenFillFields).toHaveBeenCalledTimes(1);
expect(onClick).not.toHaveBeenCalled();
});
it("keeps Enter on the Update button away from the card's open handler", () => {
const { onClick } = renderUserFields({ missingUserFields: [], hasUserFields: true });
const update = screen.getByRole("button", { name: "Update" });
expect(fireEvent.keyDown(update, { key: "Enter" }), "default activation must survive").toBe(true);
expect(onClick).not.toHaveBeenCalled();
fireEvent.keyDown(screen.getAllByRole("button")[0], { key: "Enter" });
expect(onClick).toHaveBeenCalledTimes(1);
});
it("renders no credential row for a server without per-user fields", () => {
renderUserFields({ missingUserFields: [], hasUserFields: false });
expect(screen.queryByText("Per-user credentials")).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: "Update" })).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: "Set" })).not.toBeInTheDocument();
});
});

View file

@ -21,6 +21,7 @@ interface MCPServerCardProps {
// Computed by the parent from the bulk /user-env-vars/status response, so
// the card never issues a per-row request (no N+1).
missingUserFields?: string[];
hasUserFields?: boolean;
isLoadingHealth?: boolean;
isRechecking?: boolean;
onClick: () => void;
@ -42,6 +43,7 @@ const stop = (e: MouseEvent | KeyboardEvent) => e.stopPropagation();
const MCPServerCard: FC<MCPServerCardProps> = ({
server,
missingUserFields,
hasUserFields,
isLoadingHealth,
isRechecking,
onClick,
@ -100,6 +102,7 @@ const MCPServerCard: FC<MCPServerCardProps> = ({
}
const handleKeyDown = (e: KeyboardEvent<HTMLDivElement>) => {
if (e.target !== e.currentTarget) return;
if (e.key === "Enter" || e.key === " ") {
e.preventDefault();
onClick();
@ -256,9 +259,10 @@ const MCPServerCard: FC<MCPServerCardProps> = ({
)}
</div>
{(server.is_byok || needsAttention) && (
{(server.is_byok || hasUserFields || needsAttention) && (
<div className="mt-auto flex flex-col gap-2">
{server.is_byok && <ByokRow connected={!!server.has_user_credential} onConnect={onByokConnect} />}
{hasUserFields && !needsAttention && <UserFieldsRow onUpdate={onOpenFillFields} />}
{needsAttention && (
<div className="flex items-center justify-between gap-2 text-xs">
<Tooltip>
@ -365,6 +369,29 @@ const HealthChip: FC<HealthChipProps> = ({
);
};
const UserFieldsRow: FC<{ onUpdate?: () => void }> = ({ onUpdate }) => (
<div className="flex items-center justify-between gap-2 text-xs">
<span className="text-muted-foreground">Per-user credentials</span>
<div className="flex items-center gap-2">
<Badge variant="outline">
<Check /> Set
</Badge>
{onUpdate && (
<Button
variant="link"
size="sm"
onClick={(e) => {
stop(e);
onUpdate();
}}
>
Update
</Button>
)}
</div>
</div>
);
interface ByokRowProps {
connected: boolean;
onConnect?: () => void;

View file

@ -1,5 +1,5 @@
import React from "react";
import { act, fireEvent, render, screen, waitFor } from "@testing-library/react";
import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
@ -10,6 +10,7 @@ import { MCPServer, MCPUserEnvVarsStatus } from "@/components/mcp_tools/types";
vi.mock("@/components/networking", () => ({
getMCPUserEnvVars: vi.fn(),
storeMCPUserEnvVars: vi.fn(),
clearMCPUserEnvVars: vi.fn(),
}));
const createQueryClient = () => new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } });
@ -216,6 +217,89 @@ describe("UserEnvVarsModal", () => {
expect(networking.storeMCPUserEnvVars).not.toHaveBeenCalled();
});
it("clears every stored value through the delete endpoint once the user confirms", async () => {
const user = setup();
const cleared = statusWith([{ name: "API_KEY", description: null, is_set: false }]);
vi.mocked(networking.clearMCPUserEnvVars).mockResolvedValue(cleared);
const { onSaved, onClose } = renderModal(statusWith([{ name: "API_KEY", description: null, is_set: true }]));
await fieldAfterOpen(/^API_KEY/);
await user.click(screen.getByRole("button", { name: "Clear" }));
expect(networking.clearMCPUserEnvVars).not.toHaveBeenCalled();
const confirm = await screen.findByRole("alertdialog", { name: "Clear saved credentials" });
await user.click(within(confirm).getByRole("button", { name: "Clear credentials" }));
await waitFor(() => {
expect(onSaved).toHaveBeenCalledWith(cleared);
});
expect(networking.clearMCPUserEnvVars).toHaveBeenCalledWith("sk-test", "srv-1");
expect(networking.storeMCPUserEnvVars).not.toHaveBeenCalled();
expect(onClose).toHaveBeenCalled();
});
it("keeps every stored value when the clear confirmation is cancelled", async () => {
const user = setup();
const { onSaved, onClose } = renderModal(statusWith([{ name: "API_KEY", description: null, is_set: true }]));
await fieldAfterOpen(/^API_KEY/);
await user.click(screen.getByRole("button", { name: "Clear" }));
const confirm = await screen.findByRole("alertdialog", { name: "Clear saved credentials" });
await user.click(within(confirm).getByRole("button", { name: "Cancel" }));
await waitFor(() => {
expect(screen.queryByRole("alertdialog")).not.toBeInTheDocument();
});
expect(networking.clearMCPUserEnvVars).not.toHaveBeenCalled();
expect(onSaved).not.toHaveBeenCalled();
expect(onClose).not.toHaveBeenCalled();
expect(screen.getByRole("button", { name: "Clear" })).toBeEnabled();
});
it("drops a pending clear confirmation when the modal is closed and reopened", async () => {
const user = setup();
const { onClose, setOpen } = renderModal(statusWith([{ name: "API_KEY", description: null, is_set: true }]));
await fieldAfterOpen(/^API_KEY/);
await user.click(screen.getByRole("button", { name: "Clear" }));
await screen.findByRole("alertdialog", { name: "Clear saved credentials" });
await user.click(screen.getByRole("button", { name: "Close", hidden: true }));
expect(onClose).toHaveBeenCalledTimes(1);
setOpen(false);
await waitFor(() => {
expect(screen.queryByRole("alertdialog")).not.toBeInTheDocument();
});
setOpen(true);
await fieldAfterOpen(/^API_KEY/);
expect(screen.queryByRole("alertdialog")).not.toBeInTheDocument();
expect(networking.clearMCPUserEnvVars).not.toHaveBeenCalled();
});
it("offers Clear only when a value is stored", async () => {
renderModal(statusWith([{ name: "API_KEY", description: null, is_set: false }]));
await fieldAfterOpen(/^API_KEY/);
expect(screen.queryByRole("button", { name: "Clear" })).not.toBeInTheDocument();
});
it("surfaces a clear failure without closing", async () => {
const user = setup();
vi.mocked(networking.clearMCPUserEnvVars).mockRejectedValue(new Error("boom"));
const { onSaved, onClose } = renderModal(statusWith([{ name: "API_KEY", description: null, is_set: true }]));
await fieldAfterOpen(/^API_KEY/);
await user.click(screen.getByRole("button", { name: "Clear" }));
const confirm = await screen.findByRole("alertdialog", { name: "Clear saved credentials" });
await user.click(within(confirm).getByRole("button", { name: "Clear credentials" }));
await waitFor(() => {
expect(networking.clearMCPUserEnvVars).toHaveBeenCalledTimes(1);
});
expect(onSaved).not.toHaveBeenCalled();
expect(onClose).not.toHaveBeenCalled();
});
it("surfaces a save failure without closing", async () => {
const user = setup();
vi.mocked(networking.storeMCPUserEnvVars).mockRejectedValue(new Error("boom"));

View file

@ -1,13 +1,21 @@
import React from "react";
import { CircleAlert, Info } from "lucide-react";
import { useMutation, useQuery } from "@tanstack/react-query";
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query";
import { z } from "zod/v4";
import { MCPServer, MCPUserEnvVarsStatus, MCPUserEnvVarSpec } from "@/components/mcp_tools/types";
import { getMCPUserEnvVars, storeMCPUserEnvVars } from "@/components/networking";
import { clearMCPUserEnvVars, getMCPUserEnvVars, storeMCPUserEnvVars } from "@/components/networking";
import { toast } from "@/lib/toast";
import { FieldGroup } from "@/components/ui/field";
import { FormField } from "@/components/shared/form/FormField";
import { Alert, AlertTitle } from "@/components/shared/Alert";
import {
AlertDialog,
AlertDialogContent,
AlertDialogDescription,
AlertDialogFooter,
AlertDialogHeader,
AlertDialogTitle,
} from "@/components/ui/alert-dialog";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { Badge } from "@/components/ui/badge";
import { StatusBadge } from "@/components/shared/table_cells/status_badge";
@ -28,6 +36,7 @@ interface UserEnvVarsFormProps {
required: readonly MCPUserEnvVarSpec[];
isSaving: boolean;
onCancel: () => void;
onClear?: () => void;
onSubmit: (values: Record<string, string>) => void;
}
@ -41,7 +50,7 @@ const buildSchema = (required: readonly MCPUserEnvVarSpec[]) =>
const emptyValues = (required: readonly MCPUserEnvVarSpec[]): Record<string, string> =>
Object.fromEntries(required.map((spec) => [spec.name, ""]));
const UserEnvVarsForm: React.FC<UserEnvVarsFormProps> = ({ required, isSaving, onCancel, onSubmit }) => {
const UserEnvVarsForm: React.FC<UserEnvVarsFormProps> = ({ required, isSaving, onCancel, onClear, onSubmit }) => {
const form = useZodForm(buildSchema(required), { defaultValues: emptyValues(required) });
return (
@ -73,6 +82,11 @@ const UserEnvVarsForm: React.FC<UserEnvVarsFormProps> = ({ required, isSaving, o
))}
</FieldGroup>
<div className="mt-6 flex items-center justify-end gap-2 border-t border-border pt-2">
{onClear && (
<Button type="button" variant="destructive" className="mr-auto" onClick={onClear} disabled={isSaving}>
Clear
</Button>
)}
<Button type="button" variant="outline" onClick={onCancel} disabled={isSaving}>
Cancel
</Button>
@ -93,12 +107,19 @@ const UserEnvVarsForm: React.FC<UserEnvVarsFormProps> = ({ required, isSaving, o
* description as the placeholder.
*/
const UserEnvVarsModal: React.FC<UserEnvVarsModalProps> = ({ server, open, accessToken, onClose, onSaved }) => {
const queryClient = useQueryClient();
const [confirmingClear, setConfirmingClear] = React.useState(false);
const close = () => {
setConfirmingClear(false);
onClose();
};
const queryKey = ["mcpUserEnvVars", server?.server_id];
const {
data: status,
isLoading,
isError,
} = useQuery<MCPUserEnvVarsStatus>({
queryKey: ["mcpUserEnvVars", server?.server_id],
queryKey,
queryFn: () => getMCPUserEnvVars(accessToken!, server!.server_id),
enabled: open && !!server && !!accessToken,
});
@ -106,15 +127,29 @@ const UserEnvVarsModal: React.FC<UserEnvVarsModalProps> = ({ server, open, acces
const saveMutation = useMutation({
mutationFn: (values: Record<string, string>) => storeMCPUserEnvVars(accessToken!, server!.server_id, values),
onSuccess: (saved) => {
queryClient.setQueryData(queryKey, saved);
toast.success("Credentials saved");
onSaved?.(saved);
onClose();
close();
},
onError: (err) => {
toast.fromError(`Failed to save env vars: ${err instanceof Error ? err.message : String(err)}`);
},
});
const clearMutation = useMutation({
mutationFn: () => clearMCPUserEnvVars(accessToken!, server!.server_id),
onSuccess: (cleared) => {
queryClient.setQueryData(queryKey, cleared);
toast.success("Credentials cleared");
onSaved?.(cleared);
close();
},
onError: (err) => {
toast.fromError(`Failed to clear env vars: ${err instanceof Error ? err.message : String(err)}`);
},
});
const handleSave = (values: Record<string, string>) => {
if (!server || !accessToken) return;
const trimmed: Record<string, string> = {};
@ -126,10 +161,15 @@ const UserEnvVarsModal: React.FC<UserEnvVarsModalProps> = ({ server, open, acces
const displayName = server?.server_name || server?.alias || server?.server_id || "MCP Server";
const required = status?.required ?? [];
const isSaving = saveMutation.isPending;
const isSaving = saveMutation.isPending || clearMutation.isPending;
const canClear = !!server && !!accessToken && required.some((spec) => spec.is_set);
const confirmClear = () => {
setConfirmingClear(false);
clearMutation.mutate();
};
return (
<Dialog open={open} onOpenChange={(opened) => !opened && onClose()}>
<Dialog open={open} onOpenChange={(opened) => !opened && close()}>
<DialogContent className="max-h-[calc(100dvh-2rem)] overflow-y-auto sm:max-w-[520px]">
<DialogHeader>
<div className="flex items-center gap-2">
@ -161,10 +201,35 @@ const UserEnvVarsModal: React.FC<UserEnvVarsModalProps> = ({ server, open, acces
credentials. Saved values are never shown back; leave an already-set field blank to keep it, or enter a
value to set or change it.
</span>
<UserEnvVarsForm required={required} isSaving={isSaving} onCancel={onClose} onSubmit={handleSave} />
<UserEnvVarsForm
required={required}
isSaving={isSaving}
onCancel={close}
onClear={canClear ? () => setConfirmingClear(true) : undefined}
onSubmit={handleSave}
/>
</>
)}
</div>
<AlertDialog open={confirmingClear} onOpenChange={(opened) => !opened && setConfirmingClear(false)}>
<AlertDialogContent>
<AlertDialogHeader>
<AlertDialogTitle>Clear saved credentials</AlertDialogTitle>
<AlertDialogDescription>
This deletes every per-user value you saved for {displayName}. Your next MCP request to this server
fails until you set them again.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>
<Button variant="outline" onClick={() => setConfirmingClear(false)}>
Cancel
</Button>
<Button variant="destructive" onClick={confirmClear}>
Clear credentials
</Button>
</AlertDialogFooter>
</AlertDialogContent>
</AlertDialog>
</DialogContent>
</Dialog>
);

View file

@ -241,6 +241,12 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID, i
return map;
}, [envVarStatuses]);
const serversWithUserFields = useMemo(
() =>
new Set((envVarStatuses ?? []).filter((status) => (status.required ?? []).length > 0).map((s) => s.server_id)),
[envVarStatuses],
);
// Deep-link via ?fill_env_vars=<server_id> — the link users follow from the
// friendly error the proxy returns when a per-user var is missing. The id is
// captured into state above and resolved to a server below; here we only strip
@ -730,6 +736,7 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID, i
key={server.server_id}
server={server}
missingUserFields={missingFieldsByServer[server.server_id]}
hasUserFields={serversWithUserFields.has(server.server_id)}
isLoadingHealth={isLoadingHealth}
isRechecking={recheckingServerIds?.has(server.server_id)}
onClick={() => {

View file

@ -7907,6 +7907,10 @@ export const storeMCPUserEnvVars = async (
});
};
export const clearMCPUserEnvVars = async (accessToken: string, serverId: string): Promise<MCPUserEnvVarsStatus> => {
return apiClient.delete<MCPUserEnvVarsStatus>(`/v1/mcp/server/${serverId}/user-env-vars`, { accessToken });
};
export const listMCPUserEnvVarStatus = async (accessToken: string): Promise<MCPUserEnvVarsStatus[]> => {
// Best-effort status badges: a failure here must not break the page, so fall
// back to an empty list rather than surfacing the error to the caller.