feat(provisioning-mcp): add MCP server for on-the-fly LiteLLM e2e deployments

Adds a standalone MCP server (provisioning_mcp/) that lets authenticated AI
agents provision ephemeral LiteLLM deployments from the helm/litellm chart for
end-to-end testing. Given a repo URL and git revision it derives the component
image repos, optionally stands up throwaway in-cluster Postgres/Redis, and runs
helm upgrade --install in the litellm namespace.

- OAuth 2.0 resource-server auth: validates JWT bearer tokens against the
  issuer's JWKS (signature, iss, aud, exp, required scope).
- Tools: provision / delete / status / list deployments.
- Dockerfile pins + checksum-verifies helm and kubectl; namespaced RBAC and
  Deployment manifests under deploy/.
- Unit tests cover auth, naming/registry derivation, manifests, and the
  provisioning flow (35 tests).

https://claude.ai/code/session_019WDsdJNGjNyUNrso3xigSV
This commit is contained in:
Claude 2026-05-24 00:07:15 +00:00
parent 5f73ad4fe7
commit af14aa6818
No known key found for this signature in database
22 changed files with 1887 additions and 0 deletions

5
provisioning_mcp/.gitignore vendored Normal file
View file

@ -0,0 +1,5 @@
__pycache__/
*.egg-info/
.pytest_cache/
build/
dist/

View file

@ -0,0 +1,55 @@
# LiteLLM provisioning MCP server.
#
# Build from the repository ROOT so the helm/litellm chart is in context:
# docker build -f provisioning_mcp/Dockerfile -t litellm-provisioning-mcp .
# ---- tools: fetch + checksum-verify pinned helm and kubectl ----
FROM debian:bookworm-slim AS tools
ARG HELM_VERSION=3.16.3
ARG KUBECTL_VERSION=1.31.3
# Pinned digest of the linux/amd64 kubectl binary (dl.k8s.io .../kubectl.sha256).
ARG KUBECTL_SHA256=981f6b49577068bc174275184d8ee7105d8e54f40733792c519cd85023984c0f
RUN apt-get update \
&& apt-get install -y --no-install-recommends ca-certificates curl tar \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /tmp/tools
# helm: download the release tarball and its official sha256sum sidecar, then
# verify before extracting (no piping a remote script into a shell).
RUN curl -fsSLO "https://get.helm.sh/helm-v${HELM_VERSION}-linux-amd64.tar.gz" \
&& curl -fsSLO "https://get.helm.sh/helm-v${HELM_VERSION}-linux-amd64.tar.gz.sha256sum" \
&& EXPECTED="$(awk '{print $1}' "helm-v${HELM_VERSION}-linux-amd64.tar.gz.sha256sum")" \
&& echo "${EXPECTED} helm-v${HELM_VERSION}-linux-amd64.tar.gz" | sha256sum -c - \
&& tar -xzf "helm-v${HELM_VERSION}-linux-amd64.tar.gz" \
&& install -m 0755 linux-amd64/helm /usr/local/bin/helm
# kubectl: download the binary and verify against the hardcoded digest.
RUN curl -fsSLo kubectl "https://dl.k8s.io/release/v${KUBECTL_VERSION}/bin/linux/amd64/kubectl" \
&& echo "${KUBECTL_SHA256} kubectl" | sha256sum -c - \
&& install -m 0755 kubectl /usr/local/bin/kubectl
# ---- runtime ----
FROM python:3.11-slim AS runtime
COPY --from=tools /usr/local/bin/helm /usr/local/bin/helm
COPY --from=tools /usr/local/bin/kubectl /usr/local/bin/kubectl
WORKDIR /app
COPY provisioning_mcp/pyproject.toml /app/provisioning_mcp/pyproject.toml
COPY provisioning_mcp/src /app/provisioning_mcp/src
RUN pip install --no-cache-dir /app/provisioning_mcp
# The chart this server provisions, baked into the image.
COPY helm/litellm /app/helm/litellm
ENV HELM_CHART_PATH=/app/helm/litellm
RUN useradd --create-home --uid 10001 mcp && chown -R mcp:mcp /app
USER mcp
ENV MCP_HOST=0.0.0.0 \
MCP_PORT=8080
EXPOSE 8080
ENTRYPOINT ["litellm-provisioning-mcp"]

View file

@ -0,0 +1,74 @@
apiVersion: apps/v1
kind: Deployment
metadata:
name: litellm-provisioning-mcp
namespace: litellm
labels:
app.kubernetes.io/name: litellm-provisioning-mcp
spec:
replicas: 1
selector:
matchLabels:
app.kubernetes.io/name: litellm-provisioning-mcp
template:
metadata:
labels:
app.kubernetes.io/name: litellm-provisioning-mcp
spec:
serviceAccountName: litellm-provisioning-mcp
containers:
- name: mcp
# Replace with your published image (built from provisioning_mcp/Dockerfile).
image: ghcr.io/berriai/litellm-provisioning-mcp:latest
imagePullPolicy: IfNotPresent
ports:
- name: http
containerPort: 8080
env:
# ---- required OAuth 2.0 resource-server settings (not secret) ----
- name: OAUTH_JWKS_URL
value: "https://YOUR_IDP/.well-known/jwks.json"
- name: OAUTH_ISSUER
value: "https://YOUR_IDP/"
- name: OAUTH_AUDIENCE
value: "litellm-provisioning-mcp"
- name: MCP_RESOURCE_SERVER_URL
value: "https://mcp.example.com"
# ---- optional (defaults shown) ----
- name: OAUTH_REQUIRED_SCOPE
value: "litellm:provision"
- name: LITELLM_NAMESPACE
value: "litellm"
- name: LITELLM_IMAGE_REGISTRY
value: "ghcr.io/berriai"
readinessProbe:
tcpSocket: { port: http }
initialDelaySeconds: 5
periodSeconds: 10
livenessProbe:
tcpSocket: { port: http }
initialDelaySeconds: 10
periodSeconds: 20
resources:
requests:
cpu: 100m
memory: 128Mi
limits:
cpu: 500m
memory: 256Mi
---
apiVersion: v1
kind: Service
metadata:
name: litellm-provisioning-mcp
namespace: litellm
labels:
app.kubernetes.io/name: litellm-provisioning-mcp
spec:
type: ClusterIP
selector:
app.kubernetes.io/name: litellm-provisioning-mcp
ports:
- name: http
port: 8080
targetPort: http

View file

@ -0,0 +1,6 @@
apiVersion: v1
kind: Namespace
metadata:
name: litellm
labels:
app.kubernetes.io/managed-by: litellm-provisioning-mcp

View file

@ -0,0 +1,54 @@
# Namespaced RBAC: the provisioning MCP can manage only the `litellm` namespace.
# No cluster-scoped permissions are granted, so the namespace must already exist
# (see namespace.yaml) — the server never uses `helm --create-namespace`.
apiVersion: v1
kind: ServiceAccount
metadata:
name: litellm-provisioning-mcp
namespace: litellm
---
apiVersion: rbac.authorization.k8s.io/v1
kind: Role
metadata:
name: litellm-provisioning-mcp
namespace: litellm
rules:
# Core objects created by the chart, the migration job, and the ephemeral
# datastores. Secrets cover both the generated master key / DB credentials and
# helm's own release-state storage.
- apiGroups: [""]
resources:
- pods
- pods/log
- services
- configmaps
- secrets
- serviceaccounts
- persistentvolumeclaims
verbs: ["get", "list", "watch", "create", "update", "patch", "delete"]
- apiGroups: ["apps"]
resources: ["deployments", "replicasets"]
verbs: ["get", "list", "watch", "create", "update", "patch", "delete"]
- apiGroups: ["batch"]
resources: ["jobs"]
verbs: ["get", "list", "watch", "create", "update", "patch", "delete"]
- apiGroups: ["autoscaling"]
resources: ["horizontalpodautoscalers"]
verbs: ["get", "list", "watch", "create", "update", "patch", "delete"]
- apiGroups: ["networking.k8s.io"]
resources: ["ingresses"]
verbs: ["get", "list", "watch", "create", "update", "patch", "delete"]
---
apiVersion: rbac.authorization.k8s.io/v1
kind: RoleBinding
metadata:
name: litellm-provisioning-mcp
namespace: litellm
roleRef:
apiGroup: rbac.authorization.k8s.io
kind: Role
name: litellm-provisioning-mcp
subjects:
- kind: ServiceAccount
name: litellm-provisioning-mcp
namespace: litellm

View file

@ -0,0 +1,31 @@
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[project]
name = "litellm-provisioning-mcp"
version = "0.1.0"
description = "MCP server that provisions ephemeral LiteLLM deployments from the helm/litellm chart for end-to-end testing."
requires-python = ">=3.11"
dependencies = [
"mcp>=1.26.0",
"pyjwt[crypto]>=2.8.0",
"pyyaml>=6.0",
"uvicorn>=0.30.0",
]
[project.optional-dependencies]
test = [
"pytest>=8.0.0",
"pytest-asyncio>=0.23.0",
]
[project.scripts]
litellm-provisioning-mcp = "litellm_provisioning_mcp.server:main"
[tool.hatch.build.targets.wheel]
packages = ["src/litellm_provisioning_mcp"]
[tool.pytest.ini_options]
asyncio_mode = "auto"
testpaths = ["tests"]

View file

@ -0,0 +1,5 @@
"""MCP server for provisioning ephemeral LiteLLM deployments via helm."""
__all__ = ["__version__"]
__version__ = "0.1.0"

View file

@ -0,0 +1,111 @@
"""OAuth 2.0 bearer-token validation (resource-server role).
The MCP server is a pure resource server: it never issues tokens. It validates
the signed JWT access tokens that an external authorization server issued to
the calling AI agent, checking signature (against the issuer's JWKS), issuer,
audience, expiry, and a required scope.
This module intentionally has no dependency on the ``mcp`` package so the
validation logic can be unit-tested in isolation; the thin adapter that maps a
``VerifiedToken`` onto ``mcp``'s ``AccessToken`` lives in ``server.py``.
"""
from __future__ import annotations
from dataclasses import dataclass, field
import jwt
from jwt import PyJWKClient
class TokenValidationError(Exception):
"""Raised when a presented bearer token is missing, malformed, or invalid."""
@dataclass(frozen=True)
class VerifiedToken:
subject: str
client_id: str
scopes: list[str] = field(default_factory=list)
expires_at: int | None = None
def _extract_scopes(claims: dict) -> list[str]:
"""Collect scopes across the claim shapes used by common IdPs.
OAuth standard ``scope`` is a space-delimited string; Azure AD uses ``scp``
(string) and ``roles`` (list); Auth0 uses ``permissions`` (list).
"""
scopes: set[str] = set()
for key in ("scope", "scp"):
value = claims.get(key)
if isinstance(value, str):
scopes.update(value.split())
elif isinstance(value, list):
scopes.update(str(item) for item in value)
for key in ("roles", "permissions"):
value = claims.get(key)
if isinstance(value, list):
scopes.update(str(item) for item in value)
return sorted(scopes)
class JWKSValidator:
"""Validates JWT access tokens against an issuer's published JWKS.
The underlying ``PyJWKClient`` fetches the JWKS lazily and caches signing
keys, so steady-state validation does no network IO. ``validate`` is
blocking (JWKS refresh + crypto); callers on an event loop should run it in
a worker thread.
"""
def __init__(
self,
*,
jwks_url: str,
issuer: str,
audience: str,
algorithms: tuple[str, ...],
required_scope: str | None = None,
) -> None:
self._issuer = issuer
self._audience = audience
self._algorithms = list(algorithms)
self._required_scope = required_scope
self._jwk_client = PyJWKClient(jwks_url, cache_keys=True, lifespan=600)
def validate(self, token: str) -> VerifiedToken:
if not token:
raise TokenValidationError("empty bearer token")
try:
signing_key = self._jwk_client.get_signing_key_from_jwt(token)
except Exception as exc: # PyJWK raises several distinct error types
raise TokenValidationError(f"unable to resolve signing key: {exc}") from exc
try:
claims = jwt.decode(
token,
signing_key.key,
algorithms=self._algorithms,
audience=self._audience,
issuer=self._issuer,
options={"require": ["exp", "iss", "aud"]},
)
except jwt.PyJWTError as exc:
raise TokenValidationError(f"token rejected: {exc}") from exc
scopes = _extract_scopes(claims)
if self._required_scope and self._required_scope not in scopes:
raise TokenValidationError(
f"token is missing required scope '{self._required_scope}'"
)
return VerifiedToken(
subject=str(claims.get("sub", "")),
client_id=str(
claims.get("client_id") or claims.get("azp") or claims.get("sub", "")
),
scopes=scopes,
expires_at=claims.get("exp"),
)

View file

@ -0,0 +1,64 @@
"""Async subprocess execution for the ``helm`` and ``kubectl`` CLIs.
The runner is a plain coroutine so tests can inject a fake in place of real
process execution. Commands are always invoked as an argv list (never through a
shell) to avoid command injection from caller-supplied values.
"""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
@dataclass(frozen=True)
class CommandResult:
returncode: int
stdout: str
stderr: str
class CommandError(RuntimeError):
def __init__(self, args: list[str], result: CommandResult) -> None:
self.args_list = args
self.result = result
super().__init__(
f"command failed (exit {result.returncode}): {' '.join(args)}\n{result.stderr.strip()}"
)
class CommandTimeout(RuntimeError):
def __init__(self, args: list[str], timeout: float) -> None:
self.args_list = args
self.timeout = timeout
super().__init__(f"command timed out after {timeout}s: {' '.join(args)}")
async def run_command(
args: list[str],
*,
input_text: str | None = None,
timeout: float,
) -> CommandResult:
"""Execute ``args`` and return its result. Never raises on non-zero exit."""
proc = await asyncio.create_subprocess_exec(
*args,
stdin=asyncio.subprocess.PIPE if input_text is not None else None,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
stdin_bytes = input_text.encode() if input_text is not None else None
try:
stdout, stderr = await asyncio.wait_for(
proc.communicate(input=stdin_bytes), timeout=timeout
)
except asyncio.TimeoutError as exc:
proc.kill()
await proc.wait()
raise CommandTimeout(args, timeout) from exc
return CommandResult(
returncode=proc.returncode or 0,
stdout=stdout.decode(errors="replace"),
stderr=stderr.decode(errors="replace"),
)

View file

@ -0,0 +1,93 @@
"""Environment-driven configuration for the LiteLLM provisioning MCP server.
Every setting is read once at startup from the process environment. Required
OAuth settings raise at load time so a misconfigured deployment fails fast
instead of accepting unauthenticated traffic.
"""
from __future__ import annotations
import os
from dataclasses import dataclass
class ConfigError(RuntimeError):
"""Raised when a required setting is missing or malformed."""
def _require(name: str) -> str:
value = os.environ.get(name, "").strip()
if not value:
raise ConfigError(f"environment variable {name} is required")
return value
def _optional(name: str, default: str) -> str:
value = os.environ.get(name, "").strip()
return value or default
def _int(name: str, default: int) -> int:
raw = os.environ.get(name, "").strip()
if not raw:
return default
try:
return int(raw)
except ValueError as exc:
raise ConfigError(f"environment variable {name} must be an integer") from exc
@dataclass(frozen=True)
class Settings:
# ---- OAuth 2.0 resource-server validation ----
oauth_jwks_url: str
oauth_issuer: str
oauth_audience: str
oauth_required_scope: str
oauth_algorithms: tuple[str, ...]
resource_server_url: str
# ---- provisioning targets ----
namespace: str
chart_path: str
default_image_registry: str
release_prefix: str
helm_binary: str
kubectl_binary: str
command_timeout: int
# ---- HTTP server ----
host: str
port: int
log_level: str
@classmethod
def from_env(cls) -> "Settings":
algorithms = tuple(
alg.strip()
for alg in _optional("OAUTH_ALGORITHMS", "RS256").split(",")
if alg.strip()
)
if not algorithms:
raise ConfigError("OAUTH_ALGORITHMS must list at least one algorithm")
return cls(
oauth_jwks_url=_require("OAUTH_JWKS_URL"),
oauth_issuer=_require("OAUTH_ISSUER"),
oauth_audience=_require("OAUTH_AUDIENCE"),
oauth_required_scope=_optional("OAUTH_REQUIRED_SCOPE", "litellm:provision"),
oauth_algorithms=algorithms,
resource_server_url=_require("MCP_RESOURCE_SERVER_URL"),
namespace=_optional("LITELLM_NAMESPACE", "litellm"),
chart_path=_optional("HELM_CHART_PATH", "/app/helm/litellm"),
default_image_registry=_optional(
"LITELLM_IMAGE_REGISTRY", "ghcr.io/berriai"
),
release_prefix=_optional("LITELLM_RELEASE_PREFIX", "litellm-e2e"),
helm_binary=_optional("HELM_BINARY", "helm"),
kubectl_binary=_optional("KUBECTL_BINARY", "kubectl"),
command_timeout=_int("COMMAND_TIMEOUT_SECONDS", 600),
host=_optional("MCP_HOST", "0.0.0.0"),
port=_int("MCP_PORT", 8080),
log_level=_optional("MCP_LOG_LEVEL", "INFO").upper(),
)

View file

@ -0,0 +1,67 @@
"""Thin async wrapper around the ``helm`` CLI, scoped to a single namespace."""
from __future__ import annotations
import json
from typing import Awaitable, Callable
from .commands import CommandError, CommandResult, run_command
Runner = Callable[..., Awaitable[CommandResult]]
class HelmRunner:
def __init__(
self,
*,
namespace: str,
binary: str = "helm",
runner: Runner = run_command,
wait_timeout: int = 600,
) -> None:
self._namespace = namespace
self._binary = binary
self._runner = runner
self._wait_timeout = wait_timeout
# Allow helm's --wait to run to completion before the subprocess is killed.
self._command_timeout = wait_timeout + 60
async def _run(
self, args: list[str], *, input_text: str | None = None
) -> CommandResult:
full = [self._binary, *args, "--namespace", self._namespace]
result = await self._runner(
full, input_text=input_text, timeout=self._command_timeout
)
if result.returncode != 0:
raise CommandError(full, result)
return result
async def upgrade_install(
self, *, release: str, chart_path: str, values_yaml: str
) -> CommandResult:
return await self._run(
[
"upgrade",
release,
chart_path,
"--install",
"--values",
"-", # read merged values from stdin
"--wait",
"--timeout",
f"{self._wait_timeout}s",
],
input_text=values_yaml,
)
async def uninstall(self, *, release: str) -> CommandResult:
return await self._run(["uninstall", release, "--ignore-not-found", "--wait"])
async def status(self, *, release: str) -> dict:
result = await self._run(["status", release, "--output", "json"])
return json.loads(result.stdout)
async def list_releases(self) -> list[dict]:
result = await self._run(["list", "--output", "json"])
return json.loads(result.stdout)

View file

@ -0,0 +1,60 @@
"""Thin async wrapper around the ``kubectl`` CLI, scoped to a single namespace."""
from __future__ import annotations
import json
from typing import Awaitable, Callable
from .commands import CommandError, CommandResult, run_command
Runner = Callable[..., Awaitable[CommandResult]]
class KubectlRunner:
def __init__(
self,
*,
namespace: str,
binary: str = "kubectl",
runner: Runner = run_command,
timeout: int = 180,
) -> None:
self._namespace = namespace
self._binary = binary
self._runner = runner
self._timeout = timeout
async def _run(
self, args: list[str], *, input_text: str | None = None
) -> CommandResult:
full = [self._binary, "--namespace", self._namespace, *args]
result = await self._runner(full, input_text=input_text, timeout=self._timeout)
if result.returncode != 0:
raise CommandError(full, result)
return result
async def apply(self, manifest_yaml: str) -> CommandResult:
return await self._run(["apply", "--filename", "-"], input_text=manifest_yaml)
async def wait_available(self, *, deployment: str, timeout: int) -> CommandResult:
return await self._run(
[
"wait",
f"deployment/{deployment}",
"--for=condition=Available",
f"--timeout={timeout}s",
]
)
async def delete_by_label(
self, *, selector: str, kinds: list[str]
) -> CommandResult:
return await self._run(
["delete", ",".join(kinds), "--selector", selector, "--ignore-not-found"]
)
async def get_pods(self, *, selector: str) -> list[dict]:
result = await self._run(
["get", "pods", "--selector", selector, "--output", "json"]
)
return json.loads(result.stdout).get("items", [])

View file

@ -0,0 +1,58 @@
"""Helpers for deriving release names, label values, and image repositories."""
from __future__ import annotations
import re
_NON_DNS = re.compile(r"[^a-z0-9-]+")
_GITHUB = re.compile(
r"(?:https?://|git@|ssh://git@)?github\.com[/:]([^/]+)/([^/]+?)(?:\.git)?/?$"
)
# Helm release names are used as the prefix of Kubernetes object names, which
# are RFC 1123 labels capped at 63 chars; helm itself caps the release name at
# 53 to leave room for the suffixes the chart appends.
RELEASE_MAX_LEN = 53
LABEL_MAX_LEN = 63
def sanitize_release_name(value: str) -> str:
cleaned = _NON_DNS.sub("-", value.lower()).strip("-")
cleaned = re.sub(r"-{2,}", "-", cleaned)[:RELEASE_MAX_LEN].strip("-")
if not cleaned:
raise ValueError(f"cannot derive a valid release name from {value!r}")
return cleaned
def sanitize_label(value: str) -> str:
"""Coerce ``value`` into a valid Kubernetes label value (best effort)."""
cleaned = re.sub(r"[^A-Za-z0-9_.-]+", "-", value).strip("-_.")
return cleaned[:LABEL_MAX_LEN].strip("-_.")
def registry_from_repo_url(repo_url: str) -> str | None:
"""Map a GitHub repo URL to its GHCR org, e.g. ``ghcr.io/<owner>``.
Encodes the e2e convention that a fork's CI publishes the litellm component
images to its own ``ghcr.io/<owner>`` namespace. Returns ``None`` for URLs
that are not GitHub repositories so the caller can fall back to a default.
"""
match = _GITHUB.search(repo_url.strip())
if not match:
return None
owner = match.group(1).lower()
return f"ghcr.io/{owner}"
def derive_image_repos(
*, repo_url: str, registry_override: str | None, default_registry: str
) -> dict[str, str]:
registry = (
registry_override or registry_from_repo_url(repo_url) or default_registry
).rstrip("/")
return {
"gateway": f"{registry}/litellm-gateway",
"backend": f"{registry}/litellm-backend",
"ui": f"{registry}/litellm-ui",
"migrations": f"{registry}/litellm-migrations",
}

View file

@ -0,0 +1,265 @@
"""Orchestrates ephemeral LiteLLM deployments via the ``helm/litellm`` chart.
A single ``provision`` call: mints a master-key Secret, optionally stands up a
throwaway Postgres and/or Redis, then ``helm upgrade --install``s the chart
pinned to the requested revision's images. Teardown removes both the helm
release and the auxiliary objects.
This module is free of any ``mcp`` dependency so it can be exercised directly
in unit tests with fake command runners.
"""
from __future__ import annotations
import copy
from dataclasses import dataclass, field
import yaml
from . import resources
from .config import Settings
from .helm import HelmRunner
from .kubectl import KubectlRunner
from .naming import derive_image_repos, sanitize_label, sanitize_release_name
from .resources import DatabaseConnection, RedisConnection
_DATASTORE_WAIT_SECONDS = 120
class ProvisionError(RuntimeError):
"""Raised for caller-correctable problems (bad input, missing dependency)."""
@dataclass(frozen=True)
class ProvisionRequest:
repo_url: str
revision: str
release_name: str | None = None
enable_redis: bool = False
enable_postgres: bool = True
enable_ui: bool = False
service_account: str | None = None
image_registry: str | None = None
external_database: dict | None = None
extra_values: dict = field(default_factory=dict)
def _deep_merge(base: dict, overrides: dict) -> dict:
result = copy.deepcopy(base)
for key, value in overrides.items():
if isinstance(value, dict) and isinstance(result.get(key), dict):
result[key] = _deep_merge(result[key], value)
else:
result[key] = copy.deepcopy(value)
return result
def _external_database(spec: dict) -> DatabaseConnection:
missing = [k for k in ("host", "dbname", "secret_name") if not spec.get(k)]
if missing:
raise ProvisionError(
f"external_database is missing required keys: {', '.join(missing)}"
)
return DatabaseConnection(
host=str(spec["host"]),
port=int(spec.get("port", 5432)),
dbname=str(spec["dbname"]),
secret_name=str(spec["secret_name"]),
username_key=str(spec.get("username_key", "username")),
password_key=str(spec.get("password_key", "password")),
)
def _summarize_pods(items: list[dict]) -> list[dict]:
summary = []
for pod in items:
statuses = pod.get("status", {}).get("containerStatuses", []) or []
summary.append(
{
"name": pod.get("metadata", {}).get("name", ""),
"phase": pod.get("status", {}).get("phase", "Unknown"),
"ready": bool(statuses) and all(c.get("ready") for c in statuses),
"restarts": sum(c.get("restartCount", 0) for c in statuses),
}
)
return summary
class Provisioner:
def __init__(
self, settings: Settings, *, helm: HelmRunner, kubectl: KubectlRunner
) -> None:
self._settings = settings
self._helm = helm
self._kubectl = kubectl
def _build_values(
self,
*,
release: str,
image_repos: dict[str, str],
revision: str,
master_key_secret: str,
database: DatabaseConnection,
redis_conn: RedisConnection | None,
request: ProvisionRequest,
) -> dict:
values: dict = {
"fullnameOverride": release,
"masterKey": {"secretName": master_key_secret, "secretKey": "master-key"},
"database": {
"writer": {
"host": database.host,
"port": database.port,
"dbname": database.dbname,
"passwordSecret": {
"name": database.secret_name,
"usernameKey": database.username_key,
"passwordKey": database.password_key,
},
}
},
"gateway": {
"image": {"repository": image_repos["gateway"], "tag": revision}
},
"backend": {
"image": {"repository": image_repos["backend"], "tag": revision}
},
"ui": {
"enabled": request.enable_ui,
"image": {"repository": image_repos["ui"], "tag": revision},
},
"migrationJob": {
"image": {"repository": image_repos["migrations"], "tag": revision}
},
}
if redis_conn is not None:
values["redis"] = {"host": redis_conn.host, "port": redis_conn.port}
if request.service_account:
values["serviceAccount"] = {
"create": False,
"name": request.service_account,
}
if request.extra_values:
values = _deep_merge(values, request.extra_values)
return values
async def provision(self, request: ProvisionRequest) -> dict:
if not request.repo_url.strip():
raise ProvisionError("repo_url is required")
if not request.revision.strip():
raise ProvisionError("revision is required")
release = sanitize_release_name(
request.release_name
or f"{self._settings.release_prefix}-{request.revision}"
)
image_repos = derive_image_repos(
repo_url=request.repo_url,
registry_override=request.image_registry,
default_registry=self._settings.default_image_registry,
)
mk_manifest, mk_secret = resources.master_key_secret(release)
await self._kubectl.apply(mk_manifest)
if request.enable_postgres:
pg_manifest, database = resources.postgres(release)
await self._kubectl.apply(pg_manifest)
await self._kubectl.wait_available(
deployment=f"{release}-postgres", timeout=_DATASTORE_WAIT_SECONDS
)
elif request.external_database:
database = _external_database(request.external_database)
else:
raise ProvisionError(
"a database is required: set enable_postgres=true or provide external_database"
)
redis_conn: RedisConnection | None = None
if request.enable_redis:
redis_manifest, redis_conn = resources.redis(release)
await self._kubectl.apply(redis_manifest)
await self._kubectl.wait_available(
deployment=f"{release}-redis", timeout=_DATASTORE_WAIT_SECONDS
)
values = self._build_values(
release=release,
image_repos=image_repos,
revision=sanitize_label(request.revision),
master_key_secret=mk_secret,
database=database,
redis_conn=redis_conn,
request=request,
)
await self._helm.upgrade_install(
release=release,
chart_path=self._settings.chart_path,
values_yaml=yaml.safe_dump(values, sort_keys=False),
)
pods = await self._kubectl.get_pods(
selector=f"app.kubernetes.io/instance={release}"
)
ns = self._settings.namespace
endpoints = {"gateway": f"http://{release}-gateway.{ns}.svc.cluster.local:4000"}
if request.enable_ui:
endpoints["ui"] = f"http://{release}-ui.{ns}.svc.cluster.local:3000"
endpoints["backend"] = f"http://{release}-backend.{ns}.svc.cluster.local:4001"
return {
"success": True,
"release": release,
"namespace": ns,
"revision": request.revision,
"images": image_repos,
"postgres_enabled": request.enable_postgres,
"redis_enabled": request.enable_redis,
"ui_enabled": request.enable_ui,
"endpoints": endpoints,
"pods": _summarize_pods(pods),
}
async def delete(self, release_name: str) -> dict:
release = sanitize_release_name(release_name)
await self._helm.uninstall(release=release)
await self._kubectl.delete_by_label(
selector=resources.selector_label(release),
kinds=["deployment", "service", "secret"],
)
return {
"success": True,
"release": release,
"namespace": self._settings.namespace,
}
async def status(self, release_name: str) -> dict:
release = sanitize_release_name(release_name)
helm_status = await self._helm.status(release=release)
pods = await self._kubectl.get_pods(
selector=f"app.kubernetes.io/instance={release}"
)
info = helm_status.get("info", {})
return {
"release": release,
"namespace": self._settings.namespace,
"status": info.get("status"),
"last_deployed": info.get("last_deployed"),
"pods": _summarize_pods(pods),
}
async def list_deployments(self) -> dict:
releases = await self._helm.list_releases()
return {
"namespace": self._settings.namespace,
"deployments": [
{
"release": r.get("name"),
"status": r.get("status"),
"chart": r.get("chart"),
"updated": r.get("updated"),
}
for r in releases
],
}

View file

@ -0,0 +1,237 @@
"""Builders for the auxiliary Kubernetes objects a provisioned release needs.
These are deliberately ephemeral (``emptyDir`` storage, no persistence): they
exist only to back short-lived end-to-end test deployments. Every object
carries the managed-by / release labels so it can be garbage-collected by
selector when the release is torn down.
"""
from __future__ import annotations
import secrets
from dataclasses import dataclass
import yaml
MANAGED_BY = "litellm-provisioning-mcp"
RELEASE_LABEL = "litellm.ai/release"
# Pinned, widely-mirrored images for the throwaway datastores.
POSTGRES_IMAGE = "postgres:16-alpine"
REDIS_IMAGE = "redis:7-alpine"
@dataclass(frozen=True)
class DatabaseConnection:
host: str
port: int
dbname: str
secret_name: str
username_key: str
password_key: str
@dataclass(frozen=True)
class RedisConnection:
host: str
port: int
def common_labels(release: str) -> dict[str, str]:
return {
"app.kubernetes.io/managed-by": MANAGED_BY,
# Match the chart's instance label so a single selector covers both the
# chart pods and these auxiliary datastores when reporting status.
"app.kubernetes.io/instance": release,
RELEASE_LABEL: release,
}
def selector_label(release: str) -> str:
return f"{RELEASE_LABEL}={release}"
def _dump(*objects: dict) -> str:
return yaml.safe_dump_all(objects, default_flow_style=False, sort_keys=False)
def master_key_secret(release: str) -> tuple[str, str]:
"""Return ``(manifest, secret_name)`` for a freshly generated master key."""
name = f"{release}-master-key"
master_key = "sk-" + secrets.token_urlsafe(32)
manifest = {
"apiVersion": "v1",
"kind": "Secret",
"metadata": {"name": name, "labels": common_labels(release)},
"type": "Opaque",
"stringData": {"master-key": master_key},
}
return _dump(manifest), name
def postgres(release: str) -> tuple[str, DatabaseConnection]:
name = f"{release}-postgres"
labels = {**common_labels(release), "app.kubernetes.io/component": "postgres"}
password = secrets.token_urlsafe(24)
secret = {
"apiVersion": "v1",
"kind": "Secret",
"metadata": {"name": name, "labels": labels},
"type": "Opaque",
"stringData": {
"username": "litellm",
"password": password,
"dbname": "litellm",
},
}
deployment = {
"apiVersion": "apps/v1",
"kind": "Deployment",
"metadata": {"name": name, "labels": labels},
"spec": {
"replicas": 1,
"selector": {
"matchLabels": {
RELEASE_LABEL: release,
"app.kubernetes.io/component": "postgres",
}
},
"template": {
"metadata": {"labels": labels},
"spec": {
"containers": [
{
"name": "postgres",
"image": POSTGRES_IMAGE,
"ports": [{"containerPort": 5432}],
"env": [
{
"name": "POSTGRES_USER",
"valueFrom": {
"secretKeyRef": {
"name": name,
"key": "username",
}
},
},
{
"name": "POSTGRES_PASSWORD",
"valueFrom": {
"secretKeyRef": {
"name": name,
"key": "password",
}
},
},
{"name": "POSTGRES_DB", "value": "litellm"},
{
"name": "PGDATA",
"value": "/var/lib/postgresql/data/pgdata",
},
],
"readinessProbe": {
"exec": {
"command": [
"pg_isready",
"-U",
"litellm",
"-d",
"litellm",
]
},
"initialDelaySeconds": 5,
"periodSeconds": 5,
},
"resources": {
"requests": {"cpu": "100m", "memory": "256Mi"},
"limits": {"cpu": "500m", "memory": "512Mi"},
},
"volumeMounts": [
{
"name": "data",
"mountPath": "/var/lib/postgresql/data",
}
],
}
],
"volumes": [{"name": "data", "emptyDir": {}}],
},
},
},
}
service = {
"apiVersion": "v1",
"kind": "Service",
"metadata": {"name": name, "labels": labels},
"spec": {
"selector": {
RELEASE_LABEL: release,
"app.kubernetes.io/component": "postgres",
},
"ports": [{"port": 5432, "targetPort": 5432}],
},
}
conn = DatabaseConnection(
host=name,
port=5432,
dbname="litellm",
secret_name=name,
username_key="username",
password_key="password",
)
return _dump(secret, deployment, service), conn
def redis(release: str) -> tuple[str, RedisConnection]:
name = f"{release}-redis"
labels = {**common_labels(release), "app.kubernetes.io/component": "redis"}
deployment = {
"apiVersion": "apps/v1",
"kind": "Deployment",
"metadata": {"name": name, "labels": labels},
"spec": {
"replicas": 1,
"selector": {
"matchLabels": {
RELEASE_LABEL: release,
"app.kubernetes.io/component": "redis",
}
},
"template": {
"metadata": {"labels": labels},
"spec": {
"containers": [
{
"name": "redis",
"image": REDIS_IMAGE,
"ports": [{"containerPort": 6379}],
"readinessProbe": {
"exec": {"command": ["redis-cli", "ping"]},
"initialDelaySeconds": 5,
"periodSeconds": 5,
},
"resources": {
"requests": {"cpu": "50m", "memory": "64Mi"},
"limits": {"cpu": "250m", "memory": "256Mi"},
},
}
],
},
},
},
}
service = {
"apiVersion": "v1",
"kind": "Service",
"metadata": {"name": name, "labels": labels},
"spec": {
"selector": {
RELEASE_LABEL: release,
"app.kubernetes.io/component": "redis",
},
"ports": [{"port": 6379, "targetPort": 6379}],
},
}
return _dump(deployment, service), RedisConnection(host=name, port=6379)

View file

@ -0,0 +1,191 @@
"""FastMCP server exposing LiteLLM provisioning tools over streamable HTTP.
Authentication: the server is an OAuth 2.0 resource server. Every request must
carry a ``Bearer`` access token that validates against the configured issuer's
JWKS and carries the required scope. Tokens are never issued here.
"""
from __future__ import annotations
import logging
import anyio
from mcp.server.auth.provider import AccessToken
from mcp.server.auth.settings import AuthSettings
from mcp.server.fastmcp import FastMCP
from .auth import JWKSValidator, TokenValidationError
from .commands import CommandError, CommandTimeout
from .config import Settings
from .helm import HelmRunner
from .kubectl import KubectlRunner
from .provisioner import Provisioner, ProvisionError, ProvisionRequest
logger = logging.getLogger("litellm_provisioning_mcp")
class JWTTokenVerifier:
"""Adapts :class:`JWKSValidator` to the MCP ``TokenVerifier`` protocol."""
def __init__(self, validator: JWKSValidator, *, resource: str) -> None:
self._validator = validator
self._resource = resource
async def verify_token(self, token: str) -> AccessToken | None:
try:
verified = await anyio.to_thread.run_sync(self._validator.validate, token)
except TokenValidationError as exc:
logger.warning("rejected bearer token: %s", exc)
return None
return AccessToken(
token=token,
client_id=verified.client_id,
scopes=verified.scopes,
expires_at=verified.expires_at,
resource=self._resource,
)
def _error(exc: Exception) -> dict:
if isinstance(exc, CommandError):
return {
"success": False,
"error": "command_failed",
"returncode": exc.result.returncode,
"detail": exc.result.stderr.strip() or exc.result.stdout.strip(),
}
if isinstance(exc, CommandTimeout):
return {"success": False, "error": "timeout", "detail": str(exc)}
if isinstance(exc, ProvisionError):
return {"success": False, "error": "invalid_request", "detail": str(exc)}
raise exc
def build_server(settings: Settings) -> FastMCP:
validator = JWKSValidator(
jwks_url=settings.oauth_jwks_url,
issuer=settings.oauth_issuer,
audience=settings.oauth_audience,
algorithms=settings.oauth_algorithms,
required_scope=settings.oauth_required_scope,
)
provisioner = Provisioner(
settings,
helm=HelmRunner(
namespace=settings.namespace,
binary=settings.helm_binary,
wait_timeout=settings.command_timeout,
),
kubectl=KubectlRunner(
namespace=settings.namespace, binary=settings.kubectl_binary
),
)
mcp = FastMCP(
name="litellm-provisioning",
instructions=(
"Provision and tear down ephemeral LiteLLM deployments for end-to-end "
"testing. Provide the litellm repo URL and the git revision whose images "
"should run; optionally enable a throwaway Postgres and/or Redis."
),
host=settings.host,
port=settings.port,
log_level=settings.log_level,
stateless_http=True,
json_response=True,
token_verifier=JWTTokenVerifier(
validator, resource=settings.resource_server_url
),
auth=AuthSettings(
issuer_url=settings.oauth_issuer,
resource_server_url=settings.resource_server_url,
required_scopes=[settings.oauth_required_scope],
),
)
@mcp.tool()
async def provision_litellm_deployment(
repo_url: str,
revision: str,
release_name: str | None = None,
enable_postgres: bool = True,
enable_redis: bool = False,
enable_ui: bool = False,
service_account: str | None = None,
image_registry: str | None = None,
external_database: dict | None = None,
extra_values: dict | None = None,
) -> dict:
"""Provision (or upgrade) a LiteLLM deployment in the target namespace.
Args:
repo_url: Git URL of the litellm repository under test (used to derive
the image registry, e.g. a fork's ``ghcr.io/<owner>``).
revision: Git revision; used as the container image tag for every
component, so an image with this tag must already be published.
release_name: Optional helm release name (defaults to a sanitized
``<prefix>-<revision>``). Re-using a name upgrades in place.
enable_postgres: Stand up a throwaway in-cluster Postgres (default).
enable_redis: Stand up a throwaway in-cluster Redis.
enable_ui: Also deploy the dashboard UI component.
service_account: Existing ServiceAccount the litellm pods should run as.
image_registry: Override the derived image registry base.
external_database: Use an existing DB instead of an ephemeral one. Keys:
host, dbname, secret_name (required); port, username_key, password_key.
extra_values: Helm values deep-merged last (escape hatch).
"""
request = ProvisionRequest(
repo_url=repo_url,
revision=revision,
release_name=release_name,
enable_postgres=enable_postgres,
enable_redis=enable_redis,
enable_ui=enable_ui,
service_account=service_account,
image_registry=image_registry,
external_database=external_database,
extra_values=extra_values or {},
)
try:
return await provisioner.provision(request)
except (ProvisionError, CommandError, CommandTimeout) as exc:
return _error(exc)
@mcp.tool()
async def delete_litellm_deployment(release_name: str) -> dict:
"""Tear down a deployment: uninstall the helm release and its datastores."""
try:
return await provisioner.delete(release_name)
except (ProvisionError, CommandError, CommandTimeout) as exc:
return _error(exc)
@mcp.tool()
async def get_litellm_deployment_status(release_name: str) -> dict:
"""Report helm release status and pod readiness for a deployment."""
try:
return await provisioner.status(release_name)
except (ProvisionError, CommandError, CommandTimeout) as exc:
return _error(exc)
@mcp.tool()
async def list_litellm_deployments() -> dict:
"""List the helm releases in the target namespace."""
try:
return await provisioner.list_deployments()
except (CommandError, CommandTimeout) as exc:
return _error(exc)
return mcp
def main() -> None:
settings = Settings.from_env()
logging.basicConfig(
level=settings.log_level,
format="%(asctime)s %(levelname)s %(name)s %(message)s",
)
build_server(settings).run(transport="streamable-http")
if __name__ == "__main__":
main()

View file

@ -0,0 +1,57 @@
import json
from typing import Any
from litellm_provisioning_mcp.commands import CommandResult
from litellm_provisioning_mcp.config import Settings
def make_settings(**overrides: Any) -> Settings:
base = dict(
oauth_jwks_url="https://idp/jwks",
oauth_issuer="https://idp/",
oauth_audience="litellm-provisioning-mcp",
oauth_required_scope="litellm:provision",
oauth_algorithms=("RS256",),
resource_server_url="https://mcp.example.com",
namespace="litellm",
chart_path="/app/helm/litellm",
default_image_registry="ghcr.io/berriai",
release_prefix="litellm-e2e",
helm_binary="helm",
kubectl_binary="kubectl",
command_timeout=600,
host="0.0.0.0",
port=8080,
log_level="INFO",
)
base.update(overrides)
return Settings(**base)
class FakeRunner:
"""Records invocations and returns canned results keyed by command shape."""
def __init__(self) -> None:
self.calls: list[tuple[list[str], str | None]] = []
async def __call__(self, args, *, input_text=None, timeout=None) -> CommandResult:
self.calls.append((args, input_text))
binary = args[0]
if binary.endswith("kubectl"):
if "get" in args and "pods" in args:
return CommandResult(0, json.dumps({"items": []}), "")
return CommandResult(0, "applied", "")
# helm
if "status" in args:
return CommandResult(0, json.dumps({"info": {"status": "deployed"}}), "")
if "list" in args:
return CommandResult(
0, json.dumps([{"name": "litellm-e2e-abc", "status": "deployed"}]), ""
)
return CommandResult(0, "ok", "")
def find(self, *needles: str) -> tuple[list[str], str | None]:
for args, input_text in self.calls:
if all(n in args for n in needles):
return args, input_text
raise AssertionError(f"no recorded call matching {needles}; calls={self.calls}")

View file

@ -0,0 +1,123 @@
import time
from types import SimpleNamespace
import jwt
import pytest
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from litellm_provisioning_mcp.auth import (
JWKSValidator,
TokenValidationError,
_extract_scopes,
)
ISSUER = "https://idp/"
AUDIENCE = "litellm-provisioning-mcp"
def _keypair():
priv = rsa.generate_private_key(public_exponent=65537, key_size=2048)
pem = priv.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
)
return priv, pem
def _validator(public_key, *, required_scope=None):
v = JWKSValidator(
jwks_url="https://idp/jwks",
issuer=ISSUER,
audience=AUDIENCE,
algorithms=("RS256",),
required_scope=required_scope,
)
v._jwk_client = SimpleNamespace(
get_signing_key_from_jwt=lambda token: SimpleNamespace(key=public_key)
)
return v
def _token(signing_pem, **claims):
payload = {
"iss": ISSUER,
"aud": AUDIENCE,
"exp": int(time.time()) + 3600,
"sub": "agent-1",
}
payload.update(claims)
return jwt.encode(payload, signing_pem, algorithm="RS256")
def test_valid_token_returns_subject_and_scopes():
priv, pem = _keypair()
v = _validator(priv.public_key(), required_scope="litellm:provision")
token = _token(pem, scope="litellm:provision other:scope", client_id="cli-42")
result = v.validate(token)
assert result.subject == "agent-1"
assert result.client_id == "cli-42"
assert "litellm:provision" in result.scopes
def test_missing_required_scope_is_rejected():
priv, pem = _keypair()
v = _validator(priv.public_key(), required_scope="litellm:provision")
token = _token(pem, scope="some:other")
with pytest.raises(TokenValidationError, match="missing required scope"):
v.validate(token)
def test_expired_token_is_rejected():
priv, pem = _keypair()
v = _validator(priv.public_key())
token = _token(pem, exp=int(time.time()) - 10)
with pytest.raises(TokenValidationError):
v.validate(token)
def test_wrong_audience_is_rejected():
priv, pem = _keypair()
v = _validator(priv.public_key())
token = _token(pem, aud="someone-else")
with pytest.raises(TokenValidationError):
v.validate(token)
def test_wrong_issuer_is_rejected():
priv, pem = _keypair()
v = _validator(priv.public_key())
token = _token(pem, iss="https://evil/")
with pytest.raises(TokenValidationError):
v.validate(token)
def test_bad_signature_is_rejected():
priv, _ = _keypair()
attacker_priv, attacker_pem = _keypair()
# Validator trusts `priv`'s public key, but the token is signed by attacker.
v = _validator(priv.public_key())
token = _token(attacker_pem)
with pytest.raises(TokenValidationError):
v.validate(token)
def test_empty_token_is_rejected():
priv, _ = _keypair()
v = _validator(priv.public_key())
with pytest.raises(TokenValidationError):
v.validate("")
def test_extract_scopes_across_claim_shapes():
assert _extract_scopes({"scope": "a b"}) == ["a", "b"]
assert _extract_scopes({"scp": ["c", "d"]}) == ["c", "d"]
assert _extract_scopes({"permissions": ["e"], "roles": ["f"]}) == ["e", "f"]

View file

@ -0,0 +1,38 @@
import pytest
from conftest import FakeRunner
from litellm_provisioning_mcp.commands import CommandError, CommandResult
from litellm_provisioning_mcp.helm import HelmRunner
async def test_upgrade_install_builds_command_and_passes_values_via_stdin():
runner = FakeRunner()
helm = HelmRunner(namespace="litellm", runner=runner, wait_timeout=300)
await helm.upgrade_install(
release="rel", chart_path="/chart", values_yaml="key: value\n"
)
args, stdin = runner.calls[0]
assert args[:4] == ["helm", "upgrade", "rel", "/chart"]
assert "--install" in args
assert args[-2:] == ["--namespace", "litellm"]
assert "--values" in args and "-" in args
assert "--timeout" in args and "300s" in args
assert stdin == "key: value\n"
async def test_status_parses_json():
runner = FakeRunner()
helm = HelmRunner(namespace="litellm", runner=runner)
status = await helm.status(release="rel")
assert status["info"]["status"] == "deployed"
async def test_nonzero_exit_raises_command_error():
async def runner(args, *, input_text=None, timeout=None):
return CommandResult(1, "", "release not found")
helm = HelmRunner(namespace="litellm", runner=runner)
with pytest.raises(CommandError, match="release not found"):
await helm.status(release="missing")

View file

@ -0,0 +1,71 @@
import pytest
from litellm_provisioning_mcp.naming import (
RELEASE_MAX_LEN,
derive_image_repos,
registry_from_repo_url,
sanitize_label,
sanitize_release_name,
)
def test_sanitize_release_name_lowercases_and_replaces_invalid():
assert sanitize_release_name("Feature/My_Branch") == "feature-my-branch"
def test_sanitize_release_name_truncates_and_strips():
long = "litellm-e2e-" + "a" * 80
out = sanitize_release_name(long)
assert len(out) <= RELEASE_MAX_LEN
assert not out.endswith("-")
def test_sanitize_release_name_empty_raises():
with pytest.raises(ValueError):
sanitize_release_name("///")
def test_sanitize_label_strips_invalid_edges():
assert sanitize_label("-abc/def-") == "abc-def"
@pytest.mark.parametrize(
"url,expected",
[
("https://github.com/BerriAI/litellm", "ghcr.io/berriai"),
("https://github.com/BerriAI/litellm.git", "ghcr.io/berriai"),
("git@github.com:Alice/litellm.git", "ghcr.io/alice"),
("github.com/Bob/litellm", "ghcr.io/bob"),
("https://gitlab.com/x/y", None),
],
)
def test_registry_from_repo_url(url, expected):
assert registry_from_repo_url(url) == expected
def test_derive_image_repos_override_wins():
repos = derive_image_repos(
repo_url="https://github.com/BerriAI/litellm",
registry_override="myreg.io/team",
default_registry="ghcr.io/berriai",
)
assert repos["gateway"] == "myreg.io/team/litellm-gateway"
assert repos["migrations"] == "myreg.io/team/litellm-migrations"
def test_derive_image_repos_derives_from_fork():
repos = derive_image_repos(
repo_url="https://github.com/Alice/litellm",
registry_override=None,
default_registry="ghcr.io/berriai",
)
assert repos["backend"] == "ghcr.io/alice/litellm-backend"
def test_derive_image_repos_falls_back_to_default():
repos = derive_image_repos(
repo_url="https://example.com/not-github",
registry_override=None,
default_registry="ghcr.io/berriai",
)
assert repos["ui"] == "ghcr.io/berriai/litellm-ui"

View file

@ -0,0 +1,176 @@
import pytest
import yaml
from conftest import FakeRunner, make_settings
from litellm_provisioning_mcp.commands import CommandError, CommandResult
from litellm_provisioning_mcp.helm import HelmRunner
from litellm_provisioning_mcp.kubectl import KubectlRunner
from litellm_provisioning_mcp.provisioner import (
ProvisionError,
ProvisionRequest,
Provisioner,
)
def _provisioner(runner: FakeRunner) -> Provisioner:
settings = make_settings()
return Provisioner(
settings,
helm=HelmRunner(namespace="litellm", runner=runner, wait_timeout=600),
kubectl=KubectlRunner(namespace="litellm", runner=runner),
)
def _helm_values(runner: FakeRunner) -> dict:
_, values_yaml = runner.find("upgrade")
return yaml.safe_load(values_yaml)
async def test_provision_with_ephemeral_postgres_and_redis():
runner = FakeRunner()
provisioner = _provisioner(runner)
result = await provisioner.provision(
ProvisionRequest(
repo_url="https://github.com/BerriAI/litellm",
revision="abc123",
enable_postgres=True,
enable_redis=True,
)
)
assert result["success"] is True
assert result["release"] == "litellm-e2e-abc123"
assert result["images"]["gateway"] == "ghcr.io/berriai/litellm-gateway"
values = _helm_values(runner)
assert values["fullnameOverride"] == "litellm-e2e-abc123"
assert values["gateway"]["image"]["tag"] == "abc123"
assert values["gateway"]["image"]["repository"] == "ghcr.io/berriai/litellm-gateway"
assert values["database"]["writer"]["host"] == "litellm-e2e-abc123-postgres"
assert values["redis"]["host"] == "litellm-e2e-abc123-redis"
assert values["masterKey"]["secretName"] == "litellm-e2e-abc123-master-key"
assert values["ui"]["enabled"] is False
# ephemeral datastores were applied and waited on
runner.find("apply")
runner.find("wait", "deployment/litellm-e2e-abc123-postgres")
runner.find("wait", "deployment/litellm-e2e-abc123-redis")
async def test_provision_without_db_raises():
runner = FakeRunner()
provisioner = _provisioner(runner)
with pytest.raises(ProvisionError, match="database is required"):
await provisioner.provision(
ProvisionRequest(
repo_url="https://github.com/BerriAI/litellm",
revision="abc",
enable_postgres=False,
)
)
async def test_provision_with_external_database():
runner = FakeRunner()
provisioner = _provisioner(runner)
result = await provisioner.provision(
ProvisionRequest(
repo_url="https://github.com/BerriAI/litellm",
revision="v1",
enable_postgres=False,
external_database={
"host": "pg.prod.internal",
"dbname": "litellm",
"secret_name": "prod-db",
},
)
)
assert result["success"] is True
values = _helm_values(runner)
assert values["database"]["writer"]["host"] == "pg.prod.internal"
assert values["database"]["writer"]["passwordSecret"]["name"] == "prod-db"
# no ephemeral postgres applied
with pytest.raises(AssertionError):
runner.find("wait", "deployment/litellm-e2e-v1-postgres")
async def test_external_database_missing_keys_raises():
runner = FakeRunner()
provisioner = _provisioner(runner)
with pytest.raises(ProvisionError, match="missing required keys"):
await provisioner.provision(
ProvisionRequest(
repo_url="https://github.com/BerriAI/litellm",
revision="v1",
enable_postgres=False,
external_database={"host": "pg"},
)
)
async def test_image_registry_override_and_extra_values():
runner = FakeRunner()
provisioner = _provisioner(runner)
await provisioner.provision(
ProvisionRequest(
repo_url="https://github.com/BerriAI/litellm",
revision="sha",
image_registry="myreg.io/team",
service_account="litellm-workload",
extra_values={"gateway": {"numWorkers": 4}},
)
)
values = _helm_values(runner)
assert values["gateway"]["image"]["repository"] == "myreg.io/team/litellm-gateway"
assert values["serviceAccount"] == {"create": False, "name": "litellm-workload"}
# extra_values deep-merge preserves the derived image block
assert values["gateway"]["numWorkers"] == 4
assert values["gateway"]["image"]["tag"] == "sha"
async def test_delete_uninstalls_and_cleans_datastores():
runner = FakeRunner()
provisioner = _provisioner(runner)
result = await provisioner.delete("litellm-e2e-abc")
assert result["success"] is True
runner.find("uninstall", "litellm-e2e-abc")
args, _ = runner.find("delete", "--selector")
assert "litellm.ai/release=litellm-e2e-abc" in args
async def test_status_reports_helm_and_pods():
runner = FakeRunner()
provisioner = _provisioner(runner)
result = await provisioner.status("litellm-e2e-abc")
assert result["status"] == "deployed"
assert result["pods"] == []
async def test_provision_propagates_helm_failure():
runner = FakeRunner()
async def failing(args, *, input_text=None, timeout=None):
await runner(args, input_text=input_text, timeout=timeout)
if args[0].endswith("helm") and "upgrade" in args:
return CommandResult(1, "", "boom: chart not found")
return CommandResult(0, "", "")
settings = make_settings()
provisioner = Provisioner(
settings,
helm=HelmRunner(namespace="litellm", runner=failing, wait_timeout=600),
kubectl=KubectlRunner(namespace="litellm", runner=failing),
)
with pytest.raises(CommandError, match="boom"):
await provisioner.provision(
ProvisionRequest(
repo_url="https://github.com/BerriAI/litellm", revision="x"
)
)

View file

@ -0,0 +1,46 @@
import yaml
from litellm_provisioning_mcp import resources
def _load(manifest: str) -> list[dict]:
return list(yaml.safe_load_all(manifest))
def test_master_key_secret_is_valid_and_prefixed():
manifest, name = resources.master_key_secret("rel")
docs = _load(manifest)
assert name == "rel-master-key"
assert len(docs) == 1
secret = docs[0]
assert secret["kind"] == "Secret"
assert secret["stringData"]["master-key"].startswith("sk-")
assert secret["metadata"]["labels"]["litellm.ai/release"] == "rel"
def test_postgres_manifests_and_connection():
manifest, conn = resources.postgres("rel")
docs = _load(manifest)
kinds = sorted(d["kind"] for d in docs)
assert kinds == ["Deployment", "Secret", "Service"]
assert conn.host == "rel-postgres"
assert conn.port == 5432
assert conn.dbname == "litellm"
assert conn.secret_name == "rel-postgres"
# Deployment selector is a subset of the pod labels.
deployment = next(d for d in docs if d["kind"] == "Deployment")
selector = deployment["spec"]["selector"]["matchLabels"]
pod_labels = deployment["spec"]["template"]["metadata"]["labels"]
assert selector.items() <= pod_labels.items()
def test_redis_manifests_and_connection():
manifest, conn = resources.redis("rel")
docs = _load(manifest)
assert sorted(d["kind"] for d in docs) == ["Deployment", "Service"]
assert conn.host == "rel-redis"
assert conn.port == 6379
def test_selector_label_targets_release():
assert resources.selector_label("rel") == "litellm.ai/release=rel"