mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(proxy): embed enterprise LiteAdmin MCP in LiteLLM images (#44610)
This commit is contained in:
parent
739192ec19
commit
9062fd3931
22 changed files with 1602 additions and 72 deletions
39
.github/workflows/image-scan.yml
vendored
39
.github/workflows/image-scan.yml
vendored
|
|
@ -9,10 +9,12 @@ on:
|
|||
paths:
|
||||
- Dockerfile
|
||||
- docker/Dockerfile.non_root
|
||||
- docker/Dockerfile.database
|
||||
- migrations/Dockerfile
|
||||
- migrations/run.py
|
||||
- gateway/Dockerfile
|
||||
- gateway/main.py
|
||||
- gateway/routes/allowlist.py
|
||||
- backend/Dockerfile
|
||||
- backend/main.py
|
||||
- deploy/lens/**
|
||||
|
|
@ -21,6 +23,10 @@ on:
|
|||
- docker/component_entrypoint.sh
|
||||
- docker/entrypoint.sh
|
||||
- litellm/proxy/prisma_migration.py
|
||||
- litellm/proxy/admin_mcp.py
|
||||
- litellm/proxy/proxy_server.py
|
||||
- backend/routes/allowlist.py
|
||||
- pyproject.toml
|
||||
- litellm-proxy-extras/**
|
||||
- tests/proxy_migration_tests/**
|
||||
- uv.lock
|
||||
|
|
@ -142,6 +148,9 @@ jobs:
|
|||
- name: Build runtime image
|
||||
run: docker build -f docker/Dockerfile.non_root -t litellm-image-scan:${{ github.sha }} .
|
||||
|
||||
- name: Tag the cached builder for Admin MCP schema setup
|
||||
run: docker build --target builder -f docker/Dockerfile.non_root -t litellm-admin-mcp-schema:${{ github.sha }} .
|
||||
|
||||
# The prisma bake must migrate a fresh DB with no egress as an arbitrary
|
||||
# non-root uid (OpenShift restricted-v2 / air-gapped / readOnlyRootFilesystem).
|
||||
# `docker run` as the default uid with network hides a broken bake because
|
||||
|
|
@ -155,9 +164,10 @@ jobs:
|
|||
- name: Verify offline migration as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-image-scan:${{ github.sha }}
|
||||
LITELLM_ADMIN_MCP_SCHEMA_IMAGE: litellm-admin-mcp-schema:${{ github.sha }}
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py tests/proxy_migration_tests/test_image_admin_mcp.py -v
|
||||
|
||||
# Scans the whole shipped artifact: OS/apk plus every language package
|
||||
# baked into the image, including ones no lockfile declares (e.g. prisma's
|
||||
|
|
@ -176,21 +186,28 @@ jobs:
|
|||
--output table
|
||||
|
||||
runtime-image:
|
||||
name: runtime-image
|
||||
name: runtime-image (${{ matrix.dockerfile }})
|
||||
runs-on: ubuntu-latest
|
||||
if: >-
|
||||
github.event_name != 'pull_request' ||
|
||||
github.event.pull_request.head.repo.full_name == github.repository
|
||||
timeout-minutes: 30
|
||||
timeout-minutes: 45
|
||||
permissions:
|
||||
contents: read
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
dockerfile: [Dockerfile, docker/Dockerfile.database]
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Build runtime image
|
||||
run: docker build --build-arg LITELLM_RELEASE_TAG=v0.0.0-lens-ci -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} .
|
||||
run: docker build --build-arg LITELLM_RELEASE_TAG=v0.0.0-lens-ci -f "${{ matrix.dockerfile }}" -t litellm-runtime-scan:${{ github.sha }} .
|
||||
|
||||
- name: Tag the cached builder for Admin MCP schema setup
|
||||
run: docker build --target builder -f "${{ matrix.dockerfile }}" -t litellm-admin-mcp-schema:${{ github.sha }} .
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
|
|
@ -200,11 +217,13 @@ jobs:
|
|||
- name: Verify offline migration as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
|
||||
LITELLM_ADMIN_MCP_SCHEMA_IMAGE: litellm-admin-mcp-schema:${{ github.sha }}
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
|
||||
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py tests/proxy_migration_tests/test_image_admin_mcp.py -v
|
||||
|
||||
- name: Verify the bundled Lens Compose installation and restart
|
||||
if: matrix.dockerfile == 'Dockerfile'
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
|
||||
run: bash tests/e2e/migrations/lens_compose_smoke.sh
|
||||
|
|
@ -266,9 +285,10 @@ jobs:
|
|||
env:
|
||||
LITELLM_IMAGE: litellm-gateway-scan:${{ github.sha }}
|
||||
LITELLM_COMPONENT_PORT: "4000"
|
||||
LITELLM_IMAGE_COMPONENT: gateway
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
|
||||
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py tests/proxy_migration_tests/test_image_admin_mcp.py -v
|
||||
|
||||
ui-image:
|
||||
name: ui-image
|
||||
|
|
@ -316,6 +336,9 @@ jobs:
|
|||
- name: Build backend image
|
||||
run: docker build -f backend/Dockerfile -t litellm-backend-scan:${{ github.sha }} .
|
||||
|
||||
- name: Tag the cached builder for Admin MCP schema setup
|
||||
run: docker build --target builder -f backend/Dockerfile -t litellm-admin-mcp-schema:${{ github.sha }} .
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
|
|
@ -324,7 +347,9 @@ jobs:
|
|||
- name: Verify the backend serves offline as a non-root uid
|
||||
env:
|
||||
LITELLM_IMAGE: litellm-backend-scan:${{ github.sha }}
|
||||
LITELLM_ADMIN_MCP_SCHEMA_IMAGE: litellm-admin-mcp-schema:${{ github.sha }}
|
||||
LITELLM_COMPONENT_PORT: "4001"
|
||||
LITELLM_IMAGE_COMPONENT: backend
|
||||
run: |
|
||||
python -m pip install "pytest==9.0.3"
|
||||
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py -v
|
||||
python -m pytest tests/proxy_migration_tests/test_component_image_serves_offline.py tests/proxy_migration_tests/test_image_admin_mcp.py -v
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ COPY litellm-proxy-extras/pyproject.toml litellm-proxy-extras/
|
|||
RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
@ -101,6 +102,7 @@ RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.
|
|||
RUN uv sync --frozen --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
@ -56,6 +57,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
uv sync --frozen --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
|
|||
|
|
@ -160,6 +160,7 @@ BACKEND_EXACT_PATHS: frozenset[str] = frozenset(
|
|||
|
||||
BACKEND_MOUNT_PATHS: frozenset[str] = frozenset(
|
||||
{
|
||||
"/admin",
|
||||
"/swagger", # API documentation static assets belong to the backend
|
||||
"/mcp", # lazily-mounted MCP sub-app serves on the backend component
|
||||
}
|
||||
|
|
|
|||
|
|
@ -77,6 +77,7 @@ COPY litellm-proxy-extras/pyproject.toml litellm-proxy-extras/
|
|||
RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
@ -99,6 +100,7 @@ RUN sed -i 's/\r$//' docker/build_admin_ui.sh && chmod +x docker/build_admin_ui.
|
|||
RUN uv sync --frozen --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
|
|||
|
|
@ -81,6 +81,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
@ -108,6 +109,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
uv sync --frozen --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra saml \
|
||||
|
|
|
|||
|
|
@ -70,6 +70,32 @@ To stop the running containers, use the following command:
|
|||
docker compose down
|
||||
```
|
||||
|
||||
## Embedded LiteAdmin MCP
|
||||
|
||||
Source builds containing embedded LiteAdmin MCP can serve it at `/admin/mcp` on the existing LiteLLM port. This capability is unreleased. Keep your existing database, master key, and proxy configuration, then add these settings to the serving container's environment:
|
||||
|
||||
```bash
|
||||
LITELLM_ENABLE_ADMIN_MCP=true
|
||||
LITELLM_LICENSE="your-enterprise-license"
|
||||
PROXY_BASE_URL=https://gateway.example.com
|
||||
```
|
||||
|
||||
For the unified source deployment described above, put them in its `.env` file and rebuild:
|
||||
|
||||
```bash
|
||||
docker compose up -d --build
|
||||
```
|
||||
|
||||
In componentized deployments, set the flag and license on the backend container and route `/admin/mcp` to the backend service. The gateway component excludes this endpoint. The unified, database, non-root, and backend image builds bundle the connector
|
||||
|
||||
Hosting is disabled by default. Opting in requires a valid base Enterprise license; an unlicensed opt-in or invalid flag value prevents startup. Enabling it reserves `/admin`, so rename any MCP server alias called `admin` first
|
||||
|
||||
With native key authentication, connect with a personal proxy-admin bearer key. When `enable_oauth2_proxy_auth` is enabled, the existing trusted-proxy identity headers select the user instead; the MCP bearer is required by the connector but does not select the native user. The resolved user must have the stored `proxy_admin` role, and `trusted_proxy_ranges` applies to the original caller's direct peer
|
||||
|
||||
Embedded responses default to `full`; selecting `LITELLM_ADMIN_RESPONSE_VIEW=compact` requires subsequent saved-result reads to reach the same worker process, including within a multi-worker pod
|
||||
|
||||
See the [LiteAdmin MCP guide](https://docs.litellm.ai/docs/proxy/liteadmin_mcp#run-liteadmin-mcp-inside-litellm) for client configuration, tool restrictions, and verification
|
||||
|
||||
## Hardened / Offline Testing
|
||||
|
||||
To ensure changes are safe for non-root, read-only root filesystems and restricted egress, always validate with the hardened compose file:
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
uv sync --frozen --no-install-project --no-install-workspace --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
|
|
@ -72,6 +73,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
|
|||
uv sync --frozen --no-default-groups --no-editable \
|
||||
--extra proxy \
|
||||
--extra proxy-runtime \
|
||||
--group admin-mcp \
|
||||
--extra extra_proxy \
|
||||
--extra semantic-router \
|
||||
--extra bedrock-realtime \
|
||||
|
|
|
|||
196
litellm/proxy/admin_mcp.py
Normal file
196
litellm/proxy/admin_mcp.py
Normal file
|
|
@ -0,0 +1,196 @@
|
|||
import os
|
||||
import re
|
||||
from collections.abc import AsyncGenerator, Mapping
|
||||
from contextlib import asynccontextmanager
|
||||
from contextvars import ContextVar
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from fastapi import FastAPI
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.datastructures import Headers
|
||||
from starlette.requests import Request
|
||||
from starlette.routing import Mount
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS
|
||||
from litellm.proxy._types import SpecialHeaders
|
||||
from litellm.proxy.middleware.admission_control_middleware import ADMISSION_LEASE_SCOPE_KEY
|
||||
|
||||
_REQUEST_HEADERS: Final = frozenset(
|
||||
{
|
||||
b"authorization",
|
||||
b"litellm-changed-by",
|
||||
b"cookie",
|
||||
b"content-length",
|
||||
b"content-type",
|
||||
b"transfer-encoding",
|
||||
b"connection",
|
||||
b"accept",
|
||||
b"accept-encoding",
|
||||
b"mcp-protocol-version",
|
||||
b"mcp-session-id",
|
||||
}
|
||||
)
|
||||
_CREDENTIAL_HEADERS: Final = frozenset(
|
||||
name.encode("ascii") for name in SpecialHeaders.litellm_credential_header_names()
|
||||
)
|
||||
_RESERVED_KEY_HEADERS: Final = (
|
||||
frozenset(
|
||||
{
|
||||
"host",
|
||||
"origin",
|
||||
"user-agent",
|
||||
"forwarded",
|
||||
"te",
|
||||
"trailer",
|
||||
"upgrade",
|
||||
"x-litellm-user-id",
|
||||
"x-litellm-team-id",
|
||||
"x-litellm-trace-id",
|
||||
"traceparent",
|
||||
"tracestate",
|
||||
}
|
||||
)
|
||||
| frozenset(STANDARD_CUSTOMER_ID_HEADERS)
|
||||
| frozenset(name.decode("ascii") for name in _REQUEST_HEADERS - {b"authorization"})
|
||||
)
|
||||
_SETTINGS: Final = TypeAdapter(Mapping[str, object])
|
||||
_IDENTITY_MAPPINGS: Final = TypeAdapter(tuple[dict[str, object], ...] | dict[str, object] | None)
|
||||
_OAUTH_MAPPINGS: Final = TypeAdapter(dict[str, str])
|
||||
|
||||
|
||||
def _configured_key_header() -> bytes | None:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
settings: Final = _SETTINGS.validate_python(general_settings)
|
||||
name: Final = settings.get("litellm_key_header_name")
|
||||
if name is not None and (not isinstance(name, str) or re.fullmatch(r"[!#$%&'*+\-.^_`|~0-9A-Za-z]+", name) is None):
|
||||
raise ValueError("Hosted admin MCP requires a valid litellm_key_header_name")
|
||||
raw_mappings: Final = _IDENTITY_MAPPINGS.validate_python(settings.get("user_header_mappings"))
|
||||
mappings: Final = (raw_mappings,) if isinstance(raw_mappings, dict) else raw_mappings or ()
|
||||
mapped_names: Final = tuple(mapping.get("header_name") for mapping in mappings)
|
||||
oauth_names: Final = (
|
||||
tuple(_OAUTH_MAPPINGS.validate_python(settings.get("oauth2_config_mappings") or {}).values())
|
||||
if settings.get("enable_oauth2_proxy_auth") is True
|
||||
else ()
|
||||
)
|
||||
policy_names: Final = (
|
||||
settings.get("user_header_name"),
|
||||
settings.get("mcp_client_id_header"),
|
||||
*mapped_names,
|
||||
*oauth_names,
|
||||
)
|
||||
policy_headers: Final = frozenset(value.lower() for value in policy_names if isinstance(value, str))
|
||||
overwritten_headers: Final = frozenset(value.decode("ascii") for value in _REQUEST_HEADERS | _CREDENTIAL_HEADERS)
|
||||
if policy_headers & overwritten_headers:
|
||||
raise ValueError("Hosted admin MCP cannot overwrite configured identity headers")
|
||||
if name is None:
|
||||
return None
|
||||
normalized: Final = name.lower()
|
||||
if normalized in _RESERVED_KEY_HEADERS | policy_headers or normalized.startswith("x-forwarded-"):
|
||||
raise ValueError("Hosted admin MCP litellm_key_header_name cannot replace a transport, audit, or policy header")
|
||||
return normalized.encode("ascii")
|
||||
|
||||
|
||||
def _require_enterprise_license() -> None:
|
||||
from litellm.proxy.utils import require_enterprise_license
|
||||
|
||||
require_enterprise_license("Hosted admin MCP")
|
||||
|
||||
|
||||
class _CallerContext:
|
||||
def __init__(self, app: ASGIApp, caller: ContextVar[Request]) -> None:
|
||||
self.app = app
|
||||
self.caller = caller
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
_require_enterprise_license()
|
||||
token: Final = self.caller.set(Request(scope))
|
||||
try:
|
||||
await self.app(scope, receive, send)
|
||||
finally:
|
||||
self.caller.reset(token)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def admin_mcp_lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
|
||||
enabled: Final = os.environ.get("LITELLM_ENABLE_ADMIN_MCP", "false").strip().lower()
|
||||
if enabled in ("false", "0", "off", "no", ""):
|
||||
yield
|
||||
return
|
||||
if enabled not in ("true", "1", "on", "yes"):
|
||||
raise ValueError("LITELLM_ENABLE_ADMIN_MCP must be true or false")
|
||||
_require_enterprise_license()
|
||||
_configured_key_header()
|
||||
|
||||
try:
|
||||
import httpx2
|
||||
from litellm_admin_mcp.config import ( # pyright: ignore[reportMissingTypeStubs] # upstream has no py.typed marker
|
||||
Config,
|
||||
env_bool,
|
||||
)
|
||||
from litellm_admin_mcp.gateway import ( # pyright: ignore[reportMissingTypeStubs] # upstream has no py.typed marker
|
||||
Gateway,
|
||||
)
|
||||
from litellm_admin_mcp.server import ( # pyright: ignore[reportMissingTypeStubs] # upstream has no py.typed marker
|
||||
create_http_app,
|
||||
)
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"Admin MCP requires Python 3.12+ and the admin-mcp dependency group. "
|
||||
"Use a LiteLLM image that bundles it, or run uv sync --extra proxy --group admin-mcp."
|
||||
) from exc
|
||||
|
||||
configured_url: Final = os.environ.get("LITELLM_MCP_PUBLIC_URL") or os.environ.get("PROXY_BASE_URL", "")
|
||||
public_url: Final = urlsplit(configured_url)
|
||||
config: Final = Config(
|
||||
base_url="http://localhost",
|
||||
public_url=f"{public_url.scheme}://{public_url.netloc}" if public_url.netloc else configured_url,
|
||||
read_only=env_bool("LITELLM_ADMIN_READ_ONLY"),
|
||||
allowed_tools=frozenset(
|
||||
name.strip() for name in os.environ.get("LITELLM_ADMIN_TOOLS", "").split(",") if name.strip()
|
||||
),
|
||||
response_view=os.environ.get("LITELLM_ADMIN_RESPONSE_VIEW", "full").strip(),
|
||||
schema_mode=os.environ.get("LITELLM_ADMIN_SCHEMA_MODE", "full").strip(),
|
||||
)
|
||||
caller: Final[ContextVar[Request]] = ContextVar("admin_mcp_caller")
|
||||
|
||||
async def management_api(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
request: Final = caller.get()
|
||||
configured_header: Final = _configured_key_header()
|
||||
excluded: Final = (
|
||||
_REQUEST_HEADERS
|
||||
| _CREDENTIAL_HEADERS
|
||||
| (frozenset({configured_header}) if configured_header is not None else frozenset())
|
||||
)
|
||||
caller_headers: Final = tuple(pair for pair in request.headers.raw if pair[0].lower() not in excluded)
|
||||
generated_headers: Final = Headers(scope=scope)
|
||||
api_headers: Final = tuple(pair for pair in generated_headers.raw if pair[0] in _REQUEST_HEADERS)
|
||||
configured_auth: Final = (
|
||||
((configured_header, generated_headers["authorization"].encode("ascii")),)
|
||||
if configured_header is not None and configured_header != b"authorization"
|
||||
else ()
|
||||
)
|
||||
headers: Final = list(caller_headers + api_headers + configured_auth)
|
||||
gateway_scope: Final[Scope] = {
|
||||
**scope,
|
||||
"client": request.client,
|
||||
"scheme": request.url.scheme,
|
||||
"headers": headers,
|
||||
ADMISSION_LEASE_SCOPE_KEY: request.scope.get(ADMISSION_LEASE_SCOPE_KEY),
|
||||
}
|
||||
await app(gateway_scope, receive, send)
|
||||
|
||||
async with httpx2.AsyncClient(transport=httpx2.ASGITransport(app=management_api)) as client:
|
||||
admin_app: Final = create_http_app(Gateway(config, client))
|
||||
route: Final = Mount("/admin", app=_CallerContext(admin_app, caller), name="admin_mcp")
|
||||
async with admin_app.router.lifespan_context(admin_app):
|
||||
app.router.routes.insert(0, route)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
app.router.routes[:] = [existing for existing in app.router.routes if existing is not route]
|
||||
|
|
@ -11,6 +11,8 @@ from starlette.types import ASGIApp, Receive, Scope, Send
|
|||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
ADMISSION_LEASE_SCOPE_KEY: Final = "litellm.admission_lease"
|
||||
|
||||
_EXEMPT_PATHS: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
"/health/liveliness",
|
||||
|
|
@ -130,6 +132,12 @@ class AdmissionControlState:
|
|||
return self._metrics
|
||||
|
||||
|
||||
class _AdmissionLease:
|
||||
def __init__(self, state: AdmissionControlState) -> None:
|
||||
self.state: Final = state
|
||||
self.active: bool = True
|
||||
|
||||
|
||||
class AdmissionControlMiddleware:
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -146,6 +154,15 @@ class AdmissionControlMiddleware:
|
|||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
inherited_lease: Final = scope.get(ADMISSION_LEASE_SCOPE_KEY)
|
||||
if (
|
||||
isinstance(inherited_lease, _AdmissionLease)
|
||||
and inherited_lease.state is self.state
|
||||
and inherited_lease.active
|
||||
):
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
settings: Final = self.get_settings()
|
||||
if settings is None or _get_route_path(scope) in _EXEMPT_PATHS:
|
||||
await self.app(scope, receive, send)
|
||||
|
|
@ -178,9 +195,13 @@ class AdmissionControlMiddleware:
|
|||
state.record_dequeue()
|
||||
state.record_admission()
|
||||
|
||||
lease: Final = _AdmissionLease(state)
|
||||
scope[ADMISSION_LEASE_SCOPE_KEY] = lease # rebind-ok: outer ASGI wrappers must see downstream route metadata
|
||||
try:
|
||||
await self.app(scope, receive, send)
|
||||
finally:
|
||||
lease.active = False
|
||||
scope.pop(ADMISSION_LEASE_SCOPE_KEY, None)
|
||||
semaphore.release()
|
||||
state.record_release()
|
||||
|
||||
|
|
|
|||
|
|
@ -262,7 +262,7 @@ def generate_feedback_box():
|
|||
|
||||
import contextlib
|
||||
from collections import defaultdict
|
||||
from contextlib import asynccontextmanager
|
||||
from contextlib import AsyncExitStack, asynccontextmanager
|
||||
from functools import lru_cache, partial
|
||||
|
||||
import litellm
|
||||
|
|
@ -1614,75 +1614,81 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState
|
|||
settings=tracing_settings,
|
||||
) as receiver:
|
||||
state: Final[ProxyLifespanState] = {"tracing_receiver": receiver}
|
||||
yield state
|
||||
from litellm.proxy.admin_mcp import admin_mcp_lifespan
|
||||
|
||||
if model_info_scheduler is not None and model_info_scheduler.running:
|
||||
model_info_scheduler.remove_job("refresh_model_info")
|
||||
if model_info_scheduler is not scheduler:
|
||||
model_info_scheduler.shutdown(wait=False)
|
||||
try:
|
||||
async with AsyncExitStack() as admin_mcp_stack:
|
||||
try:
|
||||
await admin_mcp_stack.enter_async_context(admin_mcp_lifespan(app))
|
||||
yield state
|
||||
finally:
|
||||
if model_info_scheduler is not None and model_info_scheduler.running:
|
||||
model_info_scheduler.remove_job("refresh_model_info")
|
||||
if model_info_scheduler is not scheduler:
|
||||
model_info_scheduler.shutdown(wait=False)
|
||||
|
||||
# Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window
|
||||
if scheduler is not None:
|
||||
pause_scheduled_jobs(scheduler)
|
||||
# Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window
|
||||
if scheduler is not None:
|
||||
pause_scheduled_jobs(scheduler)
|
||||
|
||||
# Shutdown event - drain in-flight requests before tearing down dependencies
|
||||
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
|
||||
GracefulShutdownManager.start_shutdown()
|
||||
await GracefulShutdownManager.wait_for_drain()
|
||||
# Shutdown event - drain in-flight requests before tearing down dependencies
|
||||
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
|
||||
GracefulShutdownManager.start_shutdown()
|
||||
await GracefulShutdownManager.wait_for_drain()
|
||||
finally:
|
||||
# Shutdown event - close shared aiohttp session
|
||||
if shared_aiohttp_session is not None:
|
||||
try:
|
||||
await shared_aiohttp_session.close()
|
||||
verbose_proxy_logger.info("SESSION REUSE: Closed shared aiohttp session")
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error closing shared aiohttp session: %s", e)
|
||||
|
||||
# Shutdown event - close shared aiohttp session
|
||||
if shared_aiohttp_session is not None:
|
||||
try:
|
||||
await shared_aiohttp_session.close()
|
||||
verbose_proxy_logger.info("SESSION REUSE: Closed shared aiohttp session")
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error closing shared aiohttp session: %s", e)
|
||||
# Shutdown event - stop RDS IAM token refresh background task
|
||||
if (
|
||||
prisma_client is not None
|
||||
and hasattr(prisma_client, "db")
|
||||
and hasattr(prisma_client.db, "stop_token_refresh_task")
|
||||
):
|
||||
try:
|
||||
await prisma_client.db.stop_token_refresh_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping token refresh task: %s", e)
|
||||
|
||||
# Shutdown event - stop RDS IAM token refresh background task
|
||||
if (
|
||||
prisma_client is not None
|
||||
and hasattr(prisma_client, "db")
|
||||
and hasattr(prisma_client.db, "stop_token_refresh_task")
|
||||
):
|
||||
try:
|
||||
await prisma_client.db.stop_token_refresh_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping token refresh task: %s", e)
|
||||
# Shutdown event - stop Prisma DB health watchdog task
|
||||
if prisma_client is not None and hasattr(prisma_client, "stop_db_health_watchdog_task"):
|
||||
try:
|
||||
await prisma_client.stop_db_health_watchdog_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e)
|
||||
|
||||
# Shutdown event - stop Prisma DB health watchdog task
|
||||
if prisma_client is not None and hasattr(prisma_client, "stop_db_health_watchdog_task"):
|
||||
try:
|
||||
await prisma_client.stop_db_health_watchdog_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e)
|
||||
if prisma_client is not None and hasattr(prisma_client, "stop_view_setup_task"):
|
||||
try:
|
||||
await prisma_client.stop_view_setup_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
|
||||
|
||||
if prisma_client is not None and hasattr(prisma_client, "stop_view_setup_task"):
|
||||
try:
|
||||
await prisma_client.stop_view_setup_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
|
||||
await _drain_spend_event_producer_on_shutdown()
|
||||
|
||||
await _drain_spend_event_producer_on_shutdown()
|
||||
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
|
||||
if scheduler is not None and scheduler_executor is not None:
|
||||
try:
|
||||
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
|
||||
|
||||
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
|
||||
if scheduler is not None and scheduler_executor is not None:
|
||||
try:
|
||||
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
|
||||
await flush_spend_counters_on_shutdown()
|
||||
|
||||
await flush_spend_counters_on_shutdown()
|
||||
await _flush_spend_logs_queue_on_shutdown()
|
||||
|
||||
await _flush_spend_logs_queue_on_shutdown()
|
||||
await proxy_config.stop_config_sync_subscriber()
|
||||
|
||||
await proxy_config.stop_config_sync_subscriber()
|
||||
await proxy_config.stop_auth_cache_invalidation_subscriber()
|
||||
|
||||
await proxy_config.stop_auth_cache_invalidation_subscriber()
|
||||
await proxy_shutdown_event(worker_heartbeat=worker_heartbeat)
|
||||
|
||||
await proxy_shutdown_event(worker_heartbeat=worker_heartbeat)
|
||||
|
||||
if prometheus_multiproc_dir:
|
||||
mark_worker_exit(os.getpid())
|
||||
if prometheus_multiproc_dir:
|
||||
mark_worker_exit(os.getpid())
|
||||
|
||||
|
||||
def _generate_stable_operation_id(route: "APIRoute") -> str:
|
||||
|
|
|
|||
|
|
@ -8317,7 +8317,7 @@ def handle_exception_on_proxy(e: Exception, litellm_call_id: str | None = None)
|
|||
)
|
||||
|
||||
|
||||
def _premium_user_check(feature: str | None = None):
|
||||
def require_enterprise_license(feature: str | None = None) -> None:
|
||||
"""
|
||||
Raises an HTTPException if the user is not a premium user
|
||||
"""
|
||||
|
|
@ -8337,6 +8337,9 @@ def _premium_user_check(feature: str | None = None):
|
|||
)
|
||||
|
||||
|
||||
_premium_user_check: Final = require_enterprise_license
|
||||
|
||||
|
||||
def is_known_model(model: str | None, llm_router: Router | None) -> bool:
|
||||
"""
|
||||
Returns True if the model is in the llm_router model names
|
||||
|
|
|
|||
|
|
@ -242,7 +242,11 @@ e2e-dev = [
|
|||
"psutil==7.2.2",
|
||||
"mcp>=2.2.0,<3",
|
||||
]
|
||||
admin-mcp = [
|
||||
"litellm-admin-mcp @ https://github.com/BerriAI/liteadmin-mcp/archive/d35ec9c19c117d4c50cc1eccf6ce1296aac25a1a.tar.gz#sha256=86b930f6706fb2da10d53d1798ee7fd14e44fb0afdef0da122cf0e6bb88fa2c3 ; python_version >= '3.12'",
|
||||
]
|
||||
proxy-dev = [
|
||||
{ include-group = "admin-mcp" },
|
||||
"prisma==0.11.0",
|
||||
"hypercorn==0.17.3",
|
||||
"prometheus-client==0.20.0",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
#!/usr/bin/env python3
|
||||
import configparser
|
||||
from collections.abc import Collection, Iterable, Iterator
|
||||
from dataclasses import dataclass
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
|
@ -339,6 +340,20 @@ class LicenseChecker:
|
|||
|
||||
return is_acceptable
|
||||
|
||||
def _direct_group_requirements(
|
||||
self, entries: Iterable[object], group_names: Collection[str]
|
||||
) -> Iterator[str]:
|
||||
for entry in entries:
|
||||
match entry:
|
||||
case str():
|
||||
yield entry
|
||||
case {"include-group": str(name)} if (
|
||||
len(entry) == 1 and self._normalize_package_name(name) in group_names
|
||||
):
|
||||
continue
|
||||
case _:
|
||||
raise ValueError(f"Invalid dependency group entry: {entry!r}")
|
||||
|
||||
def _load_requirements(
|
||||
self, requirements_file: Optional[Path] = None
|
||||
) -> List[Requirement]:
|
||||
|
|
@ -358,8 +373,14 @@ class LicenseChecker:
|
|||
pyproject["project"].get("optional-dependencies", {}).values()
|
||||
):
|
||||
requirement_lines.extend(extra_reqs)
|
||||
for group_reqs in pyproject.get("dependency-groups", {}).values():
|
||||
requirement_lines.extend(group_reqs)
|
||||
groups: Final = pyproject.get("dependency-groups", {})
|
||||
group_names: Final = frozenset(
|
||||
self._normalize_package_name(name) for name in groups
|
||||
)
|
||||
for group_reqs in groups.values():
|
||||
requirement_lines.extend(
|
||||
self._direct_group_requirements(group_reqs, group_names)
|
||||
)
|
||||
|
||||
lock_versions: Dict[str, List[str]] = {}
|
||||
for package in lock_data.get("package", []):
|
||||
|
|
@ -386,9 +407,10 @@ class LicenseChecker:
|
|||
requirement_lines = list(dict.fromkeys(requirement_lines))
|
||||
|
||||
return [
|
||||
Requirement(line.split("#")[0].strip())
|
||||
Requirement(requirement)
|
||||
for line in requirement_lines
|
||||
if line.split("#")[0].strip() and not line.startswith("#")
|
||||
if (requirement := re.split(r"\s+#", line, maxsplit=1)[0].strip())
|
||||
and not requirement.startswith("#")
|
||||
]
|
||||
except Exception as e:
|
||||
source = requirements_file or "pyproject.toml + uv.lock"
|
||||
|
|
@ -446,6 +468,9 @@ def main():
|
|||
|
||||
# Check requirements
|
||||
if not checker.check_requirements(req_file):
|
||||
if not checker.package_results:
|
||||
sys.exit(1)
|
||||
|
||||
# Get lists of problematic packages
|
||||
unverified = [p for p in checker.package_results if not p.license_type]
|
||||
invalid = [
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ unauthorized_licenses:
|
|||
gpl v3
|
||||
|
||||
[Authorized Packages]
|
||||
litellm-admin-mcp: ==0.1.0 # MIT, verified at https://github.com/BerriAI/liteadmin-mcp/blob/d35ec9c19c117d4c50cc1eccf6ce1296aac25a1a/LICENSE
|
||||
# Apache-2.0 https://github.com/chroma-core/hnswlib#Apache-2.0-1-ov-file
|
||||
chroma-hnswlib: >=0.7.3
|
||||
# MIT https://github.com/facebookresearch/iopath?tab=MIT-1-ov-file#readme
|
||||
|
|
|
|||
197
tests/proxy_migration_tests/test_image_admin_mcp.py
Normal file
197
tests/proxy_migration_tests/test_image_admin_mcp.py
Normal file
|
|
@ -0,0 +1,197 @@
|
|||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import test_offline_image_migration
|
||||
|
||||
offline_postgres: Final = test_offline_image_migration.offline_postgres
|
||||
|
||||
IMAGE: Final = os.getenv("LITELLM_IMAGE")
|
||||
SCHEMA_IMAGE: Final = os.getenv("LITELLM_ADMIN_MCP_SCHEMA_IMAGE")
|
||||
COMPONENT: Final = os.getenv("LITELLM_IMAGE_COMPONENT", "unified")
|
||||
PROBE: Final = """
|
||||
import asyncio
|
||||
import base64
|
||||
import importlib
|
||||
import json
|
||||
import sys
|
||||
from typing import Final
|
||||
|
||||
import httpx2
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.asymmetric import padding, rsa
|
||||
from fastapi import HTTPException
|
||||
from litellm_admin_mcp.server import create_http_app
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
module_name: Final = {
|
||||
"unified": "litellm.proxy.proxy_server",
|
||||
"backend": "backend.main",
|
||||
"gateway": "gateway.main",
|
||||
}[sys.argv[2]]
|
||||
app: Final = importlib.import_module(module_name).app
|
||||
|
||||
if sys.argv[1] == "base":
|
||||
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
message: Final = json.dumps({"expiration_date": "2999-01-01", "user_id": "image-test"}).encode()
|
||||
signature: Final = private_key.sign(
|
||||
message,
|
||||
padding.PSS(mgf=padding.MGF1(hashes.SHA256()), salt_length=padding.PSS.MAX_LENGTH),
|
||||
hashes.SHA256(),
|
||||
)
|
||||
proxy_server._license_check.public_key = private_key.public_key()
|
||||
proxy_server._license_check.license_str = base64.b64encode(message + b"." + signature).decode()
|
||||
|
||||
async def verify_admin_tools(client: httpx2.AsyncClient) -> None:
|
||||
user: Final = await client.post(
|
||||
"/user/new",
|
||||
headers={"Authorization": "Bearer sk-0123456789abcdef0123456789abcdef"},
|
||||
json={"user_id": "image-admin", "user_role": "proxy_admin", "auto_create_key": True},
|
||||
)
|
||||
assert user.status_code == 200, "Admin provisioning failed: " + str(user.status_code)
|
||||
headers: Final = {
|
||||
"Authorization": "Bearer " + user.json()["key"],
|
||||
"Accept": "application/json, text/event-stream",
|
||||
}
|
||||
discovery: Final = await client.post(
|
||||
"/admin/mcp", headers=headers, json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"}
|
||||
)
|
||||
assert discovery.status_code == 200, discovery.text
|
||||
assert {"create_team", "get_team", "delete_teams"} <= {
|
||||
tool["name"] for tool in discovery.json()["result"]["tools"]
|
||||
}
|
||||
|
||||
async def call_tool(name: str, arguments: dict[str, object]) -> str:
|
||||
response: Final = await client.post(
|
||||
"/admin/mcp", headers=headers,
|
||||
json={"jsonrpc": "2.0", "id": 2, "method": "tools/call",
|
||||
"params": {"name": name, "arguments": arguments}},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
result: Final = response.json()["result"]
|
||||
assert not result.get("isError"), response.text
|
||||
return result["content"][0]["text"]
|
||||
|
||||
team_id: Final = "image-admin-mcp-team"
|
||||
created: Final = json.loads(await call_tool(
|
||||
"create_team", {"body": {"team_id": team_id, "team_alias": team_id, "max_budget": 25}}
|
||||
))
|
||||
assert created["team_id"] == team_id and created["max_budget"] == 25, created
|
||||
read: Final = json.loads(await call_tool("get_team", {"query": {"team_id": team_id}}))
|
||||
assert read["team_info"]["team_id"] == team_id and read["team_info"]["max_budget"] == 25, read
|
||||
await call_tool("delete_teams", {"body": {"team_ids": [team_id]}})
|
||||
deleted: Final = await client.get("/team/info", headers=headers, params={"team_id": team_id})
|
||||
assert deleted.status_code == 404, deleted.text
|
||||
print("admin-mcp-tools-ok")
|
||||
|
||||
async def probe() -> None:
|
||||
try:
|
||||
async with app.router.lifespan_context(app):
|
||||
async with httpx2.AsyncClient(
|
||||
transport=httpx2.ASGITransport(app=app), base_url="http://localhost:4000"
|
||||
) as client:
|
||||
response: Final = await client.post(
|
||||
"/admin/mcp",
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"},
|
||||
headers={"Accept": "application/json, text/event-stream"},
|
||||
)
|
||||
print("admin-mcp-status=" + str(response.status_code))
|
||||
if sys.argv[3] == "tools":
|
||||
assert response.status_code == 401, response.text
|
||||
await verify_admin_tools(client)
|
||||
except HTTPException as exc:
|
||||
assert exc.status_code == 403 and "LITELLM_LICENSE" in str(exc.detail)
|
||||
assert sys.argv[1] == "none"
|
||||
print("admin-mcp-license-required")
|
||||
|
||||
asyncio.run(probe())
|
||||
"""
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(IMAGE is None, reason="requires a built image (set LITELLM_IMAGE)"),
|
||||
pytest.mark.skipif(shutil.which("docker") is None, reason="requires the docker CLI"),
|
||||
]
|
||||
|
||||
|
||||
def _run_probe(
|
||||
enabled: str | None,
|
||||
license_mode: str,
|
||||
network: str = "none",
|
||||
database_url: str | None = None,
|
||||
) -> subprocess.CompletedProcess[str]:
|
||||
assert IMAGE is not None
|
||||
enabled_args: Final = () if enabled is None else ("--env", "LITELLM_ENABLE_ADMIN_MCP=" + enabled)
|
||||
database_args: Final = () if database_url is None else ("--env", "DATABASE_URL=" + database_url)
|
||||
return subprocess.run(
|
||||
[
|
||||
"docker",
|
||||
"run",
|
||||
"--rm",
|
||||
"--network",
|
||||
network,
|
||||
"--user",
|
||||
"12345:0",
|
||||
*enabled_args,
|
||||
*database_args,
|
||||
"--env",
|
||||
"LITELLM_LOCAL_MODEL_COST_MAP=true",
|
||||
"--env",
|
||||
"LITELLM_MASTER_KEY=sk-0123456789abcdef0123456789abcdef",
|
||||
"--entrypoint",
|
||||
"python",
|
||||
IMAGE,
|
||||
"-c",
|
||||
PROBE,
|
||||
license_mode,
|
||||
COMPONENT,
|
||||
"tools" if database_url is not None else "visibility",
|
||||
],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=120,
|
||||
check=False,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"enabled,license_mode,expected",
|
||||
[
|
||||
(None, "none", "status=404"),
|
||||
("false", "base", "status=404"),
|
||||
("true", "none", "license-required"),
|
||||
("true", "base", "status=404" if COMPONENT == "gateway" else "status=401"),
|
||||
],
|
||||
)
|
||||
def test_image_admin_mcp_requires_opt_in_license_and_management_component(
|
||||
enabled: str | None, license_mode: str, expected: str
|
||||
) -> None:
|
||||
result: Final = _run_probe(enabled, license_mode)
|
||||
assert result.returncode == 0 and f"admin-mcp-{expected}" in result.stdout, (
|
||||
f"Admin MCP image probe failed with component={COMPONENT}, enabled={enabled}, license={license_mode}\n"
|
||||
f"{result.stdout}\n{result.stderr}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(COMPONENT == "gateway", reason="the gateway excludes management endpoints")
|
||||
def test_image_admin_mcp_personal_admin_manages_team(offline_postgres: tuple[str, str]) -> None:
|
||||
assert SCHEMA_IMAGE is not None, "set LITELLM_ADMIN_MCP_SCHEMA_IMAGE to the matching builder image"
|
||||
network, postgres = offline_postgres
|
||||
database_url: Final = f"postgresql://postgres:pw@{postgres}:5432/litellm"
|
||||
schema: Final = subprocess.run(
|
||||
[
|
||||
"docker", "run", "--rm", "--network", network,
|
||||
"--env", "DATABASE_URL=" + database_url,
|
||||
"--env", "HOME=/opt/prisma", "--env", "XDG_CACHE_HOME=/opt/prisma/.cache",
|
||||
"--env", "PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries",
|
||||
"--entrypoint", "prisma", SCHEMA_IMAGE,
|
||||
"db", "push", "--schema", "/app/schema.prisma", "--skip-generate", "--accept-data-loss",
|
||||
],
|
||||
capture_output=True, text=True, timeout=180, check=False,
|
||||
)
|
||||
assert schema.returncode == 0, f"Schema provisioning failed\n{schema.stdout}\n{schema.stderr}"
|
||||
result: Final = _run_probe("true", "base", network, database_url)
|
||||
assert result.returncode == 0 and "admin-mcp-tools-ok" in result.stdout, (
|
||||
f"Admin MCP management failed in {COMPONENT}\n{result.stdout}\n{result.stderr}"
|
||||
)
|
||||
|
|
@ -3,10 +3,12 @@ import json
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
from litellm.proxy.middleware.admission_control_middleware import (
|
||||
ADMISSION_LEASE_SCOPE_KEY,
|
||||
AdmissionControlMetrics,
|
||||
AdmissionControlMiddleware,
|
||||
AdmissionControlSettings,
|
||||
|
|
@ -24,9 +26,10 @@ def state() -> AdmissionControlState:
|
|||
|
||||
|
||||
async def _call(
|
||||
middleware: AdmissionControlMiddleware,
|
||||
middleware: ASGIApp,
|
||||
path: str = "/",
|
||||
root_path: str = "",
|
||||
parent_scope: Scope | None = None,
|
||||
) -> tuple[Message, ...]:
|
||||
messages: Final[list[Message]] = []
|
||||
|
||||
|
|
@ -41,7 +44,9 @@ async def _call(
|
|||
"path": path,
|
||||
"root_path": root_path,
|
||||
"method": "GET",
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
ADMISSION_LEASE_SCOPE_KEY: parent_scope.get(ADMISSION_LEASE_SCOPE_KEY) if parent_scope else None,
|
||||
}
|
||||
await middleware(scope, receive, send)
|
||||
return tuple(messages)
|
||||
|
|
@ -400,3 +405,108 @@ def test_invalid_admission_control_settings_logs_once(caplog: pytest.LogCaptureF
|
|||
if record.message.startswith("Ignoring invalid admission control settings")
|
||||
)
|
||||
assert len(messages) == 1
|
||||
|
||||
|
||||
def _single_slot() -> AdmissionControlSettings:
|
||||
return AdmissionControlSettings(1, 0, 1.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", (None, RuntimeError, asyncio.CancelledError))
|
||||
async def test_outer_wrapper_retains_route_metadata_after_admission(
|
||||
state: AdmissionControlState,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
failure: type[BaseException] | None,
|
||||
) -> None:
|
||||
monkeypatch.delenv("LITELLM_ENABLE_ADMIN_MCP", raising=False)
|
||||
app: Final = FastAPI()
|
||||
|
||||
@app.get("/items/{item_id}")
|
||||
async def item(item_id: str) -> dict[str, str]:
|
||||
assert state.get_stats().admitted == 1
|
||||
if failure is not None:
|
||||
raise failure("request interrupted")
|
||||
return {"item_id": item_id}
|
||||
|
||||
middleware: Final = AdmissionControlMiddleware(app, _single_slot, state)
|
||||
|
||||
async def outer_probe(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
try:
|
||||
await middleware(scope, receive, send)
|
||||
finally:
|
||||
assert scope["route"].path == "/items/{item_id}"
|
||||
assert scope["endpoint"] is item
|
||||
assert scope["path_params"] == {"item_id": "sample"}
|
||||
assert ADMISSION_LEASE_SCOPE_KEY not in scope
|
||||
assert state.get_stats() == AdmissionControlStats(0, 0, 0)
|
||||
|
||||
if failure is not None:
|
||||
with pytest.raises(failure, match="request interrupted"):
|
||||
await _call(outer_probe, "/items/sample")
|
||||
else:
|
||||
response: Final = await _call(outer_probe, "/items/sample")
|
||||
assert response[0]["status"] == 200
|
||||
assert json.loads(response[1]["body"]) == {"item_id": "sample"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_request_acquires_a_new_slot_after_parent_finishes(state: AdmissionControlState) -> None:
|
||||
release: Final = asyncio.Event()
|
||||
background: Final[asyncio.Future[asyncio.Task[tuple[Message, ...]]]] = asyncio.get_running_loop().create_future()
|
||||
|
||||
async def later_request(parent_scope: Scope) -> tuple[Message, ...]:
|
||||
await release.wait()
|
||||
return await _call(middleware, path="/child", parent_scope=parent_scope)
|
||||
|
||||
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["path"] == "/":
|
||||
background.set_result(asyncio.create_task(later_request(scope.copy())))
|
||||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||
await send({"type": "http.response.body", "body": str(state.get_stats().admitted).encode()})
|
||||
|
||||
middleware: Final = AdmissionControlMiddleware(handler, _single_slot, state)
|
||||
parent: Final = await _call(middleware)
|
||||
assert parent[1]["body"] == b"1"
|
||||
assert state.get_stats() == AdmissionControlStats(0, 0, 0)
|
||||
release.set()
|
||||
child: Final = await (await background)
|
||||
assert child[1]["body"] == b"1"
|
||||
assert state.get_stats() == AdmissionControlStats(0, 0, 0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("linked", [False, True])
|
||||
async def test_only_explicitly_linked_requests_share_admission(state: AdmissionControlState, linked: bool) -> None:
|
||||
async def handler(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["path"] == "/":
|
||||
nested: Final = await _call(middleware, path="/child", parent_scope=scope if linked else None)
|
||||
await send(nested[0])
|
||||
await send(nested[1])
|
||||
return
|
||||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||
await send({"type": "http.response.body", "body": str(state.get_stats().admitted).encode()})
|
||||
|
||||
middleware: Final = AdmissionControlMiddleware(handler, _single_slot, state)
|
||||
response: Final = await _call(middleware)
|
||||
assert response[0]["status"] == (200 if linked else 503)
|
||||
assert state.get_stats() == AdmissionControlStats(0, 0, 0 if linked else 1)
|
||||
if linked:
|
||||
assert response[1]["body"] == b"1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admission_lease_cannot_be_reused_by_another_worker(state: AdmissionControlState) -> None:
|
||||
def no_metrics() -> None:
|
||||
return None
|
||||
|
||||
other_state: Final = AdmissionControlState(no_metrics)
|
||||
|
||||
async def child_handler(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||
await send({"type": "http.response.body", "body": str(other_state.get_stats().admitted).encode()})
|
||||
|
||||
child: Final = AdmissionControlMiddleware(child_handler, _single_slot, other_state)
|
||||
parent: Final = AdmissionControlMiddleware(child, _single_slot, state)
|
||||
response: Final = await _call(parent)
|
||||
assert response[1]["body"] == b"1"
|
||||
assert state.get_stats() == other_state.get_stats() == AdmissionControlStats(0, 0, 0)
|
||||
|
|
|
|||
567
tests/unit/proxy/test_admin_mcp.py
Normal file
567
tests/unit/proxy/test_admin_mcp.py
Normal file
|
|
@ -0,0 +1,567 @@
|
|||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from typing import Final
|
||||
|
||||
import httpx2
|
||||
import pytest
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.routing import APIRoute
|
||||
from pydantic import BaseModel
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_UserTable, ProxyException, SpecialHeaders
|
||||
from litellm.proxy.admin_mcp import admin_mcp_lifespan
|
||||
from litellm.proxy.auth.user_api_key_auth import get_api_key, get_api_key_from_custom_header
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints import internal_user_endpoints
|
||||
from litellm.proxy.middleware.admission_control_middleware import (
|
||||
AdmissionControlMiddleware,
|
||||
AdmissionControlSettings,
|
||||
AdmissionControlState,
|
||||
AdmissionControlStats,
|
||||
)
|
||||
from litellm.proxy.middleware.per_request_root_path_middleware import PerRequestRootPathMiddleware
|
||||
|
||||
|
||||
class KeyRequest(BaseModel):
|
||||
key_alias: str
|
||||
|
||||
|
||||
def _native_credential(request: Request) -> str:
|
||||
key, _ = get_api_key(
|
||||
custom_litellm_key_header=request.headers.get("x-litellm-api-key"),
|
||||
api_key=request.headers.get("authorization", ""),
|
||||
azure_api_key_header=request.headers.get("api-key"),
|
||||
anthropic_api_key_header=request.headers.get("x-api-key"),
|
||||
google_ai_studio_api_key_header=request.headers.get("x-goog-api-key"),
|
||||
azure_apim_header=request.headers.get("ocp-apim-subscription-key"),
|
||||
pass_through_endpoints=None, route=request.url.path, request=request,
|
||||
)
|
||||
configured: Final = proxy_server.general_settings.get("litellm_key_header_name")
|
||||
return get_api_key_from_custom_header(request, configured) if configured is not None else key
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def management_app(monkeypatch: pytest.MonkeyPatch) -> FastAPI:
|
||||
monkeypatch.setattr(proxy_server, "premium_user", True)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", "true")
|
||||
monkeypatch.setenv("LITELLM_ADMIN_TOOLS", "list_keys")
|
||||
for name in (
|
||||
"PROXY_BASE_URL", "LITELLM_MCP_PUBLIC_URL", "LITELLM_ADMIN_READ_ONLY",
|
||||
"LITELLM_ADMIN_RESPONSE_VIEW", "LITELLM_ADMIN_SCHEMA_MODE",
|
||||
):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
app: Final = FastAPI(lifespan=admin_mcp_lifespan)
|
||||
|
||||
@app.get("/user/info")
|
||||
async def user_info(request: Request) -> dict[str, object]:
|
||||
await asyncio.sleep(0)
|
||||
credential: Final = _native_credential(request)
|
||||
if credential == "team-key":
|
||||
raise HTTPException(status_code=404, detail="User None not found")
|
||||
if credential not in {"admin-a", "admin-b", "admin-b-limited", "member", "viewer"}:
|
||||
raise HTTPException(status_code=401)
|
||||
user_id: Final = "admin-b" if credential == "admin-b-limited" else credential
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"user_info": {
|
||||
"user_id": user_id,
|
||||
"user_role": {"member": "internal_user", "viewer": "proxy_admin_viewer"}.get(user_id, "proxy_admin"),
|
||||
},
|
||||
}
|
||||
|
||||
@app.get("/key/list", operation_id="list_keys_key_list_get")
|
||||
async def list_keys(request: Request, large: bool = False) -> dict[str, object]:
|
||||
credential: Final = _native_credential(request)
|
||||
if credential == "admin-b-limited":
|
||||
raise HTTPException(403, "This key cannot list keys")
|
||||
if large:
|
||||
return {"keys": [{"key_alias": ("a" if credential == "admin-a" else "b") * 20000}]}
|
||||
return {
|
||||
"keys": ["owned-by-a" if credential == "admin-a" else "owned-by-b"],
|
||||
"client": request.client.host if request.client else None,
|
||||
"scheme": request.url.scheme,
|
||||
"forwarded_for": request.headers.get("x-forwarded-for"),
|
||||
"cookie": request.headers.get("cookie"),
|
||||
"caller_is_changed_by": request.headers.get("litellm-changed-by")
|
||||
== credential,
|
||||
"policy_team": request.headers.get("x-litellm-team-id"),
|
||||
"alternate_credentials_absent": not any(
|
||||
name in request.headers
|
||||
for name in SpecialHeaders.litellm_credential_header_names() - {"authorization"}
|
||||
),
|
||||
}
|
||||
|
||||
app.state.created_aliases = []
|
||||
|
||||
@app.post("/key/generate", operation_id="generate_key_fn_key_generate_post")
|
||||
async def create_key(payload: KeyRequest) -> dict[str, str]:
|
||||
app.state.created_aliases.append(payload.key_alias)
|
||||
if payload.key_alias == "fail-after-write":
|
||||
raise HTTPException(status_code=500)
|
||||
return {"key": "sk-new-key", "key_alias": payload.key_alias}
|
||||
|
||||
@app.post("/mcp")
|
||||
async def existing_mcp() -> dict[str, str]:
|
||||
return {"server": "existing"}
|
||||
|
||||
@app.post("/{server_name}/mcp")
|
||||
async def existing_namespace(server_name: str) -> dict[str, str]:
|
||||
return {"server": server_name}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", [None, "false", "0", "off", "No"])
|
||||
def test_disabled_preserves_existing_admin_namespace(monkeypatch: pytest.MonkeyPatch, enabled: str | None) -> None:
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
monkeypatch.setitem(sys.modules, "litellm_admin_mcp.config", None)
|
||||
if enabled is None:
|
||||
monkeypatch.delenv("LITELLM_ENABLE_ADMIN_MCP", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", enabled)
|
||||
app: Final = FastAPI(lifespan=admin_mcp_lifespan)
|
||||
|
||||
@app.post("/{server_name}/mcp")
|
||||
async def namespace(server_name: str) -> dict[str, str]:
|
||||
return {"server": server_name}
|
||||
|
||||
with TestClient(app) as client:
|
||||
response: Final = client.post("/admin/mcp")
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"server": "admin"}
|
||||
|
||||
|
||||
def test_enabled_without_connector_explains_installation(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(proxy_server, "premium_user", True)
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", "true")
|
||||
monkeypatch.setitem(sys.modules, "litellm_admin_mcp.config", None)
|
||||
with pytest.raises(RuntimeError, match="admin-mcp dependency group"):
|
||||
with TestClient(FastAPI(lifespan=admin_mcp_lifespan)):
|
||||
pytest.fail("Enabling the connector without its dependency must fail startup")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", ["true", "Yes", "on"])
|
||||
def test_unlicensed_opt_in_fails_before_loading_connector(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, enabled: str
|
||||
) -> None:
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", enabled)
|
||||
monkeypatch.setitem(sys.modules, "litellm_admin_mcp.config", None)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
with TestClient(management_app):
|
||||
pytest.fail("An unlicensed deployment must not serve the hosted admin connector")
|
||||
assert exc.value.status_code == 403
|
||||
assert "LITELLM_LICENSE" in str(exc.value.detail)
|
||||
assert all(route.name != "admin_mcp" for route in management_app.routes)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_losing_enterprise_status_blocks_admin_tool_calls(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_ADMIN_TOOLS", "create_key")
|
||||
headers: Final = {"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"}
|
||||
payload: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": "tools/call",
|
||||
"params": {"name": "create_key", "arguments": {"body": {"key_alias": "licensed-write"}}},
|
||||
}
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
licensed: Final = client.post("/admin/mcp", headers=headers, json=payload)
|
||||
assert licensed.status_code == 200, licensed.text
|
||||
assert licensed.json()["result"]["isError"] is False
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
denied: Final = client.post("/admin/mcp", headers=headers, json=payload)
|
||||
assert denied.status_code == 403, denied.text
|
||||
assert "LITELLM_LICENSE" in denied.text
|
||||
assert management_app.state.created_aliases == ["licensed-write"]
|
||||
|
||||
|
||||
def test_invalid_flag_fails_startup(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", "treu")
|
||||
with pytest.raises(ValueError, match="LITELLM_ENABLE_ADMIN_MCP"):
|
||||
with TestClient(FastAPI(lifespan=admin_mcp_lifespan)):
|
||||
pytest.fail("Invalid flag must fail startup")
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize(
|
||||
"authorization,status",
|
||||
[(None, 401), ("Bearer invalid", 401), ("Bearer member", 403), ("Bearer viewer", 403), ("Bearer team-key", 403)],
|
||||
)
|
||||
def test_admin_endpoint_rejects_unauthorized_callers(
|
||||
management_app: FastAPI, authorization: str | None, status: int
|
||||
) -> None:
|
||||
headers: Final = {
|
||||
**({"Authorization": authorization} if authorization else {}),
|
||||
"x-litellm-api-key": "Bearer admin-a",
|
||||
}
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
response: Final = client.post("/admin/mcp", headers=headers, json={})
|
||||
assert response.status_code == status
|
||||
assert response.headers["cache-control"] == "no-store"
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_mount_keeps_existing_mcp_and_manages_restart(management_app: FastAPI) -> None:
|
||||
for _ in range(2):
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
assert client.post("/mcp").json() == {"server": "existing"}
|
||||
assert client.post("/tools/mcp").json() == {"server": "tools"}
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert {tool["name"] for tool in response.json()["result"]["tools"]} == {
|
||||
"list_keys",
|
||||
"describe_admin_tool",
|
||||
"read_admin_result",
|
||||
}
|
||||
assert all(route.name != "admin_mcp" for route in management_app.routes)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize("root_path", ["", "/gateway"])
|
||||
@pytest.mark.parametrize("prefix_mode", ["ingress", "scalar", "multiple"])
|
||||
async def test_concurrent_calls_preserve_identity_network_context_and_strip_cookies(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, root_path: str, prefix_mode: str
|
||||
) -> None:
|
||||
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com" + root_path)
|
||||
monkeypatch.setenv("LITELLM_BASE_URL", "https://must-not-call.example.com")
|
||||
monkeypatch.setenv("LITELLM_API_KEY", "must-not-use-shared-credential")
|
||||
if prefix_mode == "scalar":
|
||||
management_app.root_path = root_path
|
||||
elif prefix_mode == "multiple":
|
||||
management_app.add_middleware(PerRequestRootPathMiddleware, root_paths=("/other", root_path))
|
||||
|
||||
async def call(credential: str) -> dict[str, object]:
|
||||
async with httpx2.AsyncClient(
|
||||
transport=httpx2.ASGITransport(
|
||||
app=management_app,
|
||||
root_path=root_path if prefix_mode == "ingress" else "",
|
||||
client=("198.51.100.7", 4567),
|
||||
),
|
||||
base_url="https://gateway.example.com",
|
||||
) as client:
|
||||
response: Final = await client.post(
|
||||
root_path + "/admin/mcp",
|
||||
headers={
|
||||
**{
|
||||
name: "Bearer member"
|
||||
for name in SpecialHeaders.litellm_credential_header_names() - {"authorization"}
|
||||
},
|
||||
"Authorization": "Bearer " + credential,
|
||||
"Accept": "application/json, text/event-stream",
|
||||
"X-Forwarded-For": "203.0.113.8",
|
||||
"Cookie": "session=must-not-forward",
|
||||
"litellm-changed-by": "forged-actor",
|
||||
"x-litellm-team-id": "policy-team",
|
||||
},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "list_keys"}},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
return response.json()
|
||||
|
||||
async with management_app.router.lifespan_context(management_app):
|
||||
results: Final = await asyncio.gather(call("admin-a"), call("admin-b"))
|
||||
|
||||
for credential, result in zip(("admin-a", "admin-b"), results):
|
||||
assert json.loads(result["result"]["content"][0]["text"]) == {
|
||||
"keys": ["owned-by-a" if credential == "admin-a" else "owned-by-b"],
|
||||
"client": "198.51.100.7",
|
||||
"scheme": "https",
|
||||
"forwarded_for": "203.0.113.8",
|
||||
"cookie": "",
|
||||
"caller_is_changed_by": True,
|
||||
"policy_team": "policy-team",
|
||||
"alternate_credentials_absent": True,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize("header_name", ["X-Custom-Key", "Authorization", "x-litellm-api-key"])
|
||||
@pytest.mark.parametrize("caller_value", [None, "Bearer member"])
|
||||
def test_configured_key_header_uses_mcp_bearer_after_settings_change(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, header_name: str, caller_value: str | None
|
||||
) -> None:
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
monkeypatch.setitem(proxy_server.general_settings, "litellm_key_header_name", header_name)
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={
|
||||
**({header_name: caller_value} if caller_value else {}),
|
||||
"Authorization": "Bearer admin-a",
|
||||
"Accept": "application/json, text/event-stream",
|
||||
},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "list_keys"}},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert json.loads(response.json()["result"]["content"][0]["text"])["keys"] == ["owned-by-a"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("header_name", [
|
||||
"", "bad name", "x-ключ", 7, "Cookie", "litellm-changed-by",
|
||||
"Content-Type", "Host", "X-Forwarded-For", "x-litellm-team-id",
|
||||
*STANDARD_CUSTOMER_ID_HEADERS,
|
||||
])
|
||||
def test_configured_key_header_rejects_invalid_or_reserved_names(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, header_name: object
|
||||
) -> None:
|
||||
monkeypatch.setitem(proxy_server.general_settings, "litellm_key_header_name", header_name)
|
||||
with pytest.raises(ValueError, match="litellm_key_header_name"):
|
||||
with TestClient(management_app):
|
||||
pytest.fail("An ambiguous or malformed credential header must fail startup")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("policy", [
|
||||
{"user_header_name": "X-Custom-Key"},
|
||||
{"user_header_mappings": {"header_name": "X-Custom-Key", "litellm_user_role": "customer"}},
|
||||
{"user_header_mappings": [{"header_name": "X-Custom-Key", "litellm_user_role": "internal_user"}]},
|
||||
{"enable_oauth2_proxy_auth": True, "oauth2_config_mappings": {"user_id": "X-Custom-Key"}},
|
||||
{"mcp_client_id_header": "X-Custom-Key"},
|
||||
])
|
||||
def test_configured_key_header_cannot_replace_configured_identity(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, policy: dict[str, object]
|
||||
) -> None:
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {**policy, "litellm_key_header_name": "x-custom-key"})
|
||||
with pytest.raises(ValueError, match="litellm_key_header_name"):
|
||||
with TestClient(management_app):
|
||||
pytest.fail("A credential header must not replace a configured identity header")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("policy", [
|
||||
{"user_header_name": "x-api-key"},
|
||||
{"user_header_mappings": {"header_name": "Cookie", "litellm_user_role": "customer"}},
|
||||
{"enable_oauth2_proxy_auth": True, "oauth2_config_mappings": {"user_id": "x-litellm-api-key"}},
|
||||
])
|
||||
def test_policy_headers_cannot_share_overwritten_slots_without_a_custom_key(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, policy: dict[str, object]
|
||||
) -> None:
|
||||
monkeypatch.setattr(proxy_server, "general_settings", policy)
|
||||
with pytest.raises(ValueError, match="configured identity headers"):
|
||||
with TestClient(management_app):
|
||||
pytest.fail("Native credentials must not overwrite configured policy headers")
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_compact_results_bind_to_bearer_despite_conflicting_alternate_keys(management_app: FastAPI) -> None:
|
||||
headers: Final = {
|
||||
"Authorization": "Bearer admin-a", "x-litellm-api-key": "Bearer admin-b",
|
||||
"Accept": "application/json, text/event-stream",
|
||||
}
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
saved: Final = client.post(
|
||||
"/admin/mcp", headers=headers,
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {
|
||||
"name": "list_keys", "arguments": {"query": {"large": True}, "response": {"view": "compact"}},
|
||||
}},
|
||||
)
|
||||
assert saved.status_code == 200, saved.text
|
||||
result_id: Final = json.loads(saved.json()["result"]["content"][0]["text"])["result_id"]
|
||||
read_payload: Final = {"jsonrpc": "2.0", "id": 2, "method": "tools/call", "params": {
|
||||
"name": "read_admin_result", "arguments": {"result_id": result_id, "view": "full"},
|
||||
}}
|
||||
replayed: Final = client.post(
|
||||
"/admin/mcp", headers={**headers, "x-litellm-api-key": "Bearer admin-b-limited"}, json=read_payload,
|
||||
)
|
||||
other_bearer: Final = client.post(
|
||||
"/admin/mcp", headers={**headers, "Authorization": "Bearer admin-b"}, json=read_payload,
|
||||
)
|
||||
limited: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin-b-limited", "Accept": headers["Accept"]},
|
||||
json={"jsonrpc": "2.0", "id": 3, "method": "tools/call", "params": {"name": "list_keys"}},
|
||||
)
|
||||
assert replayed.status_code == 200, replayed.text
|
||||
assert json.loads(replayed.json()["result"]["content"][0]["text"]) == {"keys": [{"key_alias": "a" * 20000}]}
|
||||
assert other_bearer.status_code == 200, other_bearer.text
|
||||
assert other_bearer.json()["result"]["isError"] is True
|
||||
assert limited.status_code == 200, limited.text
|
||||
assert limited.json()["result"]["isError"] is True
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize("public_url", ["gateway.example.com", "https:/broken", "http://gateway.example.com"])
|
||||
def test_invalid_public_origin_fails_startup(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, public_url: str
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_MCP_PUBLIC_URL", public_url)
|
||||
with pytest.raises(ValueError, match="HTTPS gateway origin"):
|
||||
with TestClient(management_app):
|
||||
pytest.fail("An invalid trusted public origin must fail startup")
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_public_host_and_origin_are_checked(management_app: FastAPI, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
|
||||
headers: Final = {"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"}
|
||||
with TestClient(management_app, base_url="https://gateway.example.com") as client:
|
||||
hostile_host: Final = client.post("/admin/mcp", headers={**headers, "Host": "hostile.example.com"}, json={})
|
||||
hostile_origin: Final = client.post(
|
||||
"/admin/mcp", headers={**headers, "Origin": "https://hostile.example.com"}, json={}
|
||||
)
|
||||
assert hostile_host.status_code == 421
|
||||
assert hostile_origin.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize("read_only,alias", [(False, "created"), (False, "fail-after-write"), (True, "denied")])
|
||||
def test_writes_respect_read_only_and_are_never_retried(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch, read_only: bool, alias: str
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_ADMIN_TOOLS", "list_keys,create_key")
|
||||
monkeypatch.setenv("LITELLM_ADMIN_READ_ONLY", str(read_only).lower())
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"},
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": "tools/call",
|
||||
"params": {"name": "create_key", "arguments": {"body": {"key_alias": alias}}},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert management_app.state.created_aliases == ([] if read_only else [alias])
|
||||
assert response.json()["result"]["isError"] == (read_only or alias == "fail-after-write")
|
||||
if alias == "created":
|
||||
assert json.loads(response.json()["result"]["content"][0]["text"]) == {"key": "sk-new-key", "key_alias": alias}
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_full_results_do_not_require_worker_affinity(management_app: FastAPI) -> None:
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"},
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": "tools/call",
|
||||
"params": {"name": "list_keys", "arguments": {"query": {"large": True}}},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert json.loads(response.json()["result"]["content"][0]["text"]) == {"keys": [{"key_alias": "a" * 20000}]}
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_nested_management_calls_share_one_admission_slot(management_app: FastAPI) -> None:
|
||||
def no_metrics() -> None:
|
||||
return None
|
||||
|
||||
def single_slot() -> AdmissionControlSettings:
|
||||
return AdmissionControlSettings(1, 0, 1.0)
|
||||
|
||||
state: Final = AdmissionControlState(no_metrics)
|
||||
management_app.add_middleware(
|
||||
AdmissionControlMiddleware,
|
||||
get_settings=single_slot,
|
||||
state=state,
|
||||
)
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "list_keys"}},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["result"]["isError"] is False
|
||||
assert json.loads(response.json()["result"]["content"][0]["text"])["keys"] == ["owned-by-a"]
|
||||
assert state.get_stats() == AdmissionControlStats(0, 0, 0)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
def test_oversized_mcp_body_is_rejected_before_a_write(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_ADMIN_TOOLS", "create_key")
|
||||
with TestClient(management_app, base_url="http://localhost:4000") as client:
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin-a", "Accept": "application/json, text/event-stream"},
|
||||
json={
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/call",
|
||||
"params": {"name": "create_key", "arguments": {"body": {"key_alias": "x" * 300_000}}},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 413, response.text
|
||||
assert management_app.state.created_aliases == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize("peer,mapped_user,bearer,status", [
|
||||
("10.0.0.8", "admin-a", "member", 200),
|
||||
("10.0.0.8", "member", "admin-a", 403),
|
||||
("203.0.113.8", "admin-a", "admin-a", 401),
|
||||
])
|
||||
async def test_oauth2_proxy_auth_preserves_native_identity_and_peer_trust(
|
||||
management_app: FastAPI, monkeypatch: pytest.MonkeyPatch,
|
||||
peer: str, mapped_user: str, bearer: str, status: int,
|
||||
) -> None:
|
||||
users: Final = {
|
||||
user_id: LiteLLM_UserTable(user_id=user_id, user_role=role, teams=[])
|
||||
for user_id, role in (("admin-a", "proxy_admin"), ("member", "internal_user"))
|
||||
}
|
||||
|
||||
class EmptyTable:
|
||||
async def find_many(self, **_: object) -> list[object]:
|
||||
return []
|
||||
|
||||
class DatabaseTables:
|
||||
litellm_teamtable: Final = EmptyTable()
|
||||
litellm_teammembership: Final = EmptyTable()
|
||||
|
||||
class UserInfoDatabase:
|
||||
db: Final = DatabaseTables()
|
||||
|
||||
async def get_data(
|
||||
self, *, user_id: str | None = None, table_name: str | None = None, **_: object
|
||||
) -> LiteLLM_UserTable | list[object] | None:
|
||||
return users.get(user_id) if user_id is not None and table_name is None else []
|
||||
|
||||
cache: Final = UserApiKeyCache()
|
||||
for user_id, user in users.items():
|
||||
cache.set_cache(user_id, user)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", UserInfoDatabase())
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {
|
||||
"enable_oauth2_proxy_auth": True,
|
||||
"trusted_proxy_ranges": ["10.0.0.0/24"],
|
||||
"oauth2_config_mappings": {"user_id": "X-Authenticated-User"},
|
||||
})
|
||||
management_app.router.routes[:] = [
|
||||
route for route in management_app.router.routes
|
||||
if not (isinstance(route, APIRoute) and route.path == "/user/info")
|
||||
]
|
||||
management_app.add_api_route("/user/info", internal_user_endpoints.user_info, methods=["GET"])
|
||||
management_app.add_exception_handler(ProxyException, proxy_server.openai_exception_handler)
|
||||
|
||||
async with management_app.router.lifespan_context(management_app):
|
||||
async with httpx2.AsyncClient(
|
||||
transport=httpx2.ASGITransport(app=management_app, client=(peer, 12345)),
|
||||
base_url="http://localhost:4000",
|
||||
) as client:
|
||||
response: Final = await client.post(
|
||||
"/admin/mcp",
|
||||
headers={
|
||||
"Authorization": "Bearer " + bearer,
|
||||
"X-Authenticated-User": mapped_user,
|
||||
"X-Forwarded-For": "10.0.0.8",
|
||||
"Accept": "application/json, text/event-stream",
|
||||
},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/list"},
|
||||
)
|
||||
assert response.status_code == status, response.text
|
||||
if status == 200:
|
||||
assert "list_keys" in {tool["name"] for tool in response.json()["result"]["tools"]}
|
||||
|
|
@ -32,6 +32,7 @@ from functools import partial
|
|||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse
|
||||
|
|
@ -368,3 +369,54 @@ def test_every_app_mount_is_assigned_to_a_component():
|
|||
f"Add them to GATEWAY_MOUNT_PATHS, BACKEND_MOUNT_PATHS, or serve them "
|
||||
f"from the UI container:\n " + "\n ".join(sorted(unassigned))
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize(
|
||||
"component_lifespan", (None, _gateway_lifespan, _backend_lifespan), ids=("proxy", "gateway", "backend")
|
||||
)
|
||||
@pytest.mark.parametrize("enabled", (False, True))
|
||||
def test_admin_mcp_survives_only_management_component_lifespans(
|
||||
monkeypatch: pytest.MonkeyPatch, component_lifespan: Lifespan[Starlette] | None, enabled: bool
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.admin_mcp import admin_mcp_lifespan
|
||||
|
||||
monkeypatch.setattr(proxy_server, "premium_user", True)
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", str(enabled).lower())
|
||||
for name in (
|
||||
"PROXY_BASE_URL", "LITELLM_MCP_PUBLIC_URL", "LITELLM_ADMIN_TOOLS", "LITELLM_ADMIN_READ_ONLY",
|
||||
"LITELLM_ADMIN_RESPONSE_VIEW", "LITELLM_ADMIN_SCHEMA_MODE",
|
||||
):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(application: FastAPI) -> AsyncGenerator[Mapping[str, object], None]:
|
||||
async with admin_mcp_lifespan(application):
|
||||
yield {"tracing_receiver": None}
|
||||
|
||||
application: Final = FastAPI(lifespan=lifespan)
|
||||
|
||||
@application.get("/user/info")
|
||||
async def user_info() -> dict[str, object]:
|
||||
return {"user_id": "admin", "user_info": {"user_id": "admin", "user_role": "proxy_admin"}}
|
||||
|
||||
if component_lifespan is not None:
|
||||
application.router.lifespan_context = partial(component_lifespan, lifespan=application.router.lifespan_context)
|
||||
for _ in range(2):
|
||||
with TestClient(application, base_url="http://localhost:4000") as client:
|
||||
response: Final = client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin", "Accept": "application/json, text/event-stream"},
|
||||
json={
|
||||
"jsonrpc": "2.0", "id": 1, "method": "initialize",
|
||||
"params": {
|
||||
"protocolVersion": "2025-03-26", "capabilities": {},
|
||||
"clientInfo": {"name": "component-test", "version": "1"},
|
||||
},
|
||||
},
|
||||
)
|
||||
expected_status: Final = 200 if enabled and component_lifespan is not _gateway_lifespan else 404
|
||||
assert response.status_code == expected_status, response.text
|
||||
if response.status_code == 200:
|
||||
assert response.json()["result"]["serverInfo"]["name"] == "litellm-admin-mcp"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
import os
|
||||
import sys
|
||||
import traceback
|
||||
from typing import Final
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Final, Literal
|
||||
from unittest import mock
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
|
@ -30,7 +33,7 @@ logging.basicConfig(
|
|||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi import FastAPI, HTTPException
|
||||
|
||||
# test /chat/completion request to the proxy
|
||||
from fastapi.testclient import TestClient
|
||||
|
|
@ -44,6 +47,191 @@ from litellm.proxy.proxy_server import ( # Replace with the actual module where
|
|||
)
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
@pytest.fixture
|
||||
def admin_mcp_proxy(monkeypatch: pytest.MonkeyPatch) -> FastAPI:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.litellm_license import LicenseCheck
|
||||
|
||||
monkeypatch.setenv("LITELLM_ENABLE_ADMIN_MCP", "true")
|
||||
monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-" + "1234567890abcdef" * 4)
|
||||
for name in ("WORKER_CONFIG", "CONFIG_FILE_PATH", "DATABASE_URL", "LITELLM_LICENSE"):
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
monkeypatch.setattr(proxy_server, "_license_check", LicenseCheck())
|
||||
monkeypatch.setattr(proxy_server, "premium_user", True)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {"disable_model_info_refresh": True})
|
||||
monkeypatch.setattr(proxy_server, "scheduler", None)
|
||||
return FastAPI(lifespan=proxy_server.proxy_startup_event)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
@pytest.mark.parametrize("failure_phase", ["startup", "serving", "shutdown", "cancelled", "license"])
|
||||
async def test_admin_mcp_failure_still_closes_proxy_resources(
|
||||
admin_mcp_proxy: FastAPI,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
failure_phase: Literal["startup", "serving", "shutdown", "cancelled", "license"],
|
||||
) -> None:
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from apscheduler.schedulers.base import STATE_PAUSED
|
||||
from litellm_admin_mcp import server as connector_server
|
||||
from litellm_admin_mcp.gateway import Gateway
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
|
||||
from litellm.proxy.shutdown.scheduled_jobs import AwaitableAsyncIOExecutor
|
||||
|
||||
monkeypatch.setattr(proxy_server, "premium_user", failure_phase != "license")
|
||||
executor: Final = AwaitableAsyncIOExecutor()
|
||||
scheduler: Final = AsyncIOScheduler(executors={"default": executor})
|
||||
monkeypatch.setattr(proxy_server, "scheduler", scheduler)
|
||||
monkeypatch.setattr(proxy_server, "scheduler_executor", executor)
|
||||
scheduler.start()
|
||||
|
||||
@asynccontextmanager
|
||||
async def failing_connector(_app: FastAPI) -> AsyncGenerator[None, None]:
|
||||
if failure_phase == "startup":
|
||||
raise RuntimeError("startup failed")
|
||||
yield
|
||||
assert scheduler.state == STATE_PAUSED
|
||||
assert GracefulShutdownManager.is_shutting_down()
|
||||
assert proxy_server.shared_aiohttp_session is not None
|
||||
assert not proxy_server.shared_aiohttp_session.closed
|
||||
if failure_phase == "shutdown":
|
||||
raise RuntimeError("shutdown failed")
|
||||
|
||||
def connector_app(_gateway: Gateway) -> FastAPI:
|
||||
return FastAPI(lifespan=failing_connector)
|
||||
|
||||
monkeypatch.setattr(connector_server, "create_http_app", connector_app)
|
||||
expected_error: Final = (
|
||||
HTTPException if failure_phase == "license"
|
||||
else asyncio.CancelledError if failure_phase == "cancelled"
|
||||
else RuntimeError
|
||||
)
|
||||
message: Final = "LITELLM_LICENSE" if failure_phase == "license" else f"{failure_phase} failed"
|
||||
async def run_lifespan() -> None:
|
||||
async with admin_mcp_proxy.router.lifespan_context(admin_mcp_proxy) as state:
|
||||
assert state == {"tracing_receiver": None}
|
||||
if failure_phase == "serving":
|
||||
raise RuntimeError("serving failed")
|
||||
if failure_phase == "cancelled":
|
||||
raise asyncio.CancelledError("cancelled failed")
|
||||
|
||||
with pytest.raises(expected_error, match=message):
|
||||
await run_lifespan()
|
||||
|
||||
assert proxy_server.shared_aiohttp_session is not None
|
||||
assert proxy_server.shared_aiohttp_session.closed
|
||||
assert proxy_server.master_key is None
|
||||
assert all(route.name != "admin_mcp" for route in admin_mcp_proxy.routes)
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
async def test_proxy_shutdown_drains_active_admin_tool_before_closing_connector(
|
||||
admin_mcp_proxy: FastAPI, monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
import httpx2
|
||||
|
||||
from litellm.proxy.middleware.in_flight_requests_middleware import InFlightRequestsMiddleware
|
||||
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
|
||||
|
||||
monkeypatch.setenv("LITELLM_ADMIN_TOOLS", "list_teams")
|
||||
admin_mcp_proxy.add_middleware(InFlightRequestsMiddleware)
|
||||
started: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
@admin_mcp_proxy.get("/user/info")
|
||||
async def user_info() -> dict[str, object]:
|
||||
return {"user_id": "admin", "user_info": {"user_id": "admin", "user_role": "proxy_admin"}}
|
||||
|
||||
@admin_mcp_proxy.get("/team/list", operation_id="list_team_team_list_get")
|
||||
async def list_teams() -> dict[str, object]:
|
||||
started.set()
|
||||
await release.wait()
|
||||
return {"teams": ["completed-before-shutdown"]}
|
||||
|
||||
async def complete_during_drain() -> None:
|
||||
async with asyncio.timeout(5):
|
||||
while not GracefulShutdownManager.is_shutting_down():
|
||||
await asyncio.sleep(0)
|
||||
release.set()
|
||||
|
||||
async with httpx2.AsyncClient(
|
||||
transport=httpx2.ASGITransport(app=admin_mcp_proxy), base_url="http://localhost:4000"
|
||||
) as client:
|
||||
async with admin_mcp_proxy.router.lifespan_context(admin_mcp_proxy) as state:
|
||||
assert state == {"tracing_receiver": None}
|
||||
request: Final = asyncio.create_task(client.post(
|
||||
"/admin/mcp",
|
||||
headers={"Authorization": "Bearer admin", "Accept": "application/json, text/event-stream"},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "list_teams"}},
|
||||
))
|
||||
await asyncio.wait_for(started.wait(), timeout=5)
|
||||
completion: Final = asyncio.create_task(complete_during_drain())
|
||||
await asyncio.wait_for(completion, timeout=5)
|
||||
response: Final = await asyncio.wait_for(request, timeout=5)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["result"]["isError"] is False
|
||||
assert json.loads(response.json()["result"]["content"][0]["text"]) == {
|
||||
"teams": ["completed-before-shutdown"]
|
||||
}
|
||||
assert InFlightRequestsMiddleware.get_count() == 0
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="Admin MCP requires Python 3.12+")
|
||||
async def test_proxy_shutdown_closes_admin_connector_when_drain_is_cancelled(
|
||||
admin_mcp_proxy: FastAPI,
|
||||
) -> None:
|
||||
import httpx2
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.middleware.in_flight_requests_middleware import InFlightRequestsMiddleware
|
||||
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
|
||||
|
||||
admin_mcp_proxy.add_middleware(InFlightRequestsMiddleware)
|
||||
ready: Final = asyncio.Event()
|
||||
shutdown: Final = asyncio.Event()
|
||||
started: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
|
||||
@admin_mcp_proxy.get("/hold")
|
||||
async def hold_request() -> dict[str, bool]:
|
||||
started.set()
|
||||
await release.wait()
|
||||
return {"complete": True}
|
||||
|
||||
async def serve() -> None:
|
||||
async with admin_mcp_proxy.router.lifespan_context(admin_mcp_proxy):
|
||||
ready.set()
|
||||
await shutdown.wait()
|
||||
|
||||
async with httpx2.AsyncClient(
|
||||
transport=httpx2.ASGITransport(app=admin_mcp_proxy), base_url="http://localhost:4000"
|
||||
) as client:
|
||||
serving: Final = asyncio.create_task(serve())
|
||||
await asyncio.wait_for(ready.wait(), timeout=5)
|
||||
assert any(route.name == "admin_mcp" for route in admin_mcp_proxy.routes)
|
||||
request: Final = asyncio.create_task(client.get("/hold"))
|
||||
try:
|
||||
await asyncio.wait_for(started.wait(), timeout=5)
|
||||
shutdown.set()
|
||||
async with asyncio.timeout(5):
|
||||
while not GracefulShutdownManager.is_shutting_down():
|
||||
await asyncio.sleep(0)
|
||||
serving.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await serving
|
||||
assert all(route.name != "admin_mcp" for route in admin_mcp_proxy.routes)
|
||||
assert proxy_server.shared_aiohttp_session is not None
|
||||
assert proxy_server.shared_aiohttp_session.closed
|
||||
finally:
|
||||
release.set()
|
||||
await asyncio.wait_for(request, timeout=5)
|
||||
|
||||
assert InFlightRequestsMiddleware.get_count() == 0
|
||||
|
||||
|
||||
# Your bearer token
|
||||
token = "sk-1234"
|
||||
|
||||
|
|
|
|||
|
|
@ -11,7 +11,9 @@ PyPI HTTP responses are mocked — these tests never hit the network.
|
|||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
_CODE_COVERAGE_DIR = os.path.join(
|
||||
|
|
@ -35,7 +37,7 @@ class _FakeResponse:
|
|||
return self._payload
|
||||
|
||||
|
||||
def _make_checker():
|
||||
def _make_checker() -> check_licenses.LicenseChecker:
|
||||
return check_licenses.LicenseChecker(config_file=_LICCHECK_INI)
|
||||
|
||||
|
||||
|
|
@ -280,3 +282,69 @@ def test_check_package_rejects_package_without_license(monkeypatch):
|
|||
)
|
||||
checker = _make_checker()
|
||||
assert checker.check_package("mystery-pkg", "1.0.0") is False
|
||||
|
||||
|
||||
def test_load_requirements_checks_every_dependency_group(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
_ = (tmp_path / "pyproject.toml").write_text(
|
||||
'[project]\ndependencies = ["runtime==1.0"]\n'
|
||||
"[dependency-groups]\n"
|
||||
'admin_mcp = ["connector==2.0"]\n'
|
||||
'proxy = [{include-group = "admin-mcp"}, "server==3.0"]\n'
|
||||
'dev = [{include-group = "proxy"}, "server==3.0"]\n'
|
||||
)
|
||||
_ = (tmp_path / "uv.lock").write_text("package = []\n")
|
||||
checker: Final = _make_checker()
|
||||
|
||||
assert tuple(str(req) for req in checker._load_requirements()) == (
|
||||
"runtime==1.0",
|
||||
"connector==2.0",
|
||||
"server==3.0",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entry", ('{include-group = "missing"}', '{include-group = "dev", unknown = "value"}', "123"))
|
||||
def test_load_requirements_rejects_invalid_group_entries(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, entry: str
|
||||
) -> None:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
_ = (tmp_path / "pyproject.toml").write_text(
|
||||
f"[project]\ndependencies = []\n[dependency-groups]\ndev = [{entry}]\n"
|
||||
)
|
||||
_ = (tmp_path / "uv.lock").write_text("package = []\n")
|
||||
checker: Final = _make_checker()
|
||||
|
||||
with pytest.raises(RuntimeError, match="Invalid dependency group entry"):
|
||||
checker._load_requirements()
|
||||
|
||||
|
||||
def test_load_requirements_preserves_url_hash_and_python_marker(tmp_path: Path) -> None:
|
||||
requirements: Final = tmp_path / "requirements.txt"
|
||||
_ = requirements.write_text(
|
||||
"# pinned connector\n"
|
||||
'connector @ https://example.test/connector.tar.gz#sha256=abcd ; python_version >= "3.12"\n'
|
||||
"requests==2.0 # ordinary comment\n"
|
||||
)
|
||||
checker: Final = _make_checker()
|
||||
connector, registry = checker._load_requirements(requirements)
|
||||
|
||||
assert connector.url == "https://example.test/connector.tar.gz#sha256=abcd"
|
||||
assert str(connector.marker) == 'python_version >= "3.12"'
|
||||
assert str(registry) == "requests==2.0"
|
||||
|
||||
|
||||
def test_license_cli_fails_when_requirements_cannot_be_parsed(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str]
|
||||
) -> None:
|
||||
config: Final = tmp_path / "tests/code_coverage_tests/liccheck.ini"
|
||||
config.parent.mkdir(parents=True)
|
||||
_ = config.write_text(_LICCHECK_INI.read_text())
|
||||
_ = (tmp_path / "requirements.txt").write_text("not a valid requirement\n")
|
||||
monkeypatch.chdir(tmp_path)
|
||||
monkeypatch.setattr(sys, "argv", ["check_licenses.py", "requirements.txt"])
|
||||
|
||||
with pytest.raises(SystemExit) as result:
|
||||
check_licenses.main()
|
||||
|
||||
assert result.value.code == 1
|
||||
assert "Error parsing requirements" in capsys.readouterr().out
|
||||
|
|
|
|||
30
uv.lock
generated
30
uv.lock
generated
|
|
@ -4644,6 +4644,9 @@ utils = [
|
|||
]
|
||||
|
||||
[package.dev-dependencies]
|
||||
admin-mcp = [
|
||||
{ name = "litellm-admin-mcp", marker = "python_full_version >= '3.12'" },
|
||||
]
|
||||
benchmarks = [
|
||||
{ name = "a2a-sdk" },
|
||||
{ name = "mcp" },
|
||||
|
|
@ -4738,6 +4741,7 @@ proxy-dev = [
|
|||
{ name = "a2a-sdk" },
|
||||
{ name = "azure-identity" },
|
||||
{ name = "hypercorn" },
|
||||
{ name = "litellm-admin-mcp", marker = "python_full_version >= '3.12'" },
|
||||
{ name = "opentelemetry-api" },
|
||||
{ name = "opentelemetry-exporter-otlp" },
|
||||
{ name = "opentelemetry-instrumentation-fastapi" },
|
||||
|
|
@ -4860,6 +4864,7 @@ requires-dist = [
|
|||
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "mcp", "saml", "semantic-router", "mlflow", "grpc", "stt-vertex-chirp", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
|
||||
|
||||
[package.metadata.requires-dev]
|
||||
admin-mcp = [{ name = "litellm-admin-mcp", marker = "python_full_version >= '3.12'", url = "https://github.com/BerriAI/liteadmin-mcp/archive/d35ec9c19c117d4c50cc1eccf6ce1296aac25a1a.tar.gz" }]
|
||||
benchmarks = [
|
||||
{ name = "a2a-sdk", specifier = "==1.1.0" },
|
||||
{ name = "mcp", specifier = ">=2.2.0,<3" },
|
||||
|
|
@ -4953,6 +4958,7 @@ proxy-dev = [
|
|||
{ name = "a2a-sdk", specifier = "==1.1.0" },
|
||||
{ name = "azure-identity", specifier = "==1.25.2" },
|
||||
{ name = "hypercorn", specifier = "==0.17.3" },
|
||||
{ name = "litellm-admin-mcp", marker = "python_full_version >= '3.12'", url = "https://github.com/BerriAI/liteadmin-mcp/archive/d35ec9c19c117d4c50cc1eccf6ce1296aac25a1a.tar.gz" },
|
||||
{ name = "opentelemetry-api", specifier = "==1.33.1" },
|
||||
{ name = "opentelemetry-exporter-otlp", specifier = "==1.33.1" },
|
||||
{ name = "opentelemetry-instrumentation-fastapi", specifier = "==0.54b1" },
|
||||
|
|
@ -4962,6 +4968,30 @@ proxy-dev = [
|
|||
{ name = "sentry-sdk", specifier = "==2.21.0" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-admin-mcp"
|
||||
version = "0.1.0"
|
||||
source = { url = "https://github.com/BerriAI/liteadmin-mcp/archive/d35ec9c19c117d4c50cc1eccf6ce1296aac25a1a.tar.gz" }
|
||||
dependencies = [
|
||||
{ name = "httpx2", marker = "python_full_version >= '3.12'" },
|
||||
{ name = "jsonschema", marker = "python_full_version >= '3.12'" },
|
||||
{ name = "mcp", marker = "python_full_version >= '3.12'" },
|
||||
{ name = "uvicorn", marker = "python_full_version >= '3.12'" },
|
||||
]
|
||||
sdist = { hash = "sha256:86b930f6706fb2da10d53d1798ee7fd14e44fb0afdef0da122cf0e6bb88fa2c3" }
|
||||
|
||||
[package.metadata]
|
||||
requires-dist = [
|
||||
{ name = "aiohttp", marker = "extra == 'test'", specifier = "==3.14.3" },
|
||||
{ name = "httpx2", specifier = "==2.13.0" },
|
||||
{ name = "jsonschema", specifier = "==4.26.0" },
|
||||
{ name = "mcp", specifier = "==2.2.0" },
|
||||
{ name = "pytest", marker = "extra == 'test'", specifier = "==9.1.1" },
|
||||
{ name = "pytest-asyncio", marker = "extra == 'test'", specifier = "==1.4.0" },
|
||||
{ name = "uvicorn", specifier = ">=0.30,<1" },
|
||||
]
|
||||
provides-extras = ["test"]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.73"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue