mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
ci: merge main into litellm_chore_f0413c
This commit is contained in:
commit
15fa4cd67d
337 changed files with 21104 additions and 17773 deletions
76
.github/workflows/image-scan.yml
vendored
76
.github/workflows/image-scan.yml
vendored
|
|
@ -16,7 +16,7 @@ on:
|
|||
- backend/Dockerfile
|
||||
- backend/main.py
|
||||
- deploy/lens/**
|
||||
- litellm/proxy/lens/release.py
|
||||
- litellm/proxy/lens/**
|
||||
- tests/e2e/migrations/lens_compose_smoke.sh
|
||||
- docker/component_entrypoint.sh
|
||||
- docker/entrypoint.sh
|
||||
|
|
@ -40,6 +40,80 @@ concurrency:
|
|||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
lens-worker-image:
|
||||
name: lens-worker-image (${{ matrix.arch }})
|
||||
runs-on: ${{ matrix.runner }}
|
||||
if: >-
|
||||
github.event_name != 'pull_request' ||
|
||||
github.event.pull_request.head.repo.full_name == github.repository
|
||||
timeout-minutes: 15
|
||||
permissions:
|
||||
contents: read
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- arch: amd64
|
||||
runner: ubuntu-latest
|
||||
grype_sha256: edda0968d8827daab01d32b3cd7de192ae0915005e7bbfcfef9e68e79bc43343
|
||||
- arch: arm64
|
||||
runner: ubuntu-24.04-arm
|
||||
grype_sha256: 553e4c36d9d61349830ba6034d43b8700a7f10576d3e2f4981c0fd2b96086465
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
- name: Build the release worker
|
||||
env:
|
||||
RELEASE_TAG: sha-${{ github.sha }}
|
||||
run: docker build --build-arg LITELLM_RELEASE_TAG="${RELEASE_TAG}" -f deploy/lens/Dockerfile -t lens-worker-scan .
|
||||
- name: Verify the standalone worker on a read-only filesystem
|
||||
env:
|
||||
RELEASE_TAG: sha-${{ github.sha }}
|
||||
run: |
|
||||
docker run --rm --network none --read-only --cap-drop ALL \
|
||||
--tmpfs /tmp:rw,noexec,nosuid,size=1g --security-opt no-new-privileges \
|
||||
-e EXPECTED_RELEASE_TAG="${RELEASE_TAG}" --entrypoint python lens-worker-scan -c '
|
||||
import os
|
||||
import lens.worker
|
||||
from lens.release import release_tag
|
||||
from lens.trace_store import trace_store
|
||||
assert os.getuid() == 65532
|
||||
assert release_tag() == os.environ["EXPECTED_RELEASE_TAG"]
|
||||
with trace_store() as store:
|
||||
assert store.count() == 0
|
||||
'
|
||||
- name: Reject a dependency whose hash has changed
|
||||
run: |
|
||||
docker build --target builder -f deploy/lens/Dockerfile -t lens-worker-deps .
|
||||
sed -E 's/sha256:[0-9a-f]{64}/sha256:0000000000000000000000000000000000000000000000000000000000000000/g' \
|
||||
deploy/lens/requirements.lock > "$RUNNER_TEMP/tampered.lock"
|
||||
if docker run --rm -v "$RUNNER_TEMP/tampered.lock:/tmp/tampered.lock:ro" \
|
||||
--entrypoint uv lens-worker-deps pip sync --python /app/.venv/bin/python \
|
||||
--require-hashes --only-binary :all: --reinstall --no-cache /tmp/tampered.lock \
|
||||
> "$RUNNER_TEMP/hash-check.log" 2>&1; then
|
||||
echo "::error::Dependency hash mismatch was accepted"
|
||||
exit 1
|
||||
fi
|
||||
cat "$RUNNER_TEMP/hash-check.log"
|
||||
grep -qi 'hash mismatch' "$RUNNER_TEMP/hash-check.log"
|
||||
- name: Download Grype v0.114.0
|
||||
env:
|
||||
ARCH: ${{ matrix.arch }}
|
||||
GRYPE_SHA256: ${{ matrix.grype_sha256 }}
|
||||
run: |
|
||||
curl -fsSL --retry 3 -o "$RUNNER_TEMP/grype.tar.gz" \
|
||||
"https://github.com/anchore/grype/releases/download/v0.114.0/grype_0.114.0_linux_${ARCH}.tar.gz"
|
||||
echo "${GRYPE_SHA256} $RUNNER_TEMP/grype.tar.gz" | sha256sum -c -
|
||||
tar xzf "$RUNNER_TEMP/grype.tar.gz" -C "$RUNNER_TEMP" grype
|
||||
chmod +x "$RUNNER_TEMP/grype"
|
||||
- name: Scan the worker for fixable HIGH/CRITICAL CVEs
|
||||
env:
|
||||
GRYPE_MATCH_PYTHON_USING_CPES: "true"
|
||||
run: |
|
||||
"$RUNNER_TEMP/grype" lens-worker-scan \
|
||||
--config .grype.yaml --only-fixed --fail-on high --output table
|
||||
|
||||
image-scan:
|
||||
name: image-scan
|
||||
runs-on: ubuntu-latest
|
||||
|
|
|
|||
4
.github/workflows/lens-worker.yml
vendored
4
.github/workflows/lens-worker.yml
vendored
|
|
@ -61,11 +61,11 @@ jobs:
|
|||
-v "$PWD/tests/proxy_behavior/lens/worker_storage_smoke.py:/app/storage_smoke.py:ro" \
|
||||
--entrypoint python lens-worker:${{ github.sha }} /app/storage_smoke.py
|
||||
- name: Publish versioned Lens worker
|
||||
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
|
||||
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm' && github.ref == 'refs/heads/main'
|
||||
env:
|
||||
REGISTRY_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
REGISTRY_USER: ${{ github.actor }}
|
||||
IMAGE: ghcr.io/berriai/litellm-lens-worker:sha-${{ github.sha }}
|
||||
IMAGE: ghcr.io/berriai/litellm-lens-worker-dev:sha-${{ github.sha }}
|
||||
run: |
|
||||
printf '%s' "$REGISTRY_TOKEN" | docker login ghcr.io -u "$REGISTRY_USER" --password-stdin
|
||||
docker tag lens-worker:${{ github.sha }} "$IMAGE"
|
||||
|
|
|
|||
4
Makefile
4
Makefile
|
|
@ -58,7 +58,7 @@ help:
|
|||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
|
||||
@echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)"
|
||||
@echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (ARGS=\"--seed large\", LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)"
|
||||
@echo ""
|
||||
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
|
||||
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
|
||||
|
|
@ -313,7 +313,7 @@ rust-sqlx-prepare:
|
|||
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
|
||||
|
||||
lens-dev:
|
||||
./scripts/lens_dev.sh
|
||||
./scripts/lens_dev.sh $(ARGS)
|
||||
|
||||
test: install-test-deps
|
||||
$(UV_RUN) pytest tests/
|
||||
|
|
|
|||
|
|
@ -1,9 +1,27 @@
|
|||
FROM python:3.12-slim
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
||||
FROM $LITELLM_BUILD_IMAGE AS builder
|
||||
COPY --from=uvbin /uv /usr/local/bin/uv
|
||||
RUN apk add --no-cache python-3.13
|
||||
ENV UV_PYTHON_DOWNLOADS=0 UV_LINK_MODE=copy
|
||||
WORKDIR /app
|
||||
COPY deploy/lens/requirements.lock /tmp/requirements.lock
|
||||
RUN uv venv --python python3.13 /app/.venv && \
|
||||
uv pip sync --python /app/.venv/bin/python --require-hashes --only-binary :all: /tmp/requirements.lock
|
||||
|
||||
FROM $LITELLM_RUNTIME_IMAGE AS runtime
|
||||
ARG LITELLM_RELEASE_TAG=""
|
||||
RUN : "${LITELLM_RELEASE_TAG:?Pass --build-arg LITELLM_RELEASE_TAG matching the gateway}"
|
||||
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG}
|
||||
RUN apk add --no-cache python-3.13
|
||||
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} \
|
||||
PATH="/app/.venv/bin:${PATH}" \
|
||||
PYTHONDONTWRITEBYTECODE=1
|
||||
WORKDIR /app
|
||||
RUN pip install --no-cache-dir httpx==0.28.1 pydantic==2.11.7
|
||||
COPY --from=builder /app/.venv /app/.venv
|
||||
COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py litellm/proxy/lens/release.py /app/lens/
|
||||
COPY litellm/proxy/lens/prompts/ /app/lens/prompts/
|
||||
USER 65532:65532
|
||||
|
|
|
|||
|
|
@ -1,4 +1,7 @@
|
|||
**
|
||||
!deploy/
|
||||
!deploy/lens/
|
||||
!deploy/lens/requirements.lock
|
||||
!litellm/
|
||||
!litellm/proxy/
|
||||
!litellm/proxy/lens/
|
||||
|
|
@ -7,5 +10,6 @@
|
|||
!litellm/proxy/lens/trace_store.py
|
||||
!litellm/proxy/lens/analysis.py
|
||||
!litellm/proxy/lens/worker.py
|
||||
!litellm/proxy/lens/release.py
|
||||
!litellm/proxy/lens/prompts/
|
||||
!litellm/proxy/lens/prompts/**
|
||||
|
|
|
|||
|
|
@ -2,62 +2,74 @@
|
|||
|
||||
Lens reviews recorded activity and saves evidence-linked findings in the LiteLLM dashboard under Observability, Lens (`/ui/lens/`)
|
||||
|
||||
## Install the release stack
|
||||
## Install
|
||||
|
||||
Each stable, RC, and dev release containing Lens publishes the worker at the same version on GHCR and Docker Hub. Use the [LiteLLM releases page](https://github.com/BerriAI/litellm/releases) to select a version that includes the coordinated worker release
|
||||
Build LiteLLM and its worker from the same source commit with the same release identity. The worker runs separately and connects to your gateway using a limited worker token
|
||||
|
||||
For a new local installation, install Docker with Compose, download the two release files, and create a private environment file. Replace `X.Y.Z` with the release version, without `v` (RCs use `X.Y.Z-rc.N`)
|
||||
### New local installation
|
||||
|
||||
Install Docker with Compose and Git. This builds LiteLLM and its worker from the same checkout and starts the existing local tracing stack:
|
||||
|
||||
```bash
|
||||
mkdir litellm-lens
|
||||
cd litellm-lens
|
||||
LENS_RELEASE=X.Y.Z
|
||||
curl -fSLo compose.yaml "https://raw.githubusercontent.com/BerriAI/litellm/v${LENS_RELEASE}/deploy/lens/stack.yaml"
|
||||
curl -fSLo config.yaml "https://raw.githubusercontent.com/BerriAI/litellm/v${LENS_RELEASE}/deploy/lens/config.yaml"
|
||||
umask 077
|
||||
printf 'LITELLM_VERSION=%s\nLITELLM_MASTER_KEY=sk-%s\nLITELLM_SALT_KEY=sk-%s\n' \
|
||||
"$LENS_RELEASE" "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" > .env
|
||||
printf 'POSTGRES_PASSWORD=%s\nCLICKHOUSE_PASSWORD=%s\n' \
|
||||
"$(openssl rand -hex 32)" "$(openssl rand -hex 32)" >> .env
|
||||
docker compose up -d
|
||||
git clone https://github.com/BerriAI/litellm.git
|
||||
cd litellm
|
||||
export LITELLM_RELEASE_TAG="sha-$(git rev-parse HEAD)"
|
||||
export LENS_WORKER_IMAGE="litellm-lens-worker:${LITELLM_RELEASE_TAG}"
|
||||
export OPENAI_API_KEY='sk-...'
|
||||
docker build --build-arg LITELLM_RELEASE_TAG="$LITELLM_RELEASE_TAG" \
|
||||
-f deploy/lens/Dockerfile -t "$LENS_WORKER_IMAGE" .
|
||||
docker compose -f docker/docker-compose.tracing.yml up -d --build
|
||||
```
|
||||
|
||||
Open `http://localhost:4000/ui/`, log in as `admin` with `LITELLM_MASTER_KEY` from `.env`, and add a model in the dashboard. In Lens, select **Connect worker**, choose that model and a monthly budget, then **Get install command**. Expand **Using Docker Compose or Helm?**, copy the worker token, and add `LENS_WORKER_TOKEN=<token>` to `.env`
|
||||
Open `http://localhost:4002/ui/` and sign in as `admin` with password `sk-1234`. Go to **Lens > Investigations > Connect worker**, choose a model and monthly budget, then **Get install command**. Expand **Using Docker Compose or Helm?** and copy the worker token. In the same terminal, run:
|
||||
|
||||
```bash
|
||||
docker compose --profile lens up -d
|
||||
export LITELLM_URL=http://litellm:4000
|
||||
export LENS_WORKER_TOKEN='<paste-your-worker-token>'
|
||||
docker compose -f docker/docker-compose.tracing.yml -f deploy/lens/compose.yaml up -d
|
||||
```
|
||||
|
||||
The stack starts LiteLLM, PostgreSQL, ClickHouse, and the worker from published images. The dashboard shows **Worker connected**. The worker has a limited token, no database credentials, and no provider keys. The stack exposes only the dashboard on localhost; use your normal ingress and managed databases for a public production deployment
|
||||
The worker joins the gateway's Docker network, and the dashboard shows **Worker connected**. Save the token privately for restarts and upgrades
|
||||
|
||||
Keep `.env` private and preserve its salt key. Keep both named database volumes. To upgrade, wait for active investigations to finish, stop the worker, change only `LITELLM_VERSION`, then pull and recreate the stack:
|
||||
This stack is for local evaluation: it binds to localhost and uses development database credentials. For a hosted deployment, keep your normal database, keys, networking, and deployment process. Build both images from one source revision with the same `LITELLM_RELEASE_TAG`, publish the worker to your registry, and set `LENS_WORKER_IMAGE` on LiteLLM to that image
|
||||
|
||||
### Existing LiteLLM installation
|
||||
|
||||
Keep your deployment and PostgreSQL database. A working gateway/worker pair can stay as it is until you upgrade both. For a gateway built from source, use its exact commit and `LITELLM_RELEASE_TAG`; a release version or the latest commit on `main` is not a substitute for that source identity
|
||||
|
||||
The public development package is `ghcr.io/berriai/litellm-lens-worker-dev:sha-<full-commit>`. It publishes amd64 images on Lens-related changes, so an arbitrary source commit may have no image. Check the exact image exists before using it. If it is unavailable, your gateway uses a different release identity, or you need native arm64, build the worker from the gateway's checkout:
|
||||
|
||||
```bash
|
||||
docker compose --profile lens stop lens-worker
|
||||
# Update LITELLM_VERSION in .env to the new release
|
||||
docker compose --profile lens pull
|
||||
docker compose --profile lens up -d
|
||||
export LITELLM_RELEASE_TAG='<gateway-release-identity>'
|
||||
export LENS_WORKER_IMAGE='<your-registry>/litellm-lens-worker:<your-image-tag>'
|
||||
docker build --build-arg LITELLM_RELEASE_TAG="$LITELLM_RELEASE_TAG" \
|
||||
-f deploy/lens/Dockerfile -t "$LENS_WORKER_IMAGE" .
|
||||
```
|
||||
|
||||
This preserves your investigations, findings, model credentials, and worker token. Never use `down -v` during an upgrade. If moving from an existing installation, keep its databases and add the standalone worker instead of creating an empty replacement stack
|
||||
For a remote worker host, publish that image to a registry the host can pull from. Set the gateway's `LENS_WORKER_IMAGE` to the resulting image reference, restart the gateway using its normal deployment process, then copy its install command. Prefer the published image digest for hosted installations. Do not change the gateway's release identity just to accept another worker
|
||||
|
||||
For Kubernetes or Render, run the standalone worker using `LITELLM_URL` and `LENS_WORKER_TOKEN` from setup. Keep existing databases and secrets. The worker needs no inbound port.
|
||||
|
||||
## Helm
|
||||
|
||||
The componentized `helm/litellm` chart includes an optional Lens worker. Configure PostgreSQL and ClickHouse as usual, install the chart, then obtain a limited worker token from Lens setup. Store it in a Kubernetes Secret and enable the worker in your values:
|
||||
The componentized source chart at `helm/litellm` includes an optional Lens worker. Use the chart from the same checkout as your gateway and keep your component image overrides in your values. Configure PostgreSQL and ClickHouse as usual, install the chart, then obtain a limited worker token from Lens setup. Store it in a Kubernetes Secret and enable the worker in your values:
|
||||
|
||||
```yaml
|
||||
lensWorker:
|
||||
enabled: true
|
||||
image:
|
||||
repository: <your-worker-image-repository>
|
||||
digest: sha256:<matching-worker-image-digest>
|
||||
tokenSecret:
|
||||
name: litellm-lens-worker
|
||||
key: token
|
||||
```
|
||||
|
||||
The worker image defaults to the chart's application version, and the chart connects it to the backend service. Keep these values and the Secret when upgrading the chart so the gateway and worker upgrade together. `lensWorker.replicaCount` controls simultaneous investigations. To use a private registry or external proxy, set `lensWorker.image.repository`, `lensWorker.image.tag`, and `lensWorker.url`. The dashboard uses the chart's worker image for standalone install commands too
|
||||
Set the worker repository and digest explicitly to an image built from the gateway's source commit and release identity. The chart connects the worker to the backend service. Keep these values and the Secret when upgrading the chart and update the gateway and worker image overrides together. `lensWorker.replicaCount` controls simultaneous investigations. To use a private registry or external proxy, set `lensWorker.image.repository`, `lensWorker.image.digest` (or `tag` for a source build), and `lensWorker.url`. A digest takes precedence over the tag. The dashboard uses the chart's worker image for standalone install commands too
|
||||
|
||||
## Standalone worker
|
||||
|
||||
Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL and agent tracing. Configure one ClickHouse URL for trace writes, bounded reads, and Lens queries:
|
||||
Start with a source deployment that includes Lens, PostgreSQL, and agent tracing, and prepare its matching worker as described above. Configure one ClickHouse URL for trace writes, bounded reads, and Lens queries:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
|
|
@ -74,13 +86,13 @@ Retention changes require a proxy restart. ClickHouse removes expired rows durin
|
|||
|
||||
In **Lens > Investigations**, click **Connect worker**, choose an analysis model and monthly limit, then **Get install command**. Use **Advanced options** to select an existing virtual key or change the proxy URL if the server running Docker needs a different network address. Copy the command and run it on your server. The dashboard shows **Worker connected** when the container checks in
|
||||
|
||||
The command already contains the compatible worker image and one worker token. The selected virtual key stays on the proxy; its secret is never sent to the worker. No source checkout, environment file, or second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis
|
||||
The command already contains the compatible worker image and one worker token. The selected virtual key stays on the proxy; its secret is never sent to the worker. Once the matching image is available on the worker host, no second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis
|
||||
|
||||
The dashboard selects the worker image matching the running gateway release. Release images support Linux amd64 and arm64. CI also publishes `:sha-<commit>` development images; use those only with a gateway built from the same commit and release tag
|
||||
The dashboard uses the gateway's `LENS_WORKER_IMAGE` override when set. Public `:sha-<commit>` development images must match both the gateway commit and release identity. Build from source for the worker host's native architecture
|
||||
|
||||
After upgrading the gateway, update the worker image and redeploy it while keeping its proxy URL and token. Existing containers do not update automatically. If an investigation reports a worker compatibility error, update the image before retrying
|
||||
|
||||
For deployments managed with Compose, download `compose.yaml` and provide `LITELLM_URL`, `LENS_WORKER_TOKEN`, and `LITELLM_VERSION` (without `v`) in a private environment file. To use another registry, set `LENS_WORKER_IMAGE` to the compatible image instead of setting a version:
|
||||
For deployments managed with Compose, download `compose.yaml` and provide `LITELLM_URL`, `LENS_WORKER_TOKEN`, and an explicit `LENS_WORKER_IMAGE` in a private environment file:
|
||||
|
||||
```bash
|
||||
docker compose --env-file /path/to/lens.env -f compose.yaml up -d
|
||||
|
|
@ -160,6 +172,38 @@ curl "$LITELLM_URL/lens/$LENS_ID/runs/$BATCH_ID" -H "Authorization: Bearer $LITE
|
|||
|
||||
Creation queues the first batch. Posting to `/lens/{id}/runs` queues another, or returns the existing active batch. The run response contains its ID under `jobs[0].id`. Poll the batch URL for status, findings and assessments. List responses omit large result payloads; request a batch to retrieve them. Supply an optional complete `settings` object on the runs POST for a one-off override; the saved lens stays unchanged. Selection accepts `team_id`, exact `filters`, and opaque `execution_ids` returned by `/lens/preview/sample`. Preview accepts `offset` and `as_of` to keep the time window fixed while paging. Feedback uses `PATCH /lens/{id}/findings/{finding_id}` with `status` and `reason`
|
||||
|
||||
## Local development
|
||||
|
||||
`make lens-dev ARGS=--seed` starts the full dev stack. The live dashboard is at `http://localhost:3000/ui/lens/`, with login at `http://localhost:3000/ui/login/`. Next.js forwards API requests to the proxy on port 4000, so login and navigation stay in the live UI and edits hot-reload
|
||||
|
||||
The default is Next.js dev with no production build (`LENS_DEV_BUILD_UI=0`). Set `LENS_DEV_BUILD_UI=1` when you also want a fresh static dashboard at `http://localhost:4000/ui/`. Build output goes to `.lens-dev/logs/ui-build.log`; a failed build stops startup. Both modes keep the live dashboard on port 3000. Startup checks the live login route before seeding and fails with the UI log path if Next.js exits. `LENS_DEV_STARTUP_TIMEOUT_SECONDS` controls startup readiness retries (default 300; `LENS_DEV_READINESS_REQUEST_TIMEOUT_SECONDS` caps each HTTP probe, default 5)
|
||||
|
||||
For local fixture data, run `make lens-dev ARGS=--seed`. Use `make lens-dev ARGS="--seed large"` for 2,000 fixture copies, over one million spans and linked request logs. To seed a running stack without restarting it, use `make lens-dev ARGS="--seed-only --seed large --copies 100"`. The default profile replays one copy of every checked-in capture through authenticated `/v1/traces`, including failures, retries, streaming and multiple agent frameworks. Large seeds use the same parser and compressed ClickHouse writer in batches of four copies, and write matching request logs to PostgreSQL. The first and last batches verify linked spend totals through the proxy
|
||||
|
||||
Seeds append fresh IDs on every invocation and spread copies over recent timestamps. Restarts without `SEED` do not add data. Lens excludes activity received in the last two minutes, so wait two minutes after seeding before checking investigation previews. `LENS_DEV_SEED_COPIES` overrides total copies, and `LENS_DEV_SEED_BATCH_COPIES` overrides copies per bulk insert (default 4, about 2,000 spans). Start with four or fewer on a constrained machine. Larger batches still respect the existing ClickHouse insert size limit; each capture is decoded separately within the OTLP safety budget. Large seeds test data volume and pagination, rather than concurrent ingestion throughput or review accuracy. They can use substantial disk space; adjust `--copies` for your machine. Seeding expects the generated local tracing configuration. The old `run_tracing_proxy_local.sh --seed` command forwards to Lens dev, using its ports and saved master key
|
||||
|
||||
Local ingestion limits are explicit and configurable. Set OTLP and ClickHouse variables before starting the proxy and seeder so both processes use the same settings. Invalid, zero and negative values fail instead of silently falling back. Changing these limits does not require rebuilding Rust
|
||||
|
||||
| Environment variable | Default | Controls |
|
||||
| --- | --- | --- |
|
||||
| `LENS_DEV_SEED_COPIES` | 1 default, 2000 large | Total fixture copies |
|
||||
| `LENS_DEV_SEED_BATCH_COPIES` | 4 | Copies per bulk insert |
|
||||
| `LENS_DEV_SEED_TIMEOUT_SECONDS` | 120 | Seeder HTTP timeout |
|
||||
| `OTLP_MAX_BODY_BYTES` | 16777216 | HTTP body and decompressed payload bytes |
|
||||
| `OTLP_MAX_CONCURRENT_INGESTS` | 2 | Concurrent proxy ingestion requests |
|
||||
| `OTLP_MAX_ATTRIBUTE_VALUE_BYTES` | 65536 | Stored attribute/content bytes |
|
||||
| `OTLP_MAX_DECODE_DEPTH` | 32 | Nested decode depth |
|
||||
| `OTLP_MAX_DECODE_NODES` | 65536 | JSON values or protobuf fields per export |
|
||||
| `OTLP_MAX_SPANS` | 4096 | Spans per export |
|
||||
| `OTLP_MAX_ATTRIBUTES` | 256 | Attributes per resource, scope, span, event or link |
|
||||
| `OTLP_MAX_EVENTS` | 256 | Events per span |
|
||||
| `OTLP_MAX_LINKS` | 256 | Links per span |
|
||||
| `OTLP_MAX_DECODED_SPAN_BYTES` | 16777216 | Decoded span allocation budget |
|
||||
| `CLICKHOUSE_TRACE_MAX_INSERT_BYTES` | 67108864 | Encoded trace or spend insert bytes |
|
||||
| `CLICKHOUSE_INSERT_TIMEOUT_SECONDS` | 30 | ClickHouse insert HTTP timeout |
|
||||
|
||||
The wire parsers also enforce their library recursion limits (128 levels for JSON, 100 for protobuf). Raising the configured depth does not remove those parser limits. Bulk seeding parses each capture separately, keeping the per-export limits distinct from the bulk insert limit. Use smaller batches if an insert exceeds its byte budget. For example, `LENS_DEV_SEED_COPIES=100 LENS_DEV_SEED_BATCH_COPIES=2 make lens-dev ARGS="--seed large"`
|
||||
|
||||
## Quality evaluation
|
||||
|
||||
Run the checked-in cases against a configured real model. Expected labels are used only for scoring, never passed to the model. Dev and held-out cases include missing outcomes, failed tools, recovery, handoffs, unsupported claims, repeated work, long evidence and prompt injection. The background option adds clean arithmetic traces to test rare-issue discovery at scale; those repeated synthetic cases do not establish accuracy on every production workload
|
||||
|
|
@ -187,10 +231,15 @@ Upgrades using `--use_prisma_db_push` stop before schema changes if any legacy L
|
|||
|
||||
## Release compatibility
|
||||
|
||||
Released gateway and worker images carry `LITELLM_RELEASE_TAG`. A worker announces its release and protocol before claiming an investigation. A mismatch returns HTTP 409 with the required image, leaving queued investigations untouched. During a rolling upgrade, workers wait for a gateway from their release
|
||||
Gateway and worker builds carry the same `LITELLM_RELEASE_TAG`. A worker announces its release and protocol before claiming an investigation. A mismatch returns HTTP 409 with the required image, leaving queued investigations untouched. During a rolling upgrade, workers wait for a gateway from their release
|
||||
|
||||
The dashboard reads its image from the running gateway. `LENS_WORKER_IMAGE` overrides the registry/image for private deployments. Worker-only Compose accepts `LITELLM_VERSION` (without `v`) or an explicit `LENS_WORKER_IMAGE`. Release workers are available as `ghcr.io/berriai/litellm-lens-worker:vX.Y.Z` and `docker.io/litellm/litellm-lens-worker:vX.Y.Z`, including matching RC/dev suffixes, on amd64 and arm64
|
||||
The dashboard reads its image from the running gateway. `LENS_WORKER_IMAGE` overrides the registry/image for private deployments. Set an explicit `LENS_WORKER_IMAGE` for worker-only Compose. Verify that the image exists and matches the gateway before deploying it
|
||||
|
||||
For source development, use `make lens-dev`, which gives the proxy and source worker the same commit identity. For custom containers, build both from the same checkout with `--build-arg LITELLM_RELEASE_TAG=sha-$(git rev-parse HEAD)` and set the proxy's `LENS_WORKER_IMAGE` to the worker image you built. An unlabelled custom build refuses worker setup and claims instead of guessing from the Python package version. Normal package-index installations use their installed release version
|
||||
|
||||
The hourly development pipeline pins all component images to the same selected commit and publishes its chart only after every build and worker smoke test succeeds. The public commit-tagged worker workflow publishes on Lens-related changes, so an arbitrary `main` commit may require building your own pair; do not substitute the newest available worker
|
||||
The hourly development pipeline pins all component images to the same selected commit and publishes its chart only after every build and worker smoke test succeeds. The public commit-tagged worker workflow publishes to `ghcr.io/berriai/litellm-lens-worker-dev` on Lens-related changes, so an arbitrary `main` commit may require building your own pair; do not substitute the newest available worker
|
||||
|
||||
|
||||
## Worker dependencies
|
||||
|
||||
The worker uses the same digest-pinned Wolfi base and Python version as the component images. Python dependencies and their hashes are locked in `deploy/lens/requirements.lock`. To update them, edit `deploy/lens/requirements.in`, then run `uv pip compile --universal --python-version 3.13 --generate-hashes --no-emit-index-url deploy/lens/requirements.in -o deploy/lens/requirements.lock`. The image installs only the locked wheels with hash verification. CI builds and scans both native architectures
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
services:
|
||||
lens-worker:
|
||||
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION:?Set LITELLM_VERSION to the gateway release, without the v prefix}}
|
||||
image: ${LENS_WORKER_IMAGE:-${LITELLM_VERSION:+ghcr.io/berriai/litellm-lens-worker:v}${LITELLM_VERSION:-}}
|
||||
environment:
|
||||
LITELLM_URL: ${LITELLM_URL:?Set the URL reachable from this container}
|
||||
LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:?Create a worker credential in the Lens UI}
|
||||
|
|
|
|||
2
deploy/lens/requirements.in
Normal file
2
deploy/lens/requirements.in
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
httpx==0.28.1
|
||||
pydantic==2.13.4
|
||||
172
deploy/lens/requirements.lock
Normal file
172
deploy/lens/requirements.lock
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
# This file was autogenerated by uv via the following command:
|
||||
# uv pip compile --universal --python-version 3.13 --generate-hashes --no-emit-index-url deploy/lens/requirements.in -o deploy/lens/requirements.lock
|
||||
annotated-types==0.8.0 \
|
||||
--hash=sha256:13b2beaad985e05e2d6407ee4c4f35590b11f8d693a258a561055cac8f64cab7 \
|
||||
--hash=sha256:f072f4d804ea359e4eaf198b1af7a8b0943881a87f31bb764f8bf219bb9419e0
|
||||
# via pydantic
|
||||
anyio==4.15.1 \
|
||||
--hash=sha256:6152fdbbf9a77fdec97731721bebf7c4c44f7c29b424b0065826173efc7ed101 \
|
||||
--hash=sha256:9f28306018cbd6d329e64a36d58256edff76dd996fe423bc957326e578b82a94
|
||||
# via httpx
|
||||
certifi==2026.7.22 \
|
||||
--hash=sha256:62f22742b58a1a33014a2b6b706588a8d7e2a88ae7bd1a6ebe8c992928483775 \
|
||||
--hash=sha256:741e2c3b351ddf169a738da9f2c048608ff7f2c5cc02f1ebc6b118bb090d5d55
|
||||
# via
|
||||
# httpcore
|
||||
# httpx
|
||||
h11==0.16.0 \
|
||||
--hash=sha256:4e35b956cf45792e4caa5885e69fba00bdbc6ffafbfa020300e549b208ee5ff1 \
|
||||
--hash=sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86
|
||||
# via httpcore
|
||||
httpcore==1.0.9 \
|
||||
--hash=sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55 \
|
||||
--hash=sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8
|
||||
# via httpx
|
||||
httpx==0.28.1 \
|
||||
--hash=sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc \
|
||||
--hash=sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad
|
||||
# via -r deploy/lens/requirements.in
|
||||
idna==3.20 \
|
||||
--hash=sha256:a7db850025b95ded1eae8a46181a1a6c56c92c96f0e2b005d9ff8dc0210cab44 \
|
||||
--hash=sha256:ab7ae7122974553370f0bdb919e1a960b2cd1bc1ef0276416d896db81c14582c
|
||||
# via
|
||||
# anyio
|
||||
# httpx
|
||||
pydantic==2.13.4 \
|
||||
--hash=sha256:45a282cde31d808236fd7ea9d919b128653c8b38b393d1c4ab335c62924d9aba \
|
||||
--hash=sha256:c40756b57adaa8b1efeeced5c196f3f3b7c435f90e84ea7f443901bec8099ef6
|
||||
# via -r deploy/lens/requirements.in
|
||||
pydantic-core==2.46.4 \
|
||||
--hash=sha256:00c603d540afdd6b80eb39f078f33ebd46211f02f33e34a32d9f053bba711de0 \
|
||||
--hash=sha256:0186750b482eefa11d7f435892b09c5c606193ef3375bcf94aa00ae6bfb66262 \
|
||||
--hash=sha256:041bde0a48fd37cf71cab1c9d56d3e8625a3793fef1f7dd232b3ff37e978ecda \
|
||||
--hash=sha256:0c563b08bca408dc7f65f700633d8442fffb2421fc47b8101377e9fd65051ff0 \
|
||||
--hash=sha256:0cbe8b01f948de4286c74cdd6c667aceb38f5c1e26f0693b3983d9d74887c65e \
|
||||
--hash=sha256:0ce40cd7b21210e99342afafbd4d0f76d784eb5b1d60f3bdc566be4983c6c73b \
|
||||
--hash=sha256:0e96592440881c74a213e5ad528e2b24d3d4f940de2766bed9010ab1d9e51594 \
|
||||
--hash=sha256:10e17cbb10a330363733efc4d7c4d0dd827ac0909b8f6a6542298fed1ea62f29 \
|
||||
--hash=sha256:133878133d271ade3d41d1bfb2a45ec38dbdbda40bc065921c6b04e4630127e2 \
|
||||
--hash=sha256:14d4edf427bdcf950a8a02d7cb44a08614388dd6e1bdcbf4f67504fa7887da9c \
|
||||
--hash=sha256:14f4c5d6db102bd796a627bbb3a17b4cf4574b9ae861d8b7c9a9661c6dd3362d \
|
||||
--hash=sha256:17299feefe090f2caa5b8e37222bb5f663e4935a8bfa6931d4102e5df1a9f398 \
|
||||
--hash=sha256:184c081504d17f1c1066e430e117142b2c77d9448a97f7b65c6ac9fd9aee238d \
|
||||
--hash=sha256:18e5ceec2ab67e6d5f1a9085e5a24c9c4e2ac4545730bfe668680bca05e555f3 \
|
||||
--hash=sha256:19e51f073cd3df251856a8a4189fbdf1de4012c3ebacfb1884f94f1eb406079f \
|
||||
--hash=sha256:1a7dd0b3ee80d90150e3495a3a13ac34dbcbfd4f012996a6a1d8900e91b5c0fb \
|
||||
--hash=sha256:1d8ba486450b14f3b1d63bc521d410ec7565e52f887b9fb671791886436a42f7 \
|
||||
--hash=sha256:2108ba5c1c1eca18030634489dc544844144ee36357f2f9f780b93e7ddbb44b5 \
|
||||
--hash=sha256:228ee9bae8bef5b1e97ec58302f80357c37199e0d0a99174e138d28e6957b9d9 \
|
||||
--hash=sha256:23ace664830ee0bfe014a0c7bc248b1f7f25ed7ad103852c317624a1083af462 \
|
||||
--hash=sha256:2412e734dcb48da14d4e4006b82b46b74f2518b8a26ee7e58c6844a6cd6d03c4 \
|
||||
--hash=sha256:29c61fc04a3d840155ff08e475a04809278972fe6aef51e2720554e96367e34b \
|
||||
--hash=sha256:2f84c03c8607173d16b5a854ec68a2f9079ae03237a54fb506d13af47e1d018d \
|
||||
--hash=sha256:3009f12e4e90b7f88b4f9adb1b0c4a3d58fe7820f3238c190047209d148026df \
|
||||
--hash=sha256:3245406455a5d98187ec35530fd772b1d799b26667980872c8d4614991e2c4a2 \
|
||||
--hash=sha256:3447661d99f75a3683a4cf5c87da72f2161964611864dbbeac7fbb118bb4bfc0 \
|
||||
--hash=sha256:372429a130e469c9cd698925ce5fc50940b7a1336b0d82038e63d5bbc4edc519 \
|
||||
--hash=sha256:395aebd9183f9d112f569aeb5b2214d1a10a33bec8456447f7fbdfa51d38d4cd \
|
||||
--hash=sha256:3a233125ac121aa3ffba9a2b59edfc4a985a76092dc8279586ab4b71390875e7 \
|
||||
--hash=sha256:3be77f45df024d789a672ae34f8b06fb346c4f9f46ea714956660ea4862e89ac \
|
||||
--hash=sha256:3bf92c5d0e00fefaab325a4d27828fe6b6e2a21848686b5b60d2d9eeb09d76c6 \
|
||||
--hash=sha256:3ecbc122d18468d06ca279dc26a8c2e2d5acb10943bb35e36ae92096dc3b5565 \
|
||||
--hash=sha256:3fb702cd90b0446a3a1c5e470bfa0dd23c0233b676a9099ddcc964fa6ca13898 \
|
||||
--hash=sha256:428e04521a40150c85216fc8b85e8d39fece235a9cf5e383761238c7fa9b96fb \
|
||||
--hash=sha256:432c179df7874eeb73307aad2df0755e1ae0efa61ff0ea89b93e194411ae3928 \
|
||||
--hash=sha256:4a05d69cba51d852c5c3e92758653245a50c0b646ced0cf05bd793ed592839d6 \
|
||||
--hash=sha256:4c63ebc82684aa89d9a3bcbd13d515b3be44250dc68dd3bd81526c1cb31286c3 \
|
||||
--hash=sha256:4fc73cb559bdb54b1134a706a2802a4cddd27a0633f5abb7e53056268751ac6a \
|
||||
--hash=sha256:4fcbe087dbc2068af7eda3aa87634eba216dbda64d1ae73c8684b621d33f6596 \
|
||||
--hash=sha256:56cb4851bcaf3d117eddcef4fe66afd750a50274b0da8e22be256d10e5611987 \
|
||||
--hash=sha256:5855698a4856556d86e8e6cd8434bc3ac0314ee8e12089ae0e143f64c6256e4e \
|
||||
--hash=sha256:5a4330cdbc57162e4b3aa303f588ba752257694c9c9be3e7ebb11b4aca659b5d \
|
||||
--hash=sha256:5b712b53160b79a5850310b912a5ef8e57e56947c8ad690c227f5c9d7e561712 \
|
||||
--hash=sha256:5d5902252db0d3cedf8d4a1bc68f70eeb430f7e4c7104c8c476753519b423008 \
|
||||
--hash=sha256:617d7e2ca7dcb8c5cf6bcb8c59b8832c94b36196bbf1cbd1bfb56ed341905edd \
|
||||
--hash=sha256:62f875393d7f270851f20523dd2e29f082bcc82292d66db2b64ea71f64b6e1c1 \
|
||||
--hash=sha256:633147d34cf4550417f12e2b1a0383973bdf5cdfde212cb09e9a581cf10820be \
|
||||
--hash=sha256:66ce7632c22d837c95301830e111ad0128a32b8207533b60896a96c4915192ea \
|
||||
--hash=sha256:6b3ace8194b0e5204818c92802dcdca7fc6d88aabbb799d7c795540d9cd6d292 \
|
||||
--hash=sha256:6f2eeda33a839975441c86a4119e1383c50b47faf0cbb5176985565c6bb02c33 \
|
||||
--hash=sha256:7027560ee92211647d0d34e3f7cd6f50da56399d26a9c8ad0da286d3869a53f3 \
|
||||
--hash=sha256:7283d57845ecf5a163403eb0702dfc220cc4fbdd18919cb5ccea4f95ee1cdab4 \
|
||||
--hash=sha256:7a5f930472650a82629163023e630d160863fce524c616f4e5186e5de9d9a49b \
|
||||
--hash=sha256:7bfb192b3f4b9e8a89b6277b6ce787564f62cfd272055f6e685726b111dc7826 \
|
||||
--hash=sha256:811ff8e9c313ab425368bcbb36e5c4ebd7108c2bbf4e4089cfbb0b01eff63fac \
|
||||
--hash=sha256:8233f2947cf85404441fd7e0085f53b10c93e0ee78611099b5c7237e36aacbf7 \
|
||||
--hash=sha256:82cf5301172168103724d49a1444d3378cb20cdee30b116a1bd6031236298a5d \
|
||||
--hash=sha256:8358a950c8909158e3df31538a7e4edc2d7265a7c54b47f0864d9e5bae9dcebf \
|
||||
--hash=sha256:85bb3611ff1802f3ee7fdd7dbff26b56f343fb432d57a4728fdd49b6ef35e2f4 \
|
||||
--hash=sha256:86e1a4418c6cd97d60c95c71164158eaf7324fae7b0923264016baa993eba6fc \
|
||||
--hash=sha256:8b9bab013d1c7a79d3501ff86d0bc9c31bf587db4551677b96bec07df78c6b15 \
|
||||
--hash=sha256:8c5dac79fa1614d1e06ca695109c6105923bd9c7d1d6c918d4e637b7e6b32fd3 \
|
||||
--hash=sha256:8d0820e8192167f80d88d64038e609c31452eeca865b4e1d9950a27a4609b00b \
|
||||
--hash=sha256:8daafc69c93ee8a0204506a3b6b30f586ef54028f52aeeeb5c4cfc5184fd5914 \
|
||||
--hash=sha256:9037063db01f09b09e237c282b6792bd4da634b5402c4e7f0c61effed7701a04 \
|
||||
--hash=sha256:905a0ed8ea6f2d61c1738835f99b699348d7857379083e5fc497fa0c967a407c \
|
||||
--hash=sha256:90884113d8b48f760e9587002789ddd741e76ab9f89518cd1e43b1f1a52ec44b \
|
||||
--hash=sha256:91a06d2e259ecfbd8c901d70c3c507900458498142b3026a296b7de4d1322cc9 \
|
||||
--hash=sha256:926c9541b14b12b1681dca8a0b75feb510b06c6341b70a8e500c2fdcff837cce \
|
||||
--hash=sha256:9401557acd873c3a7f3eb9383edef8ac4968f9510e340f4808d427e75667e7b4 \
|
||||
--hash=sha256:9551187363ffc0de2a00b2e47c25aeaeb1020b69b668762966df15fc5659dd5a \
|
||||
--hash=sha256:962ccbab7b642487b1d8b7df90ef677e03134cf1fd8880bf698649b22a69371f \
|
||||
--hash=sha256:97e7cf2be5c77b7d1a9713a05605d49460d02c6078d38d8bef3cbe323c548424 \
|
||||
--hash=sha256:9aa768456404a8bf48a4406685ac2bec8e72b62c69313734fa3b73cf33b3a894 \
|
||||
--hash=sha256:9bc519fbf2b7578398853d815009ae5e4d4603d12f4e3f91da8c06852d3da3e9 \
|
||||
--hash=sha256:9d56801be94b86a9da183e5f3766e6310752b99ff647e38b09a9500d88e46e76 \
|
||||
--hash=sha256:9f444c499b3eefd3a92e348059471ea0c3a6e303d9c1cec09fa748fd9f895201 \
|
||||
--hash=sha256:9fa8ae11da9e2b3126c6426f147e0fba88d96d65921799bb30c6abd1cb2c97fb \
|
||||
--hash=sha256:a0f62d0a58f4e7da165457e995725421e0064f2255d8eccebc49f41bbc23b109 \
|
||||
--hash=sha256:a396dcc17e5a0b164dbe026896245a4fa9ff402edca1dff0be3d53a517f74de4 \
|
||||
--hash=sha256:aaa2a54443eff1950ba5ddc6b6ccda0d9c84a364276a62f969bdf2a390650848 \
|
||||
--hash=sha256:ad785e92e6dc634c21555edc8bd6b64957ab844541bcb96a1366c202951ae526 \
|
||||
--hash=sha256:af8244b2bef6aaad6d92cda81372de7f8c8d36c9f0c3ea36e827c60e7d9467a0 \
|
||||
--hash=sha256:b078afbc25f3a1436c7a1d2cd3e322497ee99615ba97c563566fdf46aff1ee01 \
|
||||
--hash=sha256:b2f69dec1725e79a012d920df1707de5caf7ed5e08f3be4435e25803efc47458 \
|
||||
--hash=sha256:b8458003118a712e66286df6a707db01c52c0f52f7db8e4a38f0da1d3b94fc4e \
|
||||
--hash=sha256:bb63e0198ca18aad131c089b9204c23079c3afa95487e561f4c522d519e55aba \
|
||||
--hash=sha256:bfec22eab3c8cc2ceec0248aec886624116dc079afa027ecc8ad4a7e62010f8a \
|
||||
--hash=sha256:c1747f85cee84c26985853c6f3d9bd3e75da5212912443fa111c113b9c246f39 \
|
||||
--hash=sha256:c1b3f518abeca3aa13c712fd202306e145abf59a18b094a6bafb2d2bbf59192c \
|
||||
--hash=sha256:c50f2528cf200c5eed56faf3f4e22fcd5f38c157a8b78576e6ba3168ec35f000 \
|
||||
--hash=sha256:c68fcd102d71ea85c5b2dfac3f4f8476eff42a9e078fd5faefff6d145063536b \
|
||||
--hash=sha256:c7a7bd4e39e8e4c12c39cd480356842b6a8a06e41b23a55a5e3e191718838ddf \
|
||||
--hash=sha256:c94f0688e7b8d0a67abf40e57a7eaaecd17cc9586706a31b76c031f63df052b4 \
|
||||
--hash=sha256:cbaf13819775b7f769bf4a1f066cb6df7a28d4480081a589828ef190226881cd \
|
||||
--hash=sha256:cd2213145bcc2ba85884d0ac63d222fece9209678f77b9b4d76f054c561adb28 \
|
||||
--hash=sha256:ce5c1d2a8b27468f433ca974829c44060b8097eedc39933e3c206a90ee49c4a9 \
|
||||
--hash=sha256:d396ec2b979760aaf3218e76c24e65bd0aca24983298653b3a9d7a45f9e47b30 \
|
||||
--hash=sha256:d51026d73fcfd93610abc7b27789c26b313920fcfb20e27462d74a7f8b06e983 \
|
||||
--hash=sha256:d80ee3d731373b24cebbc10d689ca4ee1875caf0d5703a245db18efd4dd37fc1 \
|
||||
--hash=sha256:d995260fdf4e1db774581b4900e0f832abe3c7c84996726bbc161b19c8f29e76 \
|
||||
--hash=sha256:da4b951fe36dc7c3a1ccb4e3cd1747c3542b8c9ceede8fc86cae054e764485f5 \
|
||||
--hash=sha256:daa27d92c36f24388fe3ad306b174781c747627f134452e4f128ea00ce1fe8c4 \
|
||||
--hash=sha256:db06ffe51636ffe9ca531fe9023dd64bdd794be8754cb5df57c5498ae5b518a7 \
|
||||
--hash=sha256:e0d65b8c354be7fb5f720c3caa8bc940bc2d20ce749c8e06135f07f8ed95dd7c \
|
||||
--hash=sha256:e68b7a074f65a2fd746c52a7ce6142ab7006074ac269ace0c25cd8ba171f8066 \
|
||||
--hash=sha256:e739fee756ba1010f8bcccb534252e85a35fe45ae92c295a06059ce58b74ccd3 \
|
||||
--hash=sha256:e846ae7835bf0703ae43f534ab79a867146dadd59dc9ca5c8b53d5c8f7c9ef02 \
|
||||
--hash=sha256:e9c26f834c65f5752f3f06cb08cb86a913ceb7274d0db6e267808a708b46bc89 \
|
||||
--hash=sha256:ea793e075b70290d89d8142074262885d3f7da19634845135751bd6344f73b50 \
|
||||
--hash=sha256:f027324c56cd5406ca49c124b0db10e56c69064fec039acc571c29020cc87c76 \
|
||||
--hash=sha256:f13a646d65d09fbf1bc6b3a9635d30095c8e7e5cc419ff35ecc563c5fd04cd49 \
|
||||
--hash=sha256:f47286a97f0bc9b8859519809077b91b2cefe4ae47fcbf5e466a009c1c5d742b \
|
||||
--hash=sha256:f747929cf940cddb5b3668a390056ddd5ba2e5010615ea2dcf4f9c4f3ab8791d \
|
||||
--hash=sha256:f99626688942fb746e545232e7726926f3be91b5975f8b55327665fafda991c7 \
|
||||
--hash=sha256:f9fa868638bf362d3d138ea55829cefb3d5f4b0d7f142234382a15e2485dbec4 \
|
||||
--hash=sha256:fbdb89b3e1c94a30cc5edfce477c6e6a5dc4d8f84665b455c27582f211a1c72c \
|
||||
--hash=sha256:fc010ab034c8c7452522748bf937df58020d256ccae0874463d1f4d01758af8e \
|
||||
--hash=sha256:fc3e9034a63de20e15e8ade85358bc6efc614008cab72898b4b4952bea0509ff \
|
||||
--hash=sha256:fd8b3d9fd264be37976686c7f65cd52a83f5e84f4bfd2adf9c1d469676bbb6ae
|
||||
# via pydantic
|
||||
typing-extensions==4.16.0 \
|
||||
--hash=sha256:481caa481374e813c1b176ada14e97f1f67a4539ce9cfeb3f350d78d6370c2e8 \
|
||||
--hash=sha256:dc983d19a509c94dba722ee6abd33940f7c05a89e243c47e907eb4db6f1a43e5
|
||||
# via
|
||||
# anyio
|
||||
# pydantic
|
||||
# pydantic-core
|
||||
# typing-inspection
|
||||
typing-inspection==0.4.4 \
|
||||
--hash=sha256:547274fa6b0a561ccf549cc9524b999a578e737d015d8709d021f9d0d13bea47 \
|
||||
--hash=sha256:65b8397ba37ccbce054456aaccddfc91e6e3083c92824df348d96ca832f3f147
|
||||
# via pydantic
|
||||
|
|
@ -5,6 +5,8 @@ services:
|
|||
build:
|
||||
context: ..
|
||||
target: runtime
|
||||
args:
|
||||
LITELLM_RELEASE_TAG: ${LITELLM_RELEASE_TAG:-}
|
||||
command: ["--config", "/app/tracing-config.yaml", "--port", "4000"]
|
||||
environment:
|
||||
LITELLM_MASTER_KEY: sk-1234
|
||||
|
|
@ -15,6 +17,7 @@ services:
|
|||
CLICKHOUSE_URL: http://default:local-tracing@clickhouse:8123
|
||||
CLICKHOUSE_DATABASE: litellm
|
||||
OPENAI_API_KEY: ${OPENAI_API_KEY:-}
|
||||
LENS_WORKER_IMAGE: ${LENS_WORKER_IMAGE:-}
|
||||
volumes:
|
||||
- ./tracing-config.yaml:/app/tracing-config.yaml:ro
|
||||
ports:
|
||||
|
|
|
|||
|
|
@ -472,10 +472,21 @@ collector containers through an emptyDir. Empty when the sidecar is off
|
|||
or gateway.collector.address is a tcp://127.0.0.1:<port> address.
|
||||
*/}}
|
||||
{{- define "litellm.lensWorker.image" -}}
|
||||
{{- if .Values.lensWorker.image.digest -}}
|
||||
{{- if not (regexMatch "^sha256:[0-9a-f]{64}$" .Values.lensWorker.image.digest) -}}
|
||||
{{- fail "lensWorker.image.digest must be sha256 followed by 64 lowercase hex characters" -}}
|
||||
{{- end -}}
|
||||
{{- printf "%s@%s" .Values.lensWorker.image.repository .Values.lensWorker.image.digest -}}
|
||||
{{- else -}}
|
||||
{{- $backendTag := .Values.backend.image.tag | default .Chart.AppVersion -}}
|
||||
{{- $releaseTag := ternary (printf "v%s" $backendTag) $backendTag (regexMatch "^[0-9]" $backendTag) -}}
|
||||
{{- $tag := .Values.lensWorker.image.tag | default $releaseTag -}}
|
||||
{{- printf "%s:%s" .Values.lensWorker.image.repository $tag -}}
|
||||
{{- $repository := .Values.lensWorker.image.repository -}}
|
||||
{{- if and (hasPrefix "sha-" $tag) (eq $repository "ghcr.io/berriai/litellm-lens-worker") -}}
|
||||
{{- $repository = "ghcr.io/berriai/litellm-lens-worker-dev" -}}
|
||||
{{- end -}}
|
||||
{{- printf "%s:%s" $repository $tag -}}
|
||||
{{- end -}}
|
||||
{{- end -}}
|
||||
|
||||
{{- define "litellm.gateway.collectorSocketDir" -}}
|
||||
|
|
|
|||
|
|
@ -6,6 +6,66 @@ templates:
|
|||
values:
|
||||
- ./values/required.yaml
|
||||
tests:
|
||||
- it: installs the development package for a source commit
|
||||
template: lens/deployment.yaml
|
||||
set:
|
||||
backend.image.tag: sha-0123456789abcdef
|
||||
lensWorker.enabled: true
|
||||
lensWorker.tokenSecret.name: lens-credential
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].image
|
||||
value: ghcr.io/berriai/litellm-lens-worker-dev:sha-0123456789abcdef
|
||||
- it: advertises the development package for standalone source workers
|
||||
template: backend/deployment.yaml
|
||||
set:
|
||||
backend.image.tag: sha-0123456789abcdef
|
||||
asserts:
|
||||
- contains:
|
||||
path: spec.template.spec.containers[0].env
|
||||
content:
|
||||
name: LENS_WORKER_IMAGE
|
||||
value: ghcr.io/berriai/litellm-lens-worker-dev:sha-0123456789abcdef
|
||||
- it: preserves an explicit private source image repository
|
||||
template: lens/deployment.yaml
|
||||
set:
|
||||
backend.image.tag: sha-0123456789abcdef
|
||||
lensWorker.enabled: true
|
||||
lensWorker.tokenSecret.name: lens-credential
|
||||
lensWorker.image.repository: registry.example/lens-worker
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].image
|
||||
value: registry.example/lens-worker:sha-0123456789abcdef
|
||||
- it: pins the worker to its approved digest even when its tag changes
|
||||
template: lens/deployment.yaml
|
||||
set:
|
||||
lensWorker.enabled: true
|
||||
lensWorker.tokenSecret.name: lens-credential
|
||||
lensWorker.image.tag: replaced-release
|
||||
lensWorker.image.digest: sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
|
||||
asserts:
|
||||
- equal:
|
||||
path: spec.template.spec.containers[0].image
|
||||
value: ghcr.io/berriai/litellm-lens-worker@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
|
||||
- it: advertises the approved digest to standalone installers
|
||||
template: backend/deployment.yaml
|
||||
set:
|
||||
lensWorker.image.tag: replaced-release
|
||||
lensWorker.image.digest: sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
|
||||
asserts:
|
||||
- contains:
|
||||
path: spec.template.spec.containers[0].env
|
||||
content:
|
||||
name: LENS_WORKER_IMAGE
|
||||
value: ghcr.io/berriai/litellm-lens-worker@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
|
||||
- it: refuses a malformed digest instead of falling back to the tag
|
||||
template: backend/deployment.yaml
|
||||
set:
|
||||
lensWorker.image.digest: sha256:invalid
|
||||
asserts:
|
||||
- failedTemplate:
|
||||
errorMessage: lensWorker.image.digest must be sha256 followed by 64 lowercase hex characters
|
||||
- it: keeps the worker opt in
|
||||
template: lens/deployment.yaml
|
||||
asserts:
|
||||
|
|
|
|||
|
|
@ -636,6 +636,7 @@ lensWorker:
|
|||
image:
|
||||
repository: ghcr.io/berriai/litellm-lens-worker
|
||||
tag: ""
|
||||
digest: ""
|
||||
pullPolicy: IfNotPresent
|
||||
tokenSecret:
|
||||
name: ""
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import functools
|
||||
import glob
|
||||
import os
|
||||
import random
|
||||
|
|
@ -10,7 +11,8 @@ import time
|
|||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Final, Optional
|
||||
from typing import TYPE_CHECKING, Final, Optional, Union
|
||||
from urllib.parse import unquote, urlsplit
|
||||
|
||||
from litellm_proxy_extras import prisma_toolchain
|
||||
from litellm_proxy_extras._logging import logger
|
||||
|
|
@ -202,6 +204,66 @@ def _max_migration_timestamp(names) -> int:
|
|||
return max(_migration_timestamp(n) for n in names)
|
||||
|
||||
|
||||
_REDACTED: Final = "REDACTED"
|
||||
_PASSWORD_QUERY_KEYS: Final = frozenset(("password", "sslpassword"))
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _secret_shape_redactor() -> Callable[[str], str]:
|
||||
try:
|
||||
from litellm._logging import redact_secrets
|
||||
except ImportError:
|
||||
return lambda text: text
|
||||
return redact_secrets
|
||||
|
||||
|
||||
def _url_passwords(url: str) -> frozenset[str]:
|
||||
try:
|
||||
parts: Final = urlsplit(url)
|
||||
except ValueError:
|
||||
return frozenset()
|
||||
query_pairs: Final = tuple(pair.partition("=") for pair in parts.query.split("&"))
|
||||
raw_query_passwords: Final = tuple(
|
||||
value for key, separator, value in query_pairs if separator and key.lower() in _PASSWORD_QUERY_KEYS
|
||||
)
|
||||
raw_passwords: Final = ((parts.password,) if parts.password else ()) + raw_query_passwords
|
||||
return frozenset(password for password in raw_passwords + tuple(map(unquote, raw_passwords)) if password)
|
||||
|
||||
|
||||
def _configured_database_passwords() -> frozenset[str]:
|
||||
database_url: Final = os.getenv("DATABASE_URL")
|
||||
direct_url: Final = os.getenv("DIRECT_URL")
|
||||
database_passwords: Final = _url_passwords(database_url) if database_url else frozenset()
|
||||
direct_passwords: Final = _url_passwords(direct_url) if direct_url else frozenset()
|
||||
return database_passwords | direct_passwords
|
||||
|
||||
|
||||
def _redact_credentials(text: str) -> str:
|
||||
"""Mask configured database passwords before passing the text to LiteLLM redaction."""
|
||||
passwords: Final = sorted(_configured_database_passwords(), key=len, reverse=True)
|
||||
alternation: Final = "|".join(re.escape(password) for password in passwords)
|
||||
password_pattern: Final = (
|
||||
re.compile(rf"(?P<lead>:|password=)(?:{alternation})(?=@|&|$|[\s'\"\]),])", re.IGNORECASE)
|
||||
if passwords
|
||||
else None
|
||||
)
|
||||
result: Final = password_pattern.sub(rf"\g<lead>{_REDACTED}", text) if password_pattern is not None else text
|
||||
return _secret_shape_redactor()(result)
|
||||
|
||||
|
||||
def _redacted_command(command: object) -> Union[str, tuple[str, ...], list[str]]:
|
||||
if isinstance(command, tuple):
|
||||
return tuple(_redact_credentials(str(argument)) for argument in command)
|
||||
if isinstance(command, list):
|
||||
return [_redact_credentials(str(argument)) for argument in command]
|
||||
return _redact_credentials(str(command))
|
||||
|
||||
|
||||
def _redact_command_error(error: subprocess.CalledProcessError) -> str:
|
||||
redacted_command: Final = _redacted_command(error.cmd)
|
||||
return str(subprocess.CalledProcessError(error.returncode, redacted_command))
|
||||
|
||||
|
||||
def _get_prisma_command() -> str:
|
||||
"""Get the Prisma command to use, bypassing Python wrapper in offline mode."""
|
||||
if str_to_bool(os.getenv("PRISMA_OFFLINE_MODE")):
|
||||
|
|
@ -315,7 +377,8 @@ class ProxyExtrasDBManager:
|
|||
return False
|
||||
except subprocess.CalledProcessError as e:
|
||||
logger.warning(
|
||||
f"Error creating baseline migration: {e}, {e.stderr}, {e.stdout}"
|
||||
f"Error creating baseline migration: {_redact_command_error(e)}, "
|
||||
f"{_redact_credentials(str(e.stderr))}, {_redact_credentials(str(e.stdout))}"
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -1572,6 +1635,11 @@ class ProxyExtrasDBManager:
|
|||
f"Error: {stderr}"
|
||||
)
|
||||
raise
|
||||
else:
|
||||
logger.error(
|
||||
"prisma migrate deploy failed with an error the resolver does not handle: "
|
||||
f"{_redact_credentials(stderr)}"
|
||||
)
|
||||
else:
|
||||
if ProxyExtrasDBManager.spend_logs_is_partitioned():
|
||||
raise RuntimeError(PARTITIONED_SPEND_LOGS_PUSH_ERROR)
|
||||
|
|
@ -1586,7 +1654,7 @@ class ProxyExtrasDBManager:
|
|||
)
|
||||
return True
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.warning(
|
||||
logger.error(
|
||||
"Attempt %s timed out. Raise %s if this database needs longer to apply its schema.",
|
||||
attempt + 1,
|
||||
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR if use_migrate else PRISMA_COMMAND_TIMEOUT_ENV_VAR,
|
||||
|
|
@ -1599,7 +1667,12 @@ class ProxyExtrasDBManager:
|
|||
if attempts_left > 0
|
||||
else ""
|
||||
)
|
||||
logger.info(f"The process failed to execute. Details: {e}.{retry_msg}")
|
||||
stderr_detail: Final = (
|
||||
f" stderr: {_redact_credentials(str(e.stderr))}" if e.stderr else ""
|
||||
)
|
||||
logger.error(
|
||||
f"The process failed to execute. Details: {_redact_command_error(e)}.{stderr_detail}{retry_msg}"
|
||||
)
|
||||
time.sleep(random.randrange(5, 15))
|
||||
finally:
|
||||
os.chdir(original_dir)
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ fn map_error_ref(error: &Error) -> PyErr {
|
|||
| Error::InsertTooLarge
|
||||
| Error::ReadTooLarge => PyOverflowError::new_err(error.to_string()),
|
||||
Error::InvalidRow
|
||||
| Error::InvalidLimit(_)
|
||||
| Error::InvalidTable
|
||||
| Error::InvalidCursor(_)
|
||||
| Error::AmbiguousTrace
|
||||
|
|
@ -59,6 +60,7 @@ fn map_error_ref(error: &Error) -> PyErr {
|
|||
Error::Cached(source) => map_error_ref(source),
|
||||
Error::Storage(source) => match source {
|
||||
StorageError::InvalidRow
|
||||
| StorageError::InvalidLimit(_)
|
||||
| StorageError::InvalidTable
|
||||
| StorageError::InvalidSchema
|
||||
| StorageError::EmptySql
|
||||
|
|
@ -433,6 +435,13 @@ mod tests {
|
|||
|
||||
#[rstest]
|
||||
#[case::row(Error::InvalidRow, "ValueError")]
|
||||
#[case::insert_limit(Error::InvalidLimit("CLICKHOUSE_TRACE_MAX_INSERT_BYTES"), "ValueError")]
|
||||
#[case::insert_timeout(
|
||||
Error::Storage(litellm_storage_clickhouse::Error::InvalidLimit(
|
||||
"CLICKHOUSE_INSERT_TIMEOUT_SECONDS"
|
||||
)),
|
||||
"ValueError"
|
||||
)]
|
||||
#[case::insert_budget(Error::InsertTooLarge, "OverflowError")]
|
||||
#[case::scope(Error::InvalidScope, "ValueError")]
|
||||
#[case::schema(Error::SchemaFailed(503), "RuntimeError")]
|
||||
|
|
@ -482,6 +491,10 @@ mod tests {
|
|||
#[rstest]
|
||||
#[case::decode_budget(Error::Decode(litellm_traces::Error::TooLarge), "OverflowError")]
|
||||
#[case::invalid_export(Error::Decode(litellm_traces::Error::InvalidPayload), "ValueError")]
|
||||
#[case::invalid_decode_limit(
|
||||
Error::Decode(litellm_traces::Error::InvalidLimit("OTLP_MAX_SPANS")),
|
||||
"ValueError"
|
||||
)]
|
||||
#[case::cursor(Error::InvalidCursor("trace"), "ValueError")]
|
||||
#[case::ambiguous(Error::AmbiguousTrace, "ValueError")]
|
||||
#[case::changed_snapshot(Error::TraceChanged, "ValueError")]
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
#[error("{0} must be a positive integer")]
|
||||
InvalidLimit(&'static str),
|
||||
#[error("invalid ClickHouse insert table")]
|
||||
InvalidTable,
|
||||
#[error("invalid ClickHouse HTTP URL")]
|
||||
|
|
|
|||
|
|
@ -5,7 +5,19 @@ use litellm_http::Client;
|
|||
|
||||
use crate::{Connection, Error, valid_identifier};
|
||||
|
||||
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
fn insert_timeout() -> Result<Duration, Error> {
|
||||
let name = "CLICKHOUSE_INSERT_TIMEOUT_SECONDS";
|
||||
match std::env::var(name) {
|
||||
Ok(value) => value
|
||||
.parse::<u64>()
|
||||
.ok()
|
||||
.filter(|value| *value > 0)
|
||||
.map(Duration::from_secs)
|
||||
.ok_or(Error::InvalidLimit(name)),
|
||||
Err(std::env::VarError::NotPresent) => Ok(Duration::from_secs(30)),
|
||||
Err(_) => Err(Error::InvalidLimit(name)),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn insert_encoded_rows(
|
||||
client: &Client,
|
||||
|
|
@ -74,7 +86,7 @@ pub async fn insert_compressed_rows(
|
|||
.append_pair("date_time_input_format", "best_effort");
|
||||
let response = client
|
||||
.post(url)
|
||||
.timeout(INSERT_TIMEOUT)
|
||||
.timeout(insert_timeout()?)
|
||||
.header("Content-Encoding", "gzip")
|
||||
.body(body)
|
||||
.send()
|
||||
|
|
|
|||
|
|
@ -144,3 +144,50 @@ async fn server_result_limits_allow_smaller_pages_without_retrying_other_failure
|
|||
assert!(matches!(error, Error::QueryFailed(500)));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn insert_timeout_environment_controls_transport() {
|
||||
for value in ["1", "3", "0", "invalid"] {
|
||||
let result = std::process::Command::new(std::env::current_exe().unwrap())
|
||||
.args(["--exact", "insert_timeout_environment_child"])
|
||||
.env("LITELLM_TEST_INSERT_TIMEOUT", value)
|
||||
.env("CLICKHOUSE_INSERT_TIMEOUT_SECONDS", value)
|
||||
.output()
|
||||
.unwrap();
|
||||
assert!(
|
||||
result.status.success(),
|
||||
"{}",
|
||||
String::from_utf8_lossy(&result.stdout)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn insert_timeout_environment_child() {
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
|
||||
let Ok(value) = std::env::var("LITELLM_TEST_INSERT_TIMEOUT") else {
|
||||
return;
|
||||
};
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(200).set_delay(std::time::Duration::from_millis(1500)))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let result = insert_encoded_rows(
|
||||
&Client::no_redirect_for_test(),
|
||||
&Connection::parse(&server.uri()).unwrap(),
|
||||
"traces",
|
||||
"otel_traces",
|
||||
"token",
|
||||
"{}",
|
||||
)
|
||||
.await;
|
||||
match value.as_str() {
|
||||
"1" => assert!(matches!(result, Err(Error::Transport))),
|
||||
"3" => assert!(result.is_ok()),
|
||||
_ => assert!(matches!(
|
||||
result,
|
||||
Err(Error::InvalidLimit("CLICKHOUSE_INSERT_TIMEOUT_SECONDS"))
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
SELECT
|
||||
request_id, response_id, trace_id, span_id, model, spend,
|
||||
prompt_tokens, completion_tokens, status,
|
||||
JSONExtractBool(metadata, 'synthetic_spend') AS synthetic_spend
|
||||
prompt_tokens, completion_tokens, status
|
||||
FROM spend_logs FINAL
|
||||
WHERE start_time >= now() - INTERVAL 1 DAY
|
||||
ORDER BY start_time DESC, request_id
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
#[error("{0} must be a positive integer")]
|
||||
InvalidLimit(&'static str),
|
||||
#[error("invalid ClickHouse insert table")]
|
||||
InvalidTable,
|
||||
#[error("database must be a nonempty SQL identifier and retention must be positive")]
|
||||
|
|
|
|||
|
|
@ -15,7 +15,18 @@ use time::{OffsetDateTime, format_description::well_known::Rfc3339};
|
|||
use super::{Connection, Error};
|
||||
use litellm_traces::Shared;
|
||||
|
||||
const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024;
|
||||
fn max_insert_bytes() -> Result<usize, Error> {
|
||||
let name = "CLICKHOUSE_TRACE_MAX_INSERT_BYTES";
|
||||
match std::env::var(name) {
|
||||
Ok(value) => value
|
||||
.parse::<usize>()
|
||||
.ok()
|
||||
.filter(|value| *value > 0)
|
||||
.ok_or(Error::InvalidLimit(name)),
|
||||
Err(std::env::VarError::NotPresent) => Ok(64 * 1024 * 1024),
|
||||
Err(_) => Err(Error::InvalidLimit(name)),
|
||||
}
|
||||
}
|
||||
|
||||
pub type InsertRow = BTreeMap<String, Shared<Value>>;
|
||||
|
||||
|
|
@ -62,7 +73,7 @@ pub async fn insert_shared_rows(
|
|||
return Ok(());
|
||||
}
|
||||
let received_ms = (OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64;
|
||||
let (token, body) = prepare_insert(&rows, received_ms, MAX_INSERT_BYTES)?;
|
||||
let (token, body) = prepare_insert(&rows, received_ms, max_insert_bytes()?)?;
|
||||
litellm_storage_clickhouse::insert_compressed_rows(
|
||||
client,
|
||||
connection,
|
||||
|
|
|
|||
|
|
@ -184,7 +184,7 @@ Duration is nanoseconds; Timestamp has nanosecond precision, spend start_time ha
|
|||
{%- endblock %}
|
||||
|
||||
{% block missing_spend -%}
|
||||
Token usage does not establish billed spend. OTLP exports without companion spend_logs rows have unknown cost; synthetic fixture spend is marked by metadata.synthetic_spend
|
||||
Token usage does not establish billed spend. OTLP exports without companion spend_logs rows have unknown cost
|
||||
{%- endblock %}
|
||||
|
||||
{% block partial_spend -%}
|
||||
|
|
|
|||
|
|
@ -4,13 +4,11 @@ Run `cargo test -p litellm-traces-clickhouse --test queries --locked -- --test-t
|
|||
|
||||
Raw OTLP exports live in `crates/traces/tests/fixtures/query_*.json`. The seeded fixture decodes and normalizes them through `litellm_traces::decode_otlp` at test startup, then projects the decoded fields into ClickHouse columns. Team and key identities come from fixture setup rather than exporter claims. Root and child exports are inserted separately through the public insert API so materialized views process multiple blocks
|
||||
|
||||
`crates/traces/tests/fixtures/deeplite_auth_error.json` and `deeplite_swarm.json` were captured from Deeplite runs against the local proxy on 2026-10-02. The first contains a failed model call. The second contains successful model calls, searches, handoff attempts, and virtual filesystem writes. Credentials, workspace identifiers, and local user paths were redacted, and the protobuf exports were converted to OTLP JSON. Their round-trip tests check span identities, parent links, timestamps, durations, token counts, and statuses without pinning the provider's error wording
|
||||
The ClickHouse round-trip test replays the `google_adk_billed_failure`, `pydantic_ai_retry`, and `deepagents_swarm` exports and checks span identities, parent links, timestamps, durations, token counts, and statuses without pinning the provider's error wording. Exported ERROR and UNSET statuses are diagnostic and do not establish a failed execution, so the test preserves incoming statuses and checks root status separately from the count of error spans, deriving both from the decoded export. Framework-specific interpretation of control-flow exceptions belongs in the instrumentation integration
|
||||
|
||||
The swarm capture has handoff spans marked ERROR with `ParentCommand` exception events and a root with UNSET status. These are exported diagnostic statuses, which do not establish a failed execution. The tests preserve incoming statuses and check root status separately from the count of error spans, deriving both from the decoded export. They do not infer an execution outcome from exception text, framework names, successful model calls, or output presence. Framework-specific interpretation of control-flow exceptions belongs in the instrumentation integration
|
||||
For a local dashboard with linked requests and traces, run `make lens-dev ARGS=--seed` from the repository root and open `http://localhost:3000/ui/lens/`. Log in as `admin` with the master key saved in `.lens-dev/master_key`. The launcher keeps the stack running until Ctrl-C and leaves the database volumes intact
|
||||
|
||||
For a local dashboard with linked requests and traces, run `bash scripts/run_tracing_proxy_local.sh --seed` from the repository root and open `http://127.0.0.1:4002/ui/`. Log in as `admin` with password `sk-1234`, matching the UI E2E harness. The launcher keeps the proxy running until Ctrl-C and leaves the database volumes intact
|
||||
|
||||
`deeplite_swarm_spend_logs.jsonl` pairs every LLM span in the swarm export with a ClickHouse spend row. Response IDs, trace and span IDs, token counts, input, output, and timestamps come from the export. Messages and responses use the chat completion format supported by the request viewer. Spend is synthetic, set to $0.01 per request and marked in metadata, because the export does not include actual billed costs. These rows are stored here because `traces-clickhouse` owns the spend row schema
|
||||
Spend rows are stored here because `traces-clickhouse` owns the spend row schema
|
||||
|
||||
The simple and swarm exports for all twelve SDK examples were captured on 2026-10-03 against port 4002 using `openai/gpt-6-luna`. Each export has a matching `<name>_spend_logs.jsonl` with actual proxy spend, usage, request and response IDs, messages, and timestamps. Authorization headers, provider cookies, organization and project identifiers, and local paths were redacted. OTLP identifiers and enums use their canonical JSON encodings. `metadata.fixture_capture` identifies the associated export and whether model spans contain sufficient identity to join spend
|
||||
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -138,3 +138,56 @@ fn insert_encoding_preserves_timestamp_precision_and_other_fields(
|
|||
fn insert_encoding_rejects_invalid_span_timestamps(#[case] timestamp: Value) {
|
||||
assert!(encode_rows(vec![BTreeMap::from([("Timestamp".into(), timestamp)])]).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn insert_byte_limit_environment_controls_transport() {
|
||||
for value in ["1", "1024", "0", "invalid"] {
|
||||
let result = std::process::Command::new(std::env::current_exe().unwrap())
|
||||
.args(["--exact", "insert_byte_limit_environment_child"])
|
||||
.env("LITELLM_TEST_INSERT_LIMIT", value)
|
||||
.env("CLICKHOUSE_TRACE_MAX_INSERT_BYTES", value)
|
||||
.output()
|
||||
.unwrap();
|
||||
assert!(
|
||||
result.status.success(),
|
||||
"{}",
|
||||
String::from_utf8_lossy(&result.stdout)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn insert_byte_limit_environment_child() {
|
||||
let Ok(value) = std::env::var("LITELLM_TEST_INSERT_LIMIT") else {
|
||||
return;
|
||||
};
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(200))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let connection = Connection::parse(&server.uri()).unwrap();
|
||||
let result = insert_shared_rows(
|
||||
&Client::no_redirect_for_test(),
|
||||
&connection,
|
||||
"traces",
|
||||
InsertTable::OtelTraces,
|
||||
vec![BTreeMap::from([(
|
||||
"SpanId".into(),
|
||||
Shared::new(json!("test")),
|
||||
)])],
|
||||
)
|
||||
.await;
|
||||
match value.as_str() {
|
||||
"1" => assert!(matches!(result, Err(Error::InsertTooLarge))),
|
||||
"1024" => assert!(result.is_ok()),
|
||||
_ => assert!(matches!(
|
||||
result,
|
||||
Err(Error::InvalidLimit("CLICKHOUSE_TRACE_MAX_INSERT_BYTES"))
|
||||
)),
|
||||
}
|
||||
assert_eq!(
|
||||
server.received_requests().await.unwrap().len(),
|
||||
usize::from(value == "1024")
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -244,10 +244,11 @@ async fn typed_trace_cursor_returns_the_next_fixture_trace(
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::authentication_error(include_bytes!("../../traces/tests/fixtures/deeplite_auth_error.json"))]
|
||||
#[case::swarm(include_bytes!("../../traces/tests/fixtures/deeplite_swarm.json"))]
|
||||
#[case::billed_failure(include_bytes!("../../traces/tests/fixtures/google_adk_billed_failure.json"))]
|
||||
#[case::retry(include_bytes!("../../traces/tests/fixtures/pydantic_ai_retry.json"))]
|
||||
#[case::swarm(include_bytes!("../../traces/tests/fixtures/deepagents_swarm.json"))]
|
||||
#[tokio::test]
|
||||
async fn captured_deeplite_exports_round_trip_through_clickhouse(
|
||||
async fn captured_sdk_exports_round_trip_through_clickhouse(
|
||||
#[future(awt)] migrated_database: TestResult<SeededDatabase>,
|
||||
admin_access: TestResult<contracts::ReadAccessParams>,
|
||||
#[case] export: &[u8],
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
pub enum Error {
|
||||
#[error("invalid OTLP trace payload")]
|
||||
InvalidPayload,
|
||||
#[error("{0} must be a positive integer")]
|
||||
InvalidLimit(&'static str),
|
||||
#[error("OTLP trace payload exceeds the decoding budget")]
|
||||
TooLarge,
|
||||
#[error("OTLP token count is outside the storage range")]
|
||||
|
|
|
|||
|
|
@ -30,7 +30,7 @@ pub use normalize::{
|
|||
AgentMetadata, AgentType, CallEvidence, CallEvidenceKind, CallKey, Integration, NormalizedSpan,
|
||||
ObservationType,
|
||||
};
|
||||
pub use otlp::{DecodedEvent, DecodedSpan, decode_otlp};
|
||||
pub use otlp::{DecodeLimits, DecodedEvent, DecodedSpan, decode_otlp, decode_otlp_with_limits};
|
||||
pub use query::ReadQuery;
|
||||
pub use query_access::QueryScope;
|
||||
pub use resolve::{SpendLookup, iso_time, listed_summary, resolve_trace};
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ use serde::{
|
|||
ser::{SerializeMap, SerializeSeq},
|
||||
};
|
||||
|
||||
use super::limits::{Budget, MAX_ATTRIBUTES};
|
||||
use super::limits::Budget;
|
||||
use crate::Error;
|
||||
|
||||
struct AttributeWriter<'a> {
|
||||
|
|
@ -33,7 +33,7 @@ pub(super) fn attributes(
|
|||
values: Vec<KeyValue>,
|
||||
budget: &mut Budget,
|
||||
) -> Result<BTreeMap<String, String>, Error> {
|
||||
if values.len() > MAX_ATTRIBUTES {
|
||||
if values.len() > budget.limits.attributes {
|
||||
return Err(Error::TooLarge);
|
||||
}
|
||||
values
|
||||
|
|
|
|||
|
|
@ -5,14 +5,62 @@ use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor};
|
|||
|
||||
use crate::{Error, Shared};
|
||||
|
||||
pub(super) const MAX_DEPTH: usize = 32;
|
||||
pub(super) const MAX_NODES: usize = 65_536;
|
||||
pub(super) const MAX_SPANS: usize = 4_096;
|
||||
pub(super) const MAX_ATTRIBUTES: usize = 256;
|
||||
pub(super) const MAX_EVENTS: usize = 256;
|
||||
pub(super) const MAX_DECODED_SPAN_BYTES: usize = 16 * 1024 * 1024;
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
pub struct DecodeLimits {
|
||||
pub depth: usize,
|
||||
pub nodes: usize,
|
||||
pub spans: usize,
|
||||
pub attributes: usize,
|
||||
pub events: usize,
|
||||
pub links: usize,
|
||||
pub decoded_span_bytes: usize,
|
||||
}
|
||||
|
||||
pub(super) fn json_preflight(payload: &[u8]) -> Result<(), Error> {
|
||||
impl Default for DecodeLimits {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
depth: 32,
|
||||
nodes: 65_536,
|
||||
spans: 4_096,
|
||||
attributes: 256,
|
||||
events: 256,
|
||||
links: 256,
|
||||
decoded_span_bytes: 16 * 1024 * 1024,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl DecodeLimits {
|
||||
pub fn from_env() -> Result<Self, Error> {
|
||||
let defaults = Self::default();
|
||||
Ok(Self {
|
||||
depth: env_limit("OTLP_MAX_DECODE_DEPTH", defaults.depth)?,
|
||||
nodes: env_limit("OTLP_MAX_DECODE_NODES", defaults.nodes)?,
|
||||
spans: env_limit("OTLP_MAX_SPANS", defaults.spans)?,
|
||||
attributes: env_limit("OTLP_MAX_ATTRIBUTES", defaults.attributes)?,
|
||||
events: env_limit("OTLP_MAX_EVENTS", defaults.events)?,
|
||||
links: env_limit("OTLP_MAX_LINKS", defaults.links)?,
|
||||
decoded_span_bytes: env_limit(
|
||||
"OTLP_MAX_DECODED_SPAN_BYTES",
|
||||
defaults.decoded_span_bytes,
|
||||
)?,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn env_limit(name: &'static str, default: usize) -> Result<usize, Error> {
|
||||
match std::env::var(name) {
|
||||
Ok(value) => value
|
||||
.parse::<usize>()
|
||||
.ok()
|
||||
.filter(|value| *value > 0)
|
||||
.ok_or(Error::InvalidLimit(name)),
|
||||
Err(std::env::VarError::NotPresent) => Ok(default),
|
||||
Err(_) => Err(Error::InvalidLimit(name)),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn json_preflight(payload: &[u8], limits: &DecodeLimits) -> Result<(), Error> {
|
||||
let mut nodes = 0;
|
||||
let mut exceeded = false;
|
||||
let mut decoder = serde_json::Deserializer::from_slice(payload);
|
||||
|
|
@ -20,6 +68,7 @@ pub(super) fn json_preflight(payload: &[u8]) -> Result<(), Error> {
|
|||
nodes: &mut nodes,
|
||||
exceeded: &mut exceeded,
|
||||
depth: 0,
|
||||
limits,
|
||||
}
|
||||
.deserialize(&mut decoder)
|
||||
.and_then(|()| decoder.end());
|
||||
|
|
@ -33,6 +82,7 @@ struct JsonBudget<'a> {
|
|||
nodes: &'a mut usize,
|
||||
exceeded: &'a mut bool,
|
||||
depth: usize,
|
||||
limits: &'a DecodeLimits,
|
||||
}
|
||||
|
||||
impl<'de> DeserializeSeed<'de> for JsonBudget<'_> {
|
||||
|
|
@ -40,7 +90,7 @@ impl<'de> DeserializeSeed<'de> for JsonBudget<'_> {
|
|||
|
||||
fn deserialize<D: serde::Deserializer<'de>>(self, decoder: D) -> Result<(), D::Error> {
|
||||
*self.nodes += 1;
|
||||
if *self.nodes > MAX_NODES || self.depth > MAX_DEPTH {
|
||||
if *self.nodes > self.limits.nodes || self.depth > self.limits.depth {
|
||||
*self.exceeded = true;
|
||||
return Err(serde::de::Error::custom("OTLP structure exceeds budget"));
|
||||
}
|
||||
|
|
@ -79,6 +129,7 @@ impl<'de> Visitor<'de> for JsonBudget<'_> {
|
|||
nodes: self.nodes,
|
||||
exceeded: self.exceeded,
|
||||
depth: self.depth + 1,
|
||||
limits: self.limits,
|
||||
})?
|
||||
.is_some()
|
||||
{}
|
||||
|
|
@ -91,6 +142,7 @@ impl<'de> Visitor<'de> for JsonBudget<'_> {
|
|||
nodes: self.nodes,
|
||||
exceeded: self.exceeded,
|
||||
depth: self.depth + 1,
|
||||
limits: self.limits,
|
||||
})?
|
||||
.is_some()
|
||||
{
|
||||
|
|
@ -98,6 +150,7 @@ impl<'de> Visitor<'de> for JsonBudget<'_> {
|
|||
nodes: self.nodes,
|
||||
exceeded: self.exceeded,
|
||||
depth: self.depth + 1,
|
||||
limits: self.limits,
|
||||
})?;
|
||||
}
|
||||
Ok(())
|
||||
|
|
@ -146,8 +199,8 @@ impl MessageKind {
|
|||
}
|
||||
}
|
||||
|
||||
pub(super) fn protobuf_preflight(payload: &[u8]) -> Result<(), Error> {
|
||||
scan_message(payload, MessageKind::Export, 0, &mut 0)
|
||||
pub(super) fn protobuf_preflight(payload: &[u8], limits: &DecodeLimits) -> Result<(), Error> {
|
||||
scan_message(payload, MessageKind::Export, 0, &mut 0, limits)
|
||||
}
|
||||
|
||||
fn scan_message(
|
||||
|
|
@ -155,13 +208,14 @@ fn scan_message(
|
|||
kind: MessageKind,
|
||||
depth: usize,
|
||||
nodes: &mut usize,
|
||||
limits: &DecodeLimits,
|
||||
) -> Result<(), Error> {
|
||||
if depth > MAX_DEPTH {
|
||||
if depth > limits.depth {
|
||||
return Err(Error::TooLarge);
|
||||
}
|
||||
while !payload.is_empty() {
|
||||
*nodes += 1;
|
||||
if *nodes > MAX_NODES {
|
||||
if *nodes > limits.nodes {
|
||||
return Err(Error::TooLarge);
|
||||
}
|
||||
let (tag, wire) = decode_key(&mut payload).map_err(|_| Error::InvalidPayload)?;
|
||||
|
|
@ -171,7 +225,7 @@ fn scan_message(
|
|||
let (message, rest) = payload
|
||||
.split_at_checked(length)
|
||||
.ok_or(Error::InvalidPayload)?;
|
||||
scan_message(message, child, depth + 1, nodes)?;
|
||||
scan_message(message, child, depth + 1, nodes, limits)?;
|
||||
payload = rest;
|
||||
} else {
|
||||
skip_field(wire, tag, &mut payload, DecodeContext::default())
|
||||
|
|
@ -183,11 +237,15 @@ fn scan_message(
|
|||
|
||||
pub(super) struct Budget {
|
||||
remaining: usize,
|
||||
pub(super) limits: DecodeLimits,
|
||||
}
|
||||
|
||||
impl Budget {
|
||||
pub(super) fn new(remaining: usize) -> Self {
|
||||
Self { remaining }
|
||||
pub(super) fn new(limits: DecodeLimits) -> Self {
|
||||
Self {
|
||||
remaining: limits.decoded_span_bytes,
|
||||
limits,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn clone_shared<T: Clone>(
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ mod limits;
|
|||
mod span;
|
||||
mod wire;
|
||||
|
||||
pub use limits::DecodeLimits;
|
||||
|
||||
use serde::Serialize;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
|
|
@ -36,6 +38,14 @@ pub struct DecodedSpan {
|
|||
}
|
||||
|
||||
pub fn decode_otlp(body: &[u8], content_type: Option<&str>) -> Result<Vec<DecodedSpan>, Error> {
|
||||
let request = wire::decode(body, content_type)?;
|
||||
span::flatten(request)
|
||||
decode_otlp_with_limits(body, content_type, DecodeLimits::from_env()?)
|
||||
}
|
||||
|
||||
pub fn decode_otlp_with_limits(
|
||||
body: &[u8],
|
||||
content_type: Option<&str>,
|
||||
limits: DecodeLimits,
|
||||
) -> Result<Vec<DecodedSpan>, Error> {
|
||||
let request = wire::decode(body, content_type, &limits)?;
|
||||
span::flatten(request, limits)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,15 +8,18 @@ use opentelemetry_proto::tonic::{
|
|||
use super::{
|
||||
DecodedEvent, DecodedSpan,
|
||||
attributes::attributes,
|
||||
limits::{Budget, MAX_ATTRIBUTES, MAX_DECODED_SPAN_BYTES, MAX_EVENTS, MAX_SPANS},
|
||||
limits::{Budget, DecodeLimits},
|
||||
};
|
||||
use crate::{
|
||||
Error, Shared,
|
||||
normalize::{SpanContext, normalize},
|
||||
};
|
||||
|
||||
pub(super) fn flatten(request: ExportTraceServiceRequest) -> Result<Vec<DecodedSpan>, Error> {
|
||||
let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES);
|
||||
pub(super) fn flatten(
|
||||
request: ExportTraceServiceRequest,
|
||||
limits: DecodeLimits,
|
||||
) -> Result<Vec<DecodedSpan>, Error> {
|
||||
let mut budget = Budget::new(limits);
|
||||
let mut spans = Vec::new();
|
||||
for resource in request.resource_spans {
|
||||
append_resource(resource, &mut budget, &mut spans)?;
|
||||
|
|
@ -49,17 +52,17 @@ fn append_scope(
|
|||
spans: &mut Vec<DecodedSpan>,
|
||||
) -> Result<(), Error> {
|
||||
let scope = scope_spans.scope.unwrap_or_default();
|
||||
if scope.attributes.len() > MAX_ATTRIBUTES {
|
||||
if scope.attributes.len() > budget.limits.attributes {
|
||||
return Err(Error::TooLarge);
|
||||
}
|
||||
budget.consume(scope.name.len() + scope.version.len())?;
|
||||
let scope_name: Shared<String> = scope.name.into();
|
||||
let scope_version: Shared<String> = scope.version.into();
|
||||
for span in scope_spans.spans {
|
||||
if spans.len() >= MAX_SPANS {
|
||||
if spans.len() >= budget.limits.spans {
|
||||
return Err(Error::TooLarge);
|
||||
}
|
||||
validate_span(&span)?;
|
||||
validate_span(&span, &budget.limits)?;
|
||||
budget.consume(
|
||||
span.name.len()
|
||||
+ span.trace_state.len()
|
||||
|
|
@ -85,7 +88,7 @@ fn valid_id(value: &[u8], length: usize) -> bool {
|
|||
value.len() == length && value.iter().any(|byte| *byte != 0)
|
||||
}
|
||||
|
||||
fn validate_span(span: &Span) -> Result<(), Error> {
|
||||
fn validate_span(span: &Span, limits: &DecodeLimits) -> Result<(), Error> {
|
||||
if !valid_id(&span.trace_id, 16)
|
||||
|| !valid_id(&span.span_id, 8)
|
||||
|| (!span.parent_span_id.is_empty() && !valid_id(&span.parent_span_id, 8))
|
||||
|
|
@ -99,17 +102,17 @@ fn validate_span(span: &Span) -> Result<(), Error> {
|
|||
{
|
||||
return Err(Error::InvalidPayload);
|
||||
}
|
||||
if span.events.len() > MAX_EVENTS
|
||||
|| span.links.len() > MAX_EVENTS
|
||||
|| span.attributes.len() > MAX_ATTRIBUTES
|
||||
if span.events.len() > limits.events
|
||||
|| span.links.len() > limits.links
|
||||
|| span.attributes.len() > limits.attributes
|
||||
|| span
|
||||
.links
|
||||
.iter()
|
||||
.any(|link| link.attributes.len() > MAX_ATTRIBUTES)
|
||||
.any(|link| link.attributes.len() > limits.attributes)
|
||||
|| span
|
||||
.events
|
||||
.iter()
|
||||
.any(|event| event.attributes.len() > MAX_ATTRIBUTES)
|
||||
.any(|event| event.attributes.len() > limits.attributes)
|
||||
{
|
||||
return Err(Error::TooLarge);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest;
|
||||
use prost::Message;
|
||||
|
||||
use super::limits::{json_preflight, protobuf_preflight};
|
||||
use super::limits::{DecodeLimits, json_preflight, protobuf_preflight};
|
||||
use crate::Error;
|
||||
|
||||
#[derive(strum::EnumString)]
|
||||
|
|
@ -19,6 +19,7 @@ enum OtlpMediaType {
|
|||
pub(super) fn decode(
|
||||
body: &[u8],
|
||||
content_type: Option<&str>,
|
||||
limits: &DecodeLimits,
|
||||
) -> Result<ExportTraceServiceRequest, Error> {
|
||||
let media_type = content_type
|
||||
.unwrap_or("application/x-protobuf")
|
||||
|
|
@ -31,11 +32,11 @@ pub(super) fn decode(
|
|||
|
||||
let request = match media_type {
|
||||
OtlpMediaType::Json => {
|
||||
json_preflight(body)?;
|
||||
json_preflight(body, limits)?;
|
||||
serde_json::from_slice(body).map_err(|_| Error::InvalidPayload)?
|
||||
}
|
||||
OtlpMediaType::Protobuf => {
|
||||
protobuf_preflight(body)?;
|
||||
protobuf_preflight(body, limits)?;
|
||||
ExportTraceServiceRequest::decode(body).map_err(|_| Error::InvalidPayload)?
|
||||
}
|
||||
};
|
||||
|
|
|
|||
|
|
@ -330,9 +330,7 @@ fn append_response_id(document: &mut Value, trace_id: &str, span_id: &str, respo
|
|||
|
||||
#[rstest]
|
||||
fn captured_trace_cost_matches_spend_logs(
|
||||
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")]
|
||||
#[exclude("deeplite_swarm")]
|
||||
spend_logs: PathBuf,
|
||||
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")] spend_logs: PathBuf,
|
||||
) {
|
||||
let name = capture_name(&spend_logs);
|
||||
let (_, capture, rows, spends) = fixture(&spend_logs);
|
||||
|
|
@ -347,9 +345,7 @@ fn captured_trace_cost_matches_spend_logs(
|
|||
|
||||
#[rstest]
|
||||
fn unrelated_sibling_transport_leaves_cost_unchanged(
|
||||
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")]
|
||||
#[exclude("deeplite_swarm")]
|
||||
spend_logs: PathBuf,
|
||||
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")] spend_logs: PathBuf,
|
||||
) {
|
||||
let name = capture_name(&spend_logs);
|
||||
let (_, capture, rows, spends) = fixture(&spend_logs);
|
||||
|
|
@ -392,9 +388,7 @@ fn unrelated_sibling_transport_leaves_cost_unchanged(
|
|||
|
||||
#[rstest]
|
||||
fn redundant_genai_response_id_keeps_call_evidence(
|
||||
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")]
|
||||
#[exclude("deeplite_swarm")]
|
||||
spend_logs: PathBuf,
|
||||
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")] spend_logs: PathBuf,
|
||||
) {
|
||||
let (data, capture, _, _) = fixture(&spend_logs);
|
||||
let name = capture.name;
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -120,8 +120,6 @@ fn array<'a>(value: &'a Value, key: &str) -> &'a [Value] {
|
|||
#[case::crewai_swarm(include_bytes!("fixtures/crewai_swarm.json"))]
|
||||
#[case::deepagents_simple(include_bytes!("fixtures/deepagents_simple.json"))]
|
||||
#[case::deepagents_swarm(include_bytes!("fixtures/deepagents_swarm.json"))]
|
||||
#[case::deeplite_auth_error(include_bytes!("fixtures/deeplite_auth_error.json"))]
|
||||
#[case::deeplite_swarm(include_bytes!("fixtures/deeplite_swarm.json"))]
|
||||
#[case::google_adk_simple(include_bytes!("fixtures/google_adk_simple.json"))]
|
||||
#[case::google_adk_swarm(include_bytes!("fixtures/google_adk_swarm.json"))]
|
||||
#[case::langchain_simple(include_bytes!("fixtures/langchain_simple.json"))]
|
||||
|
|
|
|||
|
|
@ -1333,3 +1333,98 @@ fn resource_identity_preserves_explicit_names_and_sdk_fallbacks(
|
|||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::depth(litellm_traces::DecodeLimits { depth: 1, ..Default::default() })]
|
||||
#[case::nodes(litellm_traces::DecodeLimits { nodes: 1, ..Default::default() })]
|
||||
#[case::spans(litellm_traces::DecodeLimits { spans: 1, ..Default::default() })]
|
||||
#[case::attributes(litellm_traces::DecodeLimits { attributes: 1, ..Default::default() })]
|
||||
#[case::events(litellm_traces::DecodeLimits { events: 1, ..Default::default() })]
|
||||
#[case::links(litellm_traces::DecodeLimits { links: 1, ..Default::default() })]
|
||||
#[case::decoded_bytes(litellm_traces::DecodeLimits { decoded_span_bytes: 1, ..Default::default() })]
|
||||
fn configurable_decode_limits_apply_to_both_wire_formats(
|
||||
mut span: Span,
|
||||
#[case] limits: litellm_traces::DecodeLimits,
|
||||
) {
|
||||
use opentelemetry_proto::tonic::{
|
||||
common::v1::KeyValue,
|
||||
trace::v1::span::{Event, Link},
|
||||
};
|
||||
use prost::Message;
|
||||
span.attributes = vec![
|
||||
KeyValue {
|
||||
key: "a".into(),
|
||||
..Default::default()
|
||||
},
|
||||
KeyValue {
|
||||
key: "b".into(),
|
||||
..Default::default()
|
||||
},
|
||||
];
|
||||
span.events = vec![Event::default(), Event::default()];
|
||||
span.links = vec![
|
||||
Link {
|
||||
trace_id: vec![1; 16],
|
||||
span_id: vec![2; 8],
|
||||
..Default::default()
|
||||
};
|
||||
2
|
||||
];
|
||||
let mut request = request_with(span.clone());
|
||||
request.resource_spans[0].scope_spans[0].spans.push(span);
|
||||
for (body, content_type) in [
|
||||
(serde_json::to_vec(&request).unwrap(), "application/json"),
|
||||
(request.encode_to_vec(), "application/x-protobuf"),
|
||||
] {
|
||||
assert!(matches!(
|
||||
litellm_traces::decode_otlp_with_limits(&body, Some(content_type), limits),
|
||||
Err(litellm_traces::Error::TooLarge)
|
||||
));
|
||||
assert_eq!(
|
||||
litellm_traces::decode_otlp_with_limits(
|
||||
&body,
|
||||
Some(content_type),
|
||||
litellm_traces::DecodeLimits::default()
|
||||
)
|
||||
.unwrap()
|
||||
.len(),
|
||||
2
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_decode_limits_are_used_and_invalid_values_fail() {
|
||||
for value in ["2", "4", "0", "invalid"] {
|
||||
let result = std::process::Command::new(std::env::current_exe().unwrap())
|
||||
.args(["--exact", "environment_decode_limits_child"])
|
||||
.env("LITELLM_TEST_DECODE_LIMIT", value)
|
||||
.env("OTLP_MAX_SPANS", value)
|
||||
.output()
|
||||
.unwrap();
|
||||
assert!(
|
||||
result.status.success(),
|
||||
"{}",
|
||||
String::from_utf8_lossy(&result.stdout)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_decode_limits_child() {
|
||||
let Ok(value) = std::env::var("LITELLM_TEST_DECODE_LIMIT") else {
|
||||
return;
|
||||
};
|
||||
let result = decode_otlp(
|
||||
include_bytes!("fixtures/opentelemetry_simple.json"),
|
||||
Some("application/json"),
|
||||
);
|
||||
match value.as_str() {
|
||||
"2" => assert!(matches!(result, Err(litellm_traces::Error::TooLarge))),
|
||||
"4" => assert_eq!(result.unwrap().len(), 3),
|
||||
_ => assert!(matches!(
|
||||
result,
|
||||
Err(litellm_traces::Error::InvalidLimit("OTLP_MAX_SPANS"))
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1474,6 +1474,7 @@ from .rust_bridge import rust
|
|||
from .rag.main import *
|
||||
from .sandbox.main import *
|
||||
from .decisions.main import *
|
||||
from .tool_loop import ToolLoopMaxRoundsExceeded, arun_tool_loop, run_tool_loop
|
||||
from .search.main import *
|
||||
from .realtime_api.main import (
|
||||
_arealtime,
|
||||
|
|
|
|||
|
|
@ -2234,3 +2234,5 @@ HARNESS_SNAPSHOT_SKIP_DIRS: Final = frozenset(
|
|||
".ruff_cache",
|
||||
}
|
||||
)
|
||||
|
||||
DEFAULT_TOOL_LOOP_MAX_ROUNDS: Final = 20
|
||||
|
|
|
|||
|
|
@ -59,7 +59,7 @@ def _resolve_provider_model(model: str, custom_llm_provider: str | None) -> tupl
|
|||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
upstream_model: Final = model.removeprefix(f"{provider}/") if model.startswith(f"{provider}/") else model
|
||||
upstream_model: Final = model.removeprefix(f"{provider}/")
|
||||
if not upstream_model:
|
||||
raise litellm.BadRequestError(
|
||||
message="A model name is required for the Decisions API",
|
||||
|
|
|
|||
|
|
@ -267,11 +267,8 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
tool_specs: list[ChatCompletionToolParam] = copy.deepcopy( # mutable-ok: acompletion takes tool list
|
||||
list(self._tool_specs)
|
||||
)
|
||||
request_kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}
|
||||
}
|
||||
kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
**request_kwargs,
|
||||
**{key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}},
|
||||
"messages": messages,
|
||||
**({"tools": tool_specs} if tool_specs else {}),
|
||||
}
|
||||
|
|
@ -287,8 +284,7 @@ class ToolLoopHandler(BaseHarnessHandler):
|
|||
yield Text(content)
|
||||
tool_calls = message.tool_calls or ()
|
||||
if not tool_calls:
|
||||
final_text = content or ""
|
||||
ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.final_text = content or "" # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink
|
||||
final_message: ChatCompletionMessageParam = {
|
||||
"role": "assistant",
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ class DeepAgentsOptions:
|
|||
|
||||
@dataclass(frozen=True)
|
||||
class ToolLoopOptions:
|
||||
completion_kwargs: Mapping[str, Any] = field(default_factory=dict)
|
||||
completion_kwargs: Mapping[str, object] = field(default_factory=dict)
|
||||
|
||||
|
||||
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions | ToolLoopOptions
|
||||
|
|
|
|||
|
|
@ -125,13 +125,12 @@ class AnthropicFilesConfig(BaseFilesConfig):
|
|||
return self._finalize_headers(headers, auth_header)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_params(
|
||||
litellm_params: dict, api_base: str | None
|
||||
) -> tuple[dict | None, str | None]: # mutable-ok: mirrors the sync validate_environment contract this overrides
|
||||
def _resolve_params(litellm_params: dict, api_base: str | None) -> tuple[Mapping[str, object] | None, str | None]:
|
||||
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
|
||||
if api_base is None and params_mapping is not None:
|
||||
api_base = params_mapping.get("api_base")
|
||||
return params_mapping, api_base
|
||||
resolved_api_base: Final = (
|
||||
api_base if api_base is not None or params_mapping is None else params_mapping.get("api_base")
|
||||
)
|
||||
return params_mapping, resolved_api_base
|
||||
|
||||
@staticmethod
|
||||
def _finalize_headers(headers: dict, auth_header: Mapping[str, str] | None) -> dict: # mutable-ok: out-param
|
||||
|
|
|
|||
|
|
@ -147,6 +147,7 @@ async def anthropic_messages_with_mcp(
|
|||
|
||||
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
tool_calls=list(tool_use_blocks),
|
||||
user_api_key_auth=context.user_api_key_auth,
|
||||
mcp_auth_header=context.mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from contextlib import asynccontextmanager
|
|||
from dataclasses import dataclass, replace
|
||||
from functools import lru_cache
|
||||
from itertools import chain, groupby
|
||||
from types import MappingProxyType
|
||||
from types import EllipsisType, MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
|
||||
from urllib.parse import ParseResult, urlparse
|
||||
|
||||
|
|
@ -259,6 +259,20 @@ _user_env_vars_cache: Final[dict[tuple[str, str], tuple[dict[str, str], float]]]
|
|||
_USER_ENV_VARS_CACHE_TTL: Final = 60 # seconds
|
||||
_USER_ENV_VARS_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth
|
||||
|
||||
_ListedToolsByCaller: TypeAlias = Mapping[str | None, Mapping[str, MCPTool]]
|
||||
_LISTED_TOOLS_CALLERS_PER_SERVER: Final = 256
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ListedToolsCaller:
|
||||
"""Request inputs that select which upstream catalog a caller was shown by tools/list."""
|
||||
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None
|
||||
mcp_auth_header: str | dict[str, str] | None = None
|
||||
raw_headers: Mapping[str, str] | None = None
|
||||
oauth2_headers: Mapping[str, str] | None = None
|
||||
|
||||
|
||||
# Auth types whose upstream OAuth endpoints (protected-resource + authorization-server metadata) the
|
||||
# gateway discovers from the upstream itself: interactive oauth2 and the two client-forwarded modes.
|
||||
# OBO/M2M endpoint discovery is decided separately via _obo_needs_endpoint_discovery. Shared by the
|
||||
|
|
@ -1172,6 +1186,57 @@ def _authorization_is_litellm_admission_credential(
|
|||
return bool(user_api_key_auth and user_api_key_auth.api_key and not admission_header)
|
||||
|
||||
|
||||
def _server_auth_header_for(
|
||||
server: MCPServer,
|
||||
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""Server-specific ``x-mcp-<alias>-authorization`` header, else the deprecated global one."""
|
||||
server_specific: Final = (
|
||||
lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=server.alias,
|
||||
server_name=server.server_name,
|
||||
access_groups=server.access_groups,
|
||||
)
|
||||
if mcp_server_auth_headers
|
||||
else None
|
||||
)
|
||||
return mcp_auth_header if server_specific is None else server_specific
|
||||
|
||||
|
||||
def listed_tools_caller_for(
|
||||
server: MCPServer,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
mcp_server_auth_headers: Mapping[str, str | dict[str, str]] | None,
|
||||
raw_headers: Mapping[str, str] | None,
|
||||
oauth2_headers: Mapping[str, str] | None,
|
||||
) -> ListedToolsCaller:
|
||||
"""The caller a tools/call must look its listed entry up under: the same inputs tools/list keyed by."""
|
||||
return ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=_server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header),
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
|
||||
|
||||
def _admission_identity(
|
||||
auth: UserAPIKeyAuth, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[str | None, str | None, str | None, str | None, str | None]:
|
||||
"""The admission identity the served catalog is shaped for: the hashed key, user, team and
|
||||
organization, plus the admission credential (``x-litellm-api-key``, else ``Authorization``) of a
|
||||
caller admitted with neither a key nor a user."""
|
||||
keyless: Final = auth.api_key is None and auth.user_id is None
|
||||
credential: Final = (
|
||||
_raw_header_value(raw_headers, "x-litellm-api-key") or _raw_header_value(raw_headers, "authorization")
|
||||
if keyless
|
||||
else None
|
||||
)
|
||||
return auth.api_key, auth.user_id, auth.team_id, auth.org_id, credential
|
||||
|
||||
|
||||
def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str:
|
||||
"""Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.
|
||||
|
||||
|
|
@ -1297,6 +1362,16 @@ async def _resolve_byok_mcp_auth_header(
|
|||
return mcp_auth_header
|
||||
|
||||
|
||||
def _catalog_auth_header(
|
||||
mcp_auth_header: str | dict[str, str] | None,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType,
|
||||
) -> str | dict[str, str] | None:
|
||||
"""The header the client supplied, which keys the caller's catalog slot on both tools/list and
|
||||
tools/call. A caller that already swapped a stored BYOK credential into ``mcp_auth_header`` passes
|
||||
the client's value explicitly, since the stored credential must never be read to find the slot."""
|
||||
return mcp_auth_header if catalog_auth_header is ... else catalog_auth_header
|
||||
|
||||
|
||||
def _client_forwarded_authorization_headers(
|
||||
mcp_server: MCPServer,
|
||||
oauth2_headers: dict[str, str] | None,
|
||||
|
|
@ -1925,6 +2000,8 @@ class MCPServerManager:
|
|||
"gmail_send_email": "zapier_mcp_server",
|
||||
}
|
||||
"""
|
||||
self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list
|
||||
self._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save
|
||||
self._upstream_initialize_instructions_by_server_id: dict[str, str] = {}
|
||||
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
|
||||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
|
|
@ -3242,6 +3319,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Added MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3279,6 +3358,8 @@ class MCPServerManager:
|
|||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3754,25 +3835,14 @@ class MCPServerManager:
|
|||
verbose_logger.warning("MCP Server %s not found", server_id)
|
||||
return []
|
||||
|
||||
# Get server-specific auth header if available
|
||||
server_auth_header: str | dict[str, str] | None = None
|
||||
if mcp_server_auth_headers:
|
||||
server_auth_header = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=server.alias,
|
||||
server_name=server.server_name,
|
||||
access_groups=server.access_groups,
|
||||
)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
try:
|
||||
tools: Final = await self._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
record_listing=True,
|
||||
)
|
||||
return tools
|
||||
except Exception as e:
|
||||
|
|
@ -3852,7 +3922,7 @@ class MCPServerManager:
|
|||
def _build_stdio_env(
|
||||
self,
|
||||
server: MCPServer,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
) -> dict[str, str] | None:
|
||||
"""Resolve stdio env values, supporting header-driven placeholders."""
|
||||
|
||||
|
|
@ -4147,13 +4217,20 @@ class MCPServerManager:
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> list[MCPTool]:
|
||||
*,
|
||||
catalog_auth_header: str | dict[str, str] | None | EllipsisType = ...,
|
||||
record_listing: bool = False,
|
||||
) -> Sequence[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
||||
Args:
|
||||
server (MCPServer): The server to query tools from
|
||||
mcp_auth_header: Optional auth header for MCP server
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header``
|
||||
record_listing: Record the served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
List[MCPTool]: List of tools available on the server with prefixed names
|
||||
|
|
@ -4169,6 +4246,13 @@ class MCPServerManager:
|
|||
verbose_logger.info("_get_tools_from_server for %s...", server.name)
|
||||
|
||||
client = None
|
||||
listed_caller: Final = ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=_catalog_auth_header(mcp_auth_header, catalog_auth_header),
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
)
|
||||
listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0)
|
||||
|
||||
try:
|
||||
# Tool *listing* must not be blocked by missing per-user env vars —
|
||||
|
|
@ -4266,8 +4350,12 @@ class MCPServerManager:
|
|||
# applied (e.g. "test_petstore-getinventory"). Do NOT pass them
|
||||
# through _create_prefixed_tools — that would add the prefix a second
|
||||
# time producing "test_petstore-test_petstore-getinventory".
|
||||
unprefixed_tools: Final = guarded_openapi
|
||||
self._record_listed_tools(
|
||||
server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
if not add_prefix:
|
||||
return list(guarded_openapi)
|
||||
return unprefixed_tools
|
||||
return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi]
|
||||
else:
|
||||
tools = await self._fetch_tools_with_timeout(client, server.name)
|
||||
|
|
@ -4281,7 +4369,10 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
)
|
||||
prefixed_or_original_tools: Final = self._create_prefixed_tools(
|
||||
list(guarded_tools), server, add_prefix=add_prefix
|
||||
guarded_tools, server, add_prefix=add_prefix
|
||||
)
|
||||
self._record_listed_tools(
|
||||
server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing
|
||||
)
|
||||
|
||||
return prefixed_or_original_tools
|
||||
|
|
@ -4332,8 +4423,99 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._drop_listed_tools(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _drop_listed_tools(self, server_id: str) -> None:
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1
|
||||
|
||||
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
|
||||
"""Key the listed-tool cache by every request input that can change the served catalog.
|
||||
|
||||
The catalog is guardrail-shaped for the caller's admission identity (default-on guardrails,
|
||||
key or team selections and opt-outs), so every admitted caller gets its own slot, keyed by
|
||||
``_admission_identity``: the hashed key, user, team and organization, plus the admission
|
||||
credential of a caller admitted with neither a key nor a user (a team-only JWT). Forwarded
|
||||
headers, header-driven stdio env, the caller bearer on every server whose egress forwards it
|
||||
(``_consumes_caller_authorization``) or exchanges it as the OBO subject, and the
|
||||
server-specific auth header also reach upstream and split the slot further. Only unkeyed
|
||||
listings with none of those share the ``None`` slot.
|
||||
"""
|
||||
if caller is None:
|
||||
return None
|
||||
auth: Final = caller.user_api_key_auth
|
||||
forwarded: Final = self._forwarded_header_values(server, caller.raw_headers) or None
|
||||
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
|
||||
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
|
||||
caller_bearer: Final = (
|
||||
self._extract_subject_token(caller.oauth2_headers, caller.raw_headers, auth)
|
||||
if _consumes_caller_authorization(server) or server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
else None
|
||||
)
|
||||
identity: Final = None if auth is None else _admission_identity(auth, caller.raw_headers)
|
||||
if not (identity or caller.mcp_auth_header or forwarded or stdio_env or caller_bearer):
|
||||
return None
|
||||
material: Final = json.dumps(
|
||||
(identity, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer),
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return hashlib.sha256(material.encode()).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _forwarded_header_values(
|
||||
server: MCPServer, raw_headers: Mapping[str, str] | None
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
if not raw_headers or not server.extra_headers:
|
||||
return ()
|
||||
forwarded_names: Final = frozenset(name.lower() for name in server.extra_headers)
|
||||
return tuple(
|
||||
sorted((name.lower(), value) for name, value in raw_headers.items() if name.lower() in forwarded_names)
|
||||
)
|
||||
|
||||
def listed_tools_generation(self, server_id: str) -> int:
|
||||
return self._listed_tools_generations.get(server_id, 0)
|
||||
|
||||
def record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> None:
|
||||
self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing)
|
||||
|
||||
def _record_listed_tools(
|
||||
self,
|
||||
server: MCPServer,
|
||||
tools: Sequence[MCPTool],
|
||||
caller: ListedToolsCaller | None,
|
||||
generation: int | None = None,
|
||||
*,
|
||||
record_listing: bool = True,
|
||||
) -> None:
|
||||
"""Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation
|
||||
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
|
||||
if not record_listing:
|
||||
return
|
||||
if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0):
|
||||
return
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
|
||||
existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({}))
|
||||
shared: Final = existing.get(None)
|
||||
callers: Final = tuple((key, value) for key, value in existing.items() if key not in (None, identity))
|
||||
evicted: Final = 0 if identity is None else max(len(callers) + 1 - _LISTED_TOOLS_CALLERS_PER_SERVER, 0)
|
||||
entries: Final = (
|
||||
*(() if shared is None else ((None, shared),)),
|
||||
*callers[evicted:],
|
||||
(identity, listing),
|
||||
)
|
||||
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(dict(entries))
|
||||
|
||||
def _discovery_key(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -4343,9 +4525,11 @@ class MCPServerManager:
|
|||
stdio_env: dict[str, str] | None,
|
||||
subject_token: str | None,
|
||||
credential_fingerprint: str | None = None,
|
||||
per_caller: bool = False,
|
||||
) -> _DiscoveryKey:
|
||||
per_user: Final = (
|
||||
server.requires_per_user_auth
|
||||
per_caller
|
||||
or server.requires_per_user_auth
|
||||
or self._references_per_user_env_var(server)
|
||||
or server.delegate_auth_to_upstream
|
||||
or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag)
|
||||
|
|
@ -5253,7 +5437,12 @@ class MCPServerManager:
|
|||
{seen: kept for seen, kept in self._catalog_alert_signatures.items() if seen != key}
|
||||
)
|
||||
|
||||
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
|
||||
def _create_prefixed_tools(
|
||||
self,
|
||||
tools: Sequence[MCPTool],
|
||||
server: MCPServer,
|
||||
add_prefix: bool = True,
|
||||
) -> list[MCPTool]:
|
||||
"""
|
||||
Create prefixed tools and update tool mapping.
|
||||
|
||||
|
|
@ -5281,6 +5470,13 @@ class MCPServerManager:
|
|||
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
|
||||
return prefixed_tools
|
||||
|
||||
def get_listed_tool(self, server: MCPServer, name: str, caller: ListedToolsCaller | None = None) -> MCPTool | None:
|
||||
identity: Final = self._listed_tools_identity(server, caller)
|
||||
listed: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})).get(identity)
|
||||
if not listed:
|
||||
return None
|
||||
return listed.get(name)
|
||||
|
||||
def _create_prefixed_prompts(
|
||||
self, prompts: Sequence[Prompt], server: MCPServer, add_prefix: bool = True
|
||||
) -> list[Prompt]:
|
||||
|
|
@ -5514,6 +5710,7 @@ class MCPServerManager:
|
|||
raw_headers: dict[str, str] | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
tool: MCPTool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Run pre-call checks and guardrail hooks for an MCP tool call.
|
||||
|
|
@ -5527,6 +5724,9 @@ class MCPServerManager:
|
|||
``pre_mcp_call`` evaluation (or a block) on the spend-log row the Guardrails
|
||||
Monitor counts. It stays optional so callers that do no logging are unchanged.
|
||||
|
||||
``tool`` is the upstream tool definition when one was listed, so guardrails
|
||||
can see its description and input schema, not just the name and arguments.
|
||||
|
||||
Returns a dict that may contain:
|
||||
- "arguments": hook-modified tool arguments (only if changed)
|
||||
- "extra_headers": headers injected by pre_mcp_call guardrail hooks
|
||||
|
|
@ -5590,6 +5790,8 @@ class MCPServerManager:
|
|||
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
|
||||
"incoming_bearer_token": incoming_bearer_token,
|
||||
"headers": logging_safe_mcp_headers(raw_headers),
|
||||
"tool_description": tool.description if tool is not None else None,
|
||||
"tool_input_schema": tool.input_schema if tool is not None else None,
|
||||
}
|
||||
|
||||
# Create MCP request object for processing
|
||||
|
|
@ -5805,21 +6007,7 @@ class MCPServerManager:
|
|||
GuardrailRaisedException: If guardrails block the call
|
||||
HTTPException: If an HTTP error occurs
|
||||
"""
|
||||
# Get server-specific auth header if available (case-insensitive)
|
||||
# FIX: Added case-insensitive matching to handle auth header keys that may not match
|
||||
# the exact case of server alias/name (e.g., '1litellmagcgateway' vs '1LiteLLMAGCGateway')
|
||||
server_auth_header: dict[str, str] | str | None = None
|
||||
if mcp_server_auth_headers:
|
||||
server_auth_header = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=mcp_server.alias,
|
||||
server_name=mcp_server.server_name,
|
||||
access_groups=mcp_server.access_groups,
|
||||
)
|
||||
|
||||
# Fall back to deprecated mcp_auth_header if no server-specific header found
|
||||
if server_auth_header is None:
|
||||
server_auth_header = mcp_auth_header
|
||||
server_auth_header: Final = _server_auth_header_for(mcp_server, mcp_server_auth_headers, mcp_auth_header)
|
||||
|
||||
# Extract subject token for OAuth2 Token Exchange (OBO) and ID-JAG flows
|
||||
subject_token: str | None = None
|
||||
|
|
@ -6242,6 +6430,9 @@ class MCPServerManager:
|
|||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None | EllipsisType = ...,
|
||||
listed_tool: MCPTool | None | EllipsisType = ...,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
|
@ -6253,6 +6444,8 @@ class MCPServerManager:
|
|||
user_api_key_auth: User authentication
|
||||
mcp_auth_header: MCP auth header (deprecated)
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
catalog_auth_header: The header the client supplied, keying the caller's catalog slot;
|
||||
defaults to ``mcp_auth_header`` as received, before BYOK resolution
|
||||
proxy_logging_obj: Optional ProxyLogging object for hook integration
|
||||
litellm_logging_obj: Optional request logger the guardrail hooks record
|
||||
their evaluations onto, so MCP guardrail activity reaches the
|
||||
|
|
@ -6264,6 +6457,7 @@ class MCPServerManager:
|
|||
"""
|
||||
start_time: Final = datetime.datetime.now()
|
||||
mcp_server: Final = self._resolve_mcp_server_for_tool_call(server_name, name)
|
||||
client_auth_header: Final = _catalog_auth_header(mcp_auth_header, catalog_auth_header)
|
||||
|
||||
# Resolved before any hook runs so a missing BYOK credential (401) never
|
||||
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)
|
||||
|
|
@ -6273,6 +6467,9 @@ class MCPServerManager:
|
|||
user_api_key_auth,
|
||||
mcp_auth_header,
|
||||
)
|
||||
listed_caller: Final = listed_tools_caller_for(
|
||||
mcp_server, user_api_key_auth, client_auth_header, mcp_server_auth_headers, raw_headers, oauth2_headers
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# Pre MCP Tool Call Hook
|
||||
|
|
@ -6289,6 +6486,7 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool,
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
|
|
|||
|
|
@ -81,12 +81,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
ListedToolsCaller,
|
||||
MCPServerManager,
|
||||
_caller_authorization_fans_out,
|
||||
_client_forwarded_authorization_headers,
|
||||
_resolve_openapi_tool_auth,
|
||||
_should_strip_caller_authorization,
|
||||
global_mcp_server_manager,
|
||||
listed_tools_caller_for,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
_redact_mcp_resource_url,
|
||||
|
|
@ -954,6 +956,8 @@ async def _get_tools_from_mcp_servers(
|
|||
request_tags: list[str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
Helper method to fetch tools from MCP servers based on server filtering criteria.
|
||||
|
|
@ -964,6 +968,8 @@ async def _get_tools_from_mcp_servers(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers
|
||||
oauth2_headers: Optional dict of oauth2 headers
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from filtered servers plus each server's
|
||||
|
|
@ -1111,12 +1117,14 @@ async def _get_tools_from_mcp_servers(
|
|||
prefetched_creds=_prefetched_oauth_creds,
|
||||
)
|
||||
|
||||
catalog_auth_header: Final = server_auth_header
|
||||
if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None:
|
||||
server_auth_header = await _get_byok_credential(server, user_api_key_auth)
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
|
||||
tools: Final = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
|
|
@ -1127,6 +1135,8 @@ async def _get_tools_from_mcp_servers(
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
oauth2_headers=oauth2_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
record_listing=False,
|
||||
)
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
|
|
@ -1135,6 +1145,21 @@ async def _get_tools_from_mcp_servers(
|
|||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
global_mcp_server_manager.record_listed_tools(
|
||||
server,
|
||||
[
|
||||
tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server)})
|
||||
for tool in filtered_tools
|
||||
],
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=catalog_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
),
|
||||
listed_generation,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
if mcp_proxy_mode:
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity
|
||||
|
|
@ -1458,6 +1483,8 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source: str | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
*,
|
||||
record_listing: bool = False,
|
||||
) -> AggregateToolListing:
|
||||
"""
|
||||
List all available MCP tools.
|
||||
|
|
@ -1468,6 +1495,8 @@ async def _list_mcp_tools(
|
|||
mcp_servers: Optional list of server names/aliases to filter by
|
||||
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
|
||||
client_ip: Client IP for IP-based server access control
|
||||
record_listing: Record each served catalog into the caller's listed-tools slot; only a
|
||||
listing actually served to the caller sets it
|
||||
|
||||
Returns:
|
||||
AggregateToolListing: Combined tools from all accessible servers plus each server's
|
||||
|
|
@ -1486,6 +1515,7 @@ async def _list_mcp_tools(
|
|||
list_tools_log_source=list_tools_log_source,
|
||||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
|
||||
return listing
|
||||
|
|
@ -1798,6 +1828,7 @@ async def _list_tools_before_first_call(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
record_listing=False,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before
|
||||
verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e)
|
||||
|
|
@ -2001,6 +2032,7 @@ async def _execute_mcp_tool(
|
|||
if mcp_server is None:
|
||||
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
|
||||
client_auth_header: Final = mcp_auth_header
|
||||
if mcp_server:
|
||||
standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info")
|
||||
if litellm_logging_obj:
|
||||
|
|
@ -2071,6 +2103,18 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
mcp_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
mcp_server,
|
||||
user_api_key_auth,
|
||||
client_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
raw_headers,
|
||||
oauth2_headers,
|
||||
),
|
||||
),
|
||||
)
|
||||
# `pre_call_tool_check` may return guardrail-modified
|
||||
# arguments; honor them on the local path too.
|
||||
|
|
@ -2120,6 +2164,7 @@ async def _execute_mcp_tool(
|
|||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
catalog_auth_header=client_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
|
|
@ -2139,7 +2184,8 @@ async def _execute_mcp_tool(
|
|||
# not in the registry either, `_handle_local_mcp_tool` below reports
|
||||
# 404 and nothing runs, so demanding a server here would turn every
|
||||
# unknown tool name into a misleading 503.
|
||||
if global_mcp_tool_registry.get_tool(original_tool_name) is not None:
|
||||
registered_local_tool: Final = global_mcp_tool_registry.get_tool(original_tool_name)
|
||||
if registered_local_tool is not None:
|
||||
# `mcp_server` is None here because the tool name is not in the
|
||||
# tool -> server mapping, but the name still carries a prefix
|
||||
# that the server-level check above compared against the
|
||||
|
|
@ -2181,6 +2227,18 @@ async def _execute_mcp_tool(
|
|||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
tool=global_mcp_server_manager.get_listed_tool(
|
||||
prefix_server,
|
||||
original_tool_name,
|
||||
listed_tools_caller_for(
|
||||
prefix_server,
|
||||
user_api_key_auth,
|
||||
client_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
raw_headers,
|
||||
oauth2_headers,
|
||||
),
|
||||
),
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
|
||||
|
|
@ -2583,8 +2641,12 @@ async def _handle_managed_mcp_tool(
|
|||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
*,
|
||||
catalog_auth_header: str | None,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""Handle tool execution for managed server tools"""
|
||||
"""Handle tool execution for managed server tools. ``catalog_auth_header`` is the header the client
|
||||
supplied, which keys the caller's catalog slot; ``mcp_auth_header`` may already be the resolved
|
||||
BYOK credential."""
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
|
|
@ -2594,6 +2656,7 @@ async def _handle_managed_mcp_tool(
|
|||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
catalog_auth_header=catalog_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
|
|
@ -2712,6 +2775,7 @@ async def _execute_handle_list_tools(
|
|||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
record_listing=True,
|
||||
)
|
||||
verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools))
|
||||
if not listing.outcomes:
|
||||
|
|
|
|||
|
|
@ -223,6 +223,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
_UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES,
|
||||
ListedToolsCaller,
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
|
|
@ -701,6 +702,8 @@ if MCP_AVAILABLE:
|
|||
extra_headers: dict[str, str] | None,
|
||||
client_ip: str | None,
|
||||
proxy_logging_obj: "ProxyLogging | None",
|
||||
*,
|
||||
record_listing: bool,
|
||||
) -> list[MCPTool]:
|
||||
return await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
|
|
@ -711,11 +714,12 @@ if MCP_AVAILABLE:
|
|||
client_ip=client_ip,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
record_listing=record_listing,
|
||||
)
|
||||
|
||||
async def _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
server: MCPServer,
|
||||
server_auth_header: dict[str, str] | str | None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -731,31 +735,44 @@ if MCP_AVAILABLE:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
||||
tools = await _list_server_tools(
|
||||
server, server_auth_header, raw_headers, user_api_key_auth, extra_headers, client_ip, proxy_logging_obj
|
||||
listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id)
|
||||
tools: Final = await _list_server_tools(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers,
|
||||
user_api_key_auth,
|
||||
extra_headers,
|
||||
client_ip,
|
||||
proxy_logging_obj,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
if not apply_tool_filters:
|
||||
return _create_tool_response_objects(tools, server)
|
||||
|
||||
# Always apply allowed_tools/disallowed_tools so the blacklist is
|
||||
# enforced even when no allowlist is set (matches the SSE/HTTP path).
|
||||
tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
||||
# Filter by the key's effective tool permissions through the same
|
||||
# function the MCP protocol path uses (direct grants, toolset grants,
|
||||
# and team/agent/org ceilings), so REST listing cannot drift from it.
|
||||
# Entries here are tool names on one server, written bare by every
|
||||
# writer, and dispatch compares them bare; matching a wider set of
|
||||
# spellings would advertise a tool that tools/call then refuses
|
||||
if user_api_key_auth:
|
||||
tools = await filter_tools_by_key_team_permissions(
|
||||
tools=tools,
|
||||
server_filtered: Final = filter_tools_by_allowed_tools(tools, server) if apply_tool_filters else tools
|
||||
served_tools: Final = (
|
||||
await filter_tools_by_key_team_permissions(
|
||||
tools=server_filtered,
|
||||
server_id=server.server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
if apply_tool_filters and user_api_key_auth
|
||||
else server_filtered
|
||||
)
|
||||
if apply_tool_filters:
|
||||
# Only a listing shaped for the caller's runtime view may set their
|
||||
# listed-tools slot; the admin-only unfiltered configuration view
|
||||
# must not warm it.
|
||||
global_mcp_server_manager.record_listed_tools(
|
||||
server,
|
||||
served_tools,
|
||||
ListedToolsCaller(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=server_auth_header,
|
||||
raw_headers=raw_headers,
|
||||
),
|
||||
listed_generation,
|
||||
)
|
||||
|
||||
return _create_tool_response_objects(tools, server)
|
||||
return _create_tool_response_objects(served_tools, server)
|
||||
|
||||
async def fetch_pinnable_tool_catalog(
|
||||
server: MCPServer, request: Request, user_api_key_dict: UserAPIKeyAuth
|
||||
|
|
@ -774,6 +791,7 @@ if MCP_AVAILABLE:
|
|||
await _get_user_oauth_extra_headers(server, user_api_key_dict),
|
||||
IPAddressUtils.get_mcp_client_ip(request),
|
||||
None,
|
||||
record_listing=False,
|
||||
)
|
||||
scan: Final = await scan_tool_descriptions(
|
||||
apply_description_overrides(upstream, server), server, proxy_logging_obj, user_api_key_dict, raw_headers
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
|
|||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -61,6 +61,7 @@ _GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
|
|||
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
|
||||
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
|
||||
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
|
||||
_TOOL_INPUT_SCHEMA_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_OBO_CACHE_MAX_ENTRIES: Final = 1000
|
||||
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
|
||||
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
|
||||
|
|
@ -82,6 +83,13 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
|
|||
return ()
|
||||
|
||||
|
||||
def _parse_tool_input_schema(raw: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _TOOL_INPUT_SCHEMA_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def entra_assertion(value: object) -> str | None:
|
||||
"""``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion.
|
||||
A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``."""
|
||||
|
|
@ -100,6 +108,14 @@ class _EvaluateResponse(TypedDict, total=False):
|
|||
correlationId: ReadOnly[str]
|
||||
|
||||
|
||||
class _ToolReference(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
name: str
|
||||
description: str | None = None
|
||||
input_schema: Mapping[str, object] | None = Field(default=None, serialization_alias="inputSchema")
|
||||
|
||||
|
||||
class _UnavailableDetail(TypedDict):
|
||||
error: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
|
@ -392,8 +408,14 @@ class Agent365Guardrail(CustomGuardrail):
|
|||
arguments: Final = data.get("mcp_arguments")
|
||||
server_name: Final = str(data.get("mcp_server_name") or "litellm")
|
||||
agent_id: Final = user_api_key_dict.key_alias
|
||||
description: Final = data.get("mcp_tool_description")
|
||||
tool_reference: Final = _ToolReference(
|
||||
name=tool_name,
|
||||
description=description if isinstance(description, str) and description else None,
|
||||
input_schema=_parse_tool_input_schema(data.get("mcp_input_schema")),
|
||||
)
|
||||
payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below
|
||||
"tool": {"name": tool_name},
|
||||
"tool": tool_reference.model_dump(by_alias=True, exclude_none=True),
|
||||
"serverName": server_name,
|
||||
"conversationId": self._resolve_conversation_id(data),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,4 +32,5 @@ def worker_image() -> str:
|
|||
override: Final = os.environ.get("LENS_WORKER_IMAGE", "")
|
||||
if override:
|
||||
return override
|
||||
return f"ghcr.io/berriai/litellm-lens-worker:{tag}"
|
||||
package: Final = "litellm-lens-worker-dev" if tag.startswith("sha-") else "litellm-lens-worker"
|
||||
return f"ghcr.io/berriai/{package}:{tag}"
|
||||
|
|
|
|||
|
|
@ -1018,6 +1018,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=None,
|
||||
mcp_servers=None,
|
||||
mcp_server_auth_headers=None,
|
||||
record_listing=True,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
dumped_tools: Final = [tool.model_dump(by_alias=True) for tool in tools]
|
||||
|
|
|
|||
|
|
@ -1151,6 +1151,8 @@ def _overrides_moderation_hook(callback: CustomLogger) -> bool:
|
|||
|
||||
|
||||
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
|
||||
_MCP_TOOL_DESCRIPTION: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
|
||||
_MCP_TOOL_INPUT_SCHEMA: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -1475,9 +1477,14 @@ class ProxyLogging:
|
|||
TypeAdapter(dict[str, object]).validate_python(guardrail_context.get("metadata") or MappingProxyType({}))
|
||||
)
|
||||
|
||||
mcp_tool_description: Final = kwargs.get("mcp_tool_description")
|
||||
mcp_input_schema: Final = kwargs.get("mcp_input_schema")
|
||||
description_line: Final = f"\nDescription: {mcp_tool_description}" if mcp_tool_description else ""
|
||||
mcp_tool_description: Final = request_obj.tool_description or kwargs.get("mcp_tool_description")
|
||||
mcp_input_schema: Final = (
|
||||
request_obj.tool_input_schema
|
||||
if request_obj.tool_input_schema is not None
|
||||
else kwargs.get("mcp_input_schema")
|
||||
)
|
||||
listing_description: Final = kwargs.get("mcp_tool_description")
|
||||
description_line: Final = f"\nDescription: {listing_description}" if listing_description else ""
|
||||
tool_call_content: Final = (
|
||||
f"Tool: {request_obj.tool_name}{description_line}\nArguments: {request_obj.arguments}"
|
||||
)
|
||||
|
|
@ -1735,6 +1742,8 @@ class ProxyLogging:
|
|||
tool_name=kwargs.get("name", ""),
|
||||
arguments=kwargs.get("arguments", {}),
|
||||
server_name=kwargs.get("server_name"),
|
||||
tool_description=_MCP_TOOL_DESCRIPTION.validate_python(kwargs.get("tool_description")),
|
||||
tool_input_schema=_MCP_TOOL_INPUT_SCHEMA.validate_python(kwargs.get("tool_input_schema")),
|
||||
user_api_key_auth=user_api_key_auth_dict,
|
||||
hidden_params=HiddenParams(),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -286,6 +286,7 @@ async def aresponses_api_with_mcp(
|
|||
call_params=call_params,
|
||||
previous_response_id=previous_response_id,
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
**kwargs,
|
||||
)
|
||||
await mcp_streaming_response._create_initial_response_iterator()
|
||||
|
|
@ -339,6 +340,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -395,6 +397,7 @@ async def aresponses_api_with_mcp(
|
|||
|
||||
final_response = MCPEnhancedStreamingIterator(
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=original_mcp_tools,
|
||||
base_iterator=final_response,
|
||||
mcp_events=tool_execution_events,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
|
|
|
|||
|
|
@ -435,6 +435,7 @@ async def acompletion_with_mcp(
|
|||
# Execute tool calls
|
||||
self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
tool_calls=self.tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
@ -609,6 +610,7 @@ async def acompletion_with_mcp(
|
|||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=tool_server_map,
|
||||
tool_calls=tool_calls,
|
||||
served_tools=deduplicated_mcp_tools,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
|
|
|
|||
|
|
@ -696,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
litellm_trace_id: str | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
) -> list[MCPToolResult]:
|
||||
"""Execute tool calls and return results."""
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -860,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
listed_tool=(
|
||||
next((tool for tool in served_tools if tool.name == tool_name), None)
|
||||
if served_tools is not None
|
||||
else ...
|
||||
),
|
||||
)
|
||||
|
||||
if proxy_logging_obj:
|
||||
|
|
@ -1152,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
call_params: Mapping[str, object],
|
||||
previous_response_id: str | None,
|
||||
tool_server_map: dict[str, str],
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
"""
|
||||
|
|
@ -1181,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
base_iterator=None, # Will be created internally
|
||||
mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events
|
||||
tool_server_map=tool_server_map,
|
||||
served_tools=served_tools,
|
||||
mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy,
|
||||
user_api_key_auth=kwargs.get("user_api_key_auth")
|
||||
or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"),
|
||||
|
|
|
|||
|
|
@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None,
|
||||
user_api_key_auth: "UserAPIKeyAuth | None" = None,
|
||||
original_request_params: dict[str, Any] | None = None,
|
||||
served_tools: Sequence[MCPTool] | None = None,
|
||||
):
|
||||
# MCP setup
|
||||
self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or []
|
||||
|
|
@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
self.mcp_discovery_generated = True # Events are already generated
|
||||
self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility
|
||||
self.tool_server_map = tool_server_map
|
||||
self.served_tools = tuple(served_tools) if served_tools is not None else None
|
||||
|
||||
# Iterator references
|
||||
self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = (
|
||||
|
|
@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
# Execute the tools
|
||||
tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=self.tool_server_map,
|
||||
served_tools=self.served_tools,
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=self.user_api_key_auth,
|
||||
mcp_auth_header=self.mcp_auth_header,
|
||||
|
|
|
|||
172
litellm/tool_loop.py
Normal file
172
litellm/tool_loop.py
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
"""Client-side tool-calling loop helpers for litellm.completion."""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from typing import TypeAlias, cast
|
||||
|
||||
from typing_extensions import TypedDict, Unpack
|
||||
|
||||
from litellm.constants import DEFAULT_TOOL_LOOP_MAX_ROUNDS
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionAssistantMessage,
|
||||
ChatCompletionAssistantToolCall,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
ChatCompletionToolMessage,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Message, ModelResponse
|
||||
|
||||
ToolExecutor: TypeAlias = Callable[[ChatCompletionMessageToolCall], ChatCompletionToolMessage]
|
||||
AsyncToolExecutor: TypeAlias = Callable[[ChatCompletionMessageToolCall], Awaitable[ChatCompletionToolMessage]]
|
||||
|
||||
|
||||
class _ToolLoopCompletionKwargs(TypedDict, total=False, extra_items=object):
|
||||
"""Extra keywords forwarded verbatim to ``litellm.completion``, which owns their contract."""
|
||||
|
||||
|
||||
class ToolLoopMaxRoundsExceeded(RuntimeError):
|
||||
max_rounds: int
|
||||
|
||||
def __init__(self, max_rounds: int) -> None:
|
||||
self.max_rounds = max_rounds
|
||||
super().__init__(f"model still requested tool calls on round {max_rounds} of {max_rounds}")
|
||||
|
||||
|
||||
def _validate_tool_loop_args(max_rounds: int, completion_kwargs: _ToolLoopCompletionKwargs) -> None:
|
||||
if max_rounds < 1:
|
||||
raise ValueError(f"max_rounds must be >= 1, got {max_rounds}")
|
||||
if completion_kwargs.get("stream"):
|
||||
raise ValueError("run_tool_loop requires whole responses; stream=True is not supported")
|
||||
|
||||
|
||||
def _expect_model_response(response: object) -> ModelResponse:
|
||||
if not isinstance(response, ModelResponse):
|
||||
raise TypeError(f"run_tool_loop requires completion to return a ModelResponse, got {type(response).__name__}")
|
||||
return response
|
||||
|
||||
|
||||
def _assistant_tool_call(tool_call: ChatCompletionMessageToolCall) -> ChatCompletionAssistantToolCall:
|
||||
return ChatCompletionAssistantToolCall(
|
||||
id=tool_call.id,
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=tool_call.function.name, arguments=tool_call.function.arguments
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _assistant_message(
|
||||
message: Message, tool_calls: tuple[ChatCompletionMessageToolCall, ...]
|
||||
) -> ChatCompletionAssistantMessage:
|
||||
return cast( # cast-ok: dict literal with thinking_blocks and reasoning_items spread in only when set
|
||||
"ChatCompletionAssistantMessage",
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": message.content,
|
||||
"tool_calls": [_assistant_tool_call(tc) for tc in tool_calls],
|
||||
**{
|
||||
key: value
|
||||
for key, value in (
|
||||
("thinking_blocks", getattr(message, "thinking_blocks", None)),
|
||||
("reasoning_items", getattr(message, "reasoning_items", None)),
|
||||
)
|
||||
if value is not None
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _function_tool_call(tool_call: object) -> ChatCompletionMessageToolCall:
|
||||
if not isinstance(tool_call, ChatCompletionMessageToolCall):
|
||||
raise TypeError(
|
||||
f"run_tool_loop only executes function tool calls, got custom tool call {getattr(tool_call, 'id', None)}"
|
||||
)
|
||||
return tool_call
|
||||
|
||||
|
||||
def _function_tool_calls(message: Message) -> tuple[ChatCompletionMessageToolCall, ...]:
|
||||
return tuple(_function_tool_call(tool_call) for tool_call in message.tool_calls or ())
|
||||
|
||||
|
||||
def run_tool_loop(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
tools: Sequence[ChatCompletionToolParam],
|
||||
execute_tool: ToolExecutor,
|
||||
max_rounds: int = DEFAULT_TOOL_LOOP_MAX_ROUNDS,
|
||||
**completion_kwargs: Unpack[_ToolLoopCompletionKwargs], # kwargs-ok: forwarded verbatim to litellm.completion
|
||||
) -> str | None:
|
||||
"""Call completion, execute each requested tool, and repeat until the model answers.
|
||||
|
||||
Returns the final assistant message content. Raises ToolLoopMaxRoundsExceeded when the
|
||||
model is still requesting tools after max_rounds completions.
|
||||
|
||||
Example:
|
||||
execute = functools.partial(run_repo_tool, repository="litellm", revision="main")
|
||||
answer = litellm.run_tool_loop(
|
||||
model="anthropic/claude-sonnet-5-5", messages=messages, tools=tools, execute_tool=execute
|
||||
)
|
||||
"""
|
||||
import litellm
|
||||
|
||||
_validate_tool_loop_args(max_rounds, completion_kwargs)
|
||||
history: tuple[AllMessageValues, ...] = tuple(messages) # rebind-ok: rounds append new turns
|
||||
for round_number in range(1, max_rounds + 1):
|
||||
message = (
|
||||
_expect_model_response(
|
||||
litellm.completion(model=model, messages=list(history), tools=list(tools), **completion_kwargs)
|
||||
)
|
||||
.choices[0]
|
||||
.message
|
||||
)
|
||||
if not message.tool_calls:
|
||||
return message.content
|
||||
tool_calls = _function_tool_calls(message)
|
||||
if round_number == max_rounds:
|
||||
break
|
||||
tool_results = tuple(execute_tool(tool_call) for tool_call in tool_calls)
|
||||
history = (*history, _assistant_message(message, tool_calls), *tool_results)
|
||||
raise ToolLoopMaxRoundsExceeded(max_rounds)
|
||||
|
||||
|
||||
async def arun_tool_loop(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
tools: Sequence[ChatCompletionToolParam],
|
||||
execute_tool: AsyncToolExecutor,
|
||||
max_rounds: int = DEFAULT_TOOL_LOOP_MAX_ROUNDS,
|
||||
**completion_kwargs: Unpack[_ToolLoopCompletionKwargs], # kwargs-ok: forwarded verbatim to litellm.acompletion
|
||||
) -> str | None:
|
||||
"""Async version of run_tool_loop, awaiting each execute_tool call in order.
|
||||
|
||||
Returns the final assistant message content. Raises ToolLoopMaxRoundsExceeded when the
|
||||
model is still requesting tools after max_rounds completions.
|
||||
|
||||
Example:
|
||||
execute = functools.partial(arun_repo_tool, repository="litellm", revision="main")
|
||||
answer = await litellm.arun_tool_loop(
|
||||
model="anthropic/claude-sonnet-5-5", messages=messages, tools=tools, execute_tool=execute
|
||||
)
|
||||
"""
|
||||
import litellm
|
||||
|
||||
_validate_tool_loop_args(max_rounds, completion_kwargs)
|
||||
history: tuple[AllMessageValues, ...] = tuple(messages) # rebind-ok: rounds append new turns
|
||||
for round_number in range(1, max_rounds + 1):
|
||||
message = (
|
||||
_expect_model_response(
|
||||
await litellm.acompletion(model=model, messages=list(history), tools=list(tools), **completion_kwargs)
|
||||
)
|
||||
.choices[0]
|
||||
.message
|
||||
)
|
||||
if not message.tool_calls:
|
||||
return message.content
|
||||
tool_calls = _function_tool_calls(message)
|
||||
if round_number == max_rounds:
|
||||
break
|
||||
tool_results = tuple([await execute_tool(tool_call) for tool_call in tool_calls])
|
||||
history = (*history, _assistant_message(message, tool_calls), *tool_results)
|
||||
raise ToolLoopMaxRoundsExceeded(max_rounds)
|
||||
|
|
@ -429,6 +429,8 @@ class MCPPreCallRequestObject(BaseModel):
|
|||
tool_name: str
|
||||
arguments: dict[str, Any]
|
||||
server_name: str | None = None
|
||||
tool_description: str | None = None
|
||||
tool_input_schema: Mapping[str, object] | None = None
|
||||
user_api_key_auth: dict[str, Any] | None = None
|
||||
hidden_params: HiddenParams = HiddenParams()
|
||||
|
||||
|
|
@ -452,6 +454,8 @@ class MCPDuringCallRequestObject(BaseModel):
|
|||
tool_name: str
|
||||
arguments: dict[str, Any]
|
||||
server_name: str | None = None
|
||||
tool_description: str | None = None
|
||||
tool_input_schema: Mapping[str, object] | None = None
|
||||
start_time: float | None = None
|
||||
hidden_params: HiddenParams = HiddenParams()
|
||||
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ dependencies = [
|
|||
"pydantic-settings>=2.14.1,<3.0",
|
||||
"jsonschema>=4.0.0,<5.0",
|
||||
"boto3>=1.43.1,<2.0",
|
||||
"typing-extensions>=4.13.0,<5.0",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
|
|
|||
|
|
@ -25,6 +25,8 @@ py="${LENS_DEV_PYTHON:-$repo_root/.venv/bin/python}"
|
|||
database_url="${LENS_DEV_DATABASE_URL:-postgresql://litellm:litellm@127.0.0.1:15432/litellm}"
|
||||
clickhouse_url=http://default:local-tracing@127.0.0.1:18123
|
||||
master_key=""
|
||||
startup_timeout="${LENS_DEV_STARTUP_TIMEOUT_SECONDS:-300}"
|
||||
readiness_request_timeout="${LENS_DEV_READINESS_REQUEST_TIMEOUT_SECONDS:-5}"
|
||||
pids=()
|
||||
|
||||
die() { echo "lens-dev: $*" >&2; exit 1; }
|
||||
|
|
@ -121,6 +123,7 @@ proxy_env() {
|
|||
export CLICKHOUSE_DATABASE=litellm
|
||||
export LITELLM_LOCAL_MODEL_COST_MAP=True
|
||||
export PROXY_BASE_URL="$proxy_url"
|
||||
export LITELLM_UI_PATH="$repo_root/ui/litellm-dashboard/out"
|
||||
export UI_USERNAME=admin
|
||||
export UI_PASSWORD="$master_key"
|
||||
}
|
||||
|
|
@ -161,12 +164,26 @@ ensure_worker_token() {
|
|||
wait_for_proxy() {
|
||||
local proxy_pid="$1"
|
||||
echo "lens-dev: waiting for the proxy (log: $log_dir/proxy.log)"
|
||||
for _ in $(seq 1 300); do
|
||||
for _ in $(seq 1 "$startup_timeout"); do
|
||||
kill -0 "$proxy_pid" 2>/dev/null || die "proxy exited; see $log_dir/proxy.log"
|
||||
curl -fsS "$proxy_url/health/readiness" -H "Authorization: Bearer $master_key" >/dev/null 2>&1 && return
|
||||
curl -fsS --max-time "$readiness_request_timeout" "$proxy_url/health/readiness" -H "Authorization: Bearer $master_key" >/dev/null 2>&1 && return
|
||||
sleep 1
|
||||
done
|
||||
die "proxy not ready after 300s; see $log_dir/proxy.log"
|
||||
die "proxy not ready after ${startup_timeout}s; see $log_dir/proxy.log"
|
||||
}
|
||||
|
||||
wait_for_ui() {
|
||||
local ui_pid="$1"
|
||||
echo "lens-dev: waiting for the UI (log: $log_dir/ui.log)"
|
||||
for _ in $(seq 1 "$startup_timeout"); do
|
||||
kill -0 "$ui_pid" 2>/dev/null || die "UI exited; see $log_dir/ui.log"
|
||||
if curl -fsS --max-time "$readiness_request_timeout" "http://localhost:$ui_port/ui/login/" >/dev/null 2>&1; then
|
||||
kill -0 "$ui_pid" 2>/dev/null || die "UI exited; see $log_dir/ui.log"
|
||||
return
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
die "UI not ready after ${startup_timeout}s; see $log_dir/ui.log"
|
||||
}
|
||||
|
||||
# Children run in their own process groups (set -m), so killing -pid takes their trees too.
|
||||
|
|
@ -185,14 +202,71 @@ cleanup() {
|
|||
for pid in "${pids[@]}"; do kill -KILL -- "-$pid" 2>/dev/null || true; done
|
||||
}
|
||||
|
||||
build_dashboard() {
|
||||
local dashboard_dir="$repo_root/ui/litellm-dashboard"
|
||||
case "${LENS_DEV_BUILD_UI:-0}" in
|
||||
0) return ;;
|
||||
1) ;;
|
||||
*) die "LENS_DEV_BUILD_UI must be 0 or 1" ;;
|
||||
esac
|
||||
echo "lens-dev: building the proxy dashboard (log: $log_dir/ui-build.log)"
|
||||
(
|
||||
cd "$dashboard_dir"
|
||||
NEXT_PUBLIC_BASE_URL="" LENS_DEV_PROXY_URL="" "$repo_root/scripts/with_dashboard_node.sh" npm run build
|
||||
) > "$log_dir/ui-build.log" 2>&1 || die "UI build failed; see $log_dir/ui-build.log"
|
||||
}
|
||||
|
||||
seed_data() {
|
||||
(
|
||||
proxy_env ""
|
||||
"$py" -m scripts.seed_tracing_fixtures --profile "$seed_profile" ${seed_options[@]+"${seed_options[@]}"}
|
||||
)
|
||||
}
|
||||
|
||||
parse_args() {
|
||||
seed_profile="${LENS_DEV_SEED:-}"
|
||||
seed_only=0
|
||||
seed_options=()
|
||||
while [ "$#" -gt 0 ]; do
|
||||
case "$1" in
|
||||
--seed)
|
||||
seed_profile=default
|
||||
if [ "${2:-}" = default ] || [ "${2:-}" = large ]; then seed_profile="$2"; shift; fi
|
||||
;;
|
||||
--copies)
|
||||
[ "$#" -ge 2 ] && [[ "$2" =~ ^[1-9][0-9]*$ ]] || die "--copies requires a positive integer"
|
||||
seed_options=(--copies "$2"); shift ;;
|
||||
--seed-only) seed_only=1 ;;
|
||||
--help)
|
||||
echo "Usage: $0 [--seed [default|large]] [--copies N] [--seed-only]"
|
||||
exit 0 ;;
|
||||
*) die "unknown argument: $1 (use --help)" ;;
|
||||
esac
|
||||
shift
|
||||
done
|
||||
if [ "$seed_only" = 1 ] && [ -z "$seed_profile" ]; then seed_profile=default; fi
|
||||
case "$seed_profile" in ""|default|large) ;; *) die "seed profile must be default or large" ;; esac
|
||||
[ "${#seed_options[@]}" = 0 ] || [ -n "$seed_profile" ] || die "--copies requires --seed"
|
||||
}
|
||||
|
||||
main() {
|
||||
local config_file exports proxy_pid pid key_hint
|
||||
local config_file exports proxy_pid ui_pid pid key_hint
|
||||
parse_args "$@"
|
||||
if [ -n "${LENS_DEV_CONFIG:-}" ]; then
|
||||
[ -f "$LENS_DEV_CONFIG" ] || die "LENS_DEV_CONFIG not found: $LENS_DEV_CONFIG"
|
||||
config_file="$(cd "$(dirname "$LENS_DEV_CONFIG")" && pwd)/$(basename "$LENS_DEV_CONFIG")"
|
||||
fi
|
||||
cd "$repo_root"
|
||||
|
||||
if [ "$seed_only" = 1 ]; then
|
||||
[ -s "$key_file" ] || [ -n "${LENS_DEV_MASTER_KEY:-}" ] || die "start make lens-dev before --seed-only"
|
||||
load_master_key
|
||||
seed_data
|
||||
return
|
||||
fi
|
||||
|
||||
[[ "$startup_timeout" =~ ^[1-9][0-9]*$ ]] || die "LENS_DEV_STARTUP_TIMEOUT_SECONDS must be a positive integer"
|
||||
[[ "$readiness_request_timeout" =~ ^[1-9][0-9]*$ ]] || die "LENS_DEV_READINESS_REQUEST_TIMEOUT_SECONDS must be a positive integer"
|
||||
listening "$proxy_port" && die "port $proxy_port is in use; set LENS_DEV_PROXY_PORT"
|
||||
listening "$ui_port" && die "port $ui_port is in use; set LENS_DEV_UI_PORT"
|
||||
[ "$proxy_port" != "$ui_port" ] || die "proxy and UI ports must differ"
|
||||
|
|
@ -213,6 +287,8 @@ main() {
|
|||
(cd ui/litellm-dashboard && "$repo_root/scripts/with_dashboard_node.sh" npm ci)
|
||||
fi
|
||||
|
||||
build_dashboard
|
||||
|
||||
if [ -z "${config_file:-}" ]; then
|
||||
config_file="$state_dir/config.yaml"
|
||||
write_default_config "$config_file"
|
||||
|
|
@ -225,6 +301,7 @@ main() {
|
|||
|
||||
(
|
||||
proxy_env "$exports"
|
||||
export PROXY_BASE_URL="http://localhost:$ui_port"
|
||||
exec "$py" litellm/proxy/proxy_cli.py --config "$config_file" --host 127.0.0.1 --port "$proxy_port"
|
||||
) < /dev/null > "$log_dir/proxy.log" 2>&1 &
|
||||
proxy_pid=$!
|
||||
|
|
@ -232,12 +309,15 @@ main() {
|
|||
|
||||
(
|
||||
cd ui/litellm-dashboard
|
||||
NEXT_PUBLIC_BASE_URL="$proxy_url" exec "$repo_root/scripts/with_dashboard_node.sh" npx next dev -p "$ui_port"
|
||||
NEXT_PUBLIC_BASE_URL="" LENS_DEV_PROXY_URL="$proxy_url" exec "$repo_root/scripts/with_dashboard_node.sh" npx next dev -p "$ui_port"
|
||||
) < /dev/null > "$log_dir/ui.log" 2>&1 &
|
||||
pids+=("$!")
|
||||
ui_pid=$!
|
||||
pids+=("$ui_pid")
|
||||
|
||||
wait_for_ui "$ui_pid"
|
||||
wait_for_proxy "$proxy_pid"
|
||||
ensure_worker_token
|
||||
if [ -n "$seed_profile" ]; then seed_data; fi
|
||||
|
||||
LITELLM_RELEASE_TAG="$source_release_tag" \
|
||||
LITELLM_MODE=PRODUCTION LITELLM_URL="$proxy_url" LENS_WORKER_TOKEN="$(cat "$token_file")" \
|
||||
|
|
@ -250,12 +330,13 @@ main() {
|
|||
cat <<EOF
|
||||
|
||||
Lens dev is up. Ctrl-C stops everything.
|
||||
Log in: $proxy_url/ui/login (admin / $key_hint)
|
||||
Lens: http://localhost:$ui_port/lens
|
||||
Log in: http://localhost:$ui_port/ui/login/ (admin / $key_hint)
|
||||
Lens: http://localhost:$ui_port/ui/lens/ (hot-reloads)
|
||||
API: $proxy_url
|
||||
Logs: $log_dir/proxy.log
|
||||
$log_dir/worker.log
|
||||
$log_dir/ui.log
|
||||
Restart (Ctrl-C, make lens-dev) to pick up backend or worker edits; the UI hot-reloads.
|
||||
Restart for backend or worker edits; UI edits hot-reload.
|
||||
EOF
|
||||
|
||||
while :; do
|
||||
|
|
|
|||
|
|
@ -2,88 +2,5 @@
|
|||
set -euo pipefail
|
||||
|
||||
repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
|
||||
cd "$repo_root"
|
||||
|
||||
seed_fixtures=0
|
||||
case "${1:-}" in
|
||||
--seed) seed_fixtures=1 ;;
|
||||
"") ;;
|
||||
*) echo "Usage: $0 [--seed]" >&2; exit 2 ;;
|
||||
esac
|
||||
|
||||
if lsof -nP -iTCP:4002 -sTCP:LISTEN >/dev/null 2>&1; then
|
||||
echo "Port 4002 is already in use. Stop the existing proxy before starting this stack" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
docker compose -f docker/docker-compose.tracing.yml up -d --wait db clickhouse
|
||||
uv sync --inexact --frozen --extra proxy --group proxy-dev --no-install-project
|
||||
"$repo_root/.venv/bin/python" scripts/prisma_generate_if_needed.py
|
||||
VIRTUAL_ENV="$repo_root/.venv" uvx --from maturin==1.15.0 maturin develop \
|
||||
--release --manifest-path litellm-rust/crates/python-bridge/Cargo.toml --features extension-module
|
||||
|
||||
config_file="$(mktemp "${TMPDIR:-/tmp}/litellm-tracing-local.XXXXXX")"
|
||||
proxy_pid=""
|
||||
cleanup() {
|
||||
if [ -n "$proxy_pid" ]; then
|
||||
kill "$proxy_pid" 2>/dev/null || true
|
||||
wait "$proxy_pid" 2>/dev/null || true
|
||||
fi
|
||||
rm -f "$config_file"
|
||||
}
|
||||
trap cleanup EXIT
|
||||
trap 'exit 130' INT TERM
|
||||
cat > "$config_file" <<'EOF'
|
||||
model_list:
|
||||
- model_name: openai/gpt-6-luna
|
||||
litellm_params:
|
||||
model: openai/gpt-6-luna
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
store_prompts_in_spend_logs: true
|
||||
tracing:
|
||||
store:
|
||||
type: clickhouse
|
||||
url: os.environ/CLICKHOUSE_URL
|
||||
retention_days: 14
|
||||
EOF
|
||||
|
||||
export LITELLM_MASTER_KEY=sk-1234
|
||||
export LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true
|
||||
export LITELLM_SALT_KEY=sk-local-tracing-salt-key
|
||||
export DATABASE_URL=postgresql://litellm:litellm@127.0.0.1:15432/litellm
|
||||
export STORE_MODEL_IN_DB=True
|
||||
export CLICKHOUSE_URL=http://default:local-tracing@127.0.0.1:18123
|
||||
export CLICKHOUSE_DATABASE=litellm
|
||||
export LITELLM_LOCAL_MODEL_COST_MAP=True
|
||||
export PROXY_BASE_URL=http://127.0.0.1:4002
|
||||
|
||||
(
|
||||
cd "$repo_root/ui/litellm-dashboard"
|
||||
"$repo_root/scripts/with_dashboard_node.sh" npm ci
|
||||
NEXT_PUBLIC_BASE_URL= "$repo_root/scripts/with_dashboard_node.sh" npm run build
|
||||
)
|
||||
export LITELLM_UI_PATH="$repo_root/ui/litellm-dashboard/out"
|
||||
|
||||
printf 'Dashboard: http://127.0.0.1:4002/ui/\nProxy: http://127.0.0.1:4002\nMaster key: %s\n' "$LITELLM_MASTER_KEY"
|
||||
"$repo_root/.venv/bin/python" litellm/proxy/proxy_cli.py \
|
||||
--config "$config_file" --host 127.0.0.1 --port 4002 &
|
||||
proxy_pid=$!
|
||||
|
||||
if [ "$seed_fixtures" = "1" ]; then
|
||||
ready=0
|
||||
for attempt in $(seq 1 180); do
|
||||
kill -0 "$proxy_pid" 2>/dev/null || { echo "Proxy exited before seeding" >&2; exit 1; }
|
||||
if curl --fail --silent "$PROXY_BASE_URL/health/readiness" \
|
||||
-H "Authorization: Bearer $LITELLM_MASTER_KEY" >/dev/null; then
|
||||
ready=1
|
||||
break
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
[ "$ready" = "1" ] || { echo "Proxy did not become ready within 180 seconds" >&2; exit 1; }
|
||||
"$repo_root/.venv/bin/python" -m scripts.seed_tracing_fixtures
|
||||
fi
|
||||
|
||||
wait "$proxy_pid"
|
||||
echo "run_tracing_proxy_local.sh is deprecated; use make lens-dev" >&2
|
||||
exec "$repo_root/scripts/lens_dev.sh" "$@"
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import base64
|
||||
import binascii
|
||||
|
|
@ -10,32 +11,31 @@ import os
|
|||
import re
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from functools import cache
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
|
||||
|
||||
from litellm.rust_bridge.trace.generated.types import AllQueryScope, Trace
|
||||
from litellm.rust_bridge.trace.storage import ClickHouseStorage
|
||||
from litellm.rust_bridge.trace.storage import ClickHouseStorage, Tenant, span_rows
|
||||
from litellm.tracing.config import trace_storage_config
|
||||
from litellm.tracing.types import SpendLogRecord
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
from prisma.types import LiteLLM_SpendLogsCreateWithoutRelationsInput
|
||||
|
||||
REPO_ROOT: Final = Path(__file__).resolve().parents[1]
|
||||
TRACE_FIXTURES: Final = REPO_ROOT / "litellm-rust/crates/traces/tests/fixtures"
|
||||
SPEND_FIXTURE: Final = (
|
||||
REPO_ROOT / "litellm-rust/crates/traces-clickhouse/tests/fixtures/deeplite_swarm_spend_logs.jsonl"
|
||||
)
|
||||
SPEND_FIXTURES: Final = SPEND_FIXTURE.parent
|
||||
SPEND_FIXTURES: Final = REPO_ROOT / "litellm-rust/crates/traces-clickhouse/tests/fixtures"
|
||||
JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
SPEND_ROWS: Final = TypeAdapter(tuple[SpendLogRecord, ...])
|
||||
|
|
@ -103,7 +103,9 @@ def response_ids(rows: tuple[SpendLogRecord, ...]) -> Iterator[str]:
|
|||
|
||||
def response_pattern(rows: tuple[SpendLogRecord, ...]) -> re.Pattern[str]:
|
||||
identities: Final = sorted(
|
||||
frozenset(filter(None, chain(response_ids(rows), (row["litellm_call_id"] for row in rows)))),
|
||||
frozenset(
|
||||
identity for identity in chain(response_ids(rows), (row["litellm_call_id"] for row in rows)) if identity
|
||||
),
|
||||
key=len,
|
||||
reverse=True,
|
||||
)
|
||||
|
|
@ -118,12 +120,15 @@ def rebased_response(value: str, namespace: str, pattern: re.Pattern[str]) -> st
|
|||
return "resp_" + base64.b64encode(payload.encode()).decode()
|
||||
|
||||
|
||||
@cache
|
||||
def fixture_exports(directory: Path) -> tuple[tuple[str, JsonValue], ...]:
|
||||
return tuple((path.stem, JSON.validate_json(path.read_bytes())) for path in sorted(directory.glob("*.json")))
|
||||
|
||||
|
||||
def fixture_replays(
|
||||
directory: Path, now_ms: int, namespace: str, response_pattern: re.Pattern[str]
|
||||
) -> tuple[FixtureReplay, ...]:
|
||||
exports: Final = tuple(
|
||||
(path.stem, JSON.validate_json(path.read_bytes())) for path in sorted(directory.glob("*.json"))
|
||||
)
|
||||
exports: Final = fixture_exports(directory)
|
||||
query_latest: Final = max(
|
||||
(max(timestamps(export)) for name, export in exports if name.startswith("query_")), default=0
|
||||
)
|
||||
|
|
@ -236,15 +241,51 @@ def postgres_row(row: SpendLogRecord) -> LiteLLM_SpendLogsCreateWithoutRelations
|
|||
)
|
||||
|
||||
|
||||
async def seed() -> int:
|
||||
from prisma import Prisma
|
||||
class SeedOptions(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
profile: Literal["default", "large"]
|
||||
copies: int | None
|
||||
batch_copies: int = 4
|
||||
timeout_seconds: float = 120
|
||||
|
||||
fixtures: Final = spend_fixtures()
|
||||
spends: Final = tuple(chain.from_iterable(rows for _, rows in fixtures))
|
||||
|
||||
def seed_arguments(argv: Sequence[str] | None = None) -> SeedOptions:
|
||||
parser: Final = argparse.ArgumentParser(description="Replay tracing fixtures into a running local Lens stack")
|
||||
parser.add_argument("--profile", choices=("default", "large"), default="default")
|
||||
parser.add_argument(
|
||||
"--copies",
|
||||
type=int,
|
||||
default=os.environ.get("LENS_DEV_SEED_COPIES"),
|
||||
help="Override fixture copies (default: 1, large: 2000; env: LENS_DEV_SEED_COPIES)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch-copies",
|
||||
type=int,
|
||||
default=os.environ.get("LENS_DEV_SEED_BATCH_COPIES", "4"),
|
||||
help="Copies per bulk insert (default: 4; env: LENS_DEV_SEED_BATCH_COPIES)",
|
||||
)
|
||||
parser.add_argument("--timeout-seconds", type=float, default=os.environ.get("LENS_DEV_SEED_TIMEOUT_SECONDS", "120"))
|
||||
arguments: Final = SeedOptions.model_validate(vars(parser.parse_args(argv)))
|
||||
if arguments.copies is not None and arguments.copies < 1:
|
||||
parser.error("--copies must be positive")
|
||||
if arguments.batch_copies < 1:
|
||||
parser.error("--batch-copies must be positive")
|
||||
if not math.isfinite(arguments.timeout_seconds) or arguments.timeout_seconds <= 0:
|
||||
parser.error("--timeout-seconds must be finite and positive")
|
||||
return arguments
|
||||
|
||||
|
||||
async def seed_batch(
|
||||
client: httpx.AsyncClient,
|
||||
storage: ClickHouseStorage,
|
||||
database: Prisma,
|
||||
replays: tuple[FixtureReplay, ...],
|
||||
fixtures: tuple[tuple[str, tuple[SpendLogRecord, ...]], ...],
|
||||
pattern: re.Pattern[str],
|
||||
tenant: TenantIdentity | None,
|
||||
verify: bool,
|
||||
) -> TenantIdentity:
|
||||
by_name: Final = MappingProxyType(dict(fixtures))
|
||||
namespace: Final = uuid4().hex
|
||||
pattern: Final = response_pattern(spends)
|
||||
replays: Final = fixture_replays(TRACE_FIXTURES, time.time_ns() // 1_000_000, namespace, pattern)
|
||||
paired: Final = tuple(
|
||||
(
|
||||
replay.name,
|
||||
|
|
@ -254,45 +295,108 @@ async def seed() -> int:
|
|||
if replay.name in by_name
|
||||
)
|
||||
rebased_spends: Final = tuple(chain.from_iterable(rows for _, rows in paired))
|
||||
master_key: Final = os.environ["LITELLM_MASTER_KEY"]
|
||||
proxy_url: Final = os.environ.get("PROXY_BASE_URL", "http://127.0.0.1:4002")
|
||||
async with httpx.AsyncClient(
|
||||
base_url=proxy_url, headers={"Authorization": f"Bearer {master_key}"}, timeout=60
|
||||
) as client:
|
||||
for replay in replays:
|
||||
(
|
||||
await client.post(
|
||||
"/v1/traces", content=json.dumps(replay.export), headers={"Content-Type": "application/json"}
|
||||
)
|
||||
).raise_for_status()
|
||||
storage: Final = ClickHouseStorage(trace_storage_config({}))
|
||||
trace_id: Final = next(row["trace_id"] for row in rebased_spends if row["trace_id"])
|
||||
identity: Final = await storage.query_sql(
|
||||
"SELECT DISTINCT TeamId AS team_id, ApiKeyHash AS api_key, UserId AS user "
|
||||
f"FROM otel_traces WHERE TraceId = '{trace_id}'",
|
||||
AllQueryScope(kind="all"),
|
||||
master_key,
|
||||
)
|
||||
tenant: Final = TenantIdentity.model_validate(identity.data[0])
|
||||
stamped_spends: Final[tuple[SpendLogRecord, ...]] = tuple(
|
||||
{**row, "team_id": tenant.team_id, "api_key": tenant.api_key, "user": tenant.user} for row in rebased_spends
|
||||
)
|
||||
await storage.insert_rows("spend_logs", stamped_spends)
|
||||
async with Prisma() as database:
|
||||
await database.litellm_spendlogs.create_many(data=[postgres_row(row) for row in stamped_spends])
|
||||
resolved_tenant: Final = await ingest_replays(
|
||||
client, storage, replays, fixture_capture(*next((name, rows[0]) for name, rows in paired)).trace_id, tenant
|
||||
)
|
||||
stamped_spends: Final[tuple[SpendLogRecord, ...]] = tuple(
|
||||
{**row, "team_id": resolved_tenant.team_id, "api_key": resolved_tenant.api_key, "user": resolved_tenant.user}
|
||||
for row in rebased_spends
|
||||
)
|
||||
await storage.insert_rows("spend_logs", stamped_spends)
|
||||
await database.litellm_spendlogs.create_many(data=[postgres_row(row) for row in stamped_spends])
|
||||
if verify:
|
||||
verified: Final = tuple(await asyncio.gather(*(verify_capture(client, name, rows) for name, rows in paired)))
|
||||
sys.stdout.write(
|
||||
json.dumps(
|
||||
{
|
||||
"trace_fixtures": tuple(replay.name for replay in replays),
|
||||
"spend_rows": len(stamped_spends),
|
||||
"captures": verified,
|
||||
},
|
||||
indent=2,
|
||||
)
|
||||
+ "\n"
|
||||
sys.stdout.write(json.dumps({"spend_rows": len(stamped_spends), "captures": verified}, indent=2) + "\n")
|
||||
if not all(capture["verified"] for capture in verified):
|
||||
raise RuntimeError("Seed spend verification failed")
|
||||
return resolved_tenant
|
||||
|
||||
|
||||
async def ingest_replays(
|
||||
client: httpx.AsyncClient,
|
||||
storage: ClickHouseStorage,
|
||||
replays: tuple[FixtureReplay, ...],
|
||||
trace_id: str,
|
||||
tenant: TenantIdentity | None,
|
||||
) -> TenantIdentity:
|
||||
if tenant is not None:
|
||||
await storage.insert_rows(
|
||||
"otel_traces", bulk_span_rows(replays, Tenant(tenant.team_id, tenant.api_key, user_id=tenant.user))
|
||||
)
|
||||
return 0 if all(capture["verified"] for capture in verified) else 1
|
||||
return tenant
|
||||
for replay in replays:
|
||||
(
|
||||
await client.post(
|
||||
"/v1/traces", content=json.dumps(replay.export), headers={"Content-Type": "application/json"}
|
||||
)
|
||||
).raise_for_status()
|
||||
identity: Final = await storage.query_sql(
|
||||
"SELECT DISTINCT TeamId AS team_id, ApiKeyHash AS api_key, UserId AS user "
|
||||
f"FROM otel_traces WHERE TraceId = '{trace_id}'",
|
||||
AllQueryScope(kind="all"),
|
||||
os.environ["LITELLM_MASTER_KEY"],
|
||||
)
|
||||
return TenantIdentity.model_validate(identity.data[0])
|
||||
|
||||
|
||||
def bulk_span_rows(replays: tuple[FixtureReplay, ...], tenant: Tenant) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
chain.from_iterable(
|
||||
span_rows(json.dumps(replay.export).encode(), "application/json", tenant) for replay in replays
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def replay_batches(
|
||||
count: int, now_ms: int, namespace: str, pattern: re.Pattern[str], batch_copies: int = 4
|
||||
) -> Iterator[tuple[int, tuple[FixtureReplay, ...]]]:
|
||||
for start, stop in ((start, min(start + batch_copies, count)) for start in range(1, count, batch_copies)):
|
||||
yield (
|
||||
stop,
|
||||
tuple(
|
||||
chain.from_iterable(
|
||||
fixture_replays(TRACE_FIXTURES, now_ms - index * 1000, f"{namespace}-{index}", pattern)
|
||||
for index in range(start, stop)
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def seed(
|
||||
profile: str = "default", copies: int | None = None, batch_copies: int = 4, timeout_seconds: float = 120
|
||||
) -> int:
|
||||
from prisma import Prisma
|
||||
|
||||
fixtures: Final = spend_fixtures()
|
||||
pattern: Final = response_pattern(tuple(chain.from_iterable(rows for _, rows in fixtures)))
|
||||
count: Final = copies if copies is not None else (2000 if profile == "large" else 1)
|
||||
namespace: Final = uuid4().hex
|
||||
now_ms: Final = time.time_ns() // 1_000_000
|
||||
storage: Final = ClickHouseStorage(trace_storage_config({}))
|
||||
async with (
|
||||
httpx.AsyncClient(
|
||||
base_url=os.environ.get("PROXY_BASE_URL", "http://127.0.0.1:4002"),
|
||||
headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"},
|
||||
timeout=timeout_seconds,
|
||||
) as client,
|
||||
Prisma() as database,
|
||||
):
|
||||
tenant: Final = await seed_batch(
|
||||
client,
|
||||
storage,
|
||||
database,
|
||||
fixture_replays(TRACE_FIXTURES, now_ms, namespace + "-0", pattern),
|
||||
fixtures,
|
||||
pattern,
|
||||
None,
|
||||
True,
|
||||
)
|
||||
for stop, replays in replay_batches(count, now_ms, namespace, pattern, batch_copies):
|
||||
await seed_batch(client, storage, database, replays, fixtures, pattern, tenant, stop == count)
|
||||
sys.stdout.write(f"Seeded {stop}/{count} fixture copies\n")
|
||||
sys.stdout.flush()
|
||||
sys.stdout.write(f"Seed complete: profile={profile}, copies={count}, namespace={namespace}\n")
|
||||
return 0
|
||||
|
||||
|
||||
def fixture_capture(name: str, row: SpendLogRecord) -> FixtureCapture:
|
||||
|
|
@ -307,7 +411,7 @@ def fixture_capture(name: str, row: SpendLogRecord) -> FixtureCapture:
|
|||
|
||||
async def verify_capture(
|
||||
client: httpx.AsyncClient, name: str, rows: tuple[SpendLogRecord, ...]
|
||||
) -> dict[str, JsonValue]:
|
||||
) -> Mapping[str, JsonValue]:
|
||||
capture: Final = fixture_capture(name, rows[0])
|
||||
detail: Final = await client.get(f"/v1/traces/{capture.trace_id}")
|
||||
detail.raise_for_status()
|
||||
|
|
@ -320,9 +424,14 @@ async def verify_capture(
|
|||
"spend_rows": len(rows),
|
||||
"recorded_spend": expected,
|
||||
"trace_spend": actual,
|
||||
"verified": math.isclose(actual, expected) if actual is not None else not (capture.spend_linked and capture.spend_complete),
|
||||
"verified": math.isclose(actual, expected)
|
||||
if actual is not None
|
||||
else not (capture.spend_linked and capture.spend_complete),
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(asyncio.run(seed()))
|
||||
arguments: Final = seed_arguments()
|
||||
raise SystemExit(
|
||||
asyncio.run(seed(arguments.profile, arguments.copies, arguments.batch_copies, arguments.timeout_seconds))
|
||||
)
|
||||
|
|
|
|||
|
|
@ -177,3 +177,4 @@ hypothesis: >=6.165.10 # MPL 2.0 license
|
|||
pytest-rerunfailures: >=15.1 # MPL 2.0 license
|
||||
pytest-recording: >=0.13.4 # MIT license
|
||||
expression: >=5.6.0 # MIT License - https://github.com/cognitedata/Expression/blob/main/LICENSE
|
||||
typing-extensions: >=4.13.0 # PSF-2.0 license - https://github.com/python/typing_extensions/blob/main/LICENSE
|
||||
|
|
|
|||
|
|
@ -1,6 +1,19 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
worker_image() {
|
||||
env -u LENS_WORKER_IMAGE -u LITELLM_VERSION \
|
||||
LITELLM_URL=http://litellm:4000 LENS_WORKER_TOKEN=config-test "$@" \
|
||||
docker compose --env-file /dev/null -f deploy/lens/compose.yaml config --images
|
||||
}
|
||||
[[ "$(worker_image LENS_WORKER_IMAGE=registry.example/lens:source)" == registry.example/lens:source ]]
|
||||
[[ "$(worker_image LITELLM_VERSION=1.2.3)" == ghcr.io/berriai/litellm-lens-worker:v1.2.3 ]]
|
||||
[[ "$(worker_image LENS_WORKER_IMAGE=registry.example/lens:source LITELLM_VERSION=1.2.3)" == registry.example/lens:source ]]
|
||||
if worker_image > /dev/null 2>&1; then
|
||||
printf 'Worker Compose accepted neither an image nor a release version\n' >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
qa_dir=$(mktemp -d)
|
||||
master_key="sk-$(openssl rand -hex 32)"
|
||||
compose=(docker compose -p lens-compose-ci --env-file "$qa_dir/env" -f deploy/lens/stack.yaml)
|
||||
|
|
|
|||
|
|
@ -99,7 +99,7 @@ export const MIGRATED_E2E_PAGES: Readonly<Record<string, MigratedPage>> = {
|
|||
group: "Settings",
|
||||
content: { role: "heading", name: "UI Theme Customization" },
|
||||
},
|
||||
logs: { segment: "logs", linkName: "Logs", content: { role: "heading", name: "Request Logs" } },
|
||||
logs: { segment: "logs", linkName: "Logs", content: { role: "tab", name: "Request Logs" } },
|
||||
"admin-panel": {
|
||||
segment: "admin-panel",
|
||||
linkName: "Admin Settings",
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import threading
|
||||
|
|
|
|||
|
|
@ -216,6 +216,16 @@ JsonRpc = Mapping[str, object]
|
|||
class ScriptedTool:
|
||||
name: str
|
||||
respond: Callable[[JsonRpc], Reply | JsonRpc]
|
||||
description: str | Callable[[Mapping[str, str]], str] | None = None
|
||||
input_schema: JsonRpc = field(default_factory=lambda: {"type": "object"})
|
||||
|
||||
def listing(self, headers: Mapping[str, str]) -> JsonRpc:
|
||||
described: Final = self.description(headers) if callable(self.description) else self.description
|
||||
return {
|
||||
"name": self.name,
|
||||
"inputSchema": self.input_schema,
|
||||
**({} if described is None else {"description": described}),
|
||||
}
|
||||
|
||||
|
||||
def jsonrpc_reply(identity: object, result: JsonRpc) -> Reply:
|
||||
|
|
@ -253,9 +263,7 @@ def scripted_peer(*tools: ScriptedTool) -> Iterator[McpPeer]:
|
|||
},
|
||||
)
|
||||
if method == "tools/list":
|
||||
return jsonrpc_reply(
|
||||
identity, {"tools": [{"name": name, "inputSchema": {"type": "object"}} for name in by_name]}
|
||||
)
|
||||
return jsonrpc_reply(identity, {"tools": [tool.listing(request.headers) for tool in by_name.values()]})
|
||||
if method != "tools/call":
|
||||
return jsonrpc_error(identity, -32601, f"unsupported method {method}")
|
||||
tool: Final = by_name.get(body["params"]["name"])
|
||||
|
|
|
|||
|
|
@ -133,6 +133,7 @@ def wire_server(
|
|||
|
||||
class OwnedHTTPServer(ThreadingHTTPServer):
|
||||
daemon_threads = False
|
||||
request_queue_size = 128
|
||||
|
||||
def server_bind(self) -> None:
|
||||
super().server_bind()
|
||||
|
|
|
|||
558
tests/integration/database/test_v1_migration_error_logging.py
Normal file
558
tests/integration/database/test_v1_migration_error_logging.py
Normal file
|
|
@ -0,0 +1,558 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import socket
|
||||
import socketserver
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from contextlib import contextmanager, suppress
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal, cast
|
||||
from urllib.parse import quote, urlsplit, urlunsplit
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
from psycopg import sql
|
||||
from psycopg.types.json import Jsonb
|
||||
|
||||
from tests.integration._support.client import eventually
|
||||
from tests.integration._support.process import _free_port as free_port
|
||||
|
||||
REPO_ROOT: Final = Path(__file__).resolve().parents[3]
|
||||
MIGRATIONS_DIR: Final = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations"
|
||||
MIGRATION_NAME: Final = "20260921190000_agent_identity"
|
||||
pytestmark: Final = pytest.mark.timeout(300)
|
||||
PASSWORD: Final = "wr ong'pw9"
|
||||
FRAGMENTS: Final = ("wr ong", "wr+ong", "ong'pw9", "wr%20ong", "ong%27pw9", "pw9")
|
||||
SHIPPED_MIGRATIONS: Final = tuple(
|
||||
sorted(path.name for path in MIGRATIONS_DIR.iterdir() if path.is_dir() and path.name != "0_init")
|
||||
)
|
||||
LOG_PREFIX: Final = r"^\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2},\d{3} - [^\n]+ - "
|
||||
LOG_RECORD_START: Final = rf"{LOG_PREFIX}(?:DEBUG|INFO|WARNING|ERROR|CRITICAL) - "
|
||||
ERROR_RECORD: Final = re.compile(rf"(?ms)^({LOG_PREFIX}ERROR - .*?)(?={LOG_RECORD_START}|\Z)")
|
||||
RETRY_COUNT: Final = re.compile(r"Retrying\.\.\. \((\d+) attempts left\)")
|
||||
FRAGMENT_PATTERN: Final = re.compile(
|
||||
"|".join(re.escape(fragment) for fragment in sorted(FRAGMENTS, key=len, reverse=True))
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class MigrationResult:
|
||||
returncode: int
|
||||
output: str
|
||||
|
||||
|
||||
def _migration_log_path(test_name: str, tmp_path: Path) -> Path:
|
||||
log_directory: Final = (
|
||||
Path(os.environ["INTEGRATION_RESULTS_DIR"]) if "INTEGRATION_RESULTS_DIR" in os.environ else tmp_path
|
||||
)
|
||||
log_path: Final = log_directory / f"v1-migration-{test_name}-{uuid.uuid4().hex[:8]}.log"
|
||||
log_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
return log_path
|
||||
|
||||
|
||||
def _database_url(
|
||||
admin_url: str,
|
||||
role: str,
|
||||
password: str,
|
||||
database: str,
|
||||
*,
|
||||
encode_password: bool = True,
|
||||
) -> str:
|
||||
parsed: Final = urlsplit(admin_url)
|
||||
authority: Final = parsed.netloc.rsplit("@", 1)[-1]
|
||||
encoded_password: Final = quote(password, safe="") if encode_password else password
|
||||
netloc: Final = f"{quote(role, safe='')}:{encoded_password}@{authority}"
|
||||
return urlunsplit(parsed._replace(netloc=netloc, path=f"/{database}"))
|
||||
|
||||
|
||||
def _replace_port(database_url: str, port: int, hostname: str | None = None) -> str:
|
||||
parsed: Final = urlsplit(database_url)
|
||||
target_host: Final = hostname or parsed.hostname
|
||||
assert target_host is not None
|
||||
host: Final = f"[{target_host}]" if ":" in target_host else target_host
|
||||
userinfo: Final = parsed.netloc.rsplit("@", 1)[0]
|
||||
return urlunsplit(parsed._replace(netloc=f"{userinfo}@{host}:{port}"))
|
||||
|
||||
|
||||
def _unreachable_database_url(database_url: str) -> str:
|
||||
parsed: Final = urlsplit(database_url)
|
||||
url: Final = _database_url(
|
||||
database_url,
|
||||
parsed.username or "",
|
||||
PASSWORD,
|
||||
parsed.path.lstrip("/"),
|
||||
encode_password=False,
|
||||
)
|
||||
return _replace_port(url, free_port())
|
||||
|
||||
|
||||
@contextmanager
|
||||
def owned_database(password: str) -> Iterator[str]:
|
||||
admin_url: Final = os.environ["DATABASE_URL"]
|
||||
role: Final = f"v1_migration_{uuid.uuid4().hex}"
|
||||
database: Final = f"v1_migration_{uuid.uuid4().hex}"
|
||||
database_url: Final = _database_url(admin_url, role, password, database)
|
||||
try:
|
||||
with psycopg.connect(admin_url, autocommit=True) as admin:
|
||||
admin.execute(
|
||||
sql.SQL("CREATE ROLE {} WITH LOGIN PASSWORD {}").format(sql.Identifier(role), sql.Literal(password))
|
||||
)
|
||||
admin.execute(sql.SQL("CREATE DATABASE {} OWNER {}").format(sql.Identifier(database), sql.Identifier(role)))
|
||||
yield database_url
|
||||
finally:
|
||||
with psycopg.connect(admin_url, autocommit=True) as admin:
|
||||
admin.execute(sql.SQL("DROP DATABASE IF EXISTS {} WITH (FORCE)").format(sql.Identifier(database)))
|
||||
admin.execute(sql.SQL("DROP ROLE IF EXISTS {}").format(sql.Identifier(role)))
|
||||
|
||||
|
||||
def _migration_environment(database_url: str | None, extra_env: Mapping[str, str]) -> dict[str, str]:
|
||||
excluded_variables: Final = (
|
||||
("DIRECT_URL", "USE_V2_MIGRATION_RESOLVER")
|
||||
if database_url is not None
|
||||
else ("DATABASE_URL", "DIRECT_URL", "USE_V2_MIGRATION_RESOLVER")
|
||||
)
|
||||
inherited_environment: Final = {key: value for key, value in os.environ.items() if key not in excluded_variables}
|
||||
database_environment: Final = {"DATABASE_URL": database_url} if database_url is not None else {}
|
||||
return {
|
||||
**inherited_environment,
|
||||
**database_environment,
|
||||
"LITELLM_LOG": "ERROR",
|
||||
**extra_env,
|
||||
}
|
||||
|
||||
|
||||
def _migration_invocation(
|
||||
database_url: str | None,
|
||||
tmp_path: Path,
|
||||
extra_env: Mapping[str, str],
|
||||
resolver: Literal["legacy", "v2"] = "legacy",
|
||||
) -> tuple[tuple[str, ...], dict[str, str]]:
|
||||
config_path: Final = tmp_path / "config.yaml"
|
||||
config_path.write_text(
|
||||
"model_list:\n - model_name: integration-fake\n litellm_params:\n model: openai/integration-fake\n"
|
||||
)
|
||||
resolver_flag: Final = "--use_legacy_migration_resolver" if resolver == "legacy" else "--use_v2_migration_resolver"
|
||||
command: Final = (
|
||||
sys.executable,
|
||||
"-I",
|
||||
"-m",
|
||||
"litellm.proxy.proxy_cli",
|
||||
"--config",
|
||||
str(config_path),
|
||||
resolver_flag,
|
||||
"--skip_server_startup",
|
||||
)
|
||||
environment: Final = _migration_environment(database_url, extra_env)
|
||||
return command, environment
|
||||
|
||||
|
||||
def run_v1_migrations(
|
||||
database_url: str | None,
|
||||
tmp_path: Path,
|
||||
extra_env: Mapping[str, str],
|
||||
test_name: str,
|
||||
resolver: Literal["legacy", "v2"] = "legacy",
|
||||
) -> MigrationResult:
|
||||
command, environment = _migration_invocation(database_url, tmp_path, extra_env, resolver)
|
||||
output_path: Final = _migration_log_path(test_name, tmp_path)
|
||||
with output_path.open("w") as output_file:
|
||||
completed: Final = subprocess.run(
|
||||
command,
|
||||
cwd=REPO_ROOT,
|
||||
env=environment,
|
||||
stdout=output_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
timeout=240,
|
||||
)
|
||||
output: Final = output_path.read_text()
|
||||
return MigrationResult(completed.returncode, output)
|
||||
|
||||
|
||||
def error_lines(output: str) -> tuple[str, ...]:
|
||||
return tuple(match.group(1) for match in ERROR_RECORD.finditer(output))
|
||||
|
||||
|
||||
def _json_error_records(output: str) -> tuple[dict[str, object], ...]:
|
||||
records: Final = tuple(_parse_json_record(line, output) for line in output.splitlines() if line.startswith("{"))
|
||||
return tuple(record for record in records if record.get("level") == "ERROR")
|
||||
|
||||
|
||||
def _parse_json_record(line: str, output: str) -> dict[str, object]:
|
||||
try:
|
||||
record: Final = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
pytest.fail(_safe_output(output))
|
||||
assert isinstance(record, dict), _safe_output(output)
|
||||
return cast(dict[str, object], record)
|
||||
|
||||
|
||||
def _retry_count_texts(lines: tuple[str, ...]) -> tuple[str, ...]:
|
||||
return tuple(value for value in (_retry_count_text(line) for line in lines) if value is not None)
|
||||
|
||||
|
||||
def _retry_count_text(line: str) -> str | None:
|
||||
match: Final = RETRY_COUNT.search(line)
|
||||
return match.group(1) if match is not None else None
|
||||
|
||||
|
||||
def _safe_output(output: str) -> str:
|
||||
return FRAGMENT_PATTERN.sub("[REDACTED]", output)
|
||||
|
||||
|
||||
def _assert_no_password_fragments(output: str) -> None:
|
||||
assert [fragment for fragment in FRAGMENTS if fragment in output] == [], _safe_output(output)
|
||||
|
||||
|
||||
def _applied_migrations(database_url: str) -> tuple[str, ...]:
|
||||
with psycopg.connect(database_url) as connection:
|
||||
rows: Final = connection.execute(
|
||||
'SELECT migration_name FROM "_prisma_migrations" '
|
||||
"WHERE finished_at IS NOT NULL AND rolled_back_at IS NULL ORDER BY migration_name"
|
||||
).fetchall()
|
||||
return tuple(str(row[0]) for row in rows)
|
||||
|
||||
|
||||
def _duplicate_migrations(database_url: str) -> tuple[tuple[str, int], ...]:
|
||||
with psycopg.connect(database_url) as connection:
|
||||
rows: Final = connection.execute(
|
||||
'SELECT migration_name, COUNT(*) FROM "_prisma_migrations" '
|
||||
"GROUP BY migration_name HAVING COUNT(*) > 1 ORDER BY migration_name"
|
||||
).fetchall()
|
||||
return tuple((str(row[0]), int(row[1])) for row in rows)
|
||||
|
||||
|
||||
def _retry_p3018_migration(database_url: str) -> None:
|
||||
agent_ids: Final = (uuid.uuid4().hex, uuid.uuid4().hex)
|
||||
with psycopg.connect(database_url) as connection:
|
||||
connection.execute('DROP INDEX "LiteLLM_AgentIdentity_provider_tenant_id_client_id_key"')
|
||||
for agent_id in agent_ids:
|
||||
connection.execute(
|
||||
'INSERT INTO "LiteLLM_AgentsTable" '
|
||||
'("agent_id", "agent_name", "agent_card_params", "created_by", "updated_by") '
|
||||
"VALUES (%s, %s, %s, %s, %s)",
|
||||
(agent_id, f"audit-agent-{agent_id}", Jsonb({}), "integration", "integration"),
|
||||
)
|
||||
connection.execute(
|
||||
'INSERT INTO "LiteLLM_AgentIdentity" '
|
||||
'("agent_id", "provider", "issuer", "tenant_id", "client_id", "revision") '
|
||||
"VALUES (%s, %s, %s, %s, %s, %s)",
|
||||
(agent_id, "entra", f"https://audit.invalid/{agent_id}", "tenant-1", "client-1", uuid.uuid4().hex),
|
||||
)
|
||||
connection.execute('DELETE FROM "_prisma_migrations" WHERE migration_name = %s', (MIGRATION_NAME,))
|
||||
|
||||
|
||||
def _relay(source: socket.socket, destination: socket.socket) -> None:
|
||||
try:
|
||||
while data := source.recv(65536):
|
||||
destination.sendall(data)
|
||||
except (BrokenPipeError, ConnectionAbortedError, ConnectionResetError):
|
||||
return
|
||||
|
||||
|
||||
class _PostgresForwardingServer(socketserver.ThreadingTCPServer):
|
||||
allow_reuse_address = True
|
||||
daemon_threads = True
|
||||
request_queue_size = 64
|
||||
target: tuple[str, int]
|
||||
|
||||
def __init__(self, port: int, target: tuple[str, int]) -> None:
|
||||
self.target = target
|
||||
super().__init__(("127.0.0.1", port), _PostgresForwardingHandler)
|
||||
|
||||
|
||||
class _PostgresForwardingHandler(socketserver.BaseRequestHandler):
|
||||
def handle(self) -> None:
|
||||
server: Final = cast(_PostgresForwardingServer, self.server)
|
||||
with socket.create_connection(server.target, timeout=10) as upstream:
|
||||
reply: Final = threading.Thread(target=_relay, args=(upstream, self.request), daemon=True)
|
||||
reply.start()
|
||||
try:
|
||||
_relay(self.request, upstream)
|
||||
finally:
|
||||
with suppress(OSError):
|
||||
self.request.shutdown(socket.SHUT_WR)
|
||||
with suppress(OSError):
|
||||
upstream.shutdown(socket.SHUT_WR)
|
||||
reply.join(timeout=10)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _gated_postgres_forwarder(port: int, target: tuple[str, int]) -> Iterator[Callable[[], None]]:
|
||||
server: Final = _PostgresForwardingServer(port, target)
|
||||
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
try:
|
||||
yield thread.start
|
||||
finally:
|
||||
if thread.is_alive():
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
if thread.ident is not None:
|
||||
thread.join(timeout=10)
|
||||
|
||||
|
||||
def _stop_process(process: subprocess.Popen[str]) -> None:
|
||||
if process.poll() is None:
|
||||
with suppress(ProcessLookupError):
|
||||
os.killpg(process.pid, signal.SIGTERM)
|
||||
try:
|
||||
process.wait(timeout=30)
|
||||
except subprocess.TimeoutExpired:
|
||||
with suppress(ProcessLookupError):
|
||||
os.killpg(process.pid, signal.SIGKILL)
|
||||
process.wait(timeout=30)
|
||||
|
||||
|
||||
def _migration_log_has_p1001_or_process_exited(output_path: Path, process: subprocess.Popen[str]) -> tuple[bool, bool]:
|
||||
p1001_logged: Final = any("P1001" in line for line in error_lines(output_path.read_text()))
|
||||
process_exited: Final = process.poll() is not None
|
||||
return p1001_logged, process_exited
|
||||
|
||||
|
||||
def test_unreachable_database_emits_four_p1001_errors_without_password_fragments(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
unreachable_url: Final = _unreachable_database_url(database_url)
|
||||
completed: Final = run_v1_migrations(unreachable_url, tmp_path, {}, request.node.name)
|
||||
output: Final = completed.output
|
||||
errors: Final = error_lines(output)
|
||||
p1001_errors: Final = tuple(line for line in errors if "P1001" in line)
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 4, _safe_output(output)
|
||||
assert len(p1001_errors) == 4, _safe_output(output)
|
||||
assert _retry_count_texts(errors) == (), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_wrong_password_emits_four_p1000_errors_without_password_fragments(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(f"correct-{uuid.uuid4().hex}") as database_url:
|
||||
parsed: Final = urlsplit(database_url)
|
||||
wrong_url: Final = _database_url(
|
||||
database_url,
|
||||
parsed.username or "",
|
||||
PASSWORD,
|
||||
parsed.path.lstrip("/"),
|
||||
)
|
||||
completed: Final = run_v1_migrations(wrong_url, tmp_path, {}, request.node.name)
|
||||
output: Final = completed.output
|
||||
errors: Final = error_lines(output)
|
||||
p1000_errors: Final = tuple(line for line in errors if "P1000" in line)
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 4, _safe_output(output)
|
||||
assert len(p1000_errors) == 4, _safe_output(output)
|
||||
assert _retry_count_texts(errors) == (), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_duplicate_agent_identity_logs_the_p3018_migration_error(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
setup: Final = run_v1_migrations(database_url, tmp_path, {}, request.node.name)
|
||||
setup_output: Final = setup.output
|
||||
assert setup.returncode == 0, _safe_output(setup_output)
|
||||
_retry_p3018_migration(database_url)
|
||||
completed: Final = run_v1_migrations(database_url, tmp_path, {}, request.node.name)
|
||||
output: Final = completed.output
|
||||
errors: Final = error_lines(output)
|
||||
p3018_errors: Final = tuple(line for line in errors if "P3018" in line)
|
||||
expected_markers: Final = ((True, True), (True, True))
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 2, _safe_output(output)
|
||||
assert tuple((MIGRATION_NAME in line, "P3018" in line) for line in p3018_errors) == expected_markers, (
|
||||
_safe_output(output)
|
||||
)
|
||||
assert _retry_count_texts(errors) == ("3", "1"), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_clean_database_applies_exactly_the_shipped_migrations(tmp_path: Path, request: pytest.FixtureRequest) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
completed: Final = run_v1_migrations(database_url, tmp_path, {}, request.node.name)
|
||||
output: Final = completed.output
|
||||
assert completed.returncode == 0, _safe_output(output)
|
||||
assert error_lines(output) == (), _safe_output(output)
|
||||
assert _applied_migrations(database_url) == SHIPPED_MIGRATIONS, _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_unreachable_database_recovers_after_postgres_forwarder_starts(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
admin_url: Final = os.environ["DATABASE_URL"]
|
||||
admin: Final = urlsplit(admin_url)
|
||||
hostname: Final = admin.hostname
|
||||
port: Final = admin.port
|
||||
assert hostname is not None and port is not None
|
||||
target_host: Final = "127.0.0.1" if hostname == "localhost" else hostname
|
||||
forwarding_port: Final = free_port()
|
||||
forwarded_url: Final = _replace_port(database_url, forwarding_port, "127.0.0.1")
|
||||
command, environment = _migration_invocation(forwarded_url, tmp_path, {})
|
||||
output_path: Final = _migration_log_path(request.node.name, tmp_path)
|
||||
with _gated_postgres_forwarder(forwarding_port, (target_host, port)) as open_forwarder:
|
||||
with output_path.open("w") as output_file:
|
||||
process: Final = subprocess.Popen(
|
||||
command,
|
||||
cwd=REPO_ROOT,
|
||||
env=environment,
|
||||
stdout=output_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
start_new_session=True,
|
||||
)
|
||||
try:
|
||||
observation: Final = eventually(
|
||||
lambda: _migration_log_has_p1001_or_process_exited(output_path, process),
|
||||
lambda state: state[0] or state[1],
|
||||
seconds=240,
|
||||
)
|
||||
assert observation[0], _safe_output(output_path.read_text())
|
||||
open_forwarder()
|
||||
completed_returncode: Final = process.wait(timeout=240)
|
||||
finally:
|
||||
_stop_process(process)
|
||||
output: Final = output_path.read_text()
|
||||
errors: Final = error_lines(output)
|
||||
p1001_errors: Final = tuple(line for line in errors if "P1001" in line)
|
||||
assert completed_returncode == 0, _safe_output(output)
|
||||
assert len(p1001_errors) == 1, _safe_output(output)
|
||||
assert _retry_count_texts(errors) == (), _safe_output(output)
|
||||
assert _duplicate_migrations(database_url) == (), _safe_output(output)
|
||||
assert _applied_migrations(database_url) == SHIPPED_MIGRATIONS, _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_unreachable_database_keeps_password_masked_when_shape_redaction_is_disabled(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
unreachable_url: Final = _unreachable_database_url(database_url)
|
||||
completed: Final = run_v1_migrations(
|
||||
unreachable_url,
|
||||
tmp_path,
|
||||
{"LITELLM_DISABLE_REDACT_SECRETS": "true"},
|
||||
request.node.name,
|
||||
)
|
||||
output: Final = completed.output
|
||||
errors: Final = error_lines(output)
|
||||
p1001_errors: Final = tuple(line for line in errors if "P1001" in line)
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 4, _safe_output(output)
|
||||
assert len(p1001_errors) == 4, _safe_output(output)
|
||||
assert _retry_count_texts(errors) == (), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_component_database_env_vars_with_wrong_password_emit_four_p1000_errors_without_password_fragments(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
correct_password: Final = f"correct-{uuid.uuid4().hex}"
|
||||
with owned_database(correct_password) as database_url:
|
||||
parsed: Final = urlsplit(database_url)
|
||||
host: Final = parsed.hostname
|
||||
port: Final = parsed.port
|
||||
username: Final = parsed.username
|
||||
assert host is not None and port is not None and username is not None
|
||||
extra_env: Final = {
|
||||
"DATABASE_HOST": f"{host}:{port}",
|
||||
"DATABASE_USERNAME": username,
|
||||
"DATABASE_PASSWORD": PASSWORD,
|
||||
"DATABASE_NAME": parsed.path.lstrip("/"),
|
||||
}
|
||||
completed: Final = run_v1_migrations(None, tmp_path, extra_env, request.node.name)
|
||||
output: Final = completed.output
|
||||
errors: Final = error_lines(output)
|
||||
p1000_errors: Final = tuple(line for line in errors if "P1000" in line)
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 4, _safe_output(output)
|
||||
assert len(p1000_errors) == 4, _safe_output(output)
|
||||
assert _retry_count_texts(errors) == (), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_json_logs_emit_four_valid_json_p1001_error_records_without_password_fragments(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
unreachable_url: Final = _unreachable_database_url(database_url)
|
||||
completed: Final = run_v1_migrations(unreachable_url, tmp_path, {"JSON_LOGS": "true"}, request.node.name)
|
||||
output: Final = completed.output
|
||||
errors: Final = _json_error_records(output)
|
||||
messages: Final = tuple(record.get("message") for record in errors)
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 4, _safe_output(output)
|
||||
assert tuple(isinstance(message, str) and "P1001" in message for message in messages) == (
|
||||
True,
|
||||
True,
|
||||
True,
|
||||
True,
|
||||
), _safe_output(output)
|
||||
assert error_lines(output) == (), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_migration_job_entrypoint_emits_four_p1001_errors_without_password_fragments(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
unreachable_url: Final = _unreachable_database_url(database_url)
|
||||
command: Final = (sys.executable, "-I", "-m", "litellm.proxy.prisma_migration")
|
||||
environment: Final = _migration_environment(
|
||||
unreachable_url,
|
||||
{"USE_V2_MIGRATION_RESOLVER": "false"},
|
||||
)
|
||||
output_path: Final = _migration_log_path(request.node.name, tmp_path)
|
||||
with output_path.open("w") as output_file:
|
||||
completed: Final = subprocess.run(
|
||||
command,
|
||||
cwd=REPO_ROOT,
|
||||
env=environment,
|
||||
stdout=output_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
timeout=240,
|
||||
)
|
||||
output: Final = output_path.read_text()
|
||||
errors: Final = error_lines(output)
|
||||
p1001_errors: Final = tuple(line for line in errors if "P1001" in line)
|
||||
assert completed.returncode == 1, _safe_output(output)
|
||||
assert len(errors) == 4, _safe_output(output)
|
||||
assert len(p1001_errors) == 4, _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_v2_resolver_unreachable_database_exits_2_and_names_p1001(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
unreachable_url: Final = _unreachable_database_url(database_url)
|
||||
completed: Final = run_v1_migrations(unreachable_url, tmp_path, {}, request.node.name, resolver="v2")
|
||||
output: Final = completed.output
|
||||
assert completed.returncode == 2, _safe_output(output)
|
||||
assert "P1001" in output, _safe_output(output)
|
||||
assert error_lines(output) == (), _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
||||
|
||||
def test_v2_resolver_clean_database_applies_exactly_the_shipped_migrations(
|
||||
tmp_path: Path, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
with owned_database(PASSWORD) as database_url:
|
||||
completed: Final = run_v1_migrations(database_url, tmp_path, {}, request.node.name, resolver="v2")
|
||||
output: Final = completed.output
|
||||
assert completed.returncode == 0, _safe_output(output)
|
||||
assert error_lines(output) == (), _safe_output(output)
|
||||
assert _applied_migrations(database_url) == SHIPPED_MIGRATIONS, _safe_output(output)
|
||||
_assert_no_password_fragments(output)
|
||||
|
|
@ -1,20 +1,29 @@
|
|||
import re
|
||||
import textwrap
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.client import Gateway, gateway_from_environment
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
ENTRY_POINTS,
|
||||
EntryPoint,
|
||||
McpCaller,
|
||||
McpPeer,
|
||||
Outcome,
|
||||
PeerKind,
|
||||
ScriptedTool,
|
||||
peer_of,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.mcp_grants import SUBJECTS, Subject, grant
|
||||
from integration._support.process import owned_proxy
|
||||
|
||||
CALLABLE: Final = {"add": {"a": 1, "b": 2}, "multiply": {"a": 2, "b": 3}}
|
||||
RESULTS: Final = {"add": "3", "multiply": "6"}
|
||||
|
|
@ -135,3 +144,87 @@ def test_same_tool_name_on_two_servers_routes_by_prefix(gateway: Gateway) -> Non
|
|||
assert outcome.ok and outcome.text == "10", outcome.raw
|
||||
assert tool_calls(first.drain()) == ()
|
||||
assert [call["body"]["params"]["name"] for call in tool_calls(second.drain())] == ["add"]
|
||||
|
||||
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_ECHO: Final = "catalog-echo"
|
||||
_UNLISTED: Final = ""
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
|
||||
" return allow()\n"
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
|
||||
)
|
||||
|
||||
|
||||
_ECHO_GUARDRAIL_YAML: Final = (
|
||||
"guardrails:\n"
|
||||
" - guardrail_name: catalog-echo\n"
|
||||
" litellm_params:\n"
|
||||
" guardrail: custom_code\n"
|
||||
" mode: pre_mcp_call\n"
|
||||
" default_on: true\n"
|
||||
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("catalog-echo")
|
||||
path: Final = directory / "catalog_echo.yaml"
|
||||
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
|
||||
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
def _echoed_description(outcome: Outcome) -> str:
|
||||
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
|
||||
assert found is not None, outcome.raw
|
||||
return found.group(1)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("subject", SUBJECTS)
|
||||
def test_each_subjects_call_is_evaluated_only_against_the_catalog_its_own_listing_served(
|
||||
echo_rig: Gateway, subject: Subject
|
||||
) -> None:
|
||||
described: Final = "Adds under grant " + uuid.uuid4().hex[:8]
|
||||
tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described)
|
||||
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
|
||||
group: Final = "grp" + uuid.uuid4().hex[:8]
|
||||
alias: Final = "cat" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, mcp_access_groups=[group])
|
||||
caller: Final = grant(
|
||||
scenario, subject, (identity,), (identity,), access_group=group, allowed_tools={identity: ("add",)}
|
||||
)
|
||||
reach: Final = McpCaller(echo_rig, caller.key, "mcp", alias, caller.headers)
|
||||
assert reach.initialize().ok
|
||||
cold: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE}))
|
||||
listed: Final = reach.list_tools()
|
||||
assert listed.ok and f"{alias}-add" in listed.tools, listed.raw
|
||||
warm: Final = _echoed_description(reach.call(f"{alias}-add", {"probe": _PROBE}))
|
||||
assert (cold, warm) == (_UNLISTED, described), (cold, warm)
|
||||
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
|
||||
|
||||
|
||||
def test_end_users_of_one_key_share_its_catalog_slot_because_the_identity_excludes_the_end_user(
|
||||
echo_rig: Gateway,
|
||||
) -> None:
|
||||
described: Final = "Adds for end users " + uuid.uuid4().hex[:8]
|
||||
tool: Final = ScriptedTool("add", lambda _: text_result("3"), description=described)
|
||||
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
|
||||
alias: Final = "eu" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
granted: Final = grant(scenario, "end_user", (identity,), (identity,))
|
||||
first: Final = McpCaller(echo_rig, granted.key, "mcp", alias, granted.headers)
|
||||
second: Final = McpCaller(
|
||||
echo_rig, granted.key, "mcp", alias, {"x-litellm-end-user-id": "integration-" + uuid.uuid4().hex[:10]}
|
||||
)
|
||||
assert second.initialize().ok
|
||||
assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == _UNLISTED
|
||||
listed: Final = first.list_tools()
|
||||
assert listed.ok and f"{alias}-add" in listed.tools, listed.raw
|
||||
assert _echoed_description(second.call(f"{alias}-add", {"probe": _PROBE})) == described, (
|
||||
"the end-user header is intentionally not part of the catalog identity: one key, one slot"
|
||||
)
|
||||
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
|
||||
|
|
|
|||
|
|
@ -1,11 +1,23 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from collections.abc import Generator, Iterator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, JsonValue, Scenario, eventually
|
||||
import yaml
|
||||
from integration._support.client import (
|
||||
JSON_OBJECT,
|
||||
Gateway,
|
||||
Scenario,
|
||||
eventually,
|
||||
gateway_from_environment,
|
||||
object_value,
|
||||
)
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
ENTRY_POINTS,
|
||||
|
|
@ -13,10 +25,16 @@ from integration._support.mcp import (
|
|||
McpCaller,
|
||||
McpPeer,
|
||||
Outcome,
|
||||
ScriptedTool,
|
||||
mcp_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
DEFAULT_COST: Final = 0.25
|
||||
ADD_COST: Final = 0.5
|
||||
|
|
@ -191,3 +209,472 @@ def test_guardrail_removal_stops_blocking_without_restart(gateway: Gateway) -> N
|
|||
lambda calls: len(calls) >= 1,
|
||||
seconds=40,
|
||||
)
|
||||
|
||||
|
||||
MASK_ME: Final = "mask-integration-secret"
|
||||
MASKED: Final = "[MASKED]"
|
||||
COUNT_MISMATCH: Final = "count-mismatch-marker"
|
||||
LOOKUP_DESCRIPTION: Final = "Look up one record"
|
||||
LOOKUP_SCHEMA: Final = {
|
||||
"type": "object",
|
||||
"properties": {"record": {"type": "string", "description": "record identifier"}},
|
||||
}
|
||||
SPEND_ROW: Final = 'SELECT status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
|
||||
_RECORDER_CODE: Final = """\
|
||||
import os
|
||||
|
||||
import httpx
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
SINK = "{sink}/native"
|
||||
|
||||
|
||||
def _record(stage, data, call_type):
|
||||
logging_obj = data.get("litellm_logging_obj")
|
||||
return {{
|
||||
"stage": stage,
|
||||
"pid": os.getpid(),
|
||||
"call_type": call_type,
|
||||
"litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id,
|
||||
"messages": data.get("messages"),
|
||||
"mcp_tool_name": data.get("mcp_tool_name"),
|
||||
"mcp_arguments": data.get("mcp_arguments"),
|
||||
"mcp_tool_description": data.get("mcp_tool_description"),
|
||||
"mcp_input_schema": data.get("mcp_input_schema"),
|
||||
}}
|
||||
|
||||
|
||||
async def _post(record):
|
||||
async with httpx.AsyncClient(timeout=5) as client:
|
||||
await client.post(SINK, json=record)
|
||||
|
||||
|
||||
class HookRecorder(CustomLogger):
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
await _post(_record("pre", data, call_type))
|
||||
|
||||
async def async_moderation_hook(self, data, user_api_key_dict, call_type):
|
||||
await _post(_record("during", data, call_type))
|
||||
|
||||
async def async_post_mcp_tool_call_hook(self, kwargs, response_obj, start_time, end_time):
|
||||
await _post(
|
||||
{{
|
||||
"stage": "post",
|
||||
"pid": os.getpid(),
|
||||
"litellm_call_id": kwargs.get("litellm_call_id"),
|
||||
"tool": kwargs.get("mcp_tool_call_metadata"),
|
||||
"content": [item.model_dump() for item in response_obj.mcp_tool_call_response],
|
||||
}}
|
||||
)
|
||||
|
||||
|
||||
class SinkGuardrail(CustomGuardrail):
|
||||
def __init__(self, api_base, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.api_base = api_base
|
||||
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
payload = {{
|
||||
"pid": os.getpid(),
|
||||
"input_type": input_type,
|
||||
"litellm_call_id": None if logging_obj is None else logging_obj.litellm_call_id,
|
||||
"texts": inputs.get("texts"),
|
||||
"tools": inputs.get("tools"),
|
||||
"structured_messages": inputs.get("structured_messages"),
|
||||
"mcp_tool_name": request_data.get("mcp_tool_name"),
|
||||
}}
|
||||
async with httpx.AsyncClient(timeout=5) as client:
|
||||
verdict = (await client.post(self.api_base, json=payload)).json()
|
||||
return {{**inputs, "texts": verdict["texts"]}}
|
||||
|
||||
|
||||
recorder = HookRecorder()
|
||||
"""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Sunk:
|
||||
target: str
|
||||
body: dict[str, JsonValue]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class HooksRig:
|
||||
gateway: Gateway
|
||||
sibling: Gateway
|
||||
sink: Wire
|
||||
guardrail: str
|
||||
|
||||
def sunk(self) -> tuple[Sunk, ...]:
|
||||
return tuple(Sunk(request.target, JSON_OBJECT.validate_json(request.body)) for request in self.sink.drain())
|
||||
|
||||
|
||||
def _guardrail_sink(request: Request) -> Reply:
|
||||
if not request.target.startswith("/guardrail"):
|
||||
return Reply()
|
||||
texts: Final = JSON_OBJECT.validate_json(request.body).get("texts")
|
||||
assert isinstance(texts, list), texts
|
||||
masked: Final = [str(text).replace(MASK_ME, MASKED) for text in texts]
|
||||
extra: Final = ["extra"] if any(COUNT_MISMATCH in text for text in masked) else []
|
||||
return Reply(body=json.dumps({"texts": [*masked, *extra]}).encode())
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def hooks_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[HooksRig]:
|
||||
directory: Final = tmp_path_factory.mktemp("guardrail-payloads")
|
||||
guardrail: Final = "sink" + uuid.uuid4().hex[:8]
|
||||
with wire_server(_guardrail_sink) as sink:
|
||||
(directory / "hook_recorder.py").write_text(_RECORDER_CODE.format(sink=sink.url))
|
||||
config: Final = JSON_OBJECT.validate_python(
|
||||
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
)
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": guardrail,
|
||||
"litellm_params": {
|
||||
"guardrail": "hook_recorder.SinkGuardrail",
|
||||
"mode": ["pre_mcp_call", "post_mcp_call"],
|
||||
"default_on": True,
|
||||
"api_base": f"{sink.url}/guardrail",
|
||||
},
|
||||
}
|
||||
]
|
||||
config["litellm_settings"] = {
|
||||
**object_value(config["litellm_settings"]),
|
||||
"callbacks": ["hook_recorder.recorder"],
|
||||
}
|
||||
path: Final = directory / "config.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with (
|
||||
gateway_from_environment() as gateway,
|
||||
owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate,
|
||||
owned_proxy(gateway, directory, {}, config=path) as sibling,
|
||||
):
|
||||
yield HooksRig(candidate, sibling, sink, guardrail)
|
||||
|
||||
|
||||
def _worker(gateway: Gateway) -> int:
|
||||
response: Final = gateway.client.get("/debug/memory/summary", headers={"x-litellm-api-key": gateway.key})
|
||||
assert response.status_code == 200, response.text
|
||||
worker: Final = JSON_OBJECT.validate_json(response.content)["worker_pid"]
|
||||
assert isinstance(worker, int), response.text
|
||||
return worker
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _pinned(gateway: Gateway) -> Generator[tuple[Gateway, int], None, None]:
|
||||
limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=120)
|
||||
with httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=limits) as client:
|
||||
pinned: Final = Gateway(client, gateway.key, gateway.upstream_url)
|
||||
yield pinned, _worker(pinned)
|
||||
|
||||
|
||||
def _lookup_tool(result: str = "found") -> ScriptedTool:
|
||||
return ScriptedTool(
|
||||
"lookup", lambda _: text_result(result), description=LOOKUP_DESCRIPTION, input_schema=LOOKUP_SCHEMA
|
||||
)
|
||||
|
||||
|
||||
def _generic(sunk: tuple[Sunk, ...], input_type: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(item.body for item in sunk if item.target == "/guardrail" and item.body["input_type"] == input_type)
|
||||
|
||||
|
||||
def _native(sunk: tuple[Sunk, ...], stage: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(item.body for item in sunk if item.target == "/native" and item.body["stage"] == stage)
|
||||
|
||||
|
||||
def _only(records: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]:
|
||||
assert len(records) == 1, records
|
||||
return records[0]
|
||||
|
||||
|
||||
def _scan(sunk: tuple[Sunk, ...], call_id: JsonValue) -> dict[str, JsonValue]:
|
||||
return _only(tuple(record for record in _generic(sunk, "request") if record["litellm_call_id"] == call_id))
|
||||
|
||||
|
||||
def _texts(content: JsonValue) -> list[JsonValue]:
|
||||
assert isinstance(content, list), content
|
||||
return [object_value(item)["text"] for item in content]
|
||||
|
||||
|
||||
def _has_lookup(listing: Outcome) -> bool:
|
||||
return any(tool.endswith("lookup") for tool in listing.tools)
|
||||
|
||||
|
||||
def _spend_row(call_id: JsonValue) -> dict[str, JsonValue]:
|
||||
assert isinstance(call_id, str), call_id
|
||||
rows: Final = eventually(lambda: read_rows(SPEND_ROW, (call_id,)), lambda found: len(found) == 1, seconds=70)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _synthetic_message(name: str, arguments: Mapping[str, str]) -> list[dict[str, str]]:
|
||||
return [{"role": "user", "content": f"Tool: {name}\nArguments: {dict(arguments)}"}]
|
||||
|
||||
|
||||
def test_generic_sink_and_native_hooks_receive_listed_metadata_on_typed_keys_with_the_message_bytes_unchanged(
|
||||
hooks_rig: HooksRig,
|
||||
) -> None:
|
||||
with (
|
||||
scripted_peer(_lookup_tool()) as peer,
|
||||
_pinned(hooks_rig.gateway) as (pinned, worker),
|
||||
pinned.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "payload" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
|
||||
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
|
||||
hooks_rig.sunk()
|
||||
listed: Final = eventually(caller.list_tools, _has_lookup, seconds=30)
|
||||
name: Final = next(tool for tool in listed.tools if tool.endswith("lookup"))
|
||||
scans: Final = _generic(hooks_rig.sunk(), "request")
|
||||
assert scans and all(scan["texts"] == [LOOKUP_DESCRIPTION, "record identifier"] for scan in scans), scans
|
||||
arguments: Final = {"record": "r-1"}
|
||||
outcome: Final = caller.call(name, arguments)
|
||||
assert outcome.text == "found", outcome.raw
|
||||
assert _worker(pinned) == worker
|
||||
sunk: Final = hooks_rig.sunk()
|
||||
pre: Final = _only(_native(sunk, "pre"))
|
||||
call_id: Final = pre["litellm_call_id"]
|
||||
generic: Final = _scan(sunk, call_id)
|
||||
assert generic["texts"] == [LOOKUP_DESCRIPTION, "record identifier", "r-1"], generic
|
||||
assert generic["tools"] == [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "lookup",
|
||||
"description": LOOKUP_DESCRIPTION,
|
||||
"parameters": {**LOOKUP_SCHEMA, "additionalProperties": False},
|
||||
"strict": False,
|
||||
},
|
||||
}
|
||||
], generic
|
||||
response: Final = _only(_generic(sunk, "response"))
|
||||
assert (response["texts"], response["pid"]) == (["found"], worker), response
|
||||
during: Final = _only(_native(sunk, "during"))
|
||||
post: Final = _only(_native(sunk, "post"))
|
||||
assert all(record["pid"] == worker for record in (pre, during, post)), sunk
|
||||
assert all(record["litellm_call_id"] == call_id for record in (pre, during, post)), sunk
|
||||
assert pre["call_type"] == "call_mcp_tool" and during["call_type"] == "call_mcp_tool", sunk
|
||||
assert pre["messages"] == _synthetic_message("lookup", arguments), pre
|
||||
assert during["messages"] == _synthetic_message("lookup", arguments), during
|
||||
assert (pre["mcp_tool_name"], pre["mcp_arguments"]) == ("lookup", arguments), pre
|
||||
assert pre["mcp_tool_description"] == LOOKUP_DESCRIPTION, pre
|
||||
assert pre["mcp_input_schema"] == LOOKUP_SCHEMA, pre
|
||||
assert (during["mcp_tool_description"], during["mcp_input_schema"]) == (None, None), during
|
||||
assert _texts(post["content"]) == ["found"], post
|
||||
row: Final = _spend_row(call_id)
|
||||
assert row["status"] == "success", row
|
||||
assert _tool_metadata(row)["name"] == "lookup", row
|
||||
metadata: Final = row["metadata"]
|
||||
assert isinstance(metadata, dict) and metadata["applied_guardrails"] == [hooks_rig.guardrail], metadata
|
||||
|
||||
|
||||
def test_pre_call_mask_reaches_the_peer_and_post_call_mask_reaches_the_caller_on_one_call_id(
|
||||
hooks_rig: HooksRig,
|
||||
) -> None:
|
||||
with (
|
||||
scripted_peer(_lookup_tool(f"found {MASK_ME}")) as peer,
|
||||
_pinned(hooks_rig.gateway) as (pinned, worker),
|
||||
pinned.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "mask" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
|
||||
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
|
||||
name: Final = next(
|
||||
tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool
|
||||
)
|
||||
peer.drain()
|
||||
hooks_rig.sunk()
|
||||
outcome: Final = caller.call(name, {"record": MASK_ME})
|
||||
assert outcome.text == f"found {MASKED}", outcome.raw
|
||||
assert _worker(pinned) == worker
|
||||
reached: Final = tool_calls(peer.drain())
|
||||
assert len(reached) == 1, reached
|
||||
params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"])
|
||||
assert params["arguments"] == {"record": MASKED}, params
|
||||
sunk: Final = hooks_rig.sunk()
|
||||
generic: Final = _scan(sunk, _only(_native(sunk, "pre"))["litellm_call_id"])
|
||||
scanned: Final = generic["texts"]
|
||||
assert isinstance(scanned, list) and scanned[-1] == MASK_ME and MASKED not in scanned, generic
|
||||
assert _only(_generic(sunk, "response"))["texts"] == [f"found {MASK_ME}"], sunk
|
||||
during: Final = _only(_native(sunk, "during"))
|
||||
assert during["mcp_arguments"] == {"record": MASKED}, during
|
||||
assert during["messages"] == _synthetic_message("lookup", {"record": MASKED}), during
|
||||
assert _only(_native(sunk, "post"))["litellm_call_id"] == generic["litellm_call_id"], sunk
|
||||
assert _spend_row(generic["litellm_call_id"])["status"] == "success"
|
||||
|
||||
|
||||
def test_call_time_description_and_schema_come_from_the_catalog_of_the_worker_that_served_the_listing(
|
||||
hooks_rig: HooksRig,
|
||||
) -> None:
|
||||
with (
|
||||
scripted_peer(_lookup_tool()) as peer,
|
||||
hooks_rig.gateway.scenario() as scenario,
|
||||
_pinned(hooks_rig.sibling) as (second, second_worker),
|
||||
):
|
||||
alias: Final = "local" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
other: Final = McpCaller(second, key, "mcp", headers={"x-mcp-servers": alias})
|
||||
assert eventually(other.initialize, lambda outcome: outcome.ok, seconds=60).ok
|
||||
with _pinned(hooks_rig.gateway) as (first, first_worker):
|
||||
assert first_worker != second_worker
|
||||
lister: Final = McpCaller(first, key, "mcp", headers={"x-mcp-servers": alias})
|
||||
assert eventually(lister.initialize, lambda outcome: outcome.ok, seconds=30).ok
|
||||
name: Final = next(
|
||||
tool for tool in eventually(lister.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool
|
||||
)
|
||||
hooks_rig.sunk()
|
||||
assert other.call(name, {"record": "r-2"}).text == "found"
|
||||
assert _worker(second) == second_worker
|
||||
elsewhere: Final = hooks_rig.sunk()
|
||||
unlisted: Final = _only(_native(elsewhere, "pre"))
|
||||
assert unlisted["pid"] == second_worker, unlisted
|
||||
assert (unlisted["mcp_tool_description"], unlisted["mcp_input_schema"]) == (None, None), unlisted
|
||||
assert _scan(elsewhere, unlisted["litellm_call_id"])["texts"] == ["r-2"], elsewhere
|
||||
assert lister.call(name, {"record": "r-3"}).text == "found"
|
||||
assert _worker(first) == first_worker
|
||||
at_lister: Final = hooks_rig.sunk()
|
||||
listed: Final = _only(_native(at_lister, "pre"))
|
||||
assert listed["pid"] == first_worker, listed
|
||||
assert (listed["mcp_tool_description"], listed["mcp_input_schema"]) == (LOOKUP_DESCRIPTION, LOOKUP_SCHEMA)
|
||||
assert _scan(at_lister, listed["litellm_call_id"])["texts"] == [
|
||||
LOOKUP_DESCRIPTION,
|
||||
"record identifier",
|
||||
"r-3",
|
||||
]
|
||||
assert _has_lookup(other.list_tools())
|
||||
hooks_rig.sunk()
|
||||
assert other.call(name, {"record": "r-4"}).text == "found"
|
||||
assert _worker(second) == second_worker
|
||||
populated: Final = _only(_native(hooks_rig.sunk(), "pre"))
|
||||
assert populated["pid"] == second_worker, populated
|
||||
assert populated["mcp_tool_description"] == LOOKUP_DESCRIPTION, populated
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ("mcp", "rest"))
|
||||
@pytest.mark.parametrize("listed", (False, True))
|
||||
@pytest.mark.parametrize("record", ("r-1", MASK_ME))
|
||||
def test_long_descriptions_do_not_refuse_small_tpm_calls_or_change_masked_message_bytes(
|
||||
hooks_rig: HooksRig, entry: EntryPoint, listed: bool, record: str
|
||||
) -> None:
|
||||
description: Final = "Gateway tool metadata. " * 300
|
||||
tool: Final = ScriptedTool(
|
||||
"lookup", lambda _: text_result("found"), description=description, input_schema=LOOKUP_SCHEMA
|
||||
)
|
||||
with (
|
||||
scripted_peer(tool) as peer,
|
||||
_pinned(hooks_rig.gateway) as (pinned, worker),
|
||||
pinned.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "quota" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]}, tpm_limit=64)
|
||||
caller: Final = McpCaller(pinned, key, entry, headers={"x-mcp-servers": alias})
|
||||
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
|
||||
if listed:
|
||||
listing: Final = eventually(lambda: caller.list_tools(identity), _has_lookup, seconds=30)
|
||||
assert _has_lookup(listing), listing.raw
|
||||
hooks_rig.sunk()
|
||||
peer.drain()
|
||||
arguments: Final = {"record": record}
|
||||
outcome: Final = caller.call(f"{alias}-lookup", arguments, identity)
|
||||
assert outcome.ok and outcome.text == "found", outcome.raw
|
||||
assert _worker(pinned) == worker
|
||||
reached: Final = tool_calls(peer.drain())
|
||||
assert len(reached) == 1, reached
|
||||
params: Final = object_value(JSON_OBJECT.validate_python(reached[0]["body"])["params"])
|
||||
masked_arguments: Final = {"record": record.replace(MASK_ME, MASKED)}
|
||||
assert params["arguments"] == masked_arguments, params
|
||||
sunk: Final = hooks_rig.sunk()
|
||||
pre: Final = _only(tuple(hook for hook in _native(sunk, "pre") if hook["mcp_tool_name"] == "lookup"))
|
||||
during: Final = _only(tuple(hook for hook in _native(sunk, "during") if hook["mcp_tool_name"] == "lookup"))
|
||||
assert pre["messages"] == _synthetic_message("lookup", arguments), pre
|
||||
assert during["messages"] == _synthetic_message("lookup", masked_arguments), during
|
||||
assert _spend_row(pre["litellm_call_id"])["status"] == "success"
|
||||
|
||||
|
||||
def _nested_schema(levels: int) -> dict[str, object]:
|
||||
if levels == 0:
|
||||
return {"type": "string", "description": "deepest leaf"}
|
||||
return {"type": "object", "properties": {"a": _nested_schema(levels - 1)}}
|
||||
|
||||
|
||||
def test_a_schema_past_the_scan_depth_is_not_published_while_one_at_the_limit_is_scanned_on_the_call(
|
||||
hooks_rig: HooksRig,
|
||||
) -> None:
|
||||
shallow: Final = ScriptedTool("shallow", lambda _: text_result("found"), input_schema=_nested_schema(49))
|
||||
deep: Final = ScriptedTool("deep", lambda _: text_result("found"), input_schema=_nested_schema(50))
|
||||
with (
|
||||
scripted_peer(shallow, deep) as peer,
|
||||
_pinned(hooks_rig.gateway) as (pinned, worker),
|
||||
pinned.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "depth" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
|
||||
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
|
||||
listing: Final = eventually(caller.list_tools, lambda outcome: len(outcome.tools) > 0, seconds=30)
|
||||
assert listing.tools == (f"{alias}-shallow",), listing.raw
|
||||
assert _worker(pinned) == worker
|
||||
hooks_rig.sunk()
|
||||
assert caller.call(f"{alias}-shallow", {"record": "r-1"}).text == "found"
|
||||
scanned: Final = _only(_generic(hooks_rig.sunk(), "request"))
|
||||
assert scanned["texts"] == ["deepest leaf", "r-1"], scanned
|
||||
unpublished: Final = caller.call(f"{alias}-deep", {"record": "r-2"})
|
||||
assert _worker(pinned) == worker
|
||||
assert unpublished.error is not None and "Tool 'deep' not found" in unpublished.raw, unpublished.raw
|
||||
reached: Final = tool_calls(peer.drain())
|
||||
assert [object_value(JSON_OBJECT.validate_python(call["body"])["params"])["name"] for call in reached] == [
|
||||
"shallow"
|
||||
]
|
||||
relisted: Final = _generic(hooks_rig.sunk(), "request")
|
||||
assert all((scan["mcp_tool_name"], scan["litellm_call_id"]) == ("shallow", None) for scan in relisted), relisted
|
||||
rows: Final = _rows(key, 3)
|
||||
assert [(row["call_type"], row["status"]) for row in rows] == [
|
||||
("list_mcp_tools", "success"),
|
||||
("call_mcp_tool", "success"),
|
||||
("call_mcp_tool", "failure"),
|
||||
], rows
|
||||
assert [_tool_metadata(row)["name"] for row in rows[1:]] == ["shallow", "deep"], rows
|
||||
failed: Final = object_value(JSON_OBJECT.validate_python(rows[2]["metadata"])["error_information"])
|
||||
assert failed["error_message"] == "404: Tool 'deep' not found", failed
|
||||
|
||||
|
||||
def test_an_adapter_returning_the_wrong_number_of_texts_fails_closed_before_the_peer(hooks_rig: HooksRig) -> None:
|
||||
with (
|
||||
scripted_peer(_lookup_tool()) as peer,
|
||||
_pinned(hooks_rig.gateway) as (pinned, worker),
|
||||
pinned.scenario() as scenario,
|
||||
):
|
||||
alias: Final = "count" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(pinned, key, "mcp", headers={"x-mcp-servers": alias})
|
||||
assert eventually(caller.initialize, lambda outcome: outcome.ok, seconds=30).ok
|
||||
name: Final = next(
|
||||
tool for tool in eventually(caller.list_tools, _has_lookup, seconds=30).tools if "lookup" in tool
|
||||
)
|
||||
assert _worker(pinned) == worker
|
||||
hooks_rig.sunk()
|
||||
blocked: Final = caller.call(name, {"record": COUNT_MISMATCH})
|
||||
assert _worker(pinned) == worker
|
||||
assert blocked.error is not None, blocked.raw
|
||||
assert (
|
||||
"guardrail returned 4 texts for 3 MCP tool strings, so the redaction cannot be mapped back" in blocked.raw
|
||||
)
|
||||
assert tool_calls(peer.drain()) == (), "the blocked call reached the peer"
|
||||
sunk: Final = hooks_rig.sunk()
|
||||
assert _only(_generic(sunk, "request"))["texts"] == [LOOKUP_DESCRIPTION, "record identifier", COUNT_MISMATCH]
|
||||
assert _native(sunk, "post") == (), sunk
|
||||
rows: Final = _rows(key, 2)
|
||||
assert [(row["call_type"], row["status"]) for row in rows] == [
|
||||
("list_mcp_tools", "success"),
|
||||
("call_mcp_tool", "failure"),
|
||||
], rows
|
||||
assert _tool_metadata(rows[1])["arguments"] == {"record": COUNT_MISMATCH}, rows[1]
|
||||
|
|
|
|||
|
|
@ -1,21 +1,32 @@
|
|||
import base64
|
||||
import re
|
||||
import textwrap
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
ENTRY_POINTS,
|
||||
EntryPoint,
|
||||
McpCaller,
|
||||
McpPeer,
|
||||
Outcome,
|
||||
ScriptedTool,
|
||||
call_tool,
|
||||
mcp_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
tool_names,
|
||||
)
|
||||
from integration._support.oauth_server import oauth_server
|
||||
from integration._support.process import owned_proxy
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
ADD: Final = {"a": 2, "b": 3}
|
||||
STATIC_MODES: Final = (
|
||||
|
|
@ -192,3 +203,288 @@ def test_byok_server_uses_the_calling_users_stored_credential_and_fails_closed_w
|
|||
assert removed.status_code in (200, 204), removed.text
|
||||
eventually(lambda: call_tool(gateway, owner_key, identity, name, ADD), lambda value: value.status_code == 401)
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def _listings(peer: McpPeer) -> tuple[dict[str, object], ...]:
|
||||
return tuple(
|
||||
item
|
||||
for item in peer.drain()
|
||||
if isinstance(item.get("body"), dict) and item["body"].get("method") == "tools/list"
|
||||
)
|
||||
|
||||
|
||||
def test_oauth2_byok_listing_sends_the_minted_token_not_the_users_stored_secret(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, oauth_server() as auth, gateway.scenario() as scenario:
|
||||
alias: Final = "cc" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2",
|
||||
oauth2_flow="client_credentials",
|
||||
is_byok=True,
|
||||
token_url=auth.issuer + "/token",
|
||||
credentials={"client_id": "cc-client", "client_secret": "cc-secret-" + uuid.uuid4().hex},
|
||||
)
|
||||
owner: Final = scenario.user()
|
||||
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
|
||||
secret: Final = "byok-" + uuid.uuid4().hex
|
||||
stored: Final = gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
json={"credential": secret},
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
gateway.client.delete,
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
peer.drain()
|
||||
auth.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
|
||||
assert [request["grant_type"] for request in auth.token_requests()] == ["client_credentials"]
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
sent: Final = _header(listings[0], b"authorization")
|
||||
assert sent is not None and auth.is_live(sent.decode().removeprefix("Bearer ")), sent
|
||||
assert secret.encode() not in sent, "stored BYOK secret replaced the minted token on tools/list"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("auth_type", "header", "shape"), STATIC_MODES[:2])
|
||||
def test_byok_rest_listing_sends_the_servers_static_credential_not_the_users_stored_secret(
|
||||
gateway: Gateway, auth_type: str, header: bytes, shape: str
|
||||
) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
static: Final = "static-" + uuid.uuid4().hex
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, auth_type=auth_type, is_byok=True, credentials={"auth_value": static}
|
||||
)
|
||||
owner: Final = scenario.user()
|
||||
owner_key: Final = scenario.key(user_id=owner, object_permission={"mcp_servers": [identity]})
|
||||
secret: Final = "byok-" + uuid.uuid4().hex
|
||||
stored: Final = gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
json={"credential": secret},
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
gateway.client.delete,
|
||||
f"/v1/mcp/server/{identity}/user-credential",
|
||||
headers={"x-litellm-api-key": owner_key},
|
||||
)
|
||||
peer.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list", params={"server_id": identity}, headers={"x-litellm-api-key": owner_key}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert "add" in {tool["name"] for tool in response.json()["tools"]}, response.text
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
assert _header(listings[0], header) == shape.format(secret=static, basic="").encode(), listings[0]["headers"]
|
||||
peer.drain()
|
||||
called: Final = call_tool(gateway, owner_key, identity, f"{alias}-add", ADD)
|
||||
assert called.status_code == 200, called.text
|
||||
assert _header(_one_call(peer), header) == shape.format(secret=secret, basic="").encode()
|
||||
|
||||
|
||||
def test_deprecated_string_x_mcp_auth_lists_a_byok_server_for_a_key_without_a_user(gateway: Gateway) -> None:
|
||||
with mcp_peer() as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token", is_byok=True)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
peer.drain()
|
||||
response: Final = gateway.client.get(
|
||||
"/mcp-rest/tools/list",
|
||||
params={"server_id": identity},
|
||||
headers={"x-litellm-api-key": key, "x-mcp-auth": "Bearer hdr"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
names: Final = {tool["name"] for tool in response.json()["tools"]}
|
||||
assert "add" in names, names
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
assert listings[0]["headers"].get(b"authorization") == b"Bearer hdr"
|
||||
|
||||
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_ECHO: Final = "catalog-echo"
|
||||
_UNLISTED: Final = ""
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
|
||||
" return allow()\n"
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
|
||||
)
|
||||
|
||||
|
||||
_ECHO_GUARDRAIL_YAML: Final = (
|
||||
"guardrails:\n"
|
||||
" - guardrail_name: catalog-echo\n"
|
||||
" litellm_params:\n"
|
||||
" guardrail: custom_code\n"
|
||||
" mode: pre_mcp_call\n"
|
||||
" default_on: true\n"
|
||||
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("catalog-echo")
|
||||
path: Final = directory / "catalog_echo.yaml"
|
||||
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
|
||||
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
def _echoed_description(outcome: Outcome) -> str:
|
||||
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
|
||||
assert found is not None, outcome.raw
|
||||
return found.group(1)
|
||||
|
||||
|
||||
_PROBE_ARGUMENTS: Final = {"probe": _PROBE}
|
||||
_HEADERS: Final = TypeAdapter(dict[str, str])
|
||||
|
||||
|
||||
def _store_byok_credential(scenario: Scenario, identity: str, key: str, secret: str) -> None:
|
||||
stored: Final = scenario.gateway.client.post(
|
||||
f"/v1/mcp/server/{identity}/user-credential", json={"credential": secret}, headers={"x-litellm-api-key": key}
|
||||
)
|
||||
assert stored.status_code in (200, 201), stored.text
|
||||
scenario.cleanups.callback(
|
||||
scenario.gateway.client.delete, f"/v1/mcp/server/{identity}/user-credential", headers={"x-litellm-api-key": key}
|
||||
)
|
||||
|
||||
|
||||
def test_rotating_the_credential_drops_the_callers_listing_until_it_lists_again(echo_rig: Gateway) -> None:
|
||||
with mcp_peer() as peer, echo_rig.scenario() as scenario:
|
||||
first: Final = "cred-" + uuid.uuid4().hex
|
||||
second: Final = "cred-" + uuid.uuid4().hex
|
||||
alias: Final = "rot" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, auth_type="bearer_token", credentials={"auth_value": first}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(echo_rig, key, "mcp", alias)
|
||||
assert caller.list_tools().ok
|
||||
assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers"
|
||||
rotated: Final = echo_rig.request(
|
||||
"PUT", "/v1/mcp/server", {"server_id": identity, "credentials": {"auth_value": second}}
|
||||
)
|
||||
assert rotated.status_code == 202, rotated.text
|
||||
eventually(
|
||||
lambda: _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)), lambda seen: seen == _UNLISTED
|
||||
)
|
||||
peer.drain()
|
||||
assert caller.list_tools().ok
|
||||
assert _echoed_description(caller.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers"
|
||||
relisted: Final = _listings(peer)
|
||||
assert len(relisted) == 1, relisted
|
||||
assert _header(relisted[0], b"authorization") == f"Bearer {second}".encode(), relisted[0]["headers"]
|
||||
|
||||
|
||||
def test_byok_callers_are_evaluated_against_their_own_listing_and_the_stored_secret_never_keys_the_slot(
|
||||
echo_rig: Gateway,
|
||||
) -> None:
|
||||
with mcp_peer() as peer, echo_rig.scenario() as scenario:
|
||||
alias: Final = "byok" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, auth_type="api_key", is_byok=True)
|
||||
owner_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]})
|
||||
stranger_key: Final = scenario.key(user_id=scenario.user(), object_permission={"mcp_servers": [identity]})
|
||||
owner_secret: Final = "byok-" + uuid.uuid4().hex
|
||||
replacement: Final = "byok-" + uuid.uuid4().hex
|
||||
_store_byok_credential(scenario, identity, owner_key, owner_secret)
|
||||
_store_byok_credential(scenario, identity, stranger_key, "byok-" + uuid.uuid4().hex)
|
||||
owner: Final = McpCaller(echo_rig, owner_key, "mcp", alias)
|
||||
stranger: Final = McpCaller(echo_rig, stranger_key, "mcp", alias)
|
||||
peer.drain()
|
||||
assert owner.list_tools().ok
|
||||
listings: Final = _listings(peer)
|
||||
assert [_header(item, b"x-api-key") for item in listings] == [owner_secret.encode()], listings
|
||||
own: Final = _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS))
|
||||
other: Final = _echoed_description(stranger.call(f"{alias}-add", _PROBE_ARGUMENTS))
|
||||
assert (own, other) == ("Add two integers", _UNLISTED), (own, other)
|
||||
_store_byok_credential(scenario, identity, owner_key, replacement)
|
||||
assert _echoed_description(owner.call(f"{alias}-add", _PROBE_ARGUMENTS)) == "Add two integers", (
|
||||
"the slot is keyed by the client-supplied header, never by the stored credential"
|
||||
)
|
||||
sent: Final = eventually(
|
||||
lambda: (owner.call(f"{alias}-add", ADD).ok, tool_calls(peer.drain())),
|
||||
lambda value: any(_header(call, b"x-api-key") == replacement.encode() for call in value[1]),
|
||||
)
|
||||
assert sent[0], sent
|
||||
|
||||
|
||||
def test_callers_with_different_server_scoped_auth_headers_are_evaluated_against_their_own_listings(
|
||||
echo_rig: Gateway,
|
||||
) -> None:
|
||||
tool: Final = ScriptedTool(
|
||||
"add",
|
||||
lambda _: text_result("3"),
|
||||
description=lambda headers: "Adds for " + headers.get("authorization", "nobody"),
|
||||
)
|
||||
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
|
||||
alias: Final = "scoped" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
acme_token: Final = "acme-" + uuid.uuid4().hex
|
||||
globex_token: Final = "globex-" + uuid.uuid4().hex
|
||||
acme: Final = McpCaller(echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {acme_token}"})
|
||||
globex: Final = McpCaller(
|
||||
echo_rig, key, "mcp", alias, {f"x-mcp-{alias}-authorization": f"Bearer {globex_token}"}
|
||||
)
|
||||
assert acme.list_tools().ok and globex.list_tools().ok
|
||||
seen: Final = (
|
||||
_echoed_description(acme.call(f"{alias}-add", _PROBE_ARGUMENTS)),
|
||||
_echoed_description(globex.call(f"{alias}-add", _PROBE_ARGUMENTS)),
|
||||
)
|
||||
assert seen == (f"Adds for Bearer {acme_token}", f"Adds for Bearer {globex_token}"), seen
|
||||
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
|
||||
|
||||
|
||||
def test_deprecated_string_x_mcp_auth_callers_on_a_user_less_key_own_separate_listings(echo_rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool(
|
||||
"add",
|
||||
lambda _: text_result("3"),
|
||||
description=lambda headers: "Adds for " + headers.get("authorization", "nobody"),
|
||||
)
|
||||
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
|
||||
alias: Final = "legacy" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, auth_type="bearer_token")
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
first_token: Final = "first-" + uuid.uuid4().hex
|
||||
second_token: Final = "second-" + uuid.uuid4().hex
|
||||
first: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {first_token}"})
|
||||
second: Final = McpCaller(echo_rig, key, "mcp", alias, {"x-mcp-auth": f"Bearer {second_token}"})
|
||||
assert _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED
|
||||
peer.drain()
|
||||
assert first.list_tools().ok
|
||||
listings: Final = _listings(peer)
|
||||
assert len(listings) == 1, listings
|
||||
listed_with: Final = _HEADERS.validate_python(listings[0]["headers"])
|
||||
assert listed_with.get("authorization") == f"Bearer {first_token}", listed_with
|
||||
warm: Final = eventually(
|
||||
lambda: _echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)),
|
||||
lambda seen: seen != _UNLISTED,
|
||||
)
|
||||
assert warm == f"Adds for Bearer {first_token}", warm
|
||||
assert _echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)) == _UNLISTED
|
||||
assert second.list_tools().ok
|
||||
seen: Final = eventually(
|
||||
lambda: (
|
||||
_echoed_description(first.call(f"{alias}-add", _PROBE_ARGUMENTS)),
|
||||
_echoed_description(second.call(f"{alias}-add", _PROBE_ARGUMENTS)),
|
||||
),
|
||||
lambda pair: _UNLISTED not in pair,
|
||||
)
|
||||
assert seen == (f"Adds for Bearer {first_token}", f"Adds for Bearer {second_token}"), seen
|
||||
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
import functools
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from contextlib import ExitStack
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -14,15 +16,27 @@ from integration._support.client import Gateway, eventually
|
|||
from integration._support.database import read_rows
|
||||
from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests
|
||||
from integration._support.mcp import (
|
||||
JsonRpc,
|
||||
McpCaller,
|
||||
Outcome,
|
||||
ScriptedTool,
|
||||
call_tool,
|
||||
mcp_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
tool_names,
|
||||
)
|
||||
from integration._support.process import owned_proxy
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
_SPEND_NONCES: Final = (
|
||||
"SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce"
|
||||
' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s'
|
||||
)
|
||||
_OBJECTS: Final = TypeAdapter(Mapping[str, object])
|
||||
_STRINGS: Final = TypeAdapter(Mapping[str, str])
|
||||
|
||||
|
||||
@pytest.mark.covers("mcp.call_tool.saved_headers.reach_actual_transport")
|
||||
|
|
@ -463,9 +477,7 @@ def _update_tool_permissions(
|
|||
assert updated.status_code == 200, updated.text
|
||||
|
||||
|
||||
def _listing_on_both(
|
||||
gateway: Gateway, peer: Gateway, key: str, expected: set[str]
|
||||
) -> None:
|
||||
def _listing_on_both(gateway: Gateway, peer: Gateway, key: str, expected: set[str]) -> None:
|
||||
for worker in (gateway, peer):
|
||||
listing: Final = eventually(
|
||||
functools.partial(_granted_view, worker, key),
|
||||
|
|
@ -525,3 +537,56 @@ def test_key_update_tool_permission_widen_narrow_and_clear_apply_on_both_workers
|
|||
_listing_on_both(gateway, peer, key, all_tools)
|
||||
nulled: Final = _multiply_outcome_on_both(gateway, peer, key, alias)
|
||||
assert [call.text for call in nulled] == ["6", "6"], [call.raw for call in nulled]
|
||||
|
||||
|
||||
def _nonce_echo(params: JsonRpc) -> JsonRpc:
|
||||
return text_result(_STRINGS.validate_python(params["arguments"])["nonce"])
|
||||
|
||||
|
||||
def _call_params(call: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"])
|
||||
|
||||
|
||||
def _listed(caller: McpCaller, name: str) -> None:
|
||||
listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45)
|
||||
assert listing.error is None, (caller.gateway.client.base_url, listing.raw)
|
||||
|
||||
|
||||
def test_tool_calls_on_both_workers_stay_base_compatible_after_each_worker_lists(
|
||||
gateway: Gateway, peer: Gateway
|
||||
) -> None:
|
||||
schema: Final = {"type": "object", "properties": {"nonce": {"type": "string"}}}
|
||||
tool: Final = ScriptedTool("echo", _nonce_echo, description="Echo the nonce back", input_schema=schema)
|
||||
with scripted_peer(tool) as upstream, gateway.scenario() as scenario:
|
||||
alias: Final = "compat" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, upstream, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
name: Final = f"{alias}-echo"
|
||||
first: Final = McpCaller(gateway, key, "mcp", alias)
|
||||
second: Final = McpCaller(peer, key, "mcp", alias)
|
||||
_listed(first, name)
|
||||
_listed(second, name)
|
||||
upstream.drain()
|
||||
nonces: Final = (uuid.uuid4().hex, uuid.uuid4().hex)
|
||||
outcomes: Final = (
|
||||
first.call(name, {"nonce": nonces[0]}),
|
||||
second.call(name, {"nonce": nonces[1]}),
|
||||
first.call(name, {"nonce": nonces[0]}),
|
||||
)
|
||||
assert [outcome.text for outcome in outcomes] == [nonces[0], nonces[1], nonces[0]], [o.raw for o in outcomes]
|
||||
params: Final = [_call_params(call) for call in tool_calls(upstream.drain())]
|
||||
assert [set(entry) - {"_meta"} for entry in params] == [{"name", "arguments"}] * 3, params
|
||||
assert [entry["name"] for entry in params] == ["echo"] * 3, params
|
||||
assert [entry["arguments"] for entry in params] == [
|
||||
{"nonce": nonces[0]},
|
||||
{"nonce": nonces[1]},
|
||||
{"nonce": nonces[0]},
|
||||
], params
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")),
|
||||
lambda found: len(found) >= 3,
|
||||
seconds=70,
|
||||
)
|
||||
assert sorted((str(row["status"]), str(row["nonce"])) for row in rows) == sorted(
|
||||
("success", nonce) for nonce in (nonces[0], nonces[1], nonces[0])
|
||||
), rows
|
||||
|
|
|
|||
531
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
531
tests/integration/mcp/test_mcp_listed_tool_metadata.py
Normal file
|
|
@ -0,0 +1,531 @@
|
|||
"""pre_mcp_call guardrails are handed the tool entry ``tools/list`` served to the caller.
|
||||
|
||||
One owned proxy carries a default-on ``custom_code`` pre_mcp_call guardrail. At listing time it masks
|
||||
``SECRET`` out of every scanned text. At call time, when an argument carries the probe marker, it
|
||||
blocks and echoes the description and parameters it was handed, which is the only way to observe from outside
|
||||
what metadata the gateway attached to the hook
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Generator, Iterator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment
|
||||
from integration._support.mcp import (
|
||||
EntryPoint,
|
||||
JsonRpc,
|
||||
McpCaller,
|
||||
ScriptedTool,
|
||||
listed_tools,
|
||||
openapi_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.process import owned_proxy
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
_ECHO: Final = "catalog-echo:"
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_CALLERS_PER_SERVER: Final = 256
|
||||
_PIN_SECONDS: Final = 120
|
||||
_COLD: Final[tuple[str, JsonRpc]] = ("", {"type": "object", "properties": {}, "additionalProperties": False})
|
||||
_PID: Final = TypeAdapter(int)
|
||||
_HEADERS: Final = TypeAdapter(dict[bytes, bytes])
|
||||
_SCHEMA: Final = TypeAdapter(dict[str, object])
|
||||
_LOOKUP_SCHEMA: Final = {
|
||||
"type": "object",
|
||||
"properties": {"probe": {"type": "string", "description": "a probe marker"}},
|
||||
"additionalProperties": False,
|
||||
}
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
' texts = list(inputs.get("texts") or [])\n'
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' if "{_PROBE}" in texts:\n'
|
||||
f' return block("{_ECHO}" + json_stringify('
|
||||
'{"description": function.get("description"), "parameters": function.get("parameters")}))\n'
|
||||
' masked = [text.replace("SECRET", "[MASKED]") for text in texts]\n'
|
||||
" if masked != texts:\n"
|
||||
" return modify(texts=masked)\n"
|
||||
" return allow()\n"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("listed-tool-metadata")
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8],
|
||||
"litellm_params": {
|
||||
"guardrail": "custom_code",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"custom_code": _GUARDRAIL_CODE,
|
||||
},
|
||||
}
|
||||
]
|
||||
config["general_settings"] = {**config["general_settings"], "proxy_config_reload_interval_seconds": 1}
|
||||
path: Final = directory / "config.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with (
|
||||
gateway_from_environment() as gateway,
|
||||
owned_proxy(gateway, directory, {"KEEPALIVE_TIMEOUT": "120"}, config=path, workers=2) as candidate,
|
||||
ExitStack() as stack,
|
||||
):
|
||||
_two_workers(stack, candidate)
|
||||
yield candidate
|
||||
|
||||
|
||||
def _strings(value: object) -> Iterator[str]:
|
||||
if isinstance(value, str):
|
||||
yield value
|
||||
return
|
||||
children: Final = value.values() if isinstance(value, Mapping) else value if isinstance(value, list) else ()
|
||||
for child in children:
|
||||
yield from _strings(child)
|
||||
|
||||
|
||||
def _decoded(raw: str) -> object:
|
||||
data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:"))
|
||||
return json.loads(data[-1] if data else raw)
|
||||
|
||||
|
||||
def _echoed(raw: str) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
"""The (description, parameters) the guardrail was handed, recovered from its block reason."""
|
||||
carrier: Final = next((text for text in _strings(_decoded(raw)) if _ECHO in text), None)
|
||||
assert carrier is not None, raw
|
||||
echoed, _ = json.JSONDecoder().raw_decode(carrier.split(_ECHO, 1)[1])
|
||||
assert isinstance(echoed, dict), carrier
|
||||
return echoed.get("description"), echoed.get("parameters")
|
||||
|
||||
|
||||
def _probe(caller: McpCaller, name: str, server_id: str) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
outcome: Final = caller.call(name, {"probe": _PROBE}, server_id=server_id)
|
||||
assert outcome.error is not None, outcome.raw
|
||||
return _echoed(outcome.raw)
|
||||
|
||||
|
||||
def _worker(gateway: Gateway) -> int:
|
||||
response: Final = gateway.request("GET", "/debug/memory/summary")
|
||||
assert response.status_code == 200, response.text
|
||||
return _PID.validate_python(response.json()["worker_pid"])
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _pinned(rig: Gateway) -> Generator[Gateway, None, None]:
|
||||
"""A single keep-alive connection, so every request on it is served by the worker that accepted it."""
|
||||
limits: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1)
|
||||
with httpx.Client(base_url=rig.client.base_url, timeout=15, trust_env=False, limits=limits) as client:
|
||||
yield Gateway(client, rig.key, rig.upstream_url)
|
||||
|
||||
|
||||
def _connection(stack: ExitStack, rig: Gateway, wanted: Callable[[int], bool]) -> tuple[Gateway, int]:
|
||||
"""A pinned connection to a worker ``wanted`` accepts; a worker still starting up accepts nothing yet."""
|
||||
|
||||
def attempt() -> tuple[Gateway, int] | None:
|
||||
with ExitStack() as candidate:
|
||||
gateway: Final = candidate.enter_context(_pinned(rig))
|
||||
pid: Final = _worker(gateway)
|
||||
if not wanted(pid):
|
||||
return None
|
||||
stack.enter_context(candidate.pop_all())
|
||||
return gateway, pid
|
||||
|
||||
found: Final = eventually(attempt, lambda pair: pair is not None, seconds=_PIN_SECONDS)
|
||||
assert found is not None
|
||||
return found
|
||||
|
||||
|
||||
def _two_workers(stack: ExitStack, rig: Gateway) -> tuple[tuple[Gateway, int], tuple[Gateway, int]]:
|
||||
"""Connects while the first worker is busy answering, so the idle worker wins the accept race."""
|
||||
first: Final = _connection(stack, rig, lambda _: True)
|
||||
stop: Final = threading.Event()
|
||||
|
||||
def keep_busy() -> None:
|
||||
while not stop.is_set():
|
||||
_worker(first[0])
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
busy: Final = pool.submit(keep_busy)
|
||||
try:
|
||||
other: Final = _connection(stack, rig, lambda pid: pid != first[1])
|
||||
finally:
|
||||
stop.set()
|
||||
busy.result()
|
||||
return first, other
|
||||
|
||||
|
||||
def _served_name(gateway: Gateway, key: str, identity: str, tool: str) -> str:
|
||||
"""The prefixed name this worker lists for ``tool`` once its registry reload carries the server."""
|
||||
listing: Final = eventually(
|
||||
lambda: McpCaller(gateway, key, "rest").list_tools(server_id=identity),
|
||||
lambda value: value.ok and any(full.endswith(tool) for full in value.tools),
|
||||
)
|
||||
return next(full for full in listing.tools if full.endswith(tool))
|
||||
|
||||
|
||||
def _settled_probe(
|
||||
gateway: Gateway, key: str, name: str, identity: str
|
||||
) -> tuple[str | None, Mapping[str, object] | None]:
|
||||
"""The hook echo for a direct call, once this worker's registry reload carries the server."""
|
||||
outcome: Final = eventually(
|
||||
lambda: McpCaller(gateway, key, "rest").call(name, {"probe": _PROBE}, server_id=identity),
|
||||
lambda value: value.error is not None and _ECHO in value.raw,
|
||||
)
|
||||
return _echoed(outcome.raw)
|
||||
|
||||
|
||||
def _forwarded_tenants(observed: tuple[dict[str, object], ...]) -> frozenset[bytes]:
|
||||
return frozenset(_HEADERS.validate_python(item["headers"]).get(b"x-tenant", b"") for item in observed) - {b""}
|
||||
|
||||
|
||||
def _lookup_tool(description: str | Callable[[Mapping[str, str]], str] = "Look up one record") -> ScriptedTool:
|
||||
return ScriptedTool("lookup", lambda _: text_result("found"), description=description, input_schema=_LOOKUP_SCHEMA)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ["rest", "mcp"])
|
||||
def test_pre_call_hook_receives_the_description_and_input_schema_the_caller_was_listed(
|
||||
rig: Gateway, entry: EntryPoint
|
||||
) -> None:
|
||||
schema: Final = {"type": "object", "properties": {"probe": {"type": "string", "description": "a probe marker"}}}
|
||||
tool: Final = ScriptedTool(
|
||||
"lookup", lambda _: text_result("found"), description="Look up one record", input_schema=schema
|
||||
)
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "meta" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
caller: Final = McpCaller(rig, key, entry, headers={"x-mcp-servers": alias})
|
||||
assert caller.initialize().ok
|
||||
listed: Final = caller.list_tools(server_id=identity)
|
||||
assert listed.ok, listed.raw
|
||||
name: Final = next(full for full in listed.tools if full.endswith("lookup"))
|
||||
description, parameters = _probe(caller, name, identity)
|
||||
assert description == "Look up one record", (description, parameters)
|
||||
assert parameters is not None and parameters.get("properties") == schema["properties"], parameters
|
||||
|
||||
|
||||
def test_each_caller_is_evaluated_against_the_catalog_its_own_forwarded_headers_produced(rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool(
|
||||
"report",
|
||||
lambda _: text_result("ok"),
|
||||
description=lambda headers: f"Report for tenant {headers.get('x-tenant', 'nobody')}",
|
||||
)
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "tenant" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
acme: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "acme"})
|
||||
globex: Final = McpCaller(rig, key, "server_mcp", alias, headers={"x-tenant": "globex"})
|
||||
acme_listing: Final = acme.list_tools()
|
||||
globex_listing: Final = globex.list_tools()
|
||||
assert acme_listing.ok and globex_listing.ok, (acme_listing.raw, globex_listing.raw)
|
||||
name: Final = next(full for full in acme_listing.tools if full.endswith("report"))
|
||||
acme_seen, _ = _probe(acme, name, identity)
|
||||
globex_seen, _ = _probe(globex, name, identity)
|
||||
assert (acme_seen, globex_seen) == ("Report for tenant acme", "Report for tenant globex"), (
|
||||
"each caller's tools/call must be evaluated against the catalog its own headers listed"
|
||||
)
|
||||
|
||||
|
||||
def test_call_is_evaluated_against_the_masked_description_the_listing_served(rig: Gateway) -> None:
|
||||
tool: Final = ScriptedTool("read_note", lambda _: text_result("note"), description="Read a note")
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "note" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"read_note": "Read a SECRET note"}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(rig, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("read_note"))
|
||||
assert served[name]["description"] == "Read a [MASKED] note", served[name]
|
||||
seen, _ = _probe(McpCaller(rig, key, "rest"), name, identity)
|
||||
assert seen == "Read a [MASKED] note", "the admin override must not restore wording the listing masked"
|
||||
|
||||
|
||||
def test_openapi_call_is_evaluated_against_the_masked_override_the_listing_served(rig: Gateway) -> None:
|
||||
with openapi_peer() as peer, rig.scenario() as scenario:
|
||||
alias: Final = "pets" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(rig, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("getpet"))
|
||||
assert served[name]["description"] == "Fetch one [MASKED] pet", served[name]
|
||||
seen, parameters = _probe(McpCaller(rig, key, "rest"), name, identity)
|
||||
assert seen == "Fetch one [MASKED] pet", "the OpenAPI call path must hand hooks the entry the listing served"
|
||||
assert parameters is not None and "petId" in parameters.get("properties", {}), parameters
|
||||
assert not [call for call in peer.drain() if call["path"].startswith("/pets")], "blocked before upstream"
|
||||
|
||||
|
||||
def test_openapi_call_is_evaluated_against_the_entry_this_key_was_listed_not_the_last_listing(
|
||||
rig: Gateway,
|
||||
) -> None:
|
||||
with openapi_peer() as peer, rig.scenario() as scenario:
|
||||
alias: Final = "pets" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario, peer, alias, tool_name_to_description={"getpet": "Fetch one SECRET pet"}
|
||||
)
|
||||
guarded: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
opted_out: Final = scenario.key(
|
||||
object_permission={"mcp_servers": [identity]}, metadata={"disable_global_guardrails": True}
|
||||
)
|
||||
guarded_served: Final = listed_tools(rig, guarded, identity)
|
||||
opted_out_served: Final = listed_tools(rig, opted_out, identity)
|
||||
name: Final = next(full for full in guarded_served if full.endswith("getpet"))
|
||||
assert (guarded_served[name]["description"], opted_out_served[name]["description"]) == (
|
||||
"Fetch one [MASKED] pet",
|
||||
"Fetch one SECRET pet",
|
||||
), (guarded_served[name], opted_out_served[name])
|
||||
seen, _ = _probe(McpCaller(rig, guarded, "rest"), name, identity)
|
||||
assert seen == "Fetch one [MASKED] pet", (
|
||||
"the guarded key must be evaluated against its own listing, not the opted-out key's later one"
|
||||
)
|
||||
|
||||
|
||||
def test_direct_call_without_a_listing_hands_the_hook_no_metadata_on_either_worker(rig: Gateway) -> None:
|
||||
with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack:
|
||||
(first, first_pid), (other, other_pid) = _two_workers(stack, rig)
|
||||
with first.scenario() as scenario:
|
||||
identity: Final = register_mcp(scenario, peer, "cold" + uuid.uuid4().hex[:8])
|
||||
observer: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
name: Final = _served_name(first, observer, identity, "lookup")
|
||||
seen: Final = tuple(_settled_probe(gateway, key, name, identity) for gateway in (first, other))
|
||||
assert seen == (_COLD, _COLD), (seen, first_pid, other_pid)
|
||||
assert (_worker(first), _worker(other)) == (first_pid, other_pid)
|
||||
assert tool_calls(peer.drain()) == (), "the probe is blocked at the hook, before the upstream"
|
||||
|
||||
|
||||
def test_warm_metadata_is_local_to_the_worker_that_served_the_listing(rig: Gateway) -> None:
|
||||
with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack:
|
||||
(first, first_pid), (other, other_pid) = _two_workers(stack, rig)
|
||||
with first.scenario() as scenario:
|
||||
identity: Final = register_mcp(scenario, peer, "local" + uuid.uuid4().hex[:8])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
name: Final = _served_name(first, key, identity, "lookup")
|
||||
warm: Final = ("Look up one record", _LOOKUP_SCHEMA)
|
||||
assert _probe(McpCaller(first, key, "rest"), name, identity) == warm, first_pid
|
||||
assert _settled_probe(other, key, name, identity) == _COLD, (
|
||||
"a worker that never served this caller a listing has no catalog for it",
|
||||
other_pid,
|
||||
)
|
||||
assert _served_name(other, key, identity, "lookup") == name
|
||||
assert _probe(McpCaller(other, key, "rest"), name, identity) == warm, other_pid
|
||||
assert (_worker(first), _worker(other)) == (first_pid, other_pid)
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_listed_tools_without_a_description_still_hand_the_hook_the_schema_the_listing_served(rig: Gateway) -> None:
|
||||
open_schema: Final[JsonRpc] = {"type": "object", "properties": {}, "additionalProperties": True}
|
||||
undescribed: Final = ScriptedTool("undescribed", lambda _: text_result("ok"), input_schema=open_schema)
|
||||
blank: Final = ScriptedTool("blank", lambda _: text_result("ok"), description="", input_schema=open_schema)
|
||||
with scripted_peer(undescribed, blank) as peer, _pinned(rig) as worker, worker.scenario() as scenario:
|
||||
pid: Final = _worker(worker)
|
||||
identity: Final = register_mcp(scenario, peer, "bare" + uuid.uuid4().hex[:8])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(worker, key, identity)
|
||||
names: Final = tuple(next(full for full in served if full.endswith(tool)) for tool in ("undescribed", "blank"))
|
||||
assert tuple((served[name].get("description") or "", served[name]["inputSchema"]) for name in names) == (
|
||||
("", open_schema),
|
||||
("", open_schema),
|
||||
), served
|
||||
seen: Final = tuple(_probe(McpCaller(worker, key, "rest"), name, identity) for name in names)
|
||||
assert seen == (("", open_schema), ("", open_schema)), (seen, _COLD)
|
||||
assert _worker(worker) == pid
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_hook_receives_the_nested_schema_with_the_leaves_the_listing_masked(rig: Gateway) -> None:
|
||||
schema: Final = {
|
||||
"type": "object",
|
||||
"required": ["filter"],
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"filter": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"path": {"type": "string", "description": "SECRET path"},
|
||||
"tags": {"type": "array", "items": {"type": "string", "description": "one SECRET tag"}},
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
masked: Final = _SCHEMA.validate_python(json.loads(json.dumps(schema).replace("SECRET", "[MASKED]")))
|
||||
tool: Final = ScriptedTool(
|
||||
"search", lambda _: text_result("hit"), description="Search SECRET records", input_schema=schema
|
||||
)
|
||||
with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario:
|
||||
pid: Final = _worker(worker)
|
||||
identity: Final = register_mcp(scenario, peer, "nested" + uuid.uuid4().hex[:8])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
served: Final = listed_tools(worker, key, identity)
|
||||
name: Final = next(full for full in served if full.endswith("search"))
|
||||
assert (served[name]["description"], served[name]["inputSchema"]) == ("Search [MASKED] records", masked), served
|
||||
assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Search [MASKED] records", masked), (
|
||||
"the hook must be handed the nested schema exactly as the listing served it"
|
||||
)
|
||||
assert _worker(worker) == pid
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_a_server_definition_update_drops_the_catalog_on_every_worker_until_the_caller_lists_again(
|
||||
rig: Gateway,
|
||||
) -> None:
|
||||
with scripted_peer(_lookup_tool()) as peer, ExitStack() as stack:
|
||||
(first, first_pid), (other, other_pid) = _two_workers(stack, rig)
|
||||
with first.scenario() as scenario:
|
||||
identity: Final = register_mcp(scenario, peer, "upd" + uuid.uuid4().hex[:8])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
name: Final = _served_name(first, key, identity, "lookup")
|
||||
assert _served_name(other, key, identity, "lookup") == name
|
||||
before: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other))
|
||||
assert before == (("Look up one record", _LOOKUP_SCHEMA),) * 2, before
|
||||
updated: Final = first.request(
|
||||
"PUT",
|
||||
"/v1/mcp/server",
|
||||
{"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}},
|
||||
)
|
||||
assert updated.status_code == 202, updated.text
|
||||
assert _probe(McpCaller(first, key, "rest"), name, identity) == _COLD, (
|
||||
"the worker that applied the update must drop its catalog at once",
|
||||
first_pid,
|
||||
)
|
||||
assert (
|
||||
eventually(lambda: _probe(McpCaller(other, key, "rest"), name, identity), lambda seen: seen == _COLD)
|
||||
== _COLD
|
||||
), other_pid
|
||||
relisted: Final = tuple(
|
||||
listed_tools(gateway, key, identity)[name]["description"] for gateway in (first, other)
|
||||
)
|
||||
assert relisted == ("Audited lookup", "Audited lookup"), relisted
|
||||
after: Final = tuple(_probe(McpCaller(gateway, key, "rest"), name, identity) for gateway in (first, other))
|
||||
assert after == (("Audited lookup", _LOOKUP_SCHEMA),) * 2, after
|
||||
assert (_worker(first), _worker(other)) == (first_pid, other_pid)
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_a_listing_in_flight_across_a_server_update_does_not_resurrect_the_old_catalog(rig: Gateway) -> None:
|
||||
started: Final = threading.Event()
|
||||
release: Final = threading.Event()
|
||||
|
||||
def describe(_: Mapping[str, str]) -> str:
|
||||
started.set()
|
||||
assert release.wait(20), "the listing was never released"
|
||||
return "Look up one record"
|
||||
|
||||
with (
|
||||
scripted_peer(_lookup_tool(describe)) as peer,
|
||||
ExitStack() as stack,
|
||||
ThreadPoolExecutor(max_workers=1) as pool,
|
||||
):
|
||||
worker, pid = _connection(stack, rig, lambda _: True)
|
||||
sibling, _ = _connection(stack, rig, lambda candidate: candidate == pid)
|
||||
with worker.scenario() as scenario:
|
||||
identity: Final = register_mcp(scenario, peer, "stale" + uuid.uuid4().hex[:8])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
release.set()
|
||||
name: Final = _served_name(worker, key, identity, "lookup")
|
||||
assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Look up one record", _LOOKUP_SCHEMA)
|
||||
started.clear()
|
||||
release.clear()
|
||||
pending: Final = pool.submit(McpCaller(worker, key, "rest").list_tools, identity)
|
||||
assert started.wait(10), "the upstream never saw the in-flight listing"
|
||||
updated: Final = sibling.request(
|
||||
"PUT",
|
||||
"/v1/mcp/server",
|
||||
{"server_id": identity, "tool_name_to_description": {"lookup": "Audited lookup"}},
|
||||
)
|
||||
assert updated.status_code == 202, updated.text
|
||||
release.set()
|
||||
stale: Final = pending.result(timeout=20)
|
||||
assert stale.ok, stale.raw
|
||||
assert _probe(McpCaller(worker, key, "rest"), name, identity) == _COLD, (
|
||||
"a listing fetched before the update must not be recorded after it"
|
||||
)
|
||||
assert listed_tools(worker, key, identity)[name]["description"] == "Audited lookup"
|
||||
assert _probe(McpCaller(worker, key, "rest"), name, identity) == ("Audited lookup", _LOOKUP_SCHEMA)
|
||||
assert (_worker(worker), _worker(sibling)) == (pid, pid)
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
def test_a_server_keeps_the_newest_256_caller_catalogs_and_evicts_the_oldest(rig: Gateway) -> None:
|
||||
tool: Final = _lookup_tool(lambda headers: f"Lookup for tenant {headers.get('x-tenant', 'nobody')}")
|
||||
with scripted_peer(tool) as peer, _pinned(rig) as worker, worker.scenario() as scenario:
|
||||
pid: Final = _worker(worker)
|
||||
alias: Final = "cap" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, extra_headers=["x-tenant"])
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
tenants: Final = tuple(f"t{index}" for index in range(_CALLERS_PER_SERVER + 1))
|
||||
callers: Final = {
|
||||
tenant: McpCaller(worker, key, "server_mcp", alias, headers={"x-tenant": tenant}) for tenant in tenants
|
||||
}
|
||||
listings: Final = tuple(callers[tenant].list_tools() for tenant in tenants)
|
||||
assert all(listing.ok for listing in listings), [listing.raw for listing in listings if not listing.ok]
|
||||
name: Final = next(full for full in listings[0].tools if full.endswith("lookup"))
|
||||
|
||||
def seen(tenant: str) -> str | None:
|
||||
return _probe(callers[tenant], name, identity)[0]
|
||||
|
||||
assert (seen("t0"), seen("t1"), seen(tenants[-1])) == (
|
||||
"",
|
||||
"Lookup for tenant t1",
|
||||
f"Lookup for tenant {tenants[-1]}",
|
||||
), "the oldest of 257 callers is evicted, the newest 256 keep their own catalog"
|
||||
assert callers["t0"].list_tools().ok
|
||||
assert (seen("t0"), seen("t1")) == ("Lookup for tenant t0", ""), "relisting makes t0 newest and evicts t1"
|
||||
assert _worker(worker) == pid
|
||||
observed: Final = peer.drain()
|
||||
assert _forwarded_tenants(observed) == frozenset(tenant.encode() for tenant in tenants), len(observed)
|
||||
assert tool_calls(observed) == ()
|
||||
|
||||
|
||||
def test_an_admin_include_disabled_tools_listing_does_not_warm_the_runtime_catalog(rig: Gateway) -> None:
|
||||
"""``include_disabled_tools=true`` is the admin-only configuration view, not a listing the caller
|
||||
runs against: recording it would warm tools/call metadata no runtime listing ever served."""
|
||||
tool: Final = _lookup_tool()
|
||||
with scripted_peer(tool) as peer, rig.scenario() as scenario:
|
||||
alias: Final = "adminview" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
admin: Final = rig.key
|
||||
caller: Final = McpCaller(rig, admin, "rest")
|
||||
name: Final = eventually(
|
||||
lambda: caller.call(f"{alias}-lookup", {"probe": _PROBE}, server_id=identity),
|
||||
lambda value: value.error is not None and _ECHO in value.raw,
|
||||
)
|
||||
assert _echoed(name.raw) == _COLD, "before any listing the call is cold"
|
||||
|
||||
view: Final = eventually(
|
||||
lambda: rig.client.get(
|
||||
"/mcp-rest/tools/list",
|
||||
params={"server_id": identity, "include_disabled_tools": "true"},
|
||||
headers={"x-litellm-api-key": admin},
|
||||
),
|
||||
lambda response: response.status_code == 200
|
||||
and any(entry["name"].endswith("lookup") for entry in response.json()["tools"]),
|
||||
)
|
||||
assert view.status_code == 200, view.text
|
||||
|
||||
after_view: Final = _probe(caller, f"{alias}-lookup", identity)
|
||||
assert after_view == _COLD, (
|
||||
"the admin-only include_disabled_tools view must not record the caller's listed-tools slot"
|
||||
)
|
||||
|
||||
runtime: Final = caller.list_tools(server_id=identity)
|
||||
assert runtime.ok, runtime.raw
|
||||
assert _probe(caller, f"{alias}-lookup", identity)[0] == "Look up one record", (
|
||||
"a genuine runtime listing still warms the slot"
|
||||
)
|
||||
|
|
@ -1,16 +1,35 @@
|
|||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from collections.abc import Callable, Generator, Iterator, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Literal
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal, TypeVar
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, Scenario
|
||||
from integration._support.mcp import McpPeer, mcp_peer, register_mcp, tool_calls
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
McpPeer,
|
||||
ScriptedTool,
|
||||
mcp_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.mcp_grants import create_toolset
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
from openai.types.chat import ChatCompletionMessageParam
|
||||
from openai.types.responses.tool_param import Mcp
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
Surface = Literal["chat", "responses", "messages", "messages_bridge"]
|
||||
SURFACES: Final[tuple[Surface, ...]] = ("chat", "responses", "messages", "messages_bridge")
|
||||
|
|
@ -18,6 +37,9 @@ ADD: Final = {"a": 2, "b": 3}
|
|||
ANSWER: Final = "the sum is 5"
|
||||
GATEWAY_REF: Final = {"type": "mcp", "server_url": "litellm_proxy", "server_label": "litellm"}
|
||||
AUTO: Final = {**GATEWAY_REF, "require_approval": "never"}
|
||||
OUTAGE: Final = "bridge-outage"
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
def _json(body: Mapping[str, object]) -> Reply:
|
||||
|
|
@ -44,25 +66,94 @@ def _has_tool_result(body: Mapping[str, object]) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _model_double(tool: str) -> Callable[[Request], Reply]:
|
||||
arguments: Final = json.dumps(ADD)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Turn:
|
||||
tool: str
|
||||
arguments: str
|
||||
answer: str
|
||||
|
||||
|
||||
def _fixed_turn(tool: str) -> Callable[[Mapping[str, JsonValue]], Turn]:
|
||||
return lambda _: Turn(tool, json.dumps(ADD), ANSWER)
|
||||
|
||||
|
||||
def _echoing_turn(body: Mapping[str, JsonValue]) -> Turn:
|
||||
names: Final = _tool_names(body)
|
||||
return Turn(names[0] if names else "", json.dumps({"query": _prompt(body)}), _tool_result_text(body) or "")
|
||||
|
||||
|
||||
def _prompt(body: Mapping[str, JsonValue]) -> str:
|
||||
inputs: Final = body.get("input")
|
||||
if isinstance(inputs, str):
|
||||
return inputs
|
||||
items: Final = inputs if isinstance(inputs, list) else body.get("messages")
|
||||
first: Final = items[0] if isinstance(items, list) and items else None
|
||||
return str(first["content"]) if isinstance(first, dict) and isinstance(first.get("content"), str) else ""
|
||||
|
||||
|
||||
def _tool_result_text(body: Mapping[str, JsonValue]) -> str | None:
|
||||
inputs: Final = body.get("input")
|
||||
messages: Final = body.get("messages")
|
||||
items: Final = inputs if isinstance(inputs, list) else messages if isinstance(messages, list) else ()
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("type") == "function_call_output":
|
||||
return str(item["output"])
|
||||
if item.get("role") == "tool":
|
||||
return str(item["content"])
|
||||
if (block := _tool_result_block(item)) is not None:
|
||||
return block
|
||||
return None
|
||||
|
||||
|
||||
def _tool_result_block(item: Mapping[str, JsonValue]) -> str | None:
|
||||
content: Final = item.get("content")
|
||||
for block in content if isinstance(content, list) else ():
|
||||
if isinstance(block, dict) and block.get("type") == "tool_result":
|
||||
return str(block["content"])
|
||||
return None
|
||||
|
||||
|
||||
def _responses_stream(response: Mapping[str, JsonValue], item: Mapping[str, JsonValue]) -> Reply:
|
||||
events: Final = (
|
||||
{"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}},
|
||||
{"type": "response.in_progress", "sequence_number": 1, "response": {**response, "status": "in_progress"}},
|
||||
{"type": "response.output_item.added", "sequence_number": 2, "output_index": 0, "item": item},
|
||||
{"type": "response.output_item.done", "sequence_number": 3, "output_index": 0, "item": item},
|
||||
{"type": "response.completed", "sequence_number": 4, "response": response},
|
||||
)
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
|
||||
)
|
||||
|
||||
|
||||
def _model_double(plan: Callable[[Mapping[str, JsonValue]], Turn]) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target.endswith("/models"):
|
||||
return _json({"object": "list", "data": []})
|
||||
body: Final = json.loads(request.body)
|
||||
assert isinstance(body, dict), request.body
|
||||
body: Final = JSON_OBJECT.validate_json(request.body)
|
||||
if OUTAGE in _prompt(body):
|
||||
outage: Final = {"error": {"message": "scripted provider outage", "type": "server_error", "code": None}}
|
||||
return Reply(status=500, body=json.dumps(outage).encode())
|
||||
turn: Final = plan(body)
|
||||
done: Final = _has_tool_result(body)
|
||||
identity: Final = uuid.uuid4().hex[:12]
|
||||
usage: Final = {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
||||
if request.target.endswith("/chat/completions"):
|
||||
message: Final = (
|
||||
{"role": "assistant", "content": ANSWER}
|
||||
{"role": "assistant", "content": turn.answer}
|
||||
if done
|
||||
else {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": tool, "arguments": arguments}}
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": turn.tool, "arguments": turn.arguments},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
|
@ -74,7 +165,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
|
|||
else message
|
||||
)
|
||||
chunk: Final = {
|
||||
"id": "chatcmpl-1",
|
||||
"id": f"chatcmpl-{identity}",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
|
|
@ -89,7 +180,7 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
|
|||
)
|
||||
return _json(
|
||||
{
|
||||
"id": "chatcmpl-1",
|
||||
"id": f"chatcmpl-{identity}",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
|
|
@ -99,13 +190,13 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
|
|||
)
|
||||
if request.target.endswith("/messages"):
|
||||
content: Final = (
|
||||
[{"type": "text", "text": ANSWER}]
|
||||
[{"type": "text", "text": turn.answer}]
|
||||
if done
|
||||
else [{"type": "tool_use", "id": "toolu_1", "name": tool, "input": ADD}]
|
||||
else [{"type": "tool_use", "id": "toolu_1", "name": turn.tool, "input": json.loads(turn.arguments)}]
|
||||
)
|
||||
return _json(
|
||||
{
|
||||
"id": "msg_1",
|
||||
"id": f"msg_{identity}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude",
|
||||
|
|
@ -116,39 +207,34 @@ def _model_double(tool: str) -> Callable[[Request], Reply]:
|
|||
}
|
||||
)
|
||||
assert request.target.endswith("/responses"), request.target
|
||||
output: Final = (
|
||||
[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_1",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": ANSWER, "annotations": []}],
|
||||
}
|
||||
]
|
||||
if done
|
||||
else [
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "fc_1",
|
||||
"call_id": "call_1",
|
||||
"name": tool,
|
||||
"arguments": arguments,
|
||||
"status": "completed",
|
||||
}
|
||||
]
|
||||
)
|
||||
return _json(
|
||||
item: Final[dict[str, JsonValue]] = (
|
||||
{
|
||||
"id": "resp_1",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"type": "message",
|
||||
"id": f"msg_{identity}",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": turn.answer, "annotations": []}],
|
||||
}
|
||||
if done
|
||||
else {
|
||||
"type": "function_call",
|
||||
"id": f"fc_{identity}",
|
||||
"call_id": f"call_{identity}",
|
||||
"name": turn.tool,
|
||||
"arguments": turn.arguments,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": output,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
)
|
||||
response: Final[dict[str, JsonValue]] = {
|
||||
"id": f"resp_{identity}",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [item],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
return _responses_stream(response, item) if body.get("stream") is True else _json(response)
|
||||
|
||||
return respond
|
||||
|
||||
|
|
@ -240,7 +326,7 @@ def _rig(gateway: Gateway, surface: Surface) -> Iterator[Rig]:
|
|||
alias: Final = "llm" + uuid.uuid4().hex[:8]
|
||||
with (
|
||||
mcp_peer() as peer,
|
||||
wire_server(_model_double(f"{alias}-add")) as wire,
|
||||
wire_server(_model_double(_fixed_turn(f"{alias}-add"))) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
server_id: Final = register_mcp(scenario, peer, alias)
|
||||
|
|
@ -403,3 +489,460 @@ def test_toolset_gateway_url_gives_a_key_of_an_ungranted_team_no_tools_and_never
|
|||
assert _peer_add_calls(rig.peer) == (), "denied caller reached the peer"
|
||||
assert all(rig.tool not in names for names in rig.upstream_tools()), rig.upstream_tools()
|
||||
assert response.status_code in (200, 400, 401, 403), response.text
|
||||
|
||||
|
||||
Bridge = Literal["chat", "responses", "messages"]
|
||||
BRIDGES: Final[tuple[Bridge, ...]] = ("chat", "responses")
|
||||
Client = Literal["sync", "async"]
|
||||
HOOK_ECHO: Final = "bridge-echo:"
|
||||
HOOK_PROBE: Final = "bridge-probe"
|
||||
HOOK_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
' texts = list(inputs.get("texts") or [])\n'
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
" for text in texts:\n"
|
||||
f' if "{HOOK_PROBE}" in text:\n'
|
||||
f' return block("{HOOK_ECHO}" + json_stringify('
|
||||
'{"description": function.get("description"), "parameters": function.get("parameters")}))\n'
|
||||
" return allow()\n"
|
||||
)
|
||||
LOOKUP: Final[tuple[str, dict[str, JsonValue]]] = (
|
||||
"Look up one record",
|
||||
{"type": "object", "properties": {"query": {"type": "string"}}},
|
||||
)
|
||||
REPORT: Final[tuple[str, dict[str, JsonValue]]] = (
|
||||
"Write one report",
|
||||
{"type": "object", "properties": {"query": {"type": "string"}, "format": {}}},
|
||||
)
|
||||
COLD: Final[tuple[str, dict[str, JsonValue]]] = (
|
||||
"",
|
||||
{"type": "object", "properties": {}, "additionalProperties": False},
|
||||
)
|
||||
RELOAD_FAST: Final = {"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": "5"}
|
||||
CALL_ID: Final = "x-litellm-call-id"
|
||||
T = TypeVar("T")
|
||||
Definition = tuple[str, str, JsonValue]
|
||||
Echo = tuple[JsonValue, JsonValue]
|
||||
|
||||
|
||||
def _served(listing: tuple[str, Mapping[str, JsonValue]]) -> Echo:
|
||||
return listing[0], {**listing[1], "additionalProperties": False}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Hooked:
|
||||
proxy: Gateway
|
||||
sink: Wire
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def hooked(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Hooked]:
|
||||
directory: Final = tmp_path_factory.mktemp("bridge-hooks")
|
||||
base: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
with wire_server(lambda _: _json({"flagged": False, "session_id": "scripted"})) as sink:
|
||||
echo: Final = {"guardrail": "custom_code", "mode": "pre_mcp_call", "default_on": True, "custom_code": HOOK_CODE}
|
||||
pillar: Final = {
|
||||
"guardrail": "pillar",
|
||||
"mode": ["pre_call", "pre_mcp_call"],
|
||||
"default_on": True,
|
||||
"api_key": "sk-pillar-" + uuid.uuid4().hex,
|
||||
"api_base": sink.url,
|
||||
"on_flagged_action": "monitor",
|
||||
}
|
||||
guardrails: Final = [
|
||||
{"guardrail_name": "bridge-echo-" + uuid.uuid4().hex[:8], "litellm_params": echo},
|
||||
{"guardrail_name": "bridge-sink-" + uuid.uuid4().hex[:8], "litellm_params": pillar},
|
||||
]
|
||||
path: Final = directory / "config.yaml"
|
||||
path.write_text(yaml.safe_dump({**base, "guardrails": guardrails}))
|
||||
with (
|
||||
gateway_from_environment() as gateway,
|
||||
owned_proxy(gateway, directory, RELOAD_FAST, config=path, workers=2) as proxy,
|
||||
):
|
||||
yield Hooked(proxy, sink)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BridgeRig:
|
||||
hooked: Hooked
|
||||
scenario: Scenario
|
||||
peer: McpPeer
|
||||
wire: Wire
|
||||
alias: str
|
||||
server_id: str
|
||||
model: str
|
||||
bridge: Bridge
|
||||
|
||||
def tool(self, name: str) -> str:
|
||||
return f"{self.alias}-{name}"
|
||||
|
||||
def names(self) -> frozenset[str]:
|
||||
return frozenset(("lookup", "report"))
|
||||
|
||||
def mcp(self, name: str) -> Mcp:
|
||||
return {**AUTO_MCP, "allowed_tools": [self.tool(name)]}
|
||||
|
||||
def url(self, path: str) -> str:
|
||||
return str(self.hooked.proxy.client.base_url).rstrip("/") + path
|
||||
|
||||
def post(self, key: str, prompt: str, tools: Sequence[Mcp], **extra: object) -> httpx.Response:
|
||||
headers: Final = {"Authorization": f"Bearer {key}"}
|
||||
if self.bridge == "chat":
|
||||
body: Final = {"model": self.model, "messages": [{"role": "user", "content": prompt}], "tools": list(tools)}
|
||||
return httpx.post(self.url("/v1/chat/completions"), headers=headers, json={**body, **extra}, timeout=90)
|
||||
if self.bridge == "responses":
|
||||
body_r: Final = {"model": self.model, "input": prompt, "tools": list(tools), **extra}
|
||||
return httpx.post(self.url("/v1/responses"), headers=headers, json=body_r, timeout=90)
|
||||
body_m: Final = {
|
||||
"model": self.model,
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"tools": list(tools),
|
||||
**extra,
|
||||
}
|
||||
return httpx.post(self.url("/v1/messages"), headers=headers, json=body_m, timeout=90)
|
||||
|
||||
def upstream_by_prompt(self) -> Mapping[str, tuple[tuple[Definition, ...], ...]]:
|
||||
bodies: Final = tuple(
|
||||
JSON_OBJECT.validate_json(request.body) for request in self.wire.drain() if request.method == "POST"
|
||||
)
|
||||
prompts: Final = frozenset(_prompt(body) for body in bodies)
|
||||
return {prompt: tuple(_definitions(body) for body in bodies if _prompt(body) == prompt) for prompt in prompts}
|
||||
|
||||
def peer_calls(self) -> tuple[tuple[str, JsonValue], ...]:
|
||||
return tuple(_peer_call(call) for call in tool_calls(self.peer.drain()))
|
||||
|
||||
def hook_messages(self, marker: str) -> tuple[str, ...]:
|
||||
posted: Final = tuple(
|
||||
JSON_OBJECT.validate_json(request.body) for request in self.hooked.sink.drain() if request.method == "POST"
|
||||
)
|
||||
contents: Final = tuple(content for payload in posted for content in _contents(payload))
|
||||
return tuple(content for content in contents if _synthetic(content, marker))
|
||||
|
||||
|
||||
AUTO_MCP: Final[Mcp] = {
|
||||
"type": "mcp",
|
||||
"server_label": "litellm",
|
||||
"server_url": "litellm_proxy",
|
||||
"require_approval": "never",
|
||||
}
|
||||
|
||||
|
||||
def _peer_call(call: Mapping[str, object]) -> tuple[str, JsonValue]:
|
||||
params: Final = object_value(object_value(JSON_VALUE.validate_python(call["body"]))["params"])
|
||||
return str(params["name"]), params["arguments"]
|
||||
|
||||
|
||||
def _contents(payload: Mapping[str, JsonValue]) -> Iterator[str]:
|
||||
messages: Final = payload.get("messages")
|
||||
for message in messages if isinstance(messages, list) else ():
|
||||
if isinstance(message, dict):
|
||||
yield str(message.get("content"))
|
||||
|
||||
|
||||
def _function(tool: JsonValue) -> dict[str, JsonValue] | None:
|
||||
if not isinstance(tool, dict):
|
||||
return None
|
||||
function: Final = tool.get("function", tool)
|
||||
return function if isinstance(function, dict) else None
|
||||
|
||||
|
||||
def _definitions(body: Mapping[str, JsonValue]) -> tuple[Definition, ...]:
|
||||
tools: Final = body.get("tools")
|
||||
functions: Final = tuple(_function(tool) for tool in tools) if isinstance(tools, list) else ()
|
||||
return tuple(
|
||||
(str(function["name"]), str(function.get("description", "")), function.get("parameters"))
|
||||
for function in functions
|
||||
if function is not None
|
||||
)
|
||||
|
||||
|
||||
def _uniform(upstream: Mapping[str, tuple[tuple[Definition, ...], ...]]) -> Mapping[str, tuple[Definition, ...] | None]:
|
||||
return {
|
||||
prompt: rounds[0] if all(definitions == rounds[0] for definitions in rounds) else None
|
||||
for prompt, rounds in upstream.items()
|
||||
}
|
||||
|
||||
|
||||
def _strings(value: JsonValue) -> Iterator[str]:
|
||||
if isinstance(value, str):
|
||||
yield value
|
||||
return
|
||||
children: Final = value.values() if isinstance(value, dict) else value if isinstance(value, list) else ()
|
||||
for child in children:
|
||||
yield from _strings(child)
|
||||
|
||||
|
||||
def _echoed(value: JsonValue) -> Echo:
|
||||
carrier: Final = next((text for text in _strings(value) if HOOK_ECHO in text), None)
|
||||
assert carrier is not None, value
|
||||
payload: Final = carrier.split(HOOK_ECHO, 1)[1]
|
||||
end: Final = json.JSONDecoder().raw_decode(payload)[1]
|
||||
echoed: Final = JSON_OBJECT.validate_json(payload[:end])
|
||||
return echoed.get("description"), echoed.get("parameters")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _bridge_rig(hooked: Hooked, bridge: Bridge) -> Generator[BridgeRig, None, None]:
|
||||
alias: Final = "brg" + uuid.uuid4().hex[:8]
|
||||
lookup: Final = ScriptedTool(
|
||||
"lookup",
|
||||
lambda params: text_result("found:" + json.dumps(JSON_OBJECT.validate_python(params)["arguments"])),
|
||||
description=LOOKUP[0],
|
||||
input_schema=LOOKUP[1],
|
||||
)
|
||||
report: Final = ScriptedTool(
|
||||
"report", lambda _: text_result("reported"), description=REPORT[0], input_schema=REPORT[1]
|
||||
)
|
||||
with (
|
||||
scripted_peer(lookup, report) as peer,
|
||||
wire_server(_model_double(_echoing_turn)) as wire,
|
||||
hooked.proxy.scenario() as scenario,
|
||||
):
|
||||
server_id: Final = register_mcp(scenario, peer, alias)
|
||||
model: Final = scenario.model(model=_upstream_model(bridge), api_base=wire.url + "/v1")
|
||||
rig: Final = BridgeRig(hooked, scenario, peer, wire, alias, server_id, model, bridge)
|
||||
eventually(
|
||||
lambda: tuple(_on_worker(hooked.proxy, lambda client: _master_listing(rig, client)) for _ in range(6)),
|
||||
lambda seen: len({pid for pid, _ in seen}) >= 2 and all(names == rig.names() for _, names in seen),
|
||||
seconds=45,
|
||||
)
|
||||
peer.drain()
|
||||
hooked.sink.drain()
|
||||
yield rig
|
||||
|
||||
|
||||
def _bridge_key(rig: BridgeRig) -> str:
|
||||
return rig.scenario.key(object_permission={"mcp_servers": [rig.server_id]})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Seen:
|
||||
response_id: str
|
||||
call_id: str
|
||||
text: str
|
||||
|
||||
|
||||
def _ask_sync(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen:
|
||||
sdk: Final = OpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90)
|
||||
if rig.bridge == "responses":
|
||||
if stream:
|
||||
raw_events: Final = sdk.responses.with_raw_response.create(
|
||||
model=rig.model, input=prompt, tools=tools, stream=True
|
||||
)
|
||||
completed: Final = next(
|
||||
event.response for event in raw_events.parse() if event.type == "response.completed"
|
||||
)
|
||||
return Seen(completed.id, raw_events.headers[CALL_ID], completed.output_text)
|
||||
raw_response: Final = sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools)
|
||||
response: Final = raw_response.parse()
|
||||
return Seen(response.id, raw_response.headers[CALL_ID], response.output_text)
|
||||
messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}]
|
||||
extra: Final = {"tools": list(tools)}
|
||||
if stream:
|
||||
raw_chunks: Final = sdk.chat.completions.with_raw_response.create(
|
||||
model=rig.model, messages=messages, stream=True, extra_body=extra
|
||||
)
|
||||
parts: Final = tuple(
|
||||
(chunk.id, chunk.choices[0].delta.content or "") for chunk in raw_chunks.parse() if chunk.choices
|
||||
)
|
||||
return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts))
|
||||
raw_completion: Final = sdk.chat.completions.with_raw_response.create(
|
||||
model=rig.model, messages=messages, extra_body=extra
|
||||
)
|
||||
completion: Final = raw_completion.parse()
|
||||
return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "")
|
||||
|
||||
|
||||
async def _ask_async(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool) -> Seen:
|
||||
sdk: Final = AsyncOpenAI(base_url=rig.url("/v1"), api_key=key, max_retries=0, timeout=90)
|
||||
if rig.bridge == "responses":
|
||||
if stream:
|
||||
raw_events: Final = await sdk.responses.with_raw_response.create(
|
||||
model=rig.model, input=prompt, tools=tools, stream=True
|
||||
)
|
||||
completed: Final = [
|
||||
event.response async for event in raw_events.parse() if event.type == "response.completed"
|
||||
]
|
||||
return Seen(completed[0].id, raw_events.headers[CALL_ID], completed[0].output_text)
|
||||
raw_response: Final = await sdk.responses.with_raw_response.create(model=rig.model, input=prompt, tools=tools)
|
||||
response: Final = raw_response.parse()
|
||||
return Seen(response.id, raw_response.headers[CALL_ID], response.output_text)
|
||||
messages: Final[list[ChatCompletionMessageParam]] = [{"role": "user", "content": prompt}]
|
||||
extra: Final = {"tools": list(tools)}
|
||||
if stream:
|
||||
raw_chunks: Final = await sdk.chat.completions.with_raw_response.create(
|
||||
model=rig.model, messages=messages, stream=True, extra_body=extra
|
||||
)
|
||||
parts: Final = [
|
||||
(chunk.id, chunk.choices[0].delta.content or "") async for chunk in raw_chunks.parse() if chunk.choices
|
||||
]
|
||||
return Seen(parts[0][0], raw_chunks.headers[CALL_ID], "".join(text for _, text in parts))
|
||||
raw_completion: Final = await sdk.chat.completions.with_raw_response.create(
|
||||
model=rig.model, messages=messages, extra_body=extra
|
||||
)
|
||||
completion: Final = raw_completion.parse()
|
||||
return Seen(completion.id, raw_completion.headers[CALL_ID], completion.choices[0].message.content or "")
|
||||
|
||||
|
||||
def _ask(rig: BridgeRig, key: str, prompt: str, tools: Sequence[Mcp], stream: bool, client: Client) -> Seen:
|
||||
if client == "async":
|
||||
return asyncio.run(_ask_async(rig, key, prompt, tools, stream))
|
||||
return _ask_sync(rig, key, prompt, tools, stream)
|
||||
|
||||
|
||||
def _synthetic(content: str, marker: str) -> bool:
|
||||
return content.startswith("Tool: lookup\n") and marker in content
|
||||
|
||||
|
||||
def _on_worker(gateway: Gateway, act: Callable[[httpx.Client], T]) -> tuple[int, T]:
|
||||
with httpx.Client(base_url=str(gateway.client.base_url), timeout=30) as client:
|
||||
summary: Final = client.get("/debug/memory/summary", headers={"Authorization": f"Bearer {gateway.key}"})
|
||||
assert summary.status_code == 200, summary.text
|
||||
pid: Final = JSON_OBJECT.validate_json(summary.content)["worker_pid"]
|
||||
assert isinstance(pid, int), summary.text
|
||||
return pid, act(client)
|
||||
|
||||
|
||||
def _both_workers(gateway: Gateway) -> frozenset[int]:
|
||||
return eventually(
|
||||
lambda: frozenset(_on_worker(gateway, lambda _: None)[0] for _ in range(6)), lambda pids: len(pids) >= 2
|
||||
)
|
||||
|
||||
|
||||
def _master_listing(rig: BridgeRig, client: httpx.Client) -> frozenset[str]:
|
||||
headers: Final = {"x-litellm-api-key": rig.hooked.proxy.key}
|
||||
response: Final = client.get("/mcp-rest/tools/list", headers=headers, params={"server_id": rig.server_id})
|
||||
tools: Final = JSON_OBJECT.validate_json(response.content).get("tools") if response.status_code == 200 else None
|
||||
return frozenset(str(object_value(tool)["name"]) for tool in tools) if isinstance(tools, list) else frozenset()
|
||||
|
||||
|
||||
def _direct_probe(rig: BridgeRig, key: str, name: str, client: httpx.Client) -> Echo:
|
||||
body: Final = {"server_id": rig.server_id, "name": rig.tool(name), "arguments": {"query": HOOK_PROBE}}
|
||||
response: Final = client.post("/mcp-rest/tools/call", headers={"x-litellm-api-key": key}, json=body)
|
||||
return _echoed(JSON_VALUE.validate_json(response.content))
|
||||
|
||||
|
||||
def _spend_row(key: str, call_id: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT api_key, call_type, status, cache_hit FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)
|
||||
),
|
||||
lambda found: len(found) >= 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert rows[0]["api_key"] == sha256(key.encode()).hexdigest(), rows
|
||||
return rows[0]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", ("sync", "async"))
|
||||
@pytest.mark.parametrize("stream", (False, True), ids=("plain", "stream"))
|
||||
@pytest.mark.parametrize("bridge", BRIDGES)
|
||||
def test_bridge_hook_sees_the_definition_of_the_tool_filtered_for_that_request(
|
||||
hooked: Hooked, bridge: Bridge, stream: bool, client: Client
|
||||
) -> None:
|
||||
with _bridge_rig(hooked, bridge) as rig:
|
||||
key: Final = _bridge_key(rig)
|
||||
marker: Final = "m" + uuid.uuid4().hex
|
||||
probe: Final = f"{marker} {HOOK_PROBE}"
|
||||
found: Final = _ask(rig, key, marker, [rig.mcp("lookup")], stream, client)
|
||||
assert found.text == "found:" + json.dumps({"query": marker}), found
|
||||
blocked: Final = _ask(rig, key, probe, [rig.mcp("lookup")], stream, client)
|
||||
assert _echoed(blocked.text) == _served(LOOKUP), blocked
|
||||
expected: Final = ((rig.tool("lookup"), *_served(LOOKUP)),)
|
||||
upstream: Final = rig.upstream_by_prompt()
|
||||
assert set(upstream) == {marker, probe} and all(
|
||||
definitions == expected for definitions in upstream[marker] + upstream[probe]
|
||||
), upstream
|
||||
assert rig.peer_calls() == (("lookup", {"query": marker}),)
|
||||
assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",)
|
||||
assert _spend_row(key, found.call_id)["status"] == "success"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bridge", BRIDGES)
|
||||
def test_direct_call_after_bridge_only_discovery_stays_cold(hooked: Hooked, bridge: Bridge) -> None:
|
||||
with _bridge_rig(hooked, bridge) as rig:
|
||||
key: Final = _bridge_key(rig)
|
||||
workers: Final = _both_workers(rig.hooked.proxy)
|
||||
bridged: Final = _ask(rig, key, HOOK_PROBE, [rig.mcp("lookup")], False, "sync")
|
||||
assert _echoed(bridged.text) == _served(LOOKUP), bridged
|
||||
direct: Final = eventually(
|
||||
lambda: tuple(
|
||||
_on_worker(rig.hooked.proxy, lambda client: _direct_probe(rig, key, "lookup", client)) for _ in range(6)
|
||||
),
|
||||
lambda seen: frozenset(pid for pid, _ in seen) == workers,
|
||||
)
|
||||
assert all(echo == COLD for _, echo in direct) and frozenset(pid for pid, _ in direct) == workers, (
|
||||
direct,
|
||||
workers,
|
||||
)
|
||||
assert rig.peer_calls() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bridge", BRIDGES)
|
||||
def test_concurrent_requests_of_one_key_with_different_allowed_tools_each_see_their_own_definition(
|
||||
hooked: Hooked, bridge: Bridge
|
||||
) -> None:
|
||||
with _bridge_rig(hooked, bridge) as rig:
|
||||
key: Final = _bridge_key(rig)
|
||||
prompts: Final = {name: f"{name} {uuid.uuid4().hex} {HOOK_PROBE}" for name in ("lookup", "report")}
|
||||
|
||||
def ask(name: str) -> Seen:
|
||||
return _ask(rig, key, prompts[name], [rig.mcp(name)], False, "sync")
|
||||
|
||||
with ThreadPoolExecutor(2) as pool:
|
||||
lookup, report = pool.map(ask, ("lookup", "report"))
|
||||
assert (_echoed(lookup.text), _echoed(report.text)) == (_served(LOOKUP), _served(REPORT)), (lookup, report)
|
||||
upstream: Final = rig.upstream_by_prompt()
|
||||
assert _uniform(upstream) == {
|
||||
prompts["lookup"]: ((rig.tool("lookup"), *_served(LOOKUP)),),
|
||||
prompts["report"]: ((rig.tool("report"), *_served(REPORT)),),
|
||||
}, upstream
|
||||
assert rig.peer_calls() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bridge", BRIDGES)
|
||||
def test_provider_outage_reaches_the_caller_and_never_the_peer_or_the_hooks(hooked: Hooked, bridge: Bridge) -> None:
|
||||
with _bridge_rig(hooked, bridge) as rig:
|
||||
key: Final = _bridge_key(rig)
|
||||
prompt: Final = f"{OUTAGE} {uuid.uuid4().hex}"
|
||||
response: Final = rig.post(key, prompt, [rig.mcp("lookup")])
|
||||
assert response.status_code == 500, response.text
|
||||
upstream: Final = rig.upstream_by_prompt()
|
||||
assert set(upstream) == {prompt} and all(
|
||||
definitions == ((rig.tool("lookup"), *_served(LOOKUP)),) for definitions in upstream[prompt]
|
||||
), upstream
|
||||
assert rig.peer_calls() == ()
|
||||
assert rig.hook_messages(prompt) == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bridge", BRIDGES)
|
||||
def test_identical_nonstream_repeat_is_a_cache_hit_without_new_model_peer_or_hook_traffic(
|
||||
hooked: Hooked, bridge: Bridge
|
||||
) -> None:
|
||||
with _bridge_rig(hooked, bridge) as rig:
|
||||
key: Final = _bridge_key(rig)
|
||||
marker: Final = "m" + uuid.uuid4().hex
|
||||
first: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync")
|
||||
assert first.text == "found:" + json.dumps({"query": marker}), first
|
||||
assert set(rig.upstream_by_prompt()) == {marker} and rig.peer_calls() == (("lookup", {"query": marker}),)
|
||||
assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",)
|
||||
assert _spend_row(key, first.call_id)["cache_hit"] != "True"
|
||||
repeat: Final = _ask(rig, key, marker, [rig.mcp("lookup")], False, "sync")
|
||||
assert repeat.text == first.text, (first, repeat)
|
||||
assert (rig.upstream_by_prompt(), rig.peer_calls(), rig.hook_messages(marker)) == ({}, (), ()), repeat
|
||||
assert _spend_row(key, repeat.call_id)["cache_hit"] == "True"
|
||||
|
||||
|
||||
def test_messages_bridge_hook_keeps_the_base_shape_without_request_local_metadata(hooked: Hooked) -> None:
|
||||
with _bridge_rig(hooked, "messages") as rig:
|
||||
key: Final = _bridge_key(rig)
|
||||
marker: Final = "m" + uuid.uuid4().hex
|
||||
probe: Final = f"{marker} {HOOK_PROBE}"
|
||||
found: Final = rig.post(key, marker, [rig.mcp("lookup")])
|
||||
assert found.status_code == 200, found.text
|
||||
blocked: Final = rig.post(key, probe, [rig.mcp("lookup")])
|
||||
assert blocked.status_code == 200, blocked.text
|
||||
assert _echoed(JSON_VALUE.validate_json(blocked.content)) == COLD, blocked.text
|
||||
assert rig.peer_calls() == (("lookup", {"query": marker}),)
|
||||
assert rig.hook_messages(marker) == (f"Tool: lookup\nArguments: {dict(query=marker)}",)
|
||||
|
|
|
|||
|
|
@ -1,16 +1,20 @@
|
|||
import base64
|
||||
import hashlib
|
||||
import re
|
||||
import secrets
|
||||
import textwrap
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
ENTRY_POINTS,
|
||||
|
|
@ -27,6 +31,7 @@ from integration._support.mcp import (
|
|||
)
|
||||
from integration._support.mcp_grants import create_toolset
|
||||
from integration._support.oauth_server import AuthorizationServer, oauth_server
|
||||
from integration._support.process import owned_proxy
|
||||
|
||||
ADD: Final = {"a": 2, "b": 3}
|
||||
CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb"
|
||||
|
|
@ -503,3 +508,70 @@ def test_resource_scoped_session_bearer_opens_a_team_toolset_inside_its_server_a
|
|||
refused: Final = _toolset_rpc(gateway, bearer, outside_name, "tools/list", {})
|
||||
assert refused.status == 403, refused.raw
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_ECHO: Final = "catalog-echo"
|
||||
_UNLISTED: Final = ""
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
|
||||
" return allow()\n"
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
|
||||
)
|
||||
|
||||
|
||||
_ECHO_GUARDRAIL_YAML: Final = (
|
||||
"guardrails:\n"
|
||||
" - guardrail_name: catalog-echo\n"
|
||||
" litellm_params:\n"
|
||||
" guardrail: custom_code\n"
|
||||
" mode: pre_mcp_call\n"
|
||||
" default_on: true\n"
|
||||
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("catalog-echo")
|
||||
path: Final = directory / "catalog_echo.yaml"
|
||||
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
|
||||
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
def _echoed_description(outcome: Outcome) -> str:
|
||||
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
|
||||
assert found is not None, outcome.raw
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def test_token_exchange_callers_with_different_subject_tokens_own_separate_listings(echo_rig: Gateway) -> None:
|
||||
with mcp_peer() as peer, oauth_server() as auth, echo_rig.scenario() as scenario:
|
||||
alias: Final = "te" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2_token_exchange",
|
||||
token_exchange_endpoint=auth.issuer + "/token",
|
||||
credentials={"client_id": "te-client", "client_secret": "te-secret"},
|
||||
)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
first_subject: Final = "subject-" + uuid.uuid4().hex
|
||||
first: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": f"Bearer {first_subject}"})
|
||||
second: Final = McpCaller(echo_rig, key, "mcp", alias, {"Authorization": "Bearer subject-" + uuid.uuid4().hex})
|
||||
auth.drain()
|
||||
assert first.list_tools().ok
|
||||
assert [request["subject_token"] for request in auth.token_requests()] == [first_subject]
|
||||
probe: Final = {"probe": _PROBE}
|
||||
own: Final = _echoed_description(first.call(f"{alias}-add", probe))
|
||||
other: Final = _echoed_description(second.call(f"{alias}-add", probe))
|
||||
assert (own, other) == ("Add two integers", _UNLISTED), (
|
||||
"the caller bearer is part of the identity on a token-exchange server: one subject, one slot"
|
||||
)
|
||||
assert second.list_tools().ok
|
||||
assert _echoed_description(second.call(f"{alias}-add", probe)) == "Add two integers"
|
||||
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
|
||||
|
|
|
|||
|
|
@ -1,13 +1,23 @@
|
|||
import itertools
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import (
|
||||
ENTRY_POINTS,
|
||||
EntryPoint,
|
||||
JsonRpc,
|
||||
McpCaller,
|
||||
Outcome,
|
||||
ScriptedTool,
|
||||
disconnecting_tool,
|
||||
echo_tool,
|
||||
listed_tools,
|
||||
|
|
@ -15,8 +25,67 @@ from integration._support.mcp import (
|
|||
register_mcp,
|
||||
scripted_peer,
|
||||
slow_tool,
|
||||
text_result,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.wire import Reply
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
_ECHO: Final = "catalog-echo:"
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_DESCRIPTION: Final = "Look up one record"
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
' texts = list(inputs.get("texts") or [])\n'
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' if "{_PROBE}" in texts:\n'
|
||||
f' return block("{_ECHO}" + json_stringify({{"description": function.get("description")}}))\n'
|
||||
" return allow()\n"
|
||||
)
|
||||
_BURST: Final = 20
|
||||
_OUTAGE: Final = 6
|
||||
_SPEND_NONCES: Final = (
|
||||
"SELECT status, metadata->'mcp_tool_call_metadata'->'arguments'->>'nonce' AS nonce"
|
||||
' FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND call_type = %s'
|
||||
)
|
||||
_OBJECTS: Final = TypeAdapter(Mapping[str, object])
|
||||
_STRINGS: Final = TypeAdapter(Mapping[str, str])
|
||||
_ECHOED: Final = TypeAdapter(Mapping[str, str | None])
|
||||
|
||||
|
||||
class _Content(BaseModel):
|
||||
text: str
|
||||
|
||||
|
||||
class _Result(BaseModel):
|
||||
content: tuple[_Content, ...]
|
||||
|
||||
|
||||
class _RpcReply(BaseModel):
|
||||
id: int
|
||||
result: _Result
|
||||
|
||||
|
||||
class _SessionsReport(BaseModel):
|
||||
worker_pid: int
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def echo_config(tmp_path_factory: pytest.TempPathFactory) -> Path:
|
||||
base: Final = _OBJECTS.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
guardrail: Final = {
|
||||
"guardrail_name": "catalog-echo-" + uuid.uuid4().hex[:8],
|
||||
"litellm_params": {
|
||||
"guardrail": "custom_code",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"custom_code": _GUARDRAIL_CODE,
|
||||
},
|
||||
}
|
||||
path: Final = tmp_path_factory.mktemp("failure-recovery") / "config.yaml"
|
||||
path.write_text(yaml.safe_dump({**base, "guardrails": [guardrail]}))
|
||||
return path
|
||||
|
||||
|
||||
def _call(caller: McpCaller, name: str, arguments: dict[str, object], entry: EntryPoint, identity: str) -> Outcome:
|
||||
|
|
@ -134,3 +203,151 @@ def test_peer_restart_on_the_same_url_is_picked_up_without_gateway_restart(gatew
|
|||
)
|
||||
assert back.text == '{"a": 1, "b": 1}', back.raw
|
||||
assert len(tool_calls(replacement.drain())) >= 1
|
||||
|
||||
|
||||
def _rpc_reply(raw: str) -> _RpcReply:
|
||||
data: Final = tuple(line[5:].strip() for line in raw.splitlines() if line.startswith("data:"))
|
||||
return _RpcReply.model_validate_json(data[-1] if data else raw)
|
||||
|
||||
|
||||
def _call_params(call: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return _OBJECTS.validate_python(_OBJECTS.validate_python(call["body"])["params"])
|
||||
|
||||
|
||||
def _call_nonce(call: Mapping[str, object]) -> str:
|
||||
return _STRINGS.validate_python(_call_params(call)["arguments"])["nonce"]
|
||||
|
||||
|
||||
def _listed(caller: McpCaller, name: str) -> None:
|
||||
listing: Final = eventually(caller.list_tools, lambda outcome: name in outcome.tools, seconds=45)
|
||||
assert listing.error is None, (caller.gateway.client.base_url, listing.raw)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Worker:
|
||||
caller: McpCaller
|
||||
pid: int
|
||||
|
||||
|
||||
def _worker(proxy: Gateway, key: str, alias: str) -> _Worker:
|
||||
sessions: Final = proxy.client.get("/v1/mcp/sessions", headers={"x-litellm-api-key": proxy.key})
|
||||
assert sessions.status_code == 200, sessions.text
|
||||
return _Worker(McpCaller(proxy, key, "mcp", alias), _SessionsReport.model_validate_json(sessions.text).worker_pid)
|
||||
|
||||
|
||||
def _served(worker: _Worker, name: str, nonce: str) -> None:
|
||||
served: Final = worker.caller.call(name, {"nonce": nonce})
|
||||
assert served.text == "found", (worker.pid, served.raw)
|
||||
|
||||
|
||||
def _probed_description(worker: _Worker, name: str) -> str | None:
|
||||
"""The description the pre_mcp_call guardrail on that worker was handed, recovered from its block reason."""
|
||||
blocked: Final = worker.caller.call(name, {"nonce": _PROBE})
|
||||
assert blocked.error is not None, (worker.pid, blocked.raw)
|
||||
carrier: Final = next((item.text for item in _rpc_reply(blocked.raw).result.content if _ECHO in item.text), None)
|
||||
assert carrier is not None, (worker.pid, blocked.raw)
|
||||
return _ECHOED.validate_json(carrier.split(_ECHO, 1)[1])["description"]
|
||||
|
||||
|
||||
def _catalog_is_cold(worker: _Worker, name: str) -> bool:
|
||||
return not _probed_description(worker, name)
|
||||
|
||||
|
||||
@pytest.mark.timeout(600)
|
||||
def test_worker_restart_cools_its_listed_catalog_while_the_sibling_worker_keeps_serving(
|
||||
gateway: Gateway, echo_config: Path, tmp_path: Path
|
||||
) -> None:
|
||||
tool: Final = ScriptedTool("lookup", lambda _: text_result("found"), description=_DESCRIPTION)
|
||||
with scripted_peer(tool) as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "cold" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
name: Final = f"{alias}-lookup"
|
||||
with owned_proxy_process(gateway, tmp_path / "sibling", {}, config=echo_config) as sibling_proxy:
|
||||
sibling: Final = _worker(sibling_proxy.gateway, key, alias)
|
||||
with owned_proxy_process(gateway, tmp_path / "first", {}, config=echo_config) as first_proxy:
|
||||
first: Final = _worker(first_proxy.gateway, key, alias)
|
||||
assert first.pid != sibling.pid
|
||||
_served(first, name, "first-unlisted")
|
||||
_served(sibling, name, "sibling-unlisted")
|
||||
assert _catalog_is_cold(first, name) and _catalog_is_cold(sibling, name)
|
||||
_listed(first.caller, name)
|
||||
assert _probed_description(first, name) == _DESCRIPTION
|
||||
assert _catalog_is_cold(sibling, name), "a listing on one worker warmed its sibling"
|
||||
_listed(sibling.caller, name)
|
||||
assert _probed_description(sibling, name) == _DESCRIPTION
|
||||
_served(sibling, name, "sibling-alone")
|
||||
with owned_proxy_process(gateway, tmp_path / "restarted", {}, config=echo_config) as restarted_proxy:
|
||||
restarted: Final = _worker(restarted_proxy.gateway, key, alias)
|
||||
assert restarted.pid not in (first.pid, sibling.pid)
|
||||
assert _catalog_is_cold(restarted, name), "a restarted worker kept the old process's catalog"
|
||||
assert _probed_description(sibling, name) == _DESCRIPTION
|
||||
_served(restarted, name, "restarted-unlisted")
|
||||
_listed(restarted.caller, name)
|
||||
assert _probed_description(restarted, name) == _DESCRIPTION
|
||||
calls: Final = tool_calls(peer.drain())
|
||||
assert [_call_nonce(call) for call in calls] == [
|
||||
"first-unlisted",
|
||||
"sibling-unlisted",
|
||||
"sibling-alone",
|
||||
"restarted-unlisted",
|
||||
], calls
|
||||
assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in calls), calls
|
||||
|
||||
|
||||
def _outage_echo(name: str, failures: int) -> ScriptedTool:
|
||||
attempts: Final = itertools.count(1)
|
||||
|
||||
def respond(params: JsonRpc) -> Reply | JsonRpc:
|
||||
if next(attempts) <= failures:
|
||||
return Reply(status=503, body=b'{"error": "scripted outage"}')
|
||||
return text_result(_STRINGS.validate_python(params["arguments"])["nonce"])
|
||||
|
||||
return ScriptedTool(name, respond)
|
||||
|
||||
|
||||
def _echo_call(caller: McpCaller, name: str, nonce: str) -> Outcome:
|
||||
return caller.call(name, {"nonce": nonce})
|
||||
|
||||
|
||||
def _logged_nonces(key: str, count: int) -> tuple[tuple[str, str], ...]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(_SPEND_NONCES, (sha256(key.encode()).hexdigest(), "call_mcp_tool")),
|
||||
lambda found: len(found) >= count,
|
||||
seconds=70,
|
||||
)
|
||||
return tuple(sorted((str(row["status"]), str(row["nonce"])) for row in rows))
|
||||
|
||||
|
||||
def test_peer_outage_during_a_bounded_burst_fails_exactly_the_outage_calls_and_lands_each_call_once(
|
||||
gateway: Gateway, peer: Gateway
|
||||
) -> None:
|
||||
with scripted_peer(_outage_echo("echo", _OUTAGE)) as upstream, gateway.scenario() as scenario:
|
||||
alias: Final = "burst" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, upstream, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
name: Final = f"{alias}-echo"
|
||||
callers: Final = (McpCaller(gateway, key, "mcp", alias), McpCaller(peer, key, "mcp", alias))
|
||||
for caller in callers:
|
||||
_listed(caller, name)
|
||||
upstream.drain()
|
||||
nonces: Final = tuple(uuid.uuid4().hex for _ in range(_BURST))
|
||||
with ThreadPoolExecutor(max_workers=_BURST) as pool:
|
||||
outcomes: Final = tuple(pool.map(_echo_call, itertools.cycle(callers), itertools.repeat(name), nonces))
|
||||
raws: Final = [outcome.raw for outcome in outcomes]
|
||||
failed: Final = tuple(nonce for nonce, outcome in zip(nonces, outcomes) if outcome.error is not None)
|
||||
assert len(failed) == _OUTAGE, raws
|
||||
assert all(outcome.error is not None or outcome.text == nonce for nonce, outcome in zip(nonces, outcomes)), raws
|
||||
assert all(_rpc_reply(outcome.raw).id == 1 for outcome in outcomes), raws
|
||||
burst_calls: Final = tool_calls(upstream.drain())
|
||||
assert sorted(_call_nonce(call) for call in burst_calls) == sorted(nonces), burst_calls
|
||||
assert all(_call_params(call)["arguments"] == {"nonce": _call_nonce(call)} for call in burst_calls), burst_calls
|
||||
assert all(set(_call_params(call)) - {"_meta"} == {"name", "arguments"} for call in burst_calls), burst_calls
|
||||
recovered: Final = tuple(
|
||||
_echo_call(caller, name, nonce) for caller, nonce in zip(itertools.cycle(callers), failed)
|
||||
)
|
||||
assert [outcome.text for outcome in recovered] == list(failed), [outcome.raw for outcome in recovered]
|
||||
assert sorted(_call_nonce(call) for call in tool_calls(upstream.drain())) == sorted(failed)
|
||||
assert _logged_nonces(key, _BURST + len(failed)) == tuple(
|
||||
sorted([("success", nonce) for nonce in nonces] + [("failure", nonce) for nonce in failed])
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
import re
|
||||
import secrets
|
||||
import textwrap
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
|
@ -7,14 +10,17 @@ from typing import Final
|
|||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, object_value
|
||||
from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value
|
||||
from integration._support.mcp import (
|
||||
INITIALIZE,
|
||||
Outcome,
|
||||
ScriptedTool,
|
||||
_outcome_from_rest,
|
||||
_outcome_from_rpc,
|
||||
mcp_peer,
|
||||
register_mcp,
|
||||
scripted_peer,
|
||||
text_result,
|
||||
tool_calls,
|
||||
)
|
||||
from integration._support.mcp_grants import create_toolset
|
||||
|
|
@ -507,3 +513,63 @@ def test_a_member_of_two_teams_sees_the_union_and_each_route_stays_narrowed_to_i
|
|||
crossed: Final = _route_call(gateway, headers, first_name, f"{alias}-multiply")
|
||||
assert not crossed.ok, crossed.raw
|
||||
assert tool_calls(peer.drain()) == ()
|
||||
|
||||
|
||||
_PROBE: Final = "catalog-probe"
|
||||
_ECHO: Final = "catalog-echo"
|
||||
_UNLISTED: Final = ""
|
||||
_GUARDRAIL_CODE: Final = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
f' if "{_PROBE}" not in list(inputs.get("texts") or []):\n'
|
||||
" return allow()\n"
|
||||
' function = inputs.get("tools", [{}])[0].get("function", {})\n'
|
||||
f' return block("{_ECHO}[" + function.get("description") + "]")\n'
|
||||
)
|
||||
|
||||
|
||||
_ECHO_GUARDRAIL_YAML: Final = (
|
||||
"guardrails:\n"
|
||||
" - guardrail_name: catalog-echo\n"
|
||||
" litellm_params:\n"
|
||||
" guardrail: custom_code\n"
|
||||
" mode: pre_mcp_call\n"
|
||||
" default_on: true\n"
|
||||
" custom_code: |\n" + textwrap.indent(_GUARDRAIL_CODE, 8 * " ")
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def echo_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("catalog-echo")
|
||||
path: Final = directory / "catalog_echo.yaml"
|
||||
path.write_text((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text() + _ECHO_GUARDRAIL_YAML)
|
||||
with gateway_from_environment() as gateway, owned_proxy(gateway, directory, {}, config=path, workers=2) as rig:
|
||||
yield rig
|
||||
|
||||
|
||||
def _echoed_description(outcome: Outcome) -> str:
|
||||
found: Final = re.search(rf"{_ECHO}\[(.*?)\]", outcome.raw)
|
||||
assert found is not None, outcome.raw
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def test_a_team_keys_toolset_route_listing_feeds_its_own_calls_but_not_a_team_mates(echo_rig: Gateway) -> None:
|
||||
described: Final = "Adds for the team " + uuid.uuid4().hex[:8]
|
||||
tool: Final = ScriptedTool("add", lambda _: text_result("9"), description=described)
|
||||
with scripted_peer(tool) as peer, echo_rig.scenario() as scenario:
|
||||
alias: Final = "lit6029echo" + uuid.uuid4().hex[:6]
|
||||
server_id: Final = register_mcp(scenario, peer, alias)
|
||||
granted_id, granted_name = _toolset(scenario, server_id, "add")
|
||||
team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]})
|
||||
key: Final = scenario.key(team_id=team_id)
|
||||
team_mate: Final = scenario.key(team_id=team_id)
|
||||
_assert_team_grants_only(echo_rig, team_id, key, granted_id)
|
||||
listed: Final = _toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/list", {})
|
||||
assert listed.ok and listed.tools == (f"{alias}-add",), listed.raw
|
||||
probe: Final[dict[str, object]] = {"name": f"{alias}-add", "arguments": {"probe": _PROBE}}
|
||||
own: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(key), granted_name, "tools/call", probe))
|
||||
mate: Final = _echoed_description(_toolset_rpc(echo_rig, _bearer(team_mate), granted_name, "tools/call", probe))
|
||||
assert (own, mate) == (described, _UNLISTED), (
|
||||
"the slot is keyed by the hashed key, so a team-mate that never listed is handed nothing"
|
||||
)
|
||||
assert tool_calls(peer.drain()) == (), "a blocked probe reached the peer"
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import binascii
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import sys
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
from typing import Final, List, Optional
|
||||
import pytest
|
||||
from litellm._uuid import uuid
|
||||
import os
|
||||
|
|
@ -391,6 +391,7 @@ async def test_create_mcp_server_invalid_alias():
|
|||
@_SKIP_NO_MCP
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_redacts_credentials():
|
||||
mock_get_server: Final = mock.AsyncMock()
|
||||
with (
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE",
|
||||
|
|
@ -399,6 +400,10 @@ async def test_edit_mcp_server_redacts_credentials():
|
|||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw"
|
||||
) as mock_get_prisma,
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
new=mock_get_server,
|
||||
),
|
||||
mock.patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
|
||||
new_callable=mock.AsyncMock,
|
||||
|
|
@ -422,6 +427,18 @@ async def test_edit_mcp_server_redacts_credentials():
|
|||
mock_manager.reload_servers_from_database = mock.AsyncMock()
|
||||
|
||||
server_id = str(uuid.uuid4())
|
||||
stored_server: Final = LiteLLM_MCPServerTable(
|
||||
server_id=server_id,
|
||||
alias="Updated Server",
|
||||
url="https://updated.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
credentials={"auth_value": "secret"},
|
||||
teams=[],
|
||||
)
|
||||
mock_get_server.return_value = stored_server
|
||||
|
||||
updated_server = LiteLLM_MCPServerTable(
|
||||
server_id=server_id,
|
||||
alias="Updated Server",
|
||||
|
|
@ -458,6 +475,7 @@ async def test_edit_mcp_server_redacts_credentials():
|
|||
mock_update.assert_awaited_once()
|
||||
mock_manager.update_server.assert_called_once_with(updated_server)
|
||||
mock_manager.reload_servers_from_database.assert_awaited_once()
|
||||
mock_get_server.assert_awaited_once_with(mock_prisma, server_id)
|
||||
|
||||
|
||||
def test_validate_mcp_server_name_direct():
|
||||
|
|
|
|||
|
|
@ -490,10 +490,10 @@ def seeded_trace_api(clickhouse_url: str) -> Iterator[SeededTraceAPI]:
|
|||
rebase_spend,
|
||||
)
|
||||
|
||||
spends: Final = dict(spend_fixtures())["deeplite_swarm"]
|
||||
spends: Final = dict(spend_fixtures())["openai_agents_swarm"]
|
||||
pattern: Final = re.compile("|".join(re.escape(row["response_id"]) for row in spends))
|
||||
replays: Final = fixture_replays(TRACE_FIXTURES, time.time_ns() // 1_000_000, "query-api", pattern)
|
||||
swarm: Final = next(replay for replay in replays if replay.name == "deeplite_swarm")
|
||||
swarm: Final = next(replay for replay in replays if replay.name == "openai_agents_swarm")
|
||||
rebased: Final = rebase_spend(spends, swarm.offset_ms, swarm.namespace, pattern)
|
||||
stamped: Final[tuple[SpendLogRecord, ...]] = tuple(
|
||||
{**row, "team_id": "team-a", "api_key": "fixture-key", "user": "fixture-user"} for row in rebased
|
||||
|
|
@ -533,12 +533,11 @@ def test_fixture_backed_help_examples_execute_through_query_api(seeded_trace_api
|
|||
assert {table.name for table in api.help.tables} == {"otel_traces", "spend_logs", "agent_traces_by_key"}
|
||||
assert api.help.metadata.error is None
|
||||
assert api.help.metadata.sampled_rows == len(api.spends)
|
||||
assert any(field.path == ("synthetic_spend",) for field in api.help.metadata.fields)
|
||||
assert any(field.path == ("fixture_capture", "name") for field in api.help.metadata.fields)
|
||||
for example in api.help.examples:
|
||||
api.query_example(example.name)
|
||||
records: Final = api.query_example("Recent spend records")
|
||||
assert {str(row["request_id"]) for row in records} == {row["request_id"] for row in api.spends}
|
||||
assert all(bool(row["synthetic_spend"]) for row in records)
|
||||
total: Final = sum(row["spend"] or 0 for row in api.spends)
|
||||
recorded: Final = api.query_example("Recorded spend by trace")
|
||||
assert len(recorded) == 1
|
||||
|
|
@ -616,7 +615,7 @@ def captured_trace_api() -> Iterator[SeededTraceAPI]:
|
|||
yield from _fixture_trace_api(url, replays, stamped)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", tuple(name for name, _ in spend_fixtures() if name != "deeplite_swarm"))
|
||||
@pytest.mark.parametrize("name", tuple(name for name, _ in spend_fixtures()))
|
||||
def test_captured_sdk_cost_survives_seeding_and_is_queryable(name: str, captured_trace_api: SeededTraceAPI) -> None:
|
||||
api: Final = captured_trace_api
|
||||
rows: Final = tuple(row for row in api.spends if fixture_capture("", row).name == name)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import importlib
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
|
|
@ -12,10 +13,10 @@ import pytest
|
|||
from uvicorn.importer import import_from_string
|
||||
from uvicorn.main import main as uvicorn_main
|
||||
|
||||
import gateway.main
|
||||
from gateway.launch import GATEWAY_APP, main, pool_database_url, uvicorn_argv
|
||||
from litellm.proxy.db.db_url_settings import DatabaseURLSettings
|
||||
from litellm.proxy.db.pgbouncer import PGBOUNCER_POOLED_ENV_VAR, PgBouncerError, PgBouncerSettings
|
||||
from litellm.proxy.proxy_server import app as proxy_app
|
||||
|
||||
DB_ENV: Final = {
|
||||
"DATABASE_HOST": "db.internal",
|
||||
|
|
@ -107,8 +108,22 @@ class TestUvicornArgv:
|
|||
argv: Final = uvicorn_argv(("--timeout-keep-alive", "30"), {"KEEPALIVE_TIMEOUT": "75"})
|
||||
assert _uvicorn_params(argv)["timeout_keep_alive"] == 30
|
||||
|
||||
def test_the_app_uvicorn_is_told_to_serve_is_the_trimmed_gateway(self):
|
||||
assert import_from_string(cast(str, _uvicorn_params(uvicorn_argv((), {}))["app"])) is gateway.main.app
|
||||
def test_the_app_uvicorn_is_told_to_serve_is_the_trimmed_gateway(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(proxy_app.router, "lifespan_context", proxy_app.router.lifespan_context)
|
||||
for key in (
|
||||
"DATABASE_URL",
|
||||
"DIRECT_URL",
|
||||
"DATABASE_URL_READ_REPLICA",
|
||||
"DATABASE_HOST",
|
||||
"DATABASE_HOST_READ_REPLICA",
|
||||
"DATABASE_PASSWORD",
|
||||
"IAM_TOKEN_DB_AUTH",
|
||||
"AZURE_POSTGRESQL_AUTH",
|
||||
):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
|
||||
served: Final = import_from_string(cast(str, _uvicorn_params(uvicorn_argv((), {}))["app"]))
|
||||
assert served is importlib.import_module("gateway.main").app
|
||||
|
||||
|
||||
class TestPoolDatabaseUrl:
|
||||
|
|
|
|||
|
|
@ -1,11 +1,13 @@
|
|||
import glob
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from typing import Final, NoReturn, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -19,6 +21,8 @@ sys.path.insert(
|
|||
from litellm_proxy_extras.utils import (
|
||||
PARTITIONED_SPEND_LOGS_PUSH_ERROR,
|
||||
ProxyExtrasDBManager,
|
||||
_redact_command_error,
|
||||
_redact_credentials,
|
||||
filter_partitioned_spend_logs_diff,
|
||||
)
|
||||
|
||||
|
|
@ -1412,3 +1416,411 @@ class TestMigrationJobOwnedDrift:
|
|||
assert 'PRIMARY KEY ("request_id")' not in filtered
|
||||
assert "LiteLLM_SpendLogs_legacy" not in filtered
|
||||
assert 'ALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN "updated_by" TEXT;' in filtered
|
||||
|
||||
|
||||
_P3018_UNCLASSIFIED_STDERR: Final = (
|
||||
"Error: P3018\n\n"
|
||||
"A migration failed to apply. New migrations cannot be applied before the error is "
|
||||
"recovered from.\n\n"
|
||||
"Migration name: 20260921190000_agent_identity\n\n"
|
||||
"Database error code: 23505\n\n"
|
||||
"Database error:\n"
|
||||
'ERROR: could not create unique index "agent_identity_key"\n'
|
||||
"DETAIL: Key (agent_id)=(agent-1) is duplicated.\n"
|
||||
)
|
||||
|
||||
|
||||
_FAKE_PRISMA_PID: Final = 424242
|
||||
|
||||
|
||||
class TestV1MigrationFailuresLogAtError:
|
||||
@staticmethod
|
||||
def _run_v1_migrations(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path: Path,
|
||||
*,
|
||||
deploy_stderr: Optional[str] = None,
|
||||
deploy_timeout: bool = False,
|
||||
diff_stderr: Optional[str] = None,
|
||||
database_url: Optional[str] = None,
|
||||
) -> tuple[bool, list[list[str]], list[int]]:
|
||||
import litellm_proxy_extras.utils as utils_module
|
||||
|
||||
calls: Final[list[list[str]]] = []
|
||||
killed_pids: Final[list[int]] = []
|
||||
|
||||
class _FakePrismaPopen:
|
||||
def __init__(
|
||||
self,
|
||||
argv: tuple[str, ...],
|
||||
*,
|
||||
env: Optional[dict[str, str]] = None,
|
||||
stdout: object = None,
|
||||
stderr: object = None,
|
||||
text: object = None,
|
||||
start_new_session: object = None,
|
||||
) -> None:
|
||||
self.args: Final = argv
|
||||
self.argv: Final = argv
|
||||
self.pid: Final = _FAKE_PRISMA_PID
|
||||
self.returncode: Optional[int] = None
|
||||
calls.append(list(argv))
|
||||
|
||||
def __enter__(self) -> "_FakePrismaPopen":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def _subcommand(self) -> tuple[str, str]:
|
||||
known: Final = (
|
||||
("migrate", "deploy"),
|
||||
("migrate", "diff"),
|
||||
("migrate", "resolve"),
|
||||
("db", "execute"),
|
||||
)
|
||||
for index in range(len(self.argv) - 1):
|
||||
pair: Final = tuple(self.argv[index : index + 2])
|
||||
if pair in known:
|
||||
return pair
|
||||
return ("", "")
|
||||
|
||||
def communicate(self, timeout: Optional[float] = None) -> tuple[str, str]:
|
||||
subcommand: Final = self._subcommand()
|
||||
if subcommand == ("migrate", "deploy"):
|
||||
if deploy_timeout:
|
||||
raise subprocess.TimeoutExpired(self.argv, timeout)
|
||||
if deploy_stderr is not None:
|
||||
self.returncode = 1
|
||||
return "", deploy_stderr
|
||||
self.returncode = 0
|
||||
return "No pending migrations to apply", ""
|
||||
if subcommand == ("migrate", "diff") and diff_stderr is not None:
|
||||
self.returncode = 1
|
||||
return "", diff_stderr
|
||||
self.returncode = 0
|
||||
return "", ""
|
||||
|
||||
migration_dir: Final = tmp_path / "migration_dir"
|
||||
migration_dir.mkdir()
|
||||
if database_url is None:
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(migration_dir))
|
||||
monkeypatch.setattr(
|
||||
utils_module.prisma_toolchain.subprocess, "Popen", _FakePrismaPopen
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
utils_module.prisma_toolchain.os, "killpg", lambda pid, sig: killed_pids.append(pid)
|
||||
)
|
||||
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
|
||||
|
||||
succeeded: Final = ProxyExtrasDBManager._run_migrations(use_migrate=True, use_v2_resolver=False)
|
||||
return succeeded, calls, killed_pids
|
||||
|
||||
@staticmethod
|
||||
def _deploy_call_count(calls: list[list[str]]) -> int:
|
||||
return sum(1 for call in calls if tuple(call[-2:]) == ("migrate", "deploy"))
|
||||
|
||||
@staticmethod
|
||||
def _error_messages(caplog: pytest.LogCaptureFixture) -> list[str]:
|
||||
return [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.ERROR and record.name.startswith("litellm_proxy_extras")
|
||||
]
|
||||
|
||||
def test_an_unrecognized_prisma_error_logs_its_stderr_at_error(
|
||||
self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
stderr: Final = "Error: P1001: Can't reach database server at db:5432"
|
||||
with caplog.at_level(logging.ERROR, logger="litellm_proxy_extras"):
|
||||
succeeded, calls, _ = self._run_v1_migrations(
|
||||
monkeypatch, tmp_path, deploy_stderr=stderr
|
||||
)
|
||||
|
||||
assert succeeded is False
|
||||
assert self._deploy_call_count(calls) == 4
|
||||
assert any(stderr in message for message in self._error_messages(caplog))
|
||||
|
||||
def test_an_unclassified_p3018_logs_its_stderr_and_retry_failure_at_error(
|
||||
self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
with caplog.at_level(logging.ERROR, logger="litellm_proxy_extras"):
|
||||
succeeded, calls, _ = self._run_v1_migrations(
|
||||
monkeypatch, tmp_path, deploy_stderr=_P3018_UNCLASSIFIED_STDERR
|
||||
)
|
||||
|
||||
assert succeeded is False
|
||||
assert self._deploy_call_count(calls) == 4
|
||||
messages: Final = self._error_messages(caplog)
|
||||
assert any(
|
||||
"20260921190000_agent_identity" in message and "is duplicated" in message
|
||||
for message in messages
|
||||
)
|
||||
assert any(
|
||||
"The process failed to execute" in message and "Retrying... (3 attempts left)" in message
|
||||
for message in messages
|
||||
)
|
||||
|
||||
def test_called_process_error_with_no_command_retries_all_v1_attempts(
|
||||
self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
import litellm_proxy_extras.utils as utils_module
|
||||
|
||||
migration_dir: Final = tmp_path / "migration_dir"
|
||||
migration_dir.mkdir()
|
||||
monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(migration_dir))
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
monkeypatch.setenv("PRISMA_OFFLINE_MODE", "true")
|
||||
monkeypatch.setenv("PRISMA_CLI_PATH", sys.executable)
|
||||
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
|
||||
calls: Final[list[None]] = []
|
||||
|
||||
class _FakePrismaPopen:
|
||||
def __init__(
|
||||
self,
|
||||
argv: tuple[str, ...],
|
||||
*,
|
||||
env: Optional[dict[str, str]] = None,
|
||||
stdout: object = None,
|
||||
stderr: object = None,
|
||||
text: object = None,
|
||||
start_new_session: object = None,
|
||||
) -> None:
|
||||
self.args: Final = None
|
||||
self.returncode: Final = 1
|
||||
calls.append(None)
|
||||
|
||||
def __enter__(self) -> "_FakePrismaPopen":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def communicate(self, timeout: Optional[float] = None) -> tuple[str, str]:
|
||||
return "", "Error: P3018 unclassified"
|
||||
|
||||
monkeypatch.setattr(
|
||||
utils_module.prisma_toolchain.subprocess, "Popen", _FakePrismaPopen
|
||||
)
|
||||
|
||||
try:
|
||||
succeeded: Final = ProxyExtrasDBManager._run_migrations(
|
||||
use_migrate=True, use_v2_resolver=False
|
||||
)
|
||||
except TypeError as error:
|
||||
pytest.fail(
|
||||
f"_run_migrations raised TypeError after {len(calls)} Popen calls: {error}",
|
||||
pytrace=False,
|
||||
)
|
||||
|
||||
assert succeeded is False
|
||||
assert len(calls) == 4
|
||||
|
||||
def test_a_timeout_logs_at_error_naming_the_migrate_deploy_timeout_env_var(
|
||||
self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
from litellm_proxy_extras.prisma_toolchain import PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR
|
||||
|
||||
with caplog.at_level(logging.ERROR, logger="litellm_proxy_extras"):
|
||||
succeeded, calls, killed_pids = self._run_v1_migrations(
|
||||
monkeypatch, tmp_path, deploy_timeout=True
|
||||
)
|
||||
|
||||
assert succeeded is False
|
||||
assert self._deploy_call_count(calls) == 4
|
||||
assert killed_pids == [_FAKE_PRISMA_PID] * 4
|
||||
assert any(
|
||||
"timed out" in message and PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR in message
|
||||
for message in self._error_messages(caplog)
|
||||
)
|
||||
|
||||
def test_a_recovered_baseline_logs_nothing_at_error(
|
||||
self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
with caplog.at_level(logging.ERROR, logger="litellm_proxy_extras"):
|
||||
succeeded, calls, _ = self._run_v1_migrations(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
deploy_stderr=_P3005_STDERR,
|
||||
database_url="postgresql://user:pass@db:5432/litellm",
|
||||
)
|
||||
|
||||
assert succeeded is True
|
||||
assert tuple(calls[0][-2:]) == ("migrate", "deploy")
|
||||
assert self._error_messages(caplog) == []
|
||||
|
||||
def test_a_failed_baseline_recovery_logs_its_stderr_at_error(
|
||||
self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
database_url: Final = "postgresql://llmproxy:s3cr3t 'p\"w@db:5432/litellm"
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
with caplog.at_level(logging.DEBUG, logger="litellm_proxy_extras"):
|
||||
succeeded, calls, _ = self._run_v1_migrations(
|
||||
monkeypatch,
|
||||
tmp_path,
|
||||
deploy_stderr=_P3005_STDERR,
|
||||
diff_stderr=f"baseline diff failed: XYZ-7731 for {database_url}",
|
||||
database_url=database_url,
|
||||
)
|
||||
|
||||
assert succeeded is False
|
||||
assert self._deploy_call_count(calls) == 4
|
||||
assert [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if "s3cr3t" in record.getMessage() or 'p"w' in record.getMessage()
|
||||
] == []
|
||||
messages: Final = self._error_messages(caplog)
|
||||
assert any("postgresql://REDACTED@db:5432/litellm" in message for message in messages)
|
||||
assert any("XYZ-7731" in message for message in messages)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"database_url,direct_url,text,expected",
|
||||
(
|
||||
(
|
||||
"postgresql://u:pa ss@db:5432/litellm",
|
||||
None,
|
||||
'Error: P1000: Authentication failed against database server at "postgresql://u:pa ss@db:5432/litellm"',
|
||||
'Error: P1000: Authentication failed against database server at "postgresql://REDACTED@db:5432/litellm"',
|
||||
),
|
||||
(
|
||||
"postgresql://u:pa'ss@db:5432/litellm",
|
||||
None,
|
||||
"postgresql://u:pa'ss@db:5432/litellm",
|
||||
"postgresql://REDACTED@db:5432/litellm",
|
||||
),
|
||||
(
|
||||
'postgresql://u:pa"ss@db:5432/litellm',
|
||||
None,
|
||||
'postgresql://u:pa"ss@db:5432/litellm',
|
||||
"postgresql://REDACTED@db:5432/litellm",
|
||||
),
|
||||
(
|
||||
"postgresql://u:p@ss@db:5432/litellm",
|
||||
None,
|
||||
"postgresql://u:p@ss@db:5432/litellm",
|
||||
"postgresql://REDACTED@db:5432/litellm",
|
||||
),
|
||||
(
|
||||
"postgresql://u:p%20ss@db:5432/litellm",
|
||||
None,
|
||||
"postgresql://u:p ss@db:5432/litellm",
|
||||
"postgresql://REDACTED@db:5432/litellm",
|
||||
),
|
||||
(
|
||||
"postgresql://db/litellm?password=a b&sslmode=require",
|
||||
None,
|
||||
"postgresql://db/litellm?password=a b&sslmode=require",
|
||||
"postgresql://db/litellm?REDACTED&sslmode=require",
|
||||
),
|
||||
(
|
||||
"postgresql://db/litellm?sslpassword=zq'7x",
|
||||
None,
|
||||
"postgresql://db/litellm?sslpassword=zq'7x",
|
||||
"postgresql://db/litellm?REDACTED",
|
||||
),
|
||||
(
|
||||
None,
|
||||
"postgresql://u:pa ss@db:5432/litellm",
|
||||
"postgresql://u:pa ss@db:5432/litellm",
|
||||
"postgresql://REDACTED@db:5432/litellm",
|
||||
),
|
||||
(
|
||||
None,
|
||||
None,
|
||||
"postgresql://u:pw@db/x",
|
||||
"postgresql://REDACTED@db/x",
|
||||
),
|
||||
(
|
||||
"postgresql://u:p@db:5432/litellm",
|
||||
None,
|
||||
"Error: P1001: Can't reach database server at db:5432",
|
||||
"Error: P1001: Can't reach database server at db:5432",
|
||||
),
|
||||
(
|
||||
"postgresql://u:p@db:5432/litellm",
|
||||
None,
|
||||
"Error:P1001: Can't reach database server at db:5432",
|
||||
"Error:P1001: Can't reach database server at db:5432",
|
||||
),
|
||||
(
|
||||
None,
|
||||
None,
|
||||
"plain text with no URL",
|
||||
"plain text with no URL",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_redact_credentials_masks_passwords_in_embedded_urls(
|
||||
database_url: str | None,
|
||||
direct_url: str | None,
|
||||
text: str,
|
||||
expected: str,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
if database_url is None:
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
if direct_url is None:
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("DIRECT_URL", direct_url)
|
||||
assert _redact_credentials(text) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("password", ("zq'7x", 'zq"7x', "zq'\"7x", "zq 7x", "zq@7x"))
|
||||
def test_redact_command_error_masks_url_arguments(password: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
database_url: Final = f"postgresql://u:{password}@db:5432/litellm"
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
error: Final = subprocess.CalledProcessError(1, ["prisma", "migrate", "diff", "--to-url", database_url])
|
||||
|
||||
message: Final = _redact_command_error(error)
|
||||
|
||||
assert "zq" not in message
|
||||
assert "7x" not in message
|
||||
assert "postgresql://REDACTED@db:5432/litellm" in message
|
||||
assert "returned non-zero exit status 1" in message
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"command",
|
||||
(
|
||||
None,
|
||||
Path("/usr/bin/prisma"),
|
||||
7,
|
||||
("prisma", "migrate", "deploy"),
|
||||
["prisma", "migrate", "deploy"],
|
||||
"prisma migrate deploy",
|
||||
),
|
||||
)
|
||||
def test_redact_command_error_preserves_unredacted_command_format(
|
||||
command: object, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
error: Final = subprocess.CalledProcessError(1, command)
|
||||
|
||||
assert _redact_command_error(error) == str(error)
|
||||
|
||||
|
||||
def test_redact_command_error_masks_password_in_tuple_url_argument(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
password: Final = "zq 7x"
|
||||
database_url: Final = f"postgresql://u:{password}@db:5432/litellm"
|
||||
monkeypatch.setenv("DATABASE_URL", database_url)
|
||||
monkeypatch.delenv("DIRECT_URL", raising=False)
|
||||
error: Final = subprocess.CalledProcessError(1, ("prisma", "migrate", "deploy", "--to-url", database_url))
|
||||
|
||||
message: Final = _redact_command_error(error)
|
||||
|
||||
assert message.startswith("Command '('")
|
||||
assert password not in message
|
||||
assert "postgresql://REDACTED@db:5432/litellm" in message
|
||||
|
|
|
|||
|
|
@ -226,6 +226,76 @@ async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials(
|
|||
assert execution["guardrail_context"] == {"metadata": {"guardrails": ("block-all",)}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_with_mcp_hands_execution_the_requests_served_tools():
|
||||
"""
|
||||
Regression test: /v1/messages auto-execution must carry this request's
|
||||
resolved tool definitions into execution, matching the Responses and chat
|
||||
completions bridges.
|
||||
|
||||
Given: A request whose MCP reference resolves to a definition carrying a
|
||||
description and input schema
|
||||
When: The model asks for that tool and the gateway executes it
|
||||
Then: _execute_tool_calls receives the definition under served_tools, so
|
||||
pre_mcp_call hooks can judge the call on what the model was shown
|
||||
|
||||
Dropping it does not fail loudly; the call still runs, but the hook sees
|
||||
only the name and arguments, leaving /v1/messages permanently colder than
|
||||
the other two bridges even though all three resolve the same definitions.
|
||||
"""
|
||||
from mcp.types import Tool
|
||||
|
||||
from litellm.llms.anthropic.pass_through.messages import mcp_handler
|
||||
from litellm.responses.mcp.request_context import MCPRequestContext
|
||||
|
||||
served = [
|
||||
Tool(
|
||||
name="read_wiki_structure",
|
||||
description="Read the structure of a wiki",
|
||||
inputSchema={"type": "object", "properties": {"repoName": {"type": "string"}}},
|
||||
)
|
||||
]
|
||||
|
||||
process = AsyncMock(return_value=(served, {"read_wiki_structure": "deepwiki"}))
|
||||
execute = AsyncMock(
|
||||
return_value=[{"tool_call_id": "toolu_1", "result": "ok", "name": "read_wiki_structure"}]
|
||||
)
|
||||
responses = [
|
||||
{
|
||||
"stop_reason": "tool_use",
|
||||
"content": [{"type": "tool_use", "id": "toolu_1", "name": "read_wiki_structure", "input": {}}],
|
||||
},
|
||||
{"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]},
|
||||
]
|
||||
|
||||
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=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,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
model="claude-sonnet-4-5",
|
||||
tools=[MCP_REFERENCE],
|
||||
)
|
||||
|
||||
execution = execute.call_args.kwargs
|
||||
assert execution.get("served_tools") == served, (
|
||||
"The request's resolved tool definitions must reach execution so pre_mcp_call "
|
||||
"hooks see the listed description and input schema"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ def _bare_manager() -> MOD.MCPServerManager:
|
|||
reaches the guardrail hooks; they have their own coverage elsewhere.
|
||||
"""
|
||||
mgr = MOD.MCPServerManager.__new__(MOD.MCPServerManager)
|
||||
mgr._listed_tools_by_server_id = {}
|
||||
mgr.check_allowed_or_banned_tools = lambda name, server: True
|
||||
mgr.validate_allowed_params = lambda tool_name, arguments, server: None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,17 +1,29 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
handle_mcp_proxy_tool,
|
||||
mcp_proxy_tool_id,
|
||||
with_mcp_proxy_identity,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
AUTH = UserAPIKeyAuth(api_key="key")
|
||||
|
||||
|
|
@ -130,3 +142,46 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke
|
|||
assert hook_payload["arguments"] == arguments
|
||||
assert "raw_headers" not in hook_payload
|
||||
assert "raw-scope-secret" not in recorder.events[1][1]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_call_tool_on_a_never_listed_tool_hands_the_pre_hook_no_listed_tool() -> None:
|
||||
"""/mcp/proxy tools/list serves only the meta-tools, so the catalog call_tool reads to resolve its
|
||||
tool_id was never served: it must not fill the caller's listed-tools slot, and the pre-call hook
|
||||
must see no listed tool for the call."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="proxy-meta", name="proxy-meta", transport=MCPTransport.http, url="http://meta")
|
||||
auth = UserAPIKeyAuth(api_key="sk-proxy-meta", user_id="proxy-caller")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
served_as = with_mcp_proxy_identity(MCPTool(name="proxy-meta-echo", inputSchema={}), server.server_id)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
async def call_regular_mcp_tool(*, tasks: list[asyncio.Task[object]], **_: object) -> CallToolResult:
|
||||
await asyncio.gather(*tasks)
|
||||
return CallToolResult(content=[TextContent(type="text", text="echoed")])
|
||||
|
||||
with (
|
||||
patch.dict(manager.registry, {server.server_id: server}),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "pre_call_tool_check", pre_call_tool_check),
|
||||
patch.object(manager, "_call_regular_mcp_tool", call_regular_mcp_tool),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_proxy_tool(
|
||||
name="call_tool",
|
||||
arguments={"tool_id": mcp_proxy_tool_id(served_as), "arguments": {}},
|
||||
user_api_key_dict=auth,
|
||||
)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=auth))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "echoed"
|
||||
pre_call_tool_check.assert_awaited_once()
|
||||
assert pre_call_tool_check.await_args.kwargs["name"] == "echo"
|
||||
assert pre_call_tool_check.await_args.kwargs["tool"] is None
|
||||
assert listed is None
|
||||
|
|
|
|||
|
|
@ -963,6 +963,8 @@ async def test_get_tools_from_mcp_servers():
|
|||
user_api_key_auth=None,
|
||||
oauth2_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
catalog_auth_header=None,
|
||||
record_listing=True,
|
||||
):
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
|
|
@ -1998,6 +2000,7 @@ async def test_get_tools_for_single_server():
|
|||
client_ip=None,
|
||||
user_api_key_auth=None,
|
||||
proxy_logging_obj=ANY,
|
||||
record_listing=False,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -9,6 +9,7 @@ from types import SimpleNamespace
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp import ReadResourceResult, Resource
|
||||
|
|
@ -23,18 +24,20 @@ from mcp.types import (
|
|||
TextContent,
|
||||
TextResourceContents,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
MCPTransport,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool
|
||||
|
||||
|
||||
def test_mcp_available_on_sdk2():
|
||||
|
|
@ -85,9 +88,6 @@ def cleanup_mcp_global_state():
|
|||
yield
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def _call_tool_params(name, arguments=None):
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
|
|
@ -99,6 +99,7 @@ def _paged_params():
|
|||
|
||||
return PaginatedRequestParams()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx):
|
||||
"""Test that proxy_server_request body contains name and arguments"""
|
||||
|
|
@ -295,7 +296,9 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r
|
|||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
|
||||
with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger):
|
||||
result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"}))
|
||||
result = await mcp_server_tool_call(
|
||||
_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})
|
||||
)
|
||||
|
||||
assert result.is_error is True
|
||||
# The dedicated MCPUpstreamAuthError branch (not the generic Exception fallthrough) produces this
|
||||
|
|
@ -1167,20 +1170,32 @@ async def test_read_resource_preserves_content_metadata(_mcp_request_ctx, kind,
|
|||
else BlobResourceContents(uri=uri, blob="aGVsbG8=", mimeType="image/png", meta=metadata)
|
||||
)
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None))),
|
||||
patch.object(
|
||||
server,
|
||||
"get_or_extract_auth_context",
|
||||
AsyncMock(return_value=(caller, None, ["catalog"], None, None, None, None)),
|
||||
),
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream_server])),
|
||||
patch.object(operations.global_mcp_server_manager, "read_resource_from_server", AsyncMock(return_value=ReadResourceResult(contents=[content]))),
|
||||
patch.object(
|
||||
operations.global_mcp_server_manager,
|
||||
"read_resource_from_server",
|
||||
AsyncMock(return_value=ReadResourceResult(contents=[content])),
|
||||
),
|
||||
):
|
||||
result: Final = await server.read_resource(_mcp_request_ctx(), ReadResourceRequestParams(uri=uri))
|
||||
|
||||
assert result.model_dump(mode="json", by_alias=True, exclude_none=True) == {
|
||||
"cacheScope": "private", "resultType": "complete", "ttlMs": 0,
|
||||
"contents": [{
|
||||
"uri": uri,
|
||||
"mimeType": "text/plain" if kind == "text" else "image/png",
|
||||
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
|
||||
**({"_meta": metadata} if metadata is not None else {}),
|
||||
}],
|
||||
"cacheScope": "private",
|
||||
"resultType": "complete",
|
||||
"ttlMs": 0,
|
||||
"contents": [
|
||||
{
|
||||
"uri": uri,
|
||||
"mimeType": "text/plain" if kind == "text" else "image/png",
|
||||
"text" if kind == "text" else "blob": "hello world" if kind == "text" else "aGVsbG8=",
|
||||
**({"_meta": metadata} if metadata is not None else {}),
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -1674,7 +1689,9 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
|
|||
with (
|
||||
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
|
||||
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
|
||||
new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None),
|
||||
new=AsyncMock(
|
||||
return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None
|
||||
),
|
||||
),
|
||||
patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam
|
||||
"litellm.proxy._experimental.mcp_server.operations._list_mcp_tools",
|
||||
|
|
@ -1913,8 +1930,8 @@ async def test_streamable_http_session_manager_is_stateless():
|
|||
("DELETE", b"", False),
|
||||
),
|
||||
)
|
||||
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(_mcp_request_ctx,
|
||||
debug: bool, method: str, request_body: bytes, stateful: bool
|
||||
async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
|
||||
_mcp_request_ctx, debug: bool, method: str, request_body: bytes, stateful: bool
|
||||
) -> None:
|
||||
from starlette.requests import Request
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
|
@ -4056,7 +4073,8 @@ async def test_truncated_jsonrpc_response_with_nested_method_skips_lock(
|
|||
# parsed, with a nested "method" key in the first bytes to trip a flat
|
||||
# substring heuristic.
|
||||
response_prefix: Final = (
|
||||
'{"jsonrpc":"2.0","id":99,"' + response_field
|
||||
'{"jsonrpc":"2.0","id":99,"'
|
||||
+ response_field
|
||||
+ '":{"code":-32000,"message":"test","data":{"method":"GET","payload":"'
|
||||
).encode()
|
||||
response_body: Final = (
|
||||
|
|
@ -6564,8 +6582,12 @@ class TestGatewayCreateInitializationOptions:
|
|||
yield (None, None)
|
||||
|
||||
async def record_request(
|
||||
serving_server: object, read_stream: object, write_stream: object,
|
||||
*, lifespan_state: object, init_options: InitializationOptions,
|
||||
serving_server: object,
|
||||
read_stream: object,
|
||||
write_stream: object,
|
||||
*,
|
||||
lifespan_state: object,
|
||||
init_options: InitializationOptions,
|
||||
) -> None:
|
||||
captured["server_name"] = init_options.server_name
|
||||
|
||||
|
|
@ -6877,7 +6899,6 @@ async def test_probe_upstream_auth_surfaces_httpx_status_error():
|
|||
returning the response. The probe must catch that specifically (before the
|
||||
fail-open `except Exception`) so the auth check is not silently defeated.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth
|
||||
|
||||
|
|
@ -7412,7 +7433,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
|
|||
return_value=oauth_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7658,7 +7680,8 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
|
|||
return_value=alias_less_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7933,7 +7956,8 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req
|
|||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -7989,9 +8013,12 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
start_time = datetime.now(timezone.utc)
|
||||
litellm_logging_obj, _ = function_setup(
|
||||
|
|
@ -8042,6 +8069,348 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
assert litellm_logging_obj.model == "MCP: list_pets"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing_before_a_listing():
|
||||
"""A local-registry tools/call with no prior tools/list hands the pre-call hooks name and arguments
|
||||
only, as before this metadata existed, so a pre_mcp_call policy never scans a description the caller was
|
||||
not served. Once the caller has listed, the same call hands the entry that listing served."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
tool_name_to_description={"list_pets": "ADMIN DESC"},
|
||||
)
|
||||
schema = {"type": "object", "properties": {"limit": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-list_pets", description="List the pets", input_schema=schema, handler=lambda limit: "ok"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
pre_call_tool_check = AsyncMock(wraps=manager.pre_call_tool_check)
|
||||
|
||||
async def call() -> tuple[MCPTool | None, dict]:
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-list_pets",
|
||||
arguments={"limit": 10},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
return pre_call_tool_check.call_args.kwargs["tool"], proxy_logging.pre_call_hook.call_args.kwargs["data"]
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
never_listed_tool, never_listed_data = await call()
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
listed_tool, listed_data = await call()
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
assert never_listed_tool is None
|
||||
assert (never_listed_data.get("mcp_tool_description"), never_listed_data.get("mcp_input_schema")) == (None, None)
|
||||
assert listed_tool is not None and (listed_tool.description, listed_tool.input_schema) == ("ADMIN DESC", schema)
|
||||
assert (listed_data["mcp_tool_description"], listed_data["mcp_input_schema"]) == ("ADMIN DESC", schema)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_clients_saw():
|
||||
"""When tools/list pinned the schema and masked the description of an OpenAPI tool, the local-registry
|
||||
call path must hand the pre-call hooks that served entry, not the raw registry one."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
tool_name_to_description={"getpetbyid": "Find a SECRET pet"},
|
||||
)
|
||||
registry_schema = {"type": "object", "properties": {"petId": {"type": "integer"}, "dump_all": {"type": "boolean"}}}
|
||||
pinned_schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-getpetbyid",
|
||||
description="Find pet by ID",
|
||||
input_schema=registry_schema,
|
||||
handler=lambda petId: "ok",
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-getpetbyid",
|
||||
arguments={"petId": 1},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
handed_tool = pre_call_tool_check.call_args.kwargs["tool"]
|
||||
assert (handed_tool.description, handed_tool.input_schema) == ("Find a [MASKED] pet", pinned_schema), (
|
||||
"the pre-call policy must evaluate the entry tools/list served, not the raw registry entry"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entry():
|
||||
"""Two keys can be shown differently guarded OpenAPI catalogs. The call path must evaluate each key
|
||||
against the entry its own tools/list served, not the entry the most recent listing left behind."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
schema = {"type": "object", "properties": {"petId": {"type": "integer"}}}
|
||||
mcp_module.global_mcp_tool_registry.register_tool(
|
||||
name="petstore-getpetbyid", description="Find a SECRET pet", input_schema=schema, handler=lambda petId: "ok"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice")
|
||||
opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=guarded),
|
||||
)
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)],
|
||||
ListedToolsCaller(user_api_key_auth=opted_out),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
for caller in (guarded, opted_out):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="petstore-getpetbyid",
|
||||
arguments={"petId": 1},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=caller,
|
||||
)
|
||||
finally:
|
||||
mcp_module.global_mcp_tool_registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
handed = [call.kwargs["tool"].description for call in pre_call_tool_check.call_args_list]
|
||||
assert handed == ["Find a [MASKED] pet", "Find a SECRET pet"], (
|
||||
"each key's tools/call must be evaluated against the OpenAPI entry its own listing served"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_runs_the_longer_colliding_operation_and_hands_hooks_no_registry_metadata():
|
||||
"""An OpenAPI operation whose name starts with its own server prefix runs instead of the shorter one, and
|
||||
with no prior listing the pre-call hooks get name and arguments only, never either registry entry."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
registry = mcp_module.global_mcp_tool_registry
|
||||
registry.register_tool(name="petstore-get_pet", description="short", input_schema={}, handler=lambda: "short")
|
||||
registry.register_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
description="long",
|
||||
input_schema={"type": "object", "properties": {"petId": {"type": "integer"}}},
|
||||
handler=lambda: "long",
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice"),
|
||||
)
|
||||
finally:
|
||||
registry.unregister_tools_with_prefix("petstore-")
|
||||
|
||||
assert pre_call_tool_check.call_args.kwargs["tool"] is None
|
||||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation_named_after_a_listed_one():
|
||||
"""After the caller listed ``get_pet``, a call to the never-listed ``petstore-get_pet`` operation hands the
|
||||
pre-call hooks name and arguments only, not the listed sibling's description and schema."""
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_module
|
||||
|
||||
petstore = MCPServer(
|
||||
server_id="petstore-id",
|
||||
name="petstore",
|
||||
server_name="petstore",
|
||||
transport=MCPTransport.http,
|
||||
url=None,
|
||||
spec_path="https://example.com/petstore.yaml",
|
||||
)
|
||||
registry = mcp_module.global_mcp_tool_registry
|
||||
registry.register_tool(
|
||||
name="petstore-petstore-get_pet", description="long", input_schema={}, handler=lambda: "long"
|
||||
)
|
||||
manager = mcp_module.global_mcp_server_manager
|
||||
alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice")
|
||||
manager._record_listed_tools(
|
||||
petstore,
|
||||
[MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})],
|
||||
ListedToolsCaller(user_api_key_auth=alice),
|
||||
)
|
||||
pre_call_tool_check = AsyncMock(return_value={})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=petstore),
|
||||
patch.object(manager, "pre_call_tool_check", new=pre_call_tool_check),
|
||||
):
|
||||
result = await mcp_module.execute_mcp_tool(
|
||||
name="petstore-petstore-get_pet",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[petstore],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=alice,
|
||||
)
|
||||
finally:
|
||||
registry.unregister_tools_with_prefix("petstore-")
|
||||
manager._listed_tools_by_server_id.pop(petstore.server_id, None)
|
||||
|
||||
assert pre_call_tool_check.call_args.kwargs["tool"] is None
|
||||
assert result.content[0].text == "long"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_implicit_listing_before_the_first_call_hands_hooks_no_description():
|
||||
"""The listing tools/call runs on its own when this worker does not yet expose the tool is never served
|
||||
to the caller, so it leaves the caller's listed slot empty and the pre-call hooks still get name and
|
||||
arguments only, as on main."""
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = _never_listed_passthrough_server()
|
||||
manager.registry[server.server_id] = server
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
upstream = AsyncMock()
|
||||
upstream.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
|
||||
proxy_logging = _mock_mcp_proxy_logging()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
proxy_logging.during_call_hook = AsyncMock(return_value=None)
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Adds. FLAGWORD", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=upstream)),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging),
|
||||
):
|
||||
result = await mcp_operations.execute_mcp_tool(
|
||||
name="lazy_map-add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
raw_headers={"authorization": "Bearer caller-token"},
|
||||
)
|
||||
|
||||
assert fetch_tools.await_count == 1
|
||||
assert upstream.call_tool.await_count == 1
|
||||
assert result.content[0].text == "ok"
|
||||
hook_kwargs = proxy_logging._create_mcp_request_object_from_kwargs.call_args.args[0]
|
||||
assert (hook_kwargs["tool_description"], hook_kwargs["tool_input_schema"]) == (None, None)
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_pinnable_tool_catalog_records_no_listed_catalog_for_the_admin():
|
||||
"""The pin snapshot lists the raw upstream catalog, without the catalog guard or the admin's description
|
||||
overrides, so it must not become what the admin's own later tools/call is evaluated against."""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import fetch_pinnable_tool_catalog
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(
|
||||
server_id="pin-srv",
|
||||
name="pin_srv",
|
||||
transport=MCPTransport.http,
|
||||
url="https://up.example.com/mcp",
|
||||
tool_name_to_description={"add": "Admin wording"},
|
||||
)
|
||||
manager._listed_tools_by_server_id.pop(server.server_id, None)
|
||||
admin = UserAPIKeyAuth(api_key="sk-admin", user_id="admin")
|
||||
request = MagicMock()
|
||||
request.client.host = "10.1.2.3"
|
||||
request.headers = {"x-litellm-api-key": "sk-admin"}
|
||||
fetch_tools = AsyncMock(
|
||||
return_value=[MCPTool(name="add", description="Upstream wording", inputSchema={"type": "object"})]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", new=fetch_tools),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())),
|
||||
):
|
||||
snapshot = await fetch_pinnable_tool_catalog(server, request, admin)
|
||||
|
||||
assert snapshot == {"add": PinnedMCPTool(description="Upstream wording", input_schema={"type": "object"})}
|
||||
assert server.server_id not in manager._listed_tools_by_server_id
|
||||
assert manager.get_listed_tool(server, "add", ListedToolsCaller(user_api_key_auth=admin)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requested_server():
|
||||
"""A prefixed REST name that resolves to no tool must still dispatch to the server_id.
|
||||
|
|
@ -8098,7 +8467,8 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste
|
|||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_operations, "_handle_managed_mcp_tool",
|
||||
mcp_operations,
|
||||
"_handle_managed_mcp_tool",
|
||||
new=fake_handle_managed_mcp_tool,
|
||||
),
|
||||
patch.object(
|
||||
|
|
@ -8597,7 +8967,9 @@ class TestMCPMetaTraceCarrier:
|
|||
|
||||
assert _mcp_meta_trace_carrier(None) is None
|
||||
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
|
||||
only_progress = CallToolRequestParams.model_validate({"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False).meta
|
||||
only_progress = CallToolRequestParams.model_validate(
|
||||
{"name": "t", "_meta": {"progressToken": "p1"}}, by_name=False
|
||||
).meta
|
||||
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None
|
||||
|
||||
|
||||
|
|
@ -10394,7 +10766,9 @@ async def test_mcp_origin_admission_precedes_authentication(
|
|||
patch("litellm.proxy.proxy_server.origins", allowed_origins),
|
||||
patch.object(server, "extract_mcp_auth_context", authenticate),
|
||||
):
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.ASGITransport(app=server.app), base_url="http://gateway"
|
||||
) as client:
|
||||
response: Final = await client.request(method, path, headers=(*session_headers, *origin_headers))
|
||||
|
||||
assert response.status_code == expected_status
|
||||
|
|
@ -10477,12 +10851,15 @@ async def test_streamable_http_rejects_modern_protocol_version(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("handler_name,field", [
|
||||
("handle_list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
])
|
||||
@pytest.mark.parametrize(
|
||||
"handler_name,field",
|
||||
[
|
||||
("handle_list_tools", "tools"),
|
||||
("list_prompts", "prompts"),
|
||||
("list_resources", "resources"),
|
||||
("list_resource_templates", "resource_templates"),
|
||||
],
|
||||
)
|
||||
async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field):
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
|
|
@ -10500,7 +10877,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
auth = UserAPIKeyAuth(user_id="denied-caller")
|
||||
denial = HTTPException(status_code=403, detail="scope denied")
|
||||
logger = MagicMock()
|
||||
logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None)
|
||||
logger.post_call_failure_hook = AsyncMock(
|
||||
side_effect=RuntimeError("log unavailable") if failure_hook_raises else None
|
||||
)
|
||||
upstream = AsyncMock()
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)),
|
||||
|
|
@ -10509,7 +10888,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream),
|
||||
):
|
||||
with pytest.raises(HTTPException) as rejected:
|
||||
await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True)
|
||||
await operations._get_tools_from_mcp_servers(
|
||||
user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True
|
||||
)
|
||||
assert rejected.value is denial
|
||||
upstream.assert_not_awaited()
|
||||
logger.post_call_failure_hook.assert_awaited_once()
|
||||
|
|
@ -10521,7 +10902,9 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
@pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/")))
|
||||
@pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS))
|
||||
async def test_legacy_sse_mount_emits_message_endpoint(
|
||||
prefix: str, suffix: str, opening_protocol: str | None,
|
||||
prefix: str,
|
||||
suffix: str,
|
||||
opening_protocol: str | None,
|
||||
) -> None:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.routing import Mount
|
||||
|
|
@ -10586,16 +10969,20 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
return (await messages.get())["status"]
|
||||
|
||||
if opening_protocol is not None:
|
||||
discover: Final = json.dumps({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}},
|
||||
}).encode()
|
||||
discover: Final = json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {
|
||||
"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}
|
||||
},
|
||||
}
|
||||
).encode()
|
||||
assert await post(discover) == 202
|
||||
discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
|
||||
discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0])
|
||||
|
|
@ -10628,7 +11015,16 @@ async def test_legacy_sse_mount_emits_message_endpoint(
|
|||
patch.object(
|
||||
mcp_server,
|
||||
"extract_mcp_auth_context",
|
||||
AsyncMock(return_value=(post_auth, None, [marker], {marker: {"Authorization": marker}}, {"Authorization": marker}, {"x-request-marker": marker})),
|
||||
AsyncMock(
|
||||
return_value=(
|
||||
post_auth,
|
||||
None,
|
||||
[marker],
|
||||
{marker: {"Authorization": marker}},
|
||||
{"Authorization": marker},
|
||||
{"x-request-marker": marker},
|
||||
)
|
||||
),
|
||||
),
|
||||
patch.object(mcp_server.operations, "_get_tools_from_mcp_servers", listing),
|
||||
):
|
||||
|
|
@ -10677,7 +11073,11 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
|
|||
dispatched = AsyncMock(return_value=expected)
|
||||
auth = UserAPIKeyAuth(user_id="discover-caller")
|
||||
with (
|
||||
patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None))),
|
||||
patch.object(
|
||||
server,
|
||||
"get_or_extract_auth_context",
|
||||
AsyncMock(return_value=(auth, None, ["allowed"], None, None, None, None)),
|
||||
),
|
||||
patch.object(server.operations.GatewayOperations, "execute", dispatched),
|
||||
):
|
||||
result = await server.discover(_mcp_request_ctx(), RequestParams())
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from mcp.types import Tool
|
|||
import litellm
|
||||
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import (
|
||||
AGENT_SEARCH_TOOL_NAME,
|
||||
MCP_TOOL_CALL_TOOL_NAME,
|
||||
|
|
@ -32,12 +33,14 @@ from litellm.proxy._experimental.mcp_server.tool_search import (
|
|||
ToolSearchResult,
|
||||
coerce_top_k,
|
||||
get_virtual_tool_definitions,
|
||||
handle_mcp_tool_search,
|
||||
search_mcp_tools,
|
||||
search_tools,
|
||||
)
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.semantic_text_index import EmbeddingFailed, SemanticTextIndex, Vector
|
||||
from litellm.types.mcp import MCPToolSearchSettings
|
||||
from litellm.types.mcp import MCPToolSearchSettings, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _make_tools(specs: list[tuple[str, str]]) -> tuple[Tool, ...]:
|
||||
|
|
@ -1353,3 +1356,33 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N
|
|||
assert exc_info.value.status_code == 403
|
||||
assert "MCP server 'github'" in exc_info.value.detail["error"]
|
||||
assert "agent 'agent-123'" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_search_leaves_the_listed_tools_slot_empty(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""The search lists the whole catalog but serves only its hits, so the listing must not fill the
|
||||
caller's listed-tools slot: a later call to a tool the search never returned is not a listed tool."""
|
||||
monkeypatch.setattr(litellm, "mcp_tool_search", None)
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="search-slot", name="search-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-search-slot", user_id="searcher")
|
||||
upstream = [
|
||||
Tool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
Tool(name="delete_note", description="Delete a note", inputSchema={"type": "object"}),
|
||||
]
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
result = await handle_mcp_tool_search(query="echo", top_k=1, user_api_key_dict=user)
|
||||
caller = ListedToolsCaller(user_api_key_auth=user)
|
||||
listed = [manager.get_listed_tool(server, tool.name, caller) for tool in upstream]
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert result.is_error is False
|
||||
assert [hit["name"] for hit in json.loads(result.content[0].text)] == ["search-slot-echo"]
|
||||
assert listed == [None, None]
|
||||
|
|
|
|||
|
|
@ -42,9 +42,12 @@ async def test_openapi_local_tool_runs_pre_call_tool_check():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
|
@ -125,9 +128,12 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises():
|
|||
fake_server.server_name = "openapi-petstore"
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
fake_server.tool_name_to_description = None
|
||||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "delete_pet"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(
|
||||
side_effect=HTTPException(status_code=403, detail="not allowed")
|
||||
|
|
@ -190,6 +196,8 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable():
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_pets"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
|
@ -274,6 +282,8 @@ async def test_openapi_local_tool_injects_resolved_oauth_token():
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "get_values"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
captured: dict = {}
|
||||
|
||||
async def handle_local(_name, _arguments, _wire_compat):
|
||||
|
|
@ -620,6 +630,8 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc
|
|||
if dispatch_arm == "local_registry":
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_reports"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
with (
|
||||
patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server),
|
||||
patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool),
|
||||
|
|
@ -691,6 +703,8 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st
|
|||
|
||||
fake_tool = MagicMock()
|
||||
fake_tool.name = "list_reports"
|
||||
fake_tool.description = "test tool"
|
||||
fake_tool.input_schema = {"type": "object"}
|
||||
fake_tool.handler = raising_handler
|
||||
server = MCPServer(
|
||||
server_id="srv-openapi",
|
||||
|
|
|
|||
|
|
@ -1,15 +1,101 @@
|
|||
import asyncio
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server import rest_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller
|
||||
from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
class _CatalogHookCapture(CustomLogger):
|
||||
data: dict[str, object] | None = None
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, data: dict[str, object], call_type: str
|
||||
) -> None:
|
||||
if call_type == "call_mcp_tool":
|
||||
self.data = data.copy()
|
||||
|
||||
|
||||
async def _served_catalog_tool() -> str:
|
||||
return "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("surface", ["mcp", "rest"])
|
||||
@pytest.mark.parametrize("restriction", ["key", "server"])
|
||||
async def test_listing_records_only_tools_the_caller_received(
|
||||
monkeypatch: pytest.MonkeyPatch, surface: str, restriction: str
|
||||
) -> None:
|
||||
manager: Final = operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="served-catalog", name="served-catalog", transport=MCPTransport.http,
|
||||
spec_path="/catalog.yaml", allow_all_keys=True,
|
||||
allowed_tools=["echo"] if restriction == "server" else None,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="sk-served-catalog", user_id="lister",
|
||||
object_permission={
|
||||
"object_permission_id": "served-permission",
|
||||
"mcp_servers": [server.server_id],
|
||||
"mcp_tool_permissions": {server.server_id: ["echo"]} if restriction == "key" else None,
|
||||
},
|
||||
)
|
||||
monkeypatch.setitem(manager.registry, server.server_id, server)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "status", server.server_id)
|
||||
monkeypatch.setitem(manager.tool_name_to_mcp_server_name_mapping, "served-catalog-status", server.server_id)
|
||||
capture: Final = _CatalogHookCapture()
|
||||
monkeypatch.setattr(litellm, "callbacks", [capture])
|
||||
for name in ("echo", "status"):
|
||||
global_mcp_tool_registry.register_tool(
|
||||
name=f"served-catalog-{name}", description=f"{name} description",
|
||||
input_schema={"type": "object"}, handler=_served_catalog_tool,
|
||||
)
|
||||
try:
|
||||
if surface == "mcp":
|
||||
listing: Final = await operations._list_mcp_tools(
|
||||
user_api_key_auth=auth, mcp_servers=[server.server_id], record_listing=True,
|
||||
)
|
||||
assert [tool.name for tool in listing.tools] == ["served-catalog-echo"]
|
||||
else:
|
||||
rest_listing: Final = await rest_endpoints._get_tools_for_single_server(
|
||||
server, None, user_api_key_auth=auth,
|
||||
)
|
||||
assert [tool.name for tool in rest_listing] == ["echo"]
|
||||
granted: Final = auth.model_copy(update={"object_permission": None})
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=granted)
|
||||
assert manager.get_listed_tool(server, "status", caller) is None
|
||||
served: Final = manager.get_listed_tool(server, "echo", caller)
|
||||
assert served is not None
|
||||
assert (served.description, served.input_schema) == ("echo description", {"type": "object"})
|
||||
server.allowed_tools = None
|
||||
result: Final = await manager.call_tool(
|
||||
server_name=server.server_id, name="status", arguments={}, user_api_key_auth=granted,
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()),
|
||||
)
|
||||
assert result.is_error is False
|
||||
assert capture.data is not None
|
||||
assert capture.data["messages"] == [{"role": "user", "content": "Tool: status\nArguments: {}"}]
|
||||
assert (capture.data.get("mcp_tool_description"), capture.data.get("mcp_input_schema")) == (None, None)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
global_mcp_tool_registry.unregister_tools_with_prefix("served-catalog-")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog):
|
||||
from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user
|
||||
|
|
@ -665,3 +751,30 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
|
|||
)
|
||||
assert result.tools == []
|
||||
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
|
||||
assert listing.await_args.kwargs["record_listing"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("listing_kwargs", "recorded"), [({}, False), ({"record_listing": True}, True)])
|
||||
async def test_list_mcp_tools_records_the_catalog_only_when_asked(
|
||||
listing_kwargs: dict[str, bool], recorded: bool
|
||||
) -> None:
|
||||
"""The aggregate listing fills the caller's listed-tools slot only when asked: a listing an internal
|
||||
caller never serves must not hand a later tools/call a description the caller never saw."""
|
||||
manager = operations.global_mcp_server_manager
|
||||
server = MCPServer(server_id="listing-slot", name="listing-slot", transport=MCPTransport.http, url="http://slot")
|
||||
user = UserAPIKeyAuth(api_key="sk-listing-slot", user_id="lister")
|
||||
upstream = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})]
|
||||
with (
|
||||
patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
):
|
||||
try:
|
||||
listing = await operations._list_mcp_tools(user_api_key_auth=user, **listing_kwargs)
|
||||
listed = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user))
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
|
||||
assert (listed is not None) is recorded
|
||||
|
|
|
|||
|
|
@ -109,25 +109,6 @@ def set_salt_key(monkeypatch):
|
|||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_constants_module():
|
||||
"""Reset constants module to ensure clean state before each test"""
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
# Reload modules before test
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
yield
|
||||
|
||||
# Reload modules after test to clean up
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valid_sso_user_defined_values():
|
||||
return LiteLLM_UserTable(
|
||||
|
|
@ -875,19 +856,10 @@ def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values
|
|||
|
||||
|
||||
def test_get_cli_jwt_auth_token_custom_expiration(valid_sso_user_defined_values, monkeypatch):
|
||||
"""Test generating CLI JWT token with custom expiration via environment variable"""
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
"""Test generating a CLI JWT token with custom expiration via the configured constant"""
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
# Set custom expiration to 48 hours
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
|
||||
|
||||
# Reload the constants module to pick up the new env var
|
||||
importlib.reload(constants)
|
||||
# Also reload auth_checks to pick up the new constant value
|
||||
importlib.reload(auth_checks)
|
||||
monkeypatch.setattr(auth_checks, "CLI_JWT_EXPIRATION_HOURS", 48)
|
||||
|
||||
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
||||
|
||||
|
|
|
|||
|
|
@ -7321,15 +7321,10 @@ async def test_expired_cli_session_token_is_rejected(monkeypatch):
|
|||
on the shared validation path, not only for DB-backed keys."""
|
||||
monkeypatch.delenv("EXPERIMENTAL_UI_LOGIN", raising=False)
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-cli-test")
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "-1")
|
||||
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
monkeypatch.setattr(auth_checks, "CLI_JWT_EXPIRATION_HOURS", -1)
|
||||
|
||||
user_info = LiteLLM_UserTable(
|
||||
user_id="cli-admin",
|
||||
|
|
@ -7346,22 +7341,17 @@ async def test_expired_cli_session_token_is_rejected(monkeypatch):
|
|||
mock_request.headers = {"authorization": f"Bearer {cli_token}"}
|
||||
mock_request.query_params = {}
|
||||
|
||||
try:
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {cli_token}",
|
||||
)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await user_api_key_auth(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {cli_token}",
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.expired_key
|
||||
finally:
|
||||
monkeypatch.delenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", raising=False)
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
assert exc_info.value.type == ProxyErrorTypes.expired_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -346,6 +346,29 @@ class TestAllowFlow:
|
|||
assert evaluate_call.json["conversationId"] == "sess-123"
|
||||
assert evaluate_call.json["agentId"] == "my-agent-key"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_evaluate_payload_includes_listed_tool_metadata(self):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
schema: Final = {"type": "object", "properties": {"to": {"type": "string"}}, "required": ["to"]}
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description="Send an email", mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {
|
||||
"name": "send_email",
|
||||
"description": "Send an email",
|
||||
"inputSchema": schema,
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("description", "schema"),
|
||||
[(None, None), ("", None), (None, ["not", "a", "schema"]), (42, "type: object")],
|
||||
)
|
||||
async def test_evaluate_payload_omits_missing_or_malformed_tool_metadata(self, description, schema):
|
||||
handler: Final = FakeHandler([_token_response(), _allow_response()])
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data(mcp_tool_description=description, mcp_input_schema=schema))
|
||||
assert handler.calls[1].json["tool"] == {"name": "send_email"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_mcp_call_type_skipped(self):
|
||||
handler: Final = FakeHandler([])
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from typing import Any, Dict, Final, List, Optional
|
|||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
|
@ -21,6 +22,7 @@ from litellm.caching.in_memory_cache import InMemoryCache
|
|||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
PARALLEL_REQUEST_SLOT_TTL_SECONDS,
|
||||
ParallelSlotAcquisition,
|
||||
|
|
@ -39,6 +41,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
|||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPPreCallRequestObject
|
||||
from litellm.types.utils import (
|
||||
EmbeddingResponse,
|
||||
ModelResponse,
|
||||
|
|
@ -108,6 +111,159 @@ def test_api_key_descriptor_applies_budget_throttle(
|
|||
assert api_key_descriptor["rate_limit"]["tokens_per_unit"] == expected_tpm
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_does_not_change_admission_or_reserved_tokens(
|
||||
description: str | None, arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
|
||||
request: Final = MCPPreCallRequestObject(
|
||||
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
|
||||
)
|
||||
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
|
||||
messages: Final = data["messages"]
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-mcp-description-reservation"), tpm_limit=64)
|
||||
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert stash.reserved_tokens == 25
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
|
||||
)
|
||||
== 25
|
||||
)
|
||||
assert data["messages"] is messages
|
||||
assert data.get("mcp_tool_description") == description
|
||||
assert data["mcp_input_schema"] == schema
|
||||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"description", [None, "Gateway metadata, not caller input. " * 100], ids=["unlisted", "listed"]
|
||||
)
|
||||
@pytest.mark.parametrize("itpm_limit,otpm_limit", [(64, 4096), (4096, 64), (4096, 4096)])
|
||||
@pytest.mark.parametrize("arguments_rewritten", [False, True])
|
||||
async def test_mcp_description_preserves_project_input_and_output_reservations(
|
||||
description: str | None, itpm_limit: int, otpm_limit: int,
|
||||
arguments_rewritten: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
logger: Final = ProxyLogging(user_api_key_cache=UserApiKeyCache())
|
||||
schema: Final = {"type": "object", "properties": {"q": {"type": "string", "description": "Schema text " * 100}}}
|
||||
request: Final = MCPPreCallRequestObject(
|
||||
tool_name="echo", arguments={"q": "hello"}, tool_description=description, tool_input_schema=schema
|
||||
)
|
||||
data: Final = TypeAdapter(dict[str, object]).validate_python(logger._convert_mcp_to_llm_format(request, {}))
|
||||
messages: Final = data["messages"]
|
||||
base_data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "Tool: echo\nArguments: {'q': 'hello'}"}]
|
||||
}
|
||||
expected_input: Final = handler._estimate_precise_input_tokens(base_data, "mcp-tool-call", "call_mcp_tool")
|
||||
expected_output: Final = handler.no_max_tokens_output_floor(otpm_limit)
|
||||
expected_combined: Final = handler._estimate_tokens_for_request(
|
||||
base_data, min_configured_tpm_limit=4096, call_type="call_mcp_tool"
|
||||
)
|
||||
caller: Final = UserAPIKeyAuth(
|
||||
api_key=hash_token("sk-mcp-project-reservation"),
|
||||
tpm_limit=4096,
|
||||
project_id="mcp-project-reservation",
|
||||
project_metadata={
|
||||
"model_itpm_limit": {"mcp-tool-call": itpm_limit},
|
||||
"model_otpm_limit": {"mcp-tool-call": otpm_limit},
|
||||
},
|
||||
)
|
||||
|
||||
if arguments_rewritten:
|
||||
data["mcp_arguments"] = {"q": "Transformed arguments " * 100}
|
||||
monkeypatch.setattr(litellm, "callbacks", [handler])
|
||||
await logger.pre_call_hook(user_api_key_dict=caller, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) == (
|
||||
expected_combined,
|
||||
expected_input,
|
||||
expected_output,
|
||||
)
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys(
|
||||
"model_per_project_itpm", f"{caller.project_id}:mcp-tool-call", "tokens"
|
||||
),
|
||||
local_only=True,
|
||||
)
|
||||
== expected_input
|
||||
)
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys(
|
||||
"model_per_project_otpm", f"{caller.project_id}:mcp-tool-call", "tokens"
|
||||
),
|
||||
local_only=True,
|
||||
)
|
||||
== expected_output
|
||||
)
|
||||
assert data["messages"] is messages
|
||||
assert data.get("mcp_tool_description") == description
|
||||
assert data["mcp_input_schema"] == schema
|
||||
assert messages == [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "Tool: echo\nArguments: {'q': 'hello'}",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_llm_tpm_estimation_still_counts_messages_with_mcp_metadata() -> None:
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
data: Final[dict[str, object]] = {
|
||||
"messages": [{"role": "user", "content": "x" * 400}],
|
||||
"max_tokens": 1,
|
||||
"mcp_tool_name": "echo",
|
||||
"mcp_arguments": {},
|
||||
}
|
||||
assert handler._estimate_tokens_for_request(data, call_type="acompletion") == 101
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unconverted_mcp_request_keeps_its_reservation() -> None:
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
caller: Final = UserAPIKeyAuth(api_key=hash_token("sk-raw-mcp-request"), tpm_limit=64)
|
||||
data: Final[dict[str, object]] = {"name": "echo", "arguments": {"q": "hello"}, "server_id": "fixture"}
|
||||
|
||||
await handler.async_pre_call_hook(user_api_key_dict=caller, cache=cache, data=data, call_type="call_mcp_tool")
|
||||
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert stash.reserved_tokens == 16
|
||||
assert (
|
||||
await cache.async_get_cache(
|
||||
key=handler.create_rate_limit_keys("api_key", caller.api_key, "tokens"), local_only=True
|
||||
)
|
||||
== 16
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.flaky(reruns=3)
|
||||
@pytest.mark.asyncio
|
||||
async def test_sliding_window_rate_limit_v3(monkeypatch, time_controller):
|
||||
|
|
|
|||
|
|
@ -21,6 +21,13 @@ def test_private_registry_override_keeps_its_exact_digest(monkeypatch: pytest.Mo
|
|||
assert worker_image() == image
|
||||
|
||||
|
||||
def test_source_build_uses_the_separate_development_package(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
tag: Final = "sha-" + "a" * 40
|
||||
monkeypatch.setenv("LITELLM_RELEASE_TAG", tag)
|
||||
monkeypatch.delenv("LENS_WORKER_IMAGE", raising=False)
|
||||
assert worker_image() == f"ghcr.io/berriai/litellm-lens-worker-dev:{tag}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"installed,expected",
|
||||
(("1.2.3", "v1.2.3"), ("1.2.3rc4", "v1.2.3-rc.4"), ("1.2.3.dev5", "v1.2.3-dev.5")),
|
||||
|
|
|
|||
|
|
@ -1,52 +1,45 @@
|
|||
import os
|
||||
from typing import Final
|
||||
|
||||
import uvicorn
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
load_dotenv()
|
||||
|
||||
# Set the SERVER_ROOT_PATH environment variable to match the custom mount path
|
||||
os.environ["SERVER_ROOT_PATH"] = "/my-custom-path"
|
||||
|
||||
from litellm.proxy.proxy_server import app as litellm_app
|
||||
from litellm.proxy.proxy_server import proxy_startup_event
|
||||
|
||||
# Create main FastAPI app
|
||||
app = FastAPI(title="Custom LiteLLM Server", lifespan=proxy_startup_event)
|
||||
|
||||
# Add CORS middleware
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
custom_path = "/my-custom-path"
|
||||
|
||||
# Mount LiteLLM app at /litellm
|
||||
app.mount(custom_path, litellm_app)
|
||||
|
||||
|
||||
# Default route at /
|
||||
@app.get("/")
|
||||
async def root():
|
||||
return {
|
||||
"message": "Welcome to the API Gateway",
|
||||
"litellm_endpoint": f"{custom_path}",
|
||||
}
|
||||
def build_app() -> FastAPI:
|
||||
load_dotenv()
|
||||
os.environ["SERVER_ROOT_PATH"] = "/my-custom-path"
|
||||
|
||||
from litellm.proxy.proxy_server import app as litellm_app
|
||||
from litellm.proxy.proxy_server import proxy_startup_event
|
||||
|
||||
# Health check endpoint
|
||||
@app.get("/health")
|
||||
async def health_check():
|
||||
return {"status": "healthy"}
|
||||
app: Final = FastAPI(title="Custom LiteLLM Server", lifespan=proxy_startup_event)
|
||||
custom_path: Final = "/my-custom-path"
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.mount(custom_path, litellm_app)
|
||||
|
||||
@app.get("/")
|
||||
async def root() -> dict[str, str]:
|
||||
return {
|
||||
"message": "Welcome to the API Gateway",
|
||||
"litellm_endpoint": custom_path,
|
||||
}
|
||||
|
||||
@app.get("/health")
|
||||
async def health_check() -> dict[str, str]:
|
||||
return {"status": "healthy"}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run the server on port 8000
|
||||
uvicorn.run(app, host="0.0.0.0", port=4000, log_level="info")
|
||||
uvicorn.run(build_app(), host="0.0.0.0", port=4000, log_level="info")
|
||||
|
|
|
|||
|
|
@ -403,6 +403,26 @@ def test_create_mcp_request_object_from_kwargs_full(proxy_logging, make_user_api
|
|||
assert snapshot == {"tool_name": "calc", "arguments": {"x": 1}, "server_name": "math", "auth_user_id": "u-1"}
|
||||
|
||||
|
||||
def test_mcp_tool_metadata_flows_from_kwargs_to_synthetic_data(proxy_logging):
|
||||
schema = {"type": "object", "properties": {"x": {"type": "integer"}}}
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(
|
||||
kwargs={
|
||||
"name": "calc",
|
||||
"arguments": {"x": 1},
|
||||
"tool_description": "Adds numbers",
|
||||
"tool_input_schema": schema,
|
||||
}
|
||||
)
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
|
||||
assert (out["mcp_tool_description"], out["mcp_input_schema"]) == ("Adds numbers", schema)
|
||||
|
||||
|
||||
def test_mcp_tool_metadata_absent_when_tool_was_never_listed(proxy_logging):
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={"name": "calc", "arguments": {}})
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=obj, kwargs={})
|
||||
assert "mcp_tool_description" not in out and "mcp_input_schema" not in out
|
||||
|
||||
|
||||
def test_create_mcp_request_object_from_kwargs_empty(proxy_logging):
|
||||
obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs={})
|
||||
snapshot = {
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import types
|
||||
from typing import Any, Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from typing import Any, Final, Literal, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -12,15 +13,26 @@ from mcp.types import CallToolResult, TextContent
|
|||
from mcp.types import Tool as MCPTool
|
||||
from openai.types.responses.tool_param import Mcp
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.responses import main as responses_main
|
||||
from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
from litellm.types.responses.main import OutputFunctionToolCall
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
|
||||
|
||||
|
||||
class _DummyMCPResult:
|
||||
|
|
@ -1310,6 +1322,167 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
|
|||
assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("real_listing", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
("allowed_tools", "expected_names"),
|
||||
[
|
||||
([], ["responses_slot-echo", "responses_slot-status"]),
|
||||
(["echo"], ["responses_slot-echo"]),
|
||||
(["responses_slot-echo"], ["responses_slot-echo"]),
|
||||
(["absent"], []),
|
||||
],
|
||||
)
|
||||
async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
|
||||
monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool
|
||||
) -> None:
|
||||
manager: Final = mcp_operations.global_mcp_server_manager
|
||||
server: Final = MCPServer(
|
||||
server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http
|
||||
)
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder")
|
||||
upstream: Final = [
|
||||
MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}),
|
||||
MCPTool(name="status", description="Report status", inputSchema={"type": "object"}),
|
||||
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
|
||||
]
|
||||
fake_manager: Final = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
|
||||
fake_manager,
|
||||
)
|
||||
with (
|
||||
patch.dict(manager.tool_name_to_mcp_server_name_mapping),
|
||||
patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])),
|
||||
patch.object(manager, "_create_mcp_client", AsyncMock(return_value=object())),
|
||||
patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)),
|
||||
):
|
||||
try:
|
||||
if real_listing:
|
||||
await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True)
|
||||
caller: Final = ListedToolsCaller(user_api_key_auth=user)
|
||||
before: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert bool(before) is real_listing
|
||||
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy/mcp/responses-slot",
|
||||
"allowed_tools": allowed_tools,
|
||||
}
|
||||
],
|
||||
)
|
||||
recorded: Final = {
|
||||
tool.name: (listed.description, listed.input_schema)
|
||||
for tool in upstream
|
||||
if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None
|
||||
}
|
||||
assert recorded == before
|
||||
assert (
|
||||
manager.get_listed_tool(
|
||||
server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller"))
|
||||
)
|
||||
is None
|
||||
)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
|
||||
assert [tool.name for tool in tools] == expected_names
|
||||
|
||||
|
||||
class _BridgeMetadataGuardrail(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True)
|
||||
self.calls: tuple[tuple[object, object], ...] = ()
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Logging | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if request_data.get("mcp_arguments") == {"probe": "bridge"}:
|
||||
self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),)
|
||||
return inputs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
manager: Final = MCPServerManager()
|
||||
server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream")
|
||||
manager.registry = {server.server_id: server}
|
||||
user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user")
|
||||
upstream: Final = [
|
||||
MCPTool(
|
||||
name="echo",
|
||||
description="Echo text",
|
||||
inputSchema={"type": "object", "properties": {"text": {"type": "string"}}},
|
||||
),
|
||||
MCPTool(name="status", description="Read status", inputSchema={"type": "object"}),
|
||||
]
|
||||
client: Final = AsyncMock()
|
||||
client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")])
|
||||
manager._create_mcp_client = AsyncMock(return_value=client)
|
||||
manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream)
|
||||
guardrail: Final = _BridgeMetadataGuardrail()
|
||||
logger: Final = ProxyLogging(user_api_key_cache=DualCache())
|
||||
monkeypatch.setattr(litellm, "callbacks", [guardrail])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger)
|
||||
monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]))
|
||||
monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager)
|
||||
first_listed: Final = asyncio.Event()
|
||||
second_listed: Final = asyncio.Event()
|
||||
|
||||
async def bridge(name: str, first: bool) -> None:
|
||||
if not first:
|
||||
await first_listed.wait()
|
||||
tools, server_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=user,
|
||||
mcp_tools_with_litellm_proxy=[
|
||||
{"type": "mcp", "server_url": "litellm_proxy/mcp/bridge", "allowed_tools": [name]}
|
||||
],
|
||||
)
|
||||
(first_listed if first else second_listed).set()
|
||||
await second_listed.wait()
|
||||
result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map=server_map,
|
||||
tool_calls=[
|
||||
{"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name}
|
||||
],
|
||||
user_api_key_auth=user,
|
||||
served_tools=tools,
|
||||
)
|
||||
assert [entry["result"] for entry in result] == ["ok"]
|
||||
|
||||
try:
|
||||
await asyncio.gather(bridge("echo", True), bridge("status", False))
|
||||
assert sorted(guardrail.calls, key=str) == sorted(
|
||||
((tool.description, tool.input_schema) for tool in upstream), key=str
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (None, None)
|
||||
await manager._get_tools_from_server(
|
||||
server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True
|
||||
)
|
||||
await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger)
|
||||
assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema)
|
||||
finally:
|
||||
manager._drop_listed_tools(server.server_id)
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
|
||||
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
|
||||
return types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue