diff --git a/provisioning_mcp/src/litellm_provisioning_mcp/config.py b/provisioning_mcp/src/litellm_provisioning_mcp/config.py index cdf7128d8f6..380129629dc 100644 --- a/provisioning_mcp/src/litellm_provisioning_mcp/config.py +++ b/provisioning_mcp/src/litellm_provisioning_mcp/config.py @@ -52,6 +52,8 @@ class Settings: chart_path: str default_image_registry: str release_prefix: str + provisioning_service_account: str + allowed_service_accounts: tuple[str, ...] helm_binary: str kubectl_binary: str command_timeout: int @@ -84,6 +86,14 @@ class Settings: "LITELLM_IMAGE_REGISTRY", "ghcr.io/berriai" ), release_prefix=_optional("LITELLM_RELEASE_PREFIX", "litellm-e2e"), + provisioning_service_account=_optional( + "PROVISIONING_SERVICE_ACCOUNT", "litellm-provisioning-mcp" + ), + allowed_service_accounts=tuple( + sa.strip() + for sa in _optional("ALLOWED_SERVICE_ACCOUNTS", "").split(",") + if sa.strip() + ), helm_binary=_optional("HELM_BINARY", "helm"), kubectl_binary=_optional("KUBECTL_BINARY", "kubectl"), command_timeout=_int("COMMAND_TIMEOUT_SECONDS", 600), diff --git a/provisioning_mcp/src/litellm_provisioning_mcp/kubectl.py b/provisioning_mcp/src/litellm_provisioning_mcp/kubectl.py index 91bf5e8b9bf..3248069b4ba 100644 --- a/provisioning_mcp/src/litellm_provisioning_mcp/kubectl.py +++ b/provisioning_mcp/src/litellm_provisioning_mcp/kubectl.py @@ -58,3 +58,15 @@ class KubectlRunner: ["get", "pods", "--selector", selector, "--output", "json"] ) return json.loads(result.stdout).get("items", []) + + async def resource_exists(self, *, kind: str, name: str) -> bool: + result = await self._run( + ["get", kind, name, "--ignore-not-found", "--output", "name"] + ) + return bool(result.stdout.strip()) + + async def count_by_label(self, *, kind: str, selector: str) -> int: + result = await self._run( + ["get", kind, "--selector", selector, "--output", "json"] + ) + return len(json.loads(result.stdout).get("items", [])) diff --git a/provisioning_mcp/src/litellm_provisioning_mcp/provisioner.py b/provisioning_mcp/src/litellm_provisioning_mcp/provisioner.py index c8a64f57a24..32053a76325 100644 --- a/provisioning_mcp/src/litellm_provisioning_mcp/provisioner.py +++ b/provisioning_mcp/src/litellm_provisioning_mcp/provisioner.py @@ -21,7 +21,12 @@ 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 +from .resources import ( + MANAGED_BY, + RELEASE_LABEL, + DatabaseConnection, + RedisConnection, +) _DATASTORE_WAIT_SECONDS = 120 @@ -144,25 +149,79 @@ class Provisioner: values = _deep_merge(values, request.extra_values) return values + def _resolve_release_name(self, request: ProvisionRequest) -> str: + """Force every provisioned release into the ephemeral prefix namespace. + + Without this, a caller could pass ``release_name="litellm"`` and have + ``helm upgrade --install`` silently overwrite a real, unmanaged release + that shares the namespace. + """ + prefix = self._settings.release_prefix + base = request.release_name or request.revision + if not base.startswith(prefix): + base = f"{prefix}-{base}" + release = sanitize_release_name(base) + if release != prefix and not release.startswith(f"{prefix}-"): + raise ProvisionError( + f"release name must start with '{prefix}-'; got '{release}'" + ) + return release + + def _validate_service_account(self, service_account: str) -> None: + if service_account == self._settings.provisioning_service_account: + raise ProvisionError( + "service_account must not be the provisioning server's own " + "account (a deployed workload would inherit its permissions)" + ) + allowed = self._settings.allowed_service_accounts + if allowed and service_account not in allowed: + raise ProvisionError( + f"service_account '{service_account}' is not in the allowed list" + ) + + async def _helm_release_exists(self, release: str) -> bool: + releases = await self._helm.list_releases() + return any(r.get("name") == release for r in releases) + 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}" - ) + release = self._resolve_release_name(request) + if request.service_account: + self._validate_service_account(request.service_account) + + # Refuse to touch a release that already exists but wasn't created by + # this tool (defense in depth on top of the prefix constraint). + if await self._helm_release_exists(release) and not await self._owns_release( + release + ): + raise ProvisionError( + f"release '{release}' already exists and was not created by " + f"this tool; refusing to overwrite" + ) + 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) + # Create the master-key Secret only if it does not already exist. + # Re-provisioning the same release is an in-place upgrade; regenerating + # the key here would rotate it and invalidate every API key already + # minted against the running deployment. + mk_secret = f"{release}-master-key" + if not await self._kubectl.resource_exists(kind="secret", name=mk_secret): + mk_manifest, _ = resources.master_key_secret(release) + await self._kubectl.apply(mk_manifest) + # Note: auxiliary objects are intentionally left in place if the helm + # step below fails — they are release-named (reused on retry, so no + # accumulation) and preserved for inspecting failed e2e runs. Use + # `delete` to tear a release down. if request.enable_postgres: pg_manifest, database = resources.postgres(release) await self._kubectl.apply(pg_manifest) @@ -221,8 +280,28 @@ class Provisioner: "pods": _summarize_pods(pods), } + async def _owns_release(self, release: str) -> bool: + """Whether ``release`` was provisioned by this tool. + + Every provisioned release gets a master-key Secret carrying both the + managed-by and release labels, so the presence of such a Secret is + proof of ownership. This stops a caller from uninstalling unrelated + helm releases (e.g. a real deployment) that happen to share the + namespace. + """ + selector = ( + f"app.kubernetes.io/managed-by={MANAGED_BY},{RELEASE_LABEL}={release}" + ) + count = await self._kubectl.count_by_label(kind="secret", selector=selector) + return count > 0 + async def delete(self, release_name: str) -> dict: release = sanitize_release_name(release_name) + if not await self._owns_release(release): + raise ProvisionError( + f"release '{release}' was not created by this tool " + f"(no {MANAGED_BY} resources found); refusing to delete" + ) await self._helm.uninstall(release=release) await self._kubectl.delete_by_label( selector=resources.selector_label(release), diff --git a/provisioning_mcp/tests/conftest.py b/provisioning_mcp/tests/conftest.py index dddc617a12d..7051e22c47a 100644 --- a/provisioning_mcp/tests/conftest.py +++ b/provisioning_mcp/tests/conftest.py @@ -17,6 +17,8 @@ def make_settings(**overrides: Any) -> Settings: chart_path="/app/helm/litellm", default_image_registry="ghcr.io/berriai", release_prefix="litellm-e2e", + provisioning_service_account="litellm-provisioning-mcp", + allowed_service_accounts=(), helm_binary="helm", kubectl_binary="kubectl", command_timeout=600, @@ -33,21 +35,34 @@ class FakeRunner: def __init__(self) -> None: self.calls: list[tuple[list[str], str | None]] = [] + # Toggles driving the read-only `kubectl get` / `helm list` probes. + self.master_key_exists = False + self.owns_release = True + self.helm_releases = ["litellm-e2e-other"] 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": []}), "") + if "get" in args: + if "pods" in args: + return CommandResult(0, json.dumps({"items": []}), "") + if "--selector" in args: # count_by_label (ownership check) + items = ( + [{"metadata": {"name": "owned"}}] if self.owns_release else [] + ) + return CommandResult(0, json.dumps({"items": items}), "") + # resource_exists: `get --ignore-not-found -o name` + return CommandResult( + 0, "secret/x" if self.master_key_exists else "", "" + ) 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"}]), "" - ) + payload = [{"name": n, "status": "deployed"} for n in self.helm_releases] + return CommandResult(0, json.dumps(payload), "") return CommandResult(0, "ok", "") def find(self, *needles: str) -> tuple[list[str], str | None]: diff --git a/provisioning_mcp/tests/test_provisioner.py b/provisioning_mcp/tests/test_provisioner.py index 182af8ec0ca..cb615dedda6 100644 --- a/provisioning_mcp/tests/test_provisioner.py +++ b/provisioning_mcp/tests/test_provisioner.py @@ -133,6 +133,102 @@ async def test_image_registry_override_and_extra_values(): assert values["gateway"]["image"]["tag"] == "sha" +async def test_provision_forces_release_prefix(): + runner = FakeRunner() + provisioner = _provisioner(runner) + + result = await provisioner.provision( + ProvisionRequest( + repo_url="https://github.com/BerriAI/litellm", + revision="abc", + release_name="litellm", # would collide with a real release + ) + ) + + assert result["release"] == "litellm-e2e-litellm" + values = _helm_values(runner) + assert values["fullnameOverride"] == "litellm-e2e-litellm" + + +async def test_provision_rejects_self_service_account(): + runner = FakeRunner() + provisioner = _provisioner(runner) + with pytest.raises(ProvisionError, match="own"): + await provisioner.provision( + ProvisionRequest( + repo_url="https://github.com/BerriAI/litellm", + revision="abc", + service_account="litellm-provisioning-mcp", + ) + ) + + +async def test_provision_rejects_disallowed_service_account(): + runner = FakeRunner() + settings = make_settings(allowed_service_accounts=("litellm-workload",)) + provisioner = Provisioner( + settings, + helm=HelmRunner(namespace="litellm", runner=runner, wait_timeout=600), + kubectl=KubectlRunner(namespace="litellm", runner=runner), + ) + with pytest.raises(ProvisionError, match="not in the allowed list"): + await provisioner.provision( + ProvisionRequest( + repo_url="https://github.com/BerriAI/litellm", + revision="abc", + service_account="some-other-sa", + ) + ) + + +async def test_provision_refuses_existing_unmanaged_release(): + runner = FakeRunner() + runner.helm_releases = ["litellm-e2e-abc"] # already exists + runner.owns_release = False # but not created by us + provisioner = _provisioner(runner) + with pytest.raises(ProvisionError, match="refusing to overwrite"): + await provisioner.provision( + ProvisionRequest( + repo_url="https://github.com/BerriAI/litellm", + revision="abc", + ) + ) + with pytest.raises(AssertionError): + runner.find("upgrade") + + +async def test_provision_preserves_existing_master_key(): + runner = FakeRunner() + runner.master_key_exists = True + provisioner = _provisioner(runner) + + await provisioner.provision( + ProvisionRequest( + repo_url="https://github.com/BerriAI/litellm", + revision="abc123", + ) + ) + + # The master-key Secret already exists, so it must not be re-applied + # (re-applying would rotate the key and invalidate live API keys). + runner.find("get", "secret") # existence was probed + applied = [text for args, text in runner.calls if "apply" in args and text] + assert not any("master-key" in payload for payload in applied) + + +async def test_delete_refuses_release_not_owned_by_tool(): + runner = FakeRunner() + runner.owns_release = False + provisioner = _provisioner(runner) + + with pytest.raises(ProvisionError, match="refusing to delete"): + await provisioner.delete("some-production-release") + + # Must not have attempted a helm uninstall on an unowned release. + with pytest.raises(AssertionError): + runner.find("uninstall") + + async def test_delete_uninstalls_and_cleans_datastores(): runner = FakeRunner() provisioner = _provisioner(runner) @@ -157,10 +253,10 @@ 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) + result = 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, "", "") + return result settings = make_settings() provisioner = Provisioner(