mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/main' into litellm_langfuse_otel_metadata_keys
This commit is contained in:
commit
d0dfafbf89
27 changed files with 4103 additions and 107 deletions
|
|
@ -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
292
.circleci/tests.yml
Normal 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 "" >>
|
||||
|
|
@ -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
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
416
tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts
Normal file
416
tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts
Normal 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();
|
||||
}
|
||||
});
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
359
tests/integration/mcp/test_mcp_user_env_vars.py
Normal file
359
tests/integration/mcp/test_mcp_user_env_vars.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
60
tests/integration/spend/test_tag_budget_enforcement.py
Normal file
60
tests/integration/spend/test_tag_budget_enforcement.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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={() => {
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue