diff --git a/.github/workflows/image-scan.yml b/.github/workflows/image-scan.yml index e72b8230232..0dc629cb3b0 100644 --- a/.github/workflows/image-scan.yml +++ b/.github/workflows/image-scan.yml @@ -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 diff --git a/Dockerfile b/Dockerfile index be507f6efb4..194ac1f46d9 100644 --- a/Dockerfile +++ b/Dockerfile @@ -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 \ diff --git a/backend/Dockerfile b/backend/Dockerfile index dfff6e71a46..5213bb00c50 100644 --- a/backend/Dockerfile +++ b/backend/Dockerfile @@ -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 \ diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 2c651277dab..8a96663bd50 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -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 } diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 3309fdd5341..b850d43d8e9 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -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 \ diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index bafd1af46d1..3e3f6e0f720 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -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 \ diff --git a/docker/README.md b/docker/README.md index 376dc7b2d97..ac82c6caebf 100644 --- a/docker/README.md +++ b/docker/README.md @@ -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: diff --git a/gateway/Dockerfile b/gateway/Dockerfile index 8045a8b64cb..ffee6a9cdbb 100644 --- a/gateway/Dockerfile +++ b/gateway/Dockerfile @@ -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 \ diff --git a/litellm/proxy/admin_mcp.py b/litellm/proxy/admin_mcp.py new file mode 100644 index 00000000000..6371863f45c --- /dev/null +++ b/litellm/proxy/admin_mcp.py @@ -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] diff --git a/litellm/proxy/middleware/admission_control_middleware.py b/litellm/proxy/middleware/admission_control_middleware.py index c336b97349a..1bb94e23d15 100644 --- a/litellm/proxy/middleware/admission_control_middleware.py +++ b/litellm/proxy/middleware/admission_control_middleware.py @@ -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() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6d1a04c2b5e..d3d12df9feb 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 29f2f46f001..5c8edf34917 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index f04e66a04eb..3722dbca350 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/tests/code_coverage_tests/check_licenses.py b/tests/code_coverage_tests/check_licenses.py index a9eddc3fabb..81213e9ce79 100644 --- a/tests/code_coverage_tests/check_licenses.py +++ b/tests/code_coverage_tests/check_licenses.py @@ -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 = [ diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini index a62af1b3725..b66493386c1 100644 --- a/tests/code_coverage_tests/liccheck.ini +++ b/tests/code_coverage_tests/liccheck.ini @@ -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 diff --git a/tests/proxy_migration_tests/test_image_admin_mcp.py b/tests/proxy_migration_tests/test_image_admin_mcp.py new file mode 100644 index 00000000000..1a8b3f3008c --- /dev/null +++ b/tests/proxy_migration_tests/test_image_admin_mcp.py @@ -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}" + ) diff --git a/tests/unit/proxy/middleware/test_admission_control_middleware.py b/tests/unit/proxy/middleware/test_admission_control_middleware.py index f1ca13daa03..e449ebc83bf 100644 --- a/tests/unit/proxy/middleware/test_admission_control_middleware.py +++ b/tests/unit/proxy/middleware/test_admission_control_middleware.py @@ -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) diff --git a/tests/unit/proxy/test_admin_mcp.py b/tests/unit/proxy/test_admin_mcp.py new file mode 100644 index 00000000000..dc956038525 --- /dev/null +++ b/tests/unit/proxy/test_admin_mcp.py @@ -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"]} diff --git a/tests/unit/proxy/test_component_allowlists.py b/tests/unit/proxy/test_component_allowlists.py index 1a210fcb445..6347b326342 100644 --- a/tests/unit/proxy/test_component_allowlists.py +++ b/tests/unit/proxy/test_component_allowlists.py @@ -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" diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index 300edc8e435..d6a2de157ff 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -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" diff --git a/tests/unit/test_check_licenses.py b/tests/unit/test_check_licenses.py index 1218e44fade..b02ae436fe9 100644 --- a/tests/unit/test_check_licenses.py +++ b/tests/unit/test_check_licenses.py @@ -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 diff --git a/uv.lock b/uv.lock index a5438188d4a..ccea0afb717 100644 --- a/uv.lock +++ b/uv.lock @@ -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"