feat(proxy): embed enterprise LiteAdmin MCP in LiteLLM images (#44610)

This commit is contained in:
tin-berri 2026-10-05 15:24:34 -07:00 • committed by GitHub
parent 739192ec19
commit 9062fd3931
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 1602 additions and 72 deletions

View file

@ -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

View file

@ -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 \

View file

@ -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 \

View file

@ -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
}

View file

@ -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 \

View file

@ -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 \

View file

@ -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:

View 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
View 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]

View file

@ -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()

View file

@ -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:

View file

@ -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

View file

@ -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",

View file

@ -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 = [

View file

@ -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

View 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}"
)

View file

@ -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)

View 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"]}

View file

@ -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"

View file

@ -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"

View file

@ -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
View file

@ -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"