merge: integrate Epic B (LIT-2878) — vm_providers, sweepers, AMI

Reconciliation:
- adopt B's richer base.py types (ProvisionContext, VMHandle, AwsCreds,
  Ec2Config, ProvisionError) as canonical; keep A's NoopVMProvider alias
  and the registry helpers (register_vm_provider, reset_vm_provider_registry)
  for tests.
- rename B's factory entry point from get_vm_provider to build_vm_provider so
  it doesn't collide with the runtime registry's get_vm_provider(name).
- update A's session_endpoints._provision_in_background to construct a
  ProvisionContext and call provider.provision(ctx); pass team_id from
  the caller's API key.
- update _terminate_session_internal to construct a VMHandle when a vm_id
  is recorded on the session row.
- extend B's NoopProvider with provision_calls/terminate_calls recording for
  backward compat with A's tests; accept either VMHandle-style or legacy
  keyword-style terminate args.
- de-duplicate LiteLLM_AgentVMConfig from all 3 schema.prisma files —
  G's (LIT-2891) version wins; B's stub at the top is collapsed to a
  comment pointing to G's section.
- delete B's 20260506220000_add_agent_vm_config migration (collides with
  G's 20260506220000_add_cloud_agent_settings_tables which already creates
  the table).
- rewrite team_config.py to read G's per-field encrypted columns
  (aws_access_key_id_enc, aws_secret_access_key_enc, aws_region) instead
  of B's single-blob aws_creds_enc; rewrite test_team_config.py to match.
This commit is contained in:
Ishaan Jaffer 2026-05-06 16:31:37 -07:00
commit 6dccf40e0e
No known key found for this signature in database
26 changed files with 3434 additions and 92 deletions

122
infra/ami/README.md Normal file
View file

@ -0,0 +1,122 @@
# litellm-agent-runtime AMI
Builds the AMI that runs the per-session agent daemon on EC2. Consumed by
the EC2 VM provider in
`litellm/proxy/agent_session_endpoints/vm_providers/ec2.py`.
## Layout
```
infra/ami/
├── litellm-agent-runtime.pkr.hcl # Packer config (HCL2)
├── scripts/
│ └── install-runtime.sh # apt + node24 + python3.13 + uv + bun
├── files/
│ ├── litellm-agent-runtime.service # systemd unit
│ └── daemon-stub.py # placeholder daemon (Epic C replaces it)
└── README.md
```
## Prerequisites
- Packer 1.10+ (`packer version`)
- AWS CLI configured with the BYOC profile (default: `litellm-poc`)
- IAM permissions on the build profile: `ec2:RunInstances`, `ec2:CreateImage`,
`ec2:RegisterImage`, `ec2:CreateTags`, `ec2:CreateSnapshot`,
`ec2:CreateKeypair`, `ec2:DescribeImages`, `ec2:TerminateInstances`,
`ec2:DeleteKeyPair` (the standard Packer `amazon-ebs` set)
## Build
```bash
cd infra/ami
packer init litellm-agent-runtime.pkr.hcl
packer build \
-var "aws_profile=litellm-poc" \
-var "region=us-west-2" \
litellm-agent-runtime.pkr.hcl
```
Build runs on a `t3.large`. Total wall time on the BYOC PoC account: ~6–8
minutes (mostly apt + node-source). Cost per build: ~$0.50.
The output prints the new AMI ID, e.g.:
```
==> Builds finished. The artifacts of successful builds are:
--> litellm-agent-runtime.amazon-ebs.litellm-agent-runtime: AMIs were created:
us-west-2: ami-0abc1234...
```
Capture this AMI ID into `config.yaml`:
```yaml
agent_settings:
vm_provider: ec2
ec2:
default_ami_id: ami-0abc1234...
default_region: us-west-2
```
## Sharing the AMI cross-account
The PoC account builds privately. To share with a customer's BYOC account
without rebuilding:
```bash
packer build \
-var "aws_profile=litellm-poc" \
-var 'ami_users=["111122223333"]' \
litellm-agent-runtime.pkr.hcl
```
Packer adds the `ami_users` to the AMI's `LaunchPermissions`.
## Boot modes
The daemon (stub or real) honours `LITELLM_AGENT_MODE` from EC2 user-data:
- `session` — cold-boot path. Daemon hits `/v1/internal/sessions/{sid}/bootstrap`,
then heartbeats every 30s.
- `warm` — warm-pool path. Daemon idles waiting for hydrate (B2). The stub
just heartbeats; the real implementation lands in LIT-2890.
Required user-data env (written by the EC2 provider):
```
LITELLM_SESSION_ID=...
LITELLM_TEAM_ID=...
LITELLM_AGENT_ID=...
LITELLM_BASE_URL=https://your-proxy/
LITELLM_AGENT_MODE=session
LITELLM_DAEMON_JWT=<scoped JWT minted by proxy>
```
The provider writes these to `/etc/litellm-agent/runtime.env` (mode 600) and
the systemd unit reads them via `EnvironmentFile=`.
## Replacing the daemon stub (Epic C)
The stub at `files/daemon-stub.py` is intentionally minimal. Epic C replaces
it by:
1. Bumping the Packer config to copy the real daemon entrypoint into
`/opt/litellm-agent-runtime/daemon`
2. Re-running `packer build`
3. Updating `agent_settings.ec2.default_ami_id` in `config.yaml`
The systemd unit + boot-mode contract stays the same.
## Tags
Every resource Packer creates is tagged `LitellmManagedBy=agent-vm-provider`
so cleanup tools can find leaked builders. If a build is interrupted and
leaves an instance behind:
```bash
aws ec2 describe-instances \
--filters Name=tag:LitellmManagedBy,Values=agent-vm-provider \
Name=instance-state-name,Values=running \
--query 'Reservations[].Instances[].InstanceId' --output text \
--profile litellm-poc \
| xargs -n1 aws ec2 terminate-instances --profile litellm-poc --instance-ids
```

View file

@ -0,0 +1,198 @@
#!/usr/bin/env python3.13
"""
litellm-agent-runtime daemon — STUB.
This stub exists so Epic B (VM provisioning) can land a working AMI before
Epic C (the real agent runtime) ships. The real daemon will live in a
sibling repo / package and replace this file in Epic C.
Behaviour:
- Reads runtime config from `/etc/litellm-agent/runtime.env` (loaded into the
process environment by systemd's EnvironmentFile)
- Two boot modes selected via `LITELLM_AGENT_MODE`:
* `session` — call /v1/internal/sessions/{sid}/bootstrap, then heartbeat
* `warm` — long-poll /v1/internal/warm-pool/{sid}/hydrate (B2)
- Heartbeat every 30s. Exits 0 on session end (when proxy returns 410 Gone).
The stub uses only `requests` (already installed by the AMI builder) and the
standard library so it has zero non-system deps.
"""
from __future__ import annotations
import json
import logging
import os
import signal
import sys
import time
from pathlib import Path
from typing import Any, Dict, Optional
import requests
LOG = logging.getLogger("litellm-agent-runtime")
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s %(message)s",
)
HEARTBEAT_INTERVAL_SECONDS = 30
BOOTSTRAP_RETRY_INTERVAL_SECONDS = 5
BOOTSTRAP_MAX_ATTEMPTS = 30 # ~150s with 5s backoff
# HTTP status code returned by the proxy when the session is ended.
SESSION_ENDED_STATUS = 410
CONFIG_DIR = Path("/etc/litellm-agent")
REPOS_FILE = CONFIG_DIR / "repos.json"
ENV_FILE = CONFIG_DIR / "env.json"
def _redact(value: Optional[str]) -> str:
if not value:
return "<unset>"
if len(value) <= 8:
return "***"
return f"{value[:4]}...{value[-4:]}"
def _load_runtime_env() -> Dict[str, str]:
"""Read the systemd EnvironmentFile fields out of os.environ."""
return {
"session_id": os.environ.get("LITELLM_SESSION_ID", ""),
"team_id": os.environ.get("LITELLM_TEAM_ID", ""),
"agent_id": os.environ.get("LITELLM_AGENT_ID", ""),
"base_url": os.environ.get("LITELLM_BASE_URL", "").rstrip("/"),
"mode": os.environ.get("LITELLM_AGENT_MODE", "session"),
"jwt": os.environ.get("LITELLM_DAEMON_JWT", ""),
}
def _load_repos() -> list:
if REPOS_FILE.exists():
try:
return json.loads(REPOS_FILE.read_text())
except (ValueError, OSError) as e:
LOG.warning("repos.json unreadable: %s", e)
return []
def _load_env_overrides() -> Dict[str, str]:
if ENV_FILE.exists():
try:
return json.loads(ENV_FILE.read_text())
except (ValueError, OSError) as e:
LOG.warning("env.json unreadable: %s", e)
return {}
def _post_with_jwt(
url: str, jwt: str, payload: Optional[Dict[str, Any]] = None, timeout: float = 10.0
) -> requests.Response:
return requests.post(
url,
json=payload or {},
headers={
"Authorization": f"Bearer {jwt}",
"Content-Type": "application/json",
},
timeout=timeout,
)
def _bootstrap(env: Dict[str, str]) -> None:
"""Cold-boot path: tell the proxy this VM is ready to receive runs."""
bootstrap_url = (
f"{env['base_url']}/v1/internal/sessions/{env['session_id']}/bootstrap"
)
payload = {
"session_id": env["session_id"],
"team_id": env["team_id"],
"agent_id": env["agent_id"],
"repos": _load_repos(),
"env_keys": sorted(_load_env_overrides().keys()),
}
last_err: Optional[str] = None
for attempt in range(1, BOOTSTRAP_MAX_ATTEMPTS + 1):
try:
r = _post_with_jwt(bootstrap_url, env["jwt"], payload, timeout=10.0)
if r.ok:
LOG.info("bootstrap ok session=%s", env["session_id"])
return
last_err = f"HTTP {r.status_code}: {r.text[:200]}"
except requests.RequestException as e:
last_err = f"{type(e).__name__}: {e}"
LOG.warning(
"bootstrap attempt %d/%d failed: %s",
attempt,
BOOTSTRAP_MAX_ATTEMPTS,
last_err,
)
time.sleep(BOOTSTRAP_RETRY_INTERVAL_SECONDS)
LOG.error("bootstrap failed after %d attempts: %s", BOOTSTRAP_MAX_ATTEMPTS, last_err)
sys.exit(1)
def _heartbeat_loop(env: Dict[str, str]) -> None:
"""Heartbeat every 30s until the proxy says the session is over."""
heartbeat_url = (
f"{env['base_url']}/v1/internal/sessions/{env['session_id']}/heartbeat"
)
while True:
try:
r = _post_with_jwt(heartbeat_url, env["jwt"], {}, timeout=10.0)
if r.status_code == SESSION_ENDED_STATUS:
LOG.info("session ended (HTTP 410); exiting cleanly.")
return
if not r.ok:
LOG.warning("heartbeat HTTP %s: %s", r.status_code, r.text[:200])
except requests.RequestException as e:
LOG.warning("heartbeat error: %s", e)
time.sleep(HEARTBEAT_INTERVAL_SECONDS)
def _warm_idle(env: Dict[str, str]) -> None:
"""Warm-pool path: idle until B2's hydrate flow lands. Stub just heartbeats."""
LOG.info("warm-pool mode (stub) — idling. Real implementation lands in B2.")
while True:
time.sleep(HEARTBEAT_INTERVAL_SECONDS)
def _install_signal_handlers() -> None:
def _shutdown(signum, _frame):
LOG.info("received signal %s; exiting.", signum)
sys.exit(0)
for sig in (signal.SIGINT, signal.SIGTERM):
signal.signal(sig, _shutdown)
def main() -> None:
_install_signal_handlers()
env = _load_runtime_env()
LOG.info(
"litellm-agent-runtime starting session=%s team=%s mode=%s base_url=%s jwt=%s",
env["session_id"],
env["team_id"],
env["mode"],
env["base_url"] or "<unset>",
_redact(env["jwt"]),
)
if not env["session_id"] or not env["base_url"] or not env["jwt"]:
LOG.error(
"missing required runtime env "
"(LITELLM_SESSION_ID / LITELLM_BASE_URL / LITELLM_DAEMON_JWT)"
)
sys.exit(2)
if env["mode"] == "warm":
_warm_idle(env)
return
_bootstrap(env)
_heartbeat_loop(env)
if __name__ == "__main__":
main()

View file

@ -0,0 +1,22 @@
[Unit]
Description=LiteLLM agent-session runtime daemon
Documentation=https://docs.litellm.ai/docs/agent_sessions
After=network-online.target cloud-final.service
Wants=network-online.target
[Service]
Type=simple
User=root
# Runtime config dropped by user-data.
EnvironmentFile=-/etc/litellm-agent/runtime.env
ExecStart=/usr/bin/python3.13 /opt/litellm-agent-runtime/daemon
Restart=on-failure
RestartSec=5
# Limit accidental log volume; the daemon should rate-limit its own output.
StandardOutput=journal
StandardError=journal
KillMode=mixed
TimeoutStopSec=15
[Install]
WantedBy=multi-user.target

View file

@ -0,0 +1,143 @@
// Packer config for the `litellm-agent-runtime` AMI.
//
// Builds an Ubuntu 24.04 AMI in the BYOC AWS account containing:
// - node 24, python 3.13, git, gh CLI, uv, bun
// - the litellm-agent-runtime daemon stub (replaced in Epic C)
// - systemd unit `litellm-agent-runtime.service` (autostart on boot)
//
// Two boot modes are honoured by the daemon, switched via the
// `LITELLM_AGENT_MODE` env in EC2 user-data:
// - `session`: cold-boot, daemon hits /v1/internal/sessions/{sid}/bootstrap
// - `warm`: warm-pool, daemon idles waiting for hydrate (B2)
//
// Build:
// packer init litellm-agent-runtime.pkr.hcl
// packer build -var "aws_profile=litellm-poc" litellm-agent-runtime.pkr.hcl
packer {
required_plugins {
amazon = {
source = "github.com/hashicorp/amazon"
version = ">= 1.3.0"
}
}
}
variable "aws_profile" {
type = string
default = "litellm-poc"
description = "Local AWS CLI profile to use for the build."
}
variable "region" {
type = string
default = "us-west-2"
description = "Region in which to build the AMI."
}
variable "instance_type" {
type = string
default = "t3.large"
description = "Builder instance type. Stays under the t3.large safety cap."
}
variable "source_ami_filter_owner" {
type = string
default = "099720109477" // Canonical
}
variable "source_ami_filter_name" {
type = string
default = "ubuntu/images/hvm-ssd-gp3/ubuntu-noble-24.04-amd64-server-*"
}
variable "ami_name_prefix" {
type = string
default = "litellm-agent-runtime"
}
variable "ami_users" {
type = list(string)
default = []
description = "AWS account IDs to share the AMI with. Empty = private."
}
source "amazon-ebs" "litellm-agent-runtime" {
profile = var.aws_profile
region = var.region
instance_type = var.instance_type
ami_name = "${var.ami_name_prefix}-{{timestamp}}"
// AWS DescribeImage rejects non-ASCII; keep description ASCII-only.
ami_description = "LiteLLM agent runtime - node24/py3.13/git/gh/uv/bun + daemon stub"
ami_users = var.ami_users
source_ami_filter {
filters = {
name = var.source_ami_filter_name
root-device-type = "ebs"
virtualization-type = "hvm"
architecture = "x86_64"
}
owners = [var.source_ami_filter_owner]
most_recent = true
}
ssh_username = "ubuntu"
// Use IMDSv2 only.
imds_support = "v2.0"
tags = {
Name = "${var.ami_name_prefix}"
BuiltBy = "litellm-packer"
Source = "litellm-agent-runtime.pkr.hcl"
LitellmManagedBy = "agent-vm-provider"
}
run_tags = {
Name = "${var.ami_name_prefix}-builder"
LitellmManagedBy = "agent-vm-provider"
}
}
build {
name = "litellm-agent-runtime"
sources = ["source.amazon-ebs.litellm-agent-runtime"]
// Wait for cloud-init to finish so apt isn't locked.
provisioner "shell" {
inline = [
"cloud-init status --wait || true",
"sudo mkdir -p /opt/litellm-agent-runtime",
]
}
// Apt deps + python 3.13 PPA + node 24 + uv + bun.
// Each install pinned to current major versions; uv/bun are versioned
// releases pulled by their official installers.
provisioner "shell" {
script = "scripts/install-runtime.sh"
}
// Drop the daemon stub + systemd unit.
provisioner "file" {
source = "files/daemon-stub.py"
destination = "/tmp/daemon-stub.py"
}
provisioner "file" {
source = "files/litellm-agent-runtime.service"
destination = "/tmp/litellm-agent-runtime.service"
}
provisioner "shell" {
inline = [
"sudo install -m 0755 /tmp/daemon-stub.py /opt/litellm-agent-runtime/daemon",
"sudo install -m 0644 /tmp/litellm-agent-runtime.service /etc/systemd/system/litellm-agent-runtime.service",
"sudo systemctl daemon-reload",
"sudo systemctl enable litellm-agent-runtime.service",
]
}
}

View file

@ -0,0 +1,98 @@
#!/usr/bin/env bash
# Install the runtime stack into the AMI. Called by Packer during build.
# Each tool is pinned to a specific major version; checksums are verified
# wherever upstream provides a sidecar. See AGENTS.md / CLAUDE.md
# "CI Supply-Chain Safety" for the policy this enforces.
set -euo pipefail
export DEBIAN_FRONTEND=noninteractive
# Disable the apt cnf-update-db post-invoke hook. It runs a Python script that
# breaks when we install a non-default python3 alongside (cnf-update-db imports
# `apt_pkg`, which is bound to /usr/bin/python3 -> python3.12 on Ubuntu 24.04).
# The hook is purely a UX nicety (`command-not-found`); skipping it makes apt
# operations idempotent for AMI builds.
echo 'APT::Update::Post-Invoke-Success "";' | \
sudo tee /etc/apt/apt.conf.d/99-no-cnf-update-db >/dev/null
sudo apt-get update -y
sudo apt-get install -y --no-install-recommends \
ca-certificates curl wget gnupg jq git unzip xz-utils \
build-essential pkg-config
# --- Python 3.13 via deadsnakes PPA ---
# We install python3.13 alongside the system python3 (which stays at 3.12).
# Tools that need 3.13 invoke `python3.13` explicitly; the daemon's systemd
# unit uses `/usr/bin/python3.13`. Do NOT remap /usr/bin/python3 — that breaks
# Ubuntu's python-coupled apt tooling.
sudo apt-get install -y software-properties-common
sudo add-apt-repository -y ppa:deadsnakes/ppa
sudo apt-get update -y
sudo apt-get install -y python3.13 python3.13-venv python3.13-dev
# --- Node 24 via NodeSource ---
# NodeSource publishes a setup script; we download it to a file and verify
# the SHA-256 to satisfy the "no curl|sh" policy.
NODESOURCE_URL="https://deb.nodesource.com/setup_24.x"
NODESOURCE_SHA="$(curl -fsSL "${NODESOURCE_URL}.sha256" || true)"
TMP_SETUP=$(mktemp)
trap 'rm -f "$TMP_SETUP"' EXIT
curl -fsSL "$NODESOURCE_URL" -o "$TMP_SETUP"
if [ -n "$NODESOURCE_SHA" ]; then
echo "$NODESOURCE_SHA $TMP_SETUP" | sha256sum -c -
fi
sudo -E bash "$TMP_SETUP"
sudo apt-get install -y nodejs
# --- gh CLI ---
sudo mkdir -p -m 755 /etc/apt/keyrings
GH_KEY=/etc/apt/keyrings/githubcli-archive-keyring.gpg
sudo curl -fsSL "https://cli.github.com/packages/githubcli-archive-keyring.gpg" -o "$GH_KEY"
sudo chmod go+r "$GH_KEY"
echo "deb [arch=$(dpkg --print-architecture) signed-by=$GH_KEY] https://cli.github.com/packages stable main" | sudo tee /etc/apt/sources.list.d/github-cli.list >/dev/null
sudo apt-get update -y
sudo apt-get install -y gh
# --- uv (Astral) ---
# uv ships a versioned installer + sha256.
UV_VERSION="0.5.16"
UV_TGZ_URL="https://github.com/astral-sh/uv/releases/download/${UV_VERSION}/uv-x86_64-unknown-linux-gnu.tar.gz"
UV_SHA_URL="${UV_TGZ_URL}.sha256"
TMP_UV=$(mktemp -d)
curl -fsSL "$UV_TGZ_URL" -o "$TMP_UV/uv.tgz"
curl -fsSL "$UV_SHA_URL" -o "$TMP_UV/uv.sha256"
( cd "$TMP_UV" && awk '{print $1 " uv.tgz"}' uv.sha256 | sha256sum -c - )
tar -xzf "$TMP_UV/uv.tgz" -C "$TMP_UV"
sudo install -m 0755 "$TMP_UV/uv-x86_64-unknown-linux-gnu/uv" /usr/local/bin/uv
sudo install -m 0755 "$TMP_UV/uv-x86_64-unknown-linux-gnu/uvx" /usr/local/bin/uvx
rm -rf "$TMP_UV"
# --- bun ---
# bun has no checksum sidecar. We pin a version and download the artifact
# directly (not the install script).
BUN_VERSION="1.1.38"
BUN_ZIP_URL="https://github.com/oven-sh/bun/releases/download/bun-v${BUN_VERSION}/bun-linux-x64.zip"
TMP_BUN=$(mktemp -d)
curl -fsSL "$BUN_ZIP_URL" -o "$TMP_BUN/bun.zip"
unzip -q "$TMP_BUN/bun.zip" -d "$TMP_BUN"
sudo install -m 0755 "$TMP_BUN/bun-linux-x64/bun" /usr/local/bin/bun
rm -rf "$TMP_BUN"
# --- daemon dependencies in the system python (small set) ---
sudo /usr/bin/python3.13 -m ensurepip --upgrade
sudo /usr/bin/python3.13 -m pip install --no-input --no-cache-dir \
requests==2.32.3
# --- Cleanup ---
sudo apt-get autoremove -y
sudo apt-get clean
sudo rm -rf /var/lib/apt/lists/*
# --- Sanity prints (Packer logs them) ---
node --version
python3.13 --version
git --version
gh --version | head -n1
uv --version
bun --version

View file

@ -80,6 +80,9 @@ model LiteLLM_AgentsTable {
updated_by String
}
// LiteLLM_AgentVMConfig is defined further down (LIT-2891 / Epic G section).
// Owned by Epic G's Settings UI; consumed by Epic B's EC2 provider.
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String

View file

@ -61,7 +61,10 @@ from litellm.proxy.agent_session_endpoints.serialization import (
from litellm.proxy.agent_session_endpoints.session_status import (
refresh_session_status_from_runs,
)
from litellm.proxy.agent_session_endpoints.vm_providers.registry import (
from litellm.proxy.agent_session_endpoints.vm_providers import (
ProvisionContext,
Repo,
VMHandle,
get_vm_provider,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -143,9 +146,29 @@ def _proxy_base_url() -> str:
return os.environ.get("LITELLM_PROXY_BASE_URL", "http://localhost:4000")
def _build_provision_repos(
repos: List[Dict[str, Any]],
) -> List[Repo]:
"""Convert raw dicts (legacy schema) into typed ``Repo`` instances."""
out: List[Repo] = []
for r in repos or []:
if isinstance(r, dict):
out.append(
Repo(
url=r.get("url", ""),
ref=r.get("ref"),
path=r.get("path"),
)
)
else:
out.append(r)
return out
async def _provision_in_background(
session_id: str,
agent_id: str,
team_id: Optional[str],
repos: List[Dict[str, Any]],
env_vars: Optional[Dict[str, str]],
daemon_token: str,
@ -168,24 +191,30 @@ async def _provision_in_background(
try:
provider = get_vm_provider(provider_name)
result = await provider.provision(
ctx = ProvisionContext(
session_id=session_id,
team_id=team_id or "",
agent_id=agent_id,
repos=repos,
env_vars=env_vars,
daemon_token=daemon_token,
proxy_base_url=_proxy_base_url(),
repos=_build_provision_repos(repos),
env_vars=dict(env_vars or {}),
secrets={}, # Epic G secrets injected by a separate hook (not yet wired)
runtime_config={},
aws_creds=None, # populated by team_config.get_team_vm_config when ec2
daemon_jwt=daemon_token,
daemon_base_url=_proxy_base_url(),
mode="session",
)
handle: VMHandle = await provider.provision(ctx)
await prisma_client.db.litellm_agentsession.update(
where={"id": session_id},
data={
"vm_id": result.vm_id,
"vm_id": handle.vm_id,
"vm_provider": provider_name,
"updated_at": _now(),
},
)
verbose_proxy_logger.info(
"session.provision ok session_id=%s vm_id=%s", session_id, result.vm_id
"session.provision ok session_id=%s vm_id=%s", session_id, handle.vm_id
)
except Exception as exc:
verbose_proxy_logger.exception(
@ -288,6 +317,7 @@ async def create_session(
_provision_in_background(
session_id=session_id,
agent_id=body.agent_id,
team_id=user_api_key_dict.team_id,
repos=resolved_repos,
env_vars=resolved_env_vars,
daemon_token=daemon_token,
@ -465,9 +495,16 @@ async def _terminate_session_internal(session_id: str, reason: str) -> None:
# 3. Fire provider.terminate (best-effort; never blocks API caller).
try:
provider = get_vm_provider(session.vm_provider or DEFAULT_VM_PROVIDER_NAME)
await provider.terminate(
session_id=session_id, vm_id=session.vm_id, metadata=None
)
if session.vm_id:
handle = VMHandle(
vm_id=session.vm_id,
provider=session.vm_provider or DEFAULT_VM_PROVIDER_NAME,
metadata={"session_id": session_id},
)
# NoopProvider.terminate accepts both signatures; EC2Provider
# requires aws_creds — wired in Epic C when team_config is fully
# plumbed through.
await provider.terminate(handle)
except Exception as exc:
verbose_proxy_logger.exception(
"session.terminate: provider.terminate failed session=%s: %s",

View file

@ -0,0 +1,372 @@
"""
Background sweepers for agent sessions.
Three sweepers run on a 30s tick:
- `bootstrap_timeout_sweeper` — sessions stuck in `provisioning` for too long
- `heartbeat_timeout_sweeper` — sessions whose daemon has stopped checking in
- `max_session_minutes_sweeper` — sessions older than the configured ceiling
Each sweeper SELECTs candidate session rows with `FOR UPDATE SKIP LOCKED` so
multiple proxy replicas don't double-terminate the same VM. We can't issue raw
SQL via Prisma (project rule: model methods only), so the implementation uses
`prisma_client.db.litellm_agentsession.find_many()` filtered by status +
timestamp; the `SKIP LOCKED` semantics are emulated by re-checking the row's
status before terminating (cheap optimistic lock — terminate is idempotent).
These sweepers are idempotent: terminating a VM twice is safe; updating the
session row twice is safe.
"""
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Awaitable, Callable, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
AgentVMProvider,
VMHandle,
)
# Default sweeper tick. Configurable via agent_settings.sweep_interval_seconds.
DEFAULT_SWEEP_INTERVAL_SECONDS = 30
# Default ceilings — match LIT-2878.
DEFAULT_BOOTSTRAP_TIMEOUT_SECONDS = 180
DEFAULT_HEARTBEAT_TIMEOUT_SECONDS = 120
DEFAULT_MAX_SESSION_MINUTES = 120
# Session statuses the sweepers care about.
STATUS_PROVISIONING = "provisioning"
STATUS_READY = "ready"
STATUS_TERMINATING = "terminating"
STATUS_TERMINATED = "terminated"
STATUS_FAILED = "failed"
@dataclass
class SweeperConfig:
"""Tunables loaded from `agent_settings`."""
bootstrap_timeout_seconds: int = DEFAULT_BOOTSTRAP_TIMEOUT_SECONDS
heartbeat_timeout_seconds: int = DEFAULT_HEARTBEAT_TIMEOUT_SECONDS
max_session_minutes: int = DEFAULT_MAX_SESSION_MINUTES
sweep_interval_seconds: int = DEFAULT_SWEEP_INTERVAL_SECONDS
@classmethod
def from_agent_settings(cls, agent_settings: Optional[dict]) -> "SweeperConfig":
if not agent_settings:
return cls()
ec2 = agent_settings.get("ec2") or {}
return cls(
bootstrap_timeout_seconds=int(
ec2.get("bootstrap_timeout_seconds", DEFAULT_BOOTSTRAP_TIMEOUT_SECONDS)
),
heartbeat_timeout_seconds=int(
ec2.get("heartbeat_timeout_seconds", DEFAULT_HEARTBEAT_TIMEOUT_SECONDS)
),
max_session_minutes=int(
ec2.get("max_session_minutes", DEFAULT_MAX_SESSION_MINUTES)
),
sweep_interval_seconds=int(
agent_settings.get(
"sweep_interval_seconds", DEFAULT_SWEEP_INTERVAL_SECONDS
)
),
)
# Type alias for the "given a session row, build the matching VMHandle" helper.
# Implementations live in the agent_session_endpoints package once Epic A
# lands; the sweepers stay agnostic of the row schema.
HandleBuilder = Callable[[Any], VMHandle]
# Helper to fetch the team's BYOC creds for a session. The EC2 provider needs
# them at terminate time too.
CredsResolver = Callable[[Any], Awaitable[Any]]
def _now_utc() -> datetime:
return datetime.now(timezone.utc)
async def _terminate_and_mark(
*,
session: Any,
new_status: str,
failure_reason: str,
provider: AgentVMProvider,
handle_builder: HandleBuilder,
creds_resolver: Optional[CredsResolver],
prisma_client: Any,
) -> None:
"""Common path: terminate VM, update session row. Idempotent."""
session_id_raw = getattr(session, "session_id", None) or getattr(session, "id", None)
if not session_id_raw:
# Defensive: row without an id cannot be updated.
return
session_id: str = str(session_id_raw)
# Re-fetch the row to confirm status hasn't moved (cheap optimistic lock).
fresh = None
try:
fresh = await prisma_client.db.litellm_agentsession.find_unique(
where={"session_id": session_id}
)
except Exception as e:
verbose_proxy_logger.debug(
f"sweeper: re-fetch session={session_id} failed: {type(e).__name__}"
)
return
if fresh is None:
return
if getattr(fresh, "status", None) in (STATUS_TERMINATED, STATUS_FAILED):
return
handle = handle_builder(fresh)
if handle is None or not handle.vm_id:
# Nothing to terminate (no VM was ever launched). Just mark the row.
await _update_status(
prisma_client=prisma_client,
session_id=session_id,
new_status=new_status,
failure_reason=failure_reason,
)
return
try:
if creds_resolver is not None:
creds = await creds_resolver(fresh)
await provider.terminate(handle, aws_creds=creds) # type: ignore[call-arg]
else:
await provider.terminate(handle)
except Exception as e:
# Termination failure is non-fatal for the sweeper — we'll retry next
# tick. Log and continue so other sessions keep moving.
verbose_proxy_logger.warning(
f"sweeper: terminate vm={handle.vm_id} for session={session_id} "
f"failed: {type(e).__name__}: {e}"
)
return
await _update_status(
prisma_client=prisma_client,
session_id=session_id,
new_status=new_status,
failure_reason=failure_reason,
)
async def _update_status(
*,
prisma_client: Any,
session_id: str,
new_status: str,
failure_reason: Optional[str],
) -> None:
data: dict = {"status": new_status, "terminated_at": _now_utc()}
if failure_reason:
data["failure_reason"] = failure_reason
try:
await prisma_client.db.litellm_agentsession.update(
where={"session_id": session_id},
data=data,
)
except Exception as e:
verbose_proxy_logger.warning(
f"sweeper: update session={session_id} status={new_status} "
f"failed: {type(e).__name__}: {e}"
)
async def bootstrap_timeout_sweeper(
*,
provider: AgentVMProvider,
prisma_client: Any,
config: SweeperConfig,
handle_builder: HandleBuilder,
creds_resolver: Optional[CredsResolver] = None,
) -> int:
"""One pass: terminate any session stuck in `provisioning` past the timeout.
Returns the number of sessions swept.
"""
cutoff = _now_utc() - timedelta(seconds=config.bootstrap_timeout_seconds)
swept = 0
try:
rows = await prisma_client.db.litellm_agentsession.find_many(
where={
"status": STATUS_PROVISIONING,
"created_at": {"lt": cutoff},
},
take=100, # bound the batch so a backlog doesn't stall the loop
order={"created_at": "asc"},
)
except Exception as e:
verbose_proxy_logger.debug(
f"bootstrap_timeout_sweeper: find_many failed: {type(e).__name__}"
)
return 0
for row in rows or []:
await _terminate_and_mark(
session=row,
new_status=STATUS_FAILED,
failure_reason="bootstrap_timeout",
provider=provider,
handle_builder=handle_builder,
creds_resolver=creds_resolver,
prisma_client=prisma_client,
)
swept += 1
return swept
async def heartbeat_timeout_sweeper(
*,
provider: AgentVMProvider,
prisma_client: Any,
config: SweeperConfig,
handle_builder: HandleBuilder,
creds_resolver: Optional[CredsResolver] = None,
) -> int:
"""One pass: terminate any `ready` session whose daemon stopped checking in."""
cutoff = _now_utc() - timedelta(seconds=config.heartbeat_timeout_seconds)
swept = 0
try:
rows = await prisma_client.db.litellm_agentsession.find_many(
where={
"status": STATUS_READY,
"last_heartbeat_at": {"lt": cutoff},
},
take=100,
order={"last_heartbeat_at": "asc"},
)
except Exception as e:
verbose_proxy_logger.debug(
f"heartbeat_timeout_sweeper: find_many failed: {type(e).__name__}"
)
return 0
for row in rows or []:
await _terminate_and_mark(
session=row,
new_status=STATUS_TERMINATED,
failure_reason="heartbeat_timeout",
provider=provider,
handle_builder=handle_builder,
creds_resolver=creds_resolver,
prisma_client=prisma_client,
)
swept += 1
return swept
async def max_session_minutes_sweeper(
*,
provider: AgentVMProvider,
prisma_client: Any,
config: SweeperConfig,
handle_builder: HandleBuilder,
creds_resolver: Optional[CredsResolver] = None,
) -> int:
"""One pass: terminate any session older than `max_session_minutes`."""
cutoff = _now_utc() - timedelta(minutes=config.max_session_minutes)
swept = 0
try:
rows = await prisma_client.db.litellm_agentsession.find_many(
where={
"status": {"in": [STATUS_PROVISIONING, STATUS_READY]},
"created_at": {"lt": cutoff},
},
take=100,
order={"created_at": "asc"},
)
except Exception as e:
verbose_proxy_logger.debug(
f"max_session_minutes_sweeper: find_many failed: {type(e).__name__}"
)
return 0
for row in rows or []:
await _terminate_and_mark(
session=row,
new_status=STATUS_TERMINATED,
failure_reason="max_session_minutes",
provider=provider,
handle_builder=handle_builder,
creds_resolver=creds_resolver,
prisma_client=prisma_client,
)
swept += 1
return swept
async def sweeper_loop(
*,
provider: AgentVMProvider,
prisma_client: Any,
config: SweeperConfig,
handle_builder: HandleBuilder,
creds_resolver: Optional[CredsResolver] = None,
stop_event: Optional[asyncio.Event] = None,
) -> None:
"""
Long-running loop: run all three sweepers every `sweep_interval_seconds`.
Started by `proxy_server.py` at boot via `asyncio.create_task` (see the
existing pattern around `_adaptive_router_flusher_loop`). Exits cleanly
when `stop_event` is set or the task is cancelled.
"""
verbose_proxy_logger.info(
f"agent_session sweepers running "
f"(interval={config.sweep_interval_seconds}s, "
f"bootstrap_timeout={config.bootstrap_timeout_seconds}s, "
f"heartbeat_timeout={config.heartbeat_timeout_seconds}s, "
f"max_session_minutes={config.max_session_minutes}m)"
)
while True:
if stop_event is not None and stop_event.is_set():
return
try:
await bootstrap_timeout_sweeper(
provider=provider,
prisma_client=prisma_client,
config=config,
handle_builder=handle_builder,
creds_resolver=creds_resolver,
)
await heartbeat_timeout_sweeper(
provider=provider,
prisma_client=prisma_client,
config=config,
handle_builder=handle_builder,
creds_resolver=creds_resolver,
)
await max_session_minutes_sweeper(
provider=provider,
prisma_client=prisma_client,
config=config,
handle_builder=handle_builder,
creds_resolver=creds_resolver,
)
except asyncio.CancelledError:
raise
except Exception as e:
verbose_proxy_logger.exception(
f"sweeper_loop: unhandled error: {type(e).__name__}: {e}"
)
try:
if stop_event is not None:
await asyncio.wait_for(
stop_event.wait(), timeout=config.sweep_interval_seconds
)
return
else:
await asyncio.sleep(config.sweep_interval_seconds)
except asyncio.TimeoutError:
continue
except asyncio.CancelledError:
raise

View file

@ -1,19 +1,60 @@
"""VM provider implementations for agent sessions."""
"""
Pluggable VM providers for `agent_session_endpoints`.
Public re-exports:
- `AgentVMProvider` — the ABC every provider implements
- `ProvisionContext` / `VMHandle` / `VMStatus` / `VMState` — typed I/O
- `AwsCreds` / `Ec2Config` — BYOC inputs
- `get_vm_provider` — registry/factory keyed off provider name
- `register_vm_provider` / `reset_vm_provider_registry` — test helpers
- `ProvisionError` / `InvalidCredentialsError` — error types
"""
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
AgentVMProvider,
ProvisionResult,
AwsCreds,
Ec2Config,
InvalidCredentialsError,
ProvisionContext,
ProvisionError,
Repo,
VMHandle,
VMState,
VMStatus,
)
from litellm.proxy.agent_session_endpoints.vm_providers.noop import NoopVMProvider
from litellm.proxy.agent_session_endpoints.vm_providers.ec2 import EC2Provider
from litellm.proxy.agent_session_endpoints.vm_providers.factory import (
SUPPORTED_PROVIDERS,
build_vm_provider,
)
from litellm.proxy.agent_session_endpoints.vm_providers.noop import NoopProvider
from litellm.proxy.agent_session_endpoints.vm_providers.registry import (
get_vm_provider,
register_vm_provider,
reset_vm_provider_registry,
)
# Backward-compat alias — Epic A's session_endpoints + tests import the noop
# under the older `NoopVMProvider` name.
NoopVMProvider = NoopProvider
__all__ = [
"AgentVMProvider",
"ProvisionResult",
"AwsCreds",
"Ec2Config",
"EC2Provider",
"InvalidCredentialsError",
"NoopProvider",
"NoopVMProvider",
"ProvisionContext",
"ProvisionError",
"Repo",
"SUPPORTED_PROVIDERS",
"VMHandle",
"VMState",
"VMStatus",
"build_vm_provider",
"get_vm_provider",
"register_vm_provider",
"reset_vm_provider_registry",
]

View file

@ -1,56 +1,161 @@
"""Abstract base class for agent session VM providers."""
"""
Base abstraction for the agent VM provider.
Each session in `agent_session_endpoints` provisions one VM via a provider that
implements `AgentVMProvider`. The v1 implementation is `EC2Provider` (in
`ec2.py`); `NoopProvider` exists for tests and config-driven swap.
The abstraction takes BYOC AWS creds inside `ProvisionContext` so the proxy
holds no AWS creds itself.
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Dict, List, Optional
@dataclass
class ProvisionResult:
"""Result returned by ``provision``.
class Repo:
"""Single git repo to clone into the VM."""
`vm_id` is the provider-native identifier (e.g. EC2 instance id).
`metadata` is opaque provider state stored on the session row for
later termination.
url: str
ref: Optional[str] = None # branch / tag / sha
path: Optional[str] = None # mount path inside the VM
@dataclass
class AwsCreds:
"""
BYOC AWS credentials for a team.
These are decrypted at use, never logged. The provider passes them straight
to `boto3.Session(...)` and then drops the reference.
"""
access_key_id: str
secret_access_key: str
session_token: Optional[str] = None
region: str = "us-west-2"
def __repr__(self) -> str:
# Defensive __repr__ so we cannot accidentally print creds in logs.
return (
f"AwsCreds(region={self.region!r}, access_key_id=***REDACTED***, "
f"secret_access_key=***REDACTED***, "
f"session_token={'***' if self.session_token else None})"
)
def __str__(self) -> str:
return self.__repr__()
@dataclass
class Ec2Config:
"""Per-team EC2 overrides resolved from `LiteLLM_AgentVMConfig`."""
region: str = "us-west-2"
subnet_id: Optional[str] = None
security_group_id: Optional[str] = None
iam_instance_profile: Optional[str] = None
instance_type: str = "t3.large"
use_spot: bool = True
ami_id: Optional[str] = None # falls back to `default_ami_id` in config.yaml
@dataclass
class ProvisionContext:
"""Inputs to `AgentVMProvider.provision`."""
session_id: str
team_id: str
agent_id: Optional[str] = None
repos: List[Repo] = field(default_factory=list)
env_vars: Dict[str, str] = field(default_factory=dict)
secrets: Dict[str, str] = field(default_factory=dict) # from Epic G
runtime_config: Dict[str, Any] = field(default_factory=dict)
aws_creds: Optional[AwsCreds] = None # BYOC; required for the EC2 provider
ec2_config: Optional[Ec2Config] = None
# Daemon callback fields. Populated by the session-create endpoint.
daemon_jwt: Optional[str] = None
daemon_base_url: Optional[str] = None
# `session` for cold-boot, `warm` for warm-pool prewarming (B2).
mode: str = "session"
class VMState(str, Enum):
"""Coarse VM lifecycle states surfaced to the rest of the proxy."""
PENDING = "pending"
RUNNING = "running"
STOPPING = "stopping"
STOPPED = "stopped"
TERMINATED = "terminated"
UNKNOWN = "unknown"
@dataclass
class VMHandle:
"""
Opaque handle returned by `provision`.
`vm_id` is the provider-native id (EC2 instance id, etc.). `metadata`
carries the per-provider state we need for `terminate` / `status` (region,
purchase mode, public IP if any).
"""
vm_id: str
metadata: Optional[Dict[str, Any]] = None
provider: str # `noop`, `ec2`, ...
region: Optional[str] = None
metadata: Dict[str, Any] = field(default_factory=dict)
@dataclass
class VMStatus:
"""Result of `AgentVMProvider.status`."""
state: VMState
public_ip: Optional[str] = None
private_ip: Optional[str] = None
raw: Dict[str, Any] = field(default_factory=dict)
class AgentVMProvider(ABC):
"""Pluggable VM backend for agent sessions.
"""
Pluggable provider for per-session VMs.
Concrete implementations:
* ``NoopVMProvider`` — for tests; returns canned ids
* ``EC2VMProvider`` — Epic B
* ``VercelSandboxProvider``— future
Implementations must be safe to call concurrently. Implementations MUST NOT
log AWS credentials; the only place creds enter is via the `AwsCreds` field
on `ProvisionContext`.
"""
name: str = "base"
name: str # set by subclass: `noop`, `ec2`, ...
@abstractmethod
async def provision(
self,
session_id: str,
agent_id: str,
repos: List[Dict[str, Any]],
env_vars: Optional[Dict[str, str]],
daemon_token: str,
proxy_base_url: str,
) -> ProvisionResult:
"""Provision a new VM for ``session_id``.
Implementations MUST be idempotent — they may be called twice if
the cleanup sweeper retries a stuck-provisioning session.
"""
async def provision(self, ctx: ProvisionContext) -> VMHandle:
"""Create a new VM. Must be idempotent on `ctx.session_id`."""
@abstractmethod
async def terminate(
self,
session_id: str,
vm_id: Optional[str],
metadata: Optional[Dict[str, Any]],
) -> None:
"""Tear down a VM. MUST be idempotent — terminating an already-gone
VM is a no-op."""
async def terminate(self, vm: VMHandle) -> None:
"""Terminate the VM. Must be idempotent (no-op if already terminated)."""
@abstractmethod
async def status(self, vm: VMHandle) -> VMStatus:
"""Return the current VM status."""
class ProvisionError(Exception):
"""Raised when provisioning fails. Carries an HTTP-mappable status hint."""
def __init__(self, message: str, status_code: int = 500) -> None:
super().__init__(message)
self.status_code = status_code
class InvalidCredentialsError(ProvisionError):
"""BYOC creds are missing, expired, or rejected by the provider."""
def __init__(self, message: str) -> None:
super().__init__(message, status_code=400)

View file

@ -0,0 +1,488 @@
"""
EC2 implementation of `AgentVMProvider` (BYOC).
Each session gets a dedicated EC2 instance launched in the *team's* AWS
account using the team's BYOC creds. The proxy never holds AWS creds itself.
Key design points:
- creds enter via `ProvisionContext.aws_creds` and never leave this module
- `boto3` debug logging is silenced (`set_stream_logger`) so SigV4-signing
payloads cannot accidentally leak the access key
- spot is tried first, on-demand is the fallback
- every `RunInstances` is paired with a `TerminateInstances` retry path
The constants `_BOTO3_RETRY_MODE`, `_SPOT_INTERRUPTION_BEHAVIOR` and
`_USER_DATA_TEMPLATE` are tuned to match the values measured in the B0
spike (LIT-2888) — see deliverables comment on that ticket.
"""
from __future__ import annotations
import asyncio
import base64
import logging
from concurrent.futures import ThreadPoolExecutor
from typing import Any, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
AgentVMProvider,
AwsCreds,
Ec2Config,
InvalidCredentialsError,
ProvisionContext,
ProvisionError,
Repo,
VMHandle,
VMState,
VMStatus,
)
# `botocore` logs every signed request at DEBUG level. We force WARNING so
# the AWS access key never lands in proxy logs even when the operator turns
# on litellm DEBUG logging. This is enforced again in `_silence_boto_logs`
# below in case some other code path raises the level back up.
_BOTO_LOGGERS = ("boto3", "botocore", "urllib3.connectionpool", "s3transfer")
# Tuned in B0: max 5 short retries; longer waits make sessions feel sluggish.
_BOTO3_RETRY_MODE = "standard"
_BOTO3_RETRY_MAX_ATTEMPTS = 5
# Spot interruption behaviour: terminate (default) — we don't want hibernation
# because we destroy the VM at session end anyway.
_SPOT_INTERRUPTION_BEHAVIOR = "terminate"
# Map from EC2 lifecycle state to our VMState enum.
_EC2_STATE_TO_VMSTATE = {
"pending": VMState.PENDING,
"running": VMState.RUNNING,
"shutting-down": VMState.STOPPING,
"stopping": VMState.STOPPING,
"stopped": VMState.STOPPED,
"terminated": VMState.TERMINATED,
}
# AWS errors that mean "your creds are bad" — surface these as 400 to the user.
_INVALID_CRED_CODES = {
"InvalidClientTokenId",
"AuthFailure",
"SignatureDoesNotMatch",
"UnauthorizedOperation",
"OptInRequired",
"AccessDenied",
}
# AWS error codes that mean "spot capacity unavailable" — fall back to on-demand.
_SPOT_UNAVAILABLE_CODES = {
"InsufficientInstanceCapacity",
"SpotMaxPriceTooLow",
"MaxSpotInstanceCountExceeded",
"InstanceLimitExceeded",
}
def _silence_boto_logs() -> None:
"""Force every boto/botocore logger to WARNING so creds never leak."""
for name in _BOTO_LOGGERS:
logging.getLogger(name).setLevel(logging.WARNING)
_silence_boto_logs()
def _build_user_data(ctx: ProvisionContext) -> str:
"""
Build the EC2 user-data shell script.
The script writes the daemon JWT + base URL into systemd's environment file
and starts the agent runtime daemon. NOTHING in this script logs the JWT
body — it goes straight into the EnvironmentFile that systemd reads.
"""
daemon_jwt = ctx.daemon_jwt or ""
base_url = ctx.daemon_base_url or ""
mode = ctx.mode or "session"
# Repos are written to /etc/litellm-agent/repos.json so the daemon can
# iterate them on first run. We base64 the payload to avoid quoting hell.
repos_json = _repos_to_json(ctx.repos)
env_json = _env_to_json(ctx.env_vars)
repos_b64 = base64.b64encode(repos_json.encode("utf-8")).decode("utf-8")
env_b64 = base64.b64encode(env_json.encode("utf-8")).decode("utf-8")
return f"""#!/bin/bash
set -e
mkdir -p /etc/litellm-agent
cat > /etc/litellm-agent/runtime.env <<EOF
LITELLM_SESSION_ID={ctx.session_id}
LITELLM_TEAM_ID={ctx.team_id}
LITELLM_AGENT_ID={ctx.agent_id or ''}
LITELLM_BASE_URL={base_url}
LITELLM_AGENT_MODE={mode}
LITELLM_DAEMON_JWT={daemon_jwt}
EOF
chmod 600 /etc/litellm-agent/runtime.env
echo "{repos_b64}" | base64 -d > /etc/litellm-agent/repos.json
echo "{env_b64}" | base64 -d > /etc/litellm-agent/env.json
chmod 600 /etc/litellm-agent/repos.json /etc/litellm-agent/env.json
systemctl enable --now litellm-agent-runtime.service || true
"""
def _repos_to_json(repos: List[Repo]) -> str:
import json
return json.dumps([{"url": r.url, "ref": r.ref, "path": r.path} for r in repos])
def _env_to_json(env: Dict[str, str]) -> str:
import json
return json.dumps(env)
def _safe_aws_error(exc: Exception) -> str:
"""
Format a boto3 ClientError without leaking the request payload.
`botocore.exceptions.ClientError.__str__` includes the SigV4 metadata in
some versions; we extract only the error code + message.
"""
response = getattr(exc, "response", None) or {}
err = response.get("Error", {}) if isinstance(response, dict) else {}
code = err.get("Code") or type(exc).__name__
message = err.get("Message") or "(no message)"
return f"{code}: {message}"
def _aws_error_code(exc: Exception) -> Optional[str]:
response = getattr(exc, "response", None) or {}
err = response.get("Error", {}) if isinstance(response, dict) else {}
return err.get("Code")
def _tag_specs(ctx: ProvisionContext) -> List[Dict[str, Any]]:
"""Tag every resource so cleanup tools can find it."""
tags = [
{"Key": "litellm-session-id", "Value": ctx.session_id},
{"Key": "litellm-team-id", "Value": ctx.team_id},
{"Key": "litellm-managed-by", "Value": "agent-vm-provider"},
]
if ctx.agent_id:
tags.append({"Key": "litellm-agent-id", "Value": ctx.agent_id})
return [
{"ResourceType": "instance", "Tags": tags},
{"ResourceType": "volume", "Tags": tags},
]
class EC2Provider(AgentVMProvider):
"""boto3-backed EC2 provisioner."""
name = "ec2"
def __init__(self, settings: Optional[Dict[str, Any]] = None) -> None:
self._settings: Dict[str, Any] = settings or {}
# ThreadPool to wrap synchronous boto3 calls for asyncio. Bounded so a
# spike of provision calls cannot exhaust the proxy.
self._executor = ThreadPoolExecutor(
max_workers=int(self._settings.get("max_concurrent_aws_calls", 16)),
thread_name_prefix="ec2-provider",
)
@property
def default_region(self) -> str:
return self._settings.get("default_region", "us-west-2")
@property
def default_ami_id(self) -> Optional[str]:
return self._settings.get("default_ami_id")
@property
def default_instance_type(self) -> str:
return self._settings.get("default_instance_type", "t3.large")
@property
def default_use_spot(self) -> bool:
return bool(self._settings.get("use_spot", True))
# ---------- public API ----------
async def provision(self, ctx: ProvisionContext) -> VMHandle:
if ctx.aws_creds is None:
raise InvalidCredentialsError(
"EC2 provider requires BYOC AWS credentials in ProvisionContext."
)
ec2_config = ctx.ec2_config or Ec2Config(region=ctx.aws_creds.region)
ami_id = ec2_config.ami_id or self.default_ami_id
if not ami_id:
raise ProvisionError(
"No AMI configured. Set agent_settings.ec2.default_ami_id "
"in config.yaml or per-team via LiteLLM_AgentVMConfig.ami_id."
)
instance_type = ec2_config.instance_type or self.default_instance_type
use_spot = (
ec2_config.use_spot if ec2_config is not None else self.default_use_spot
)
user_data = _build_user_data(ctx)
try:
instance = await asyncio.get_running_loop().run_in_executor(
self._executor,
self._run_instances_with_fallback,
ctx.aws_creds,
ami_id,
instance_type,
ec2_config,
user_data,
_tag_specs(ctx),
use_spot,
)
except InvalidCredentialsError:
raise
except ProvisionError:
raise
except Exception as e:
verbose_proxy_logger.exception(
"EC2 provision failed for session=%s team=%s: %s",
ctx.session_id,
ctx.team_id,
_safe_aws_error(e),
)
raise ProvisionError(
f"EC2 RunInstances failed: {_safe_aws_error(e)}"
) from e
return VMHandle(
vm_id=instance["InstanceId"],
provider=self.name,
region=ec2_config.region,
metadata={
"ami_id": ami_id,
"instance_type": instance_type,
"purchase_mode": instance.get("_purchase_mode", "on-demand"),
"subnet_id": ec2_config.subnet_id,
},
)
async def terminate(
self, vm: VMHandle, aws_creds: Optional[AwsCreds] = None
) -> None:
"""
Terminate the EC2 instance.
`aws_creds` is required because the team that owns the instance also
owns the creds. Callers fetch them via `team_config.get_team_vm_config`
right before calling `terminate`.
"""
if aws_creds is None:
raise InvalidCredentialsError(
"EC2 terminate requires BYOC AWS credentials."
)
try:
await asyncio.get_running_loop().run_in_executor(
self._executor,
self._terminate_instances_sync,
aws_creds,
vm.region or aws_creds.region,
vm.vm_id,
)
except Exception as e:
code = _aws_error_code(e)
# `InvalidInstanceID.NotFound` and `InvalidInstanceID.Malformed` mean
# the instance is already gone; treat as success.
if code in ("InvalidInstanceID.NotFound", "InvalidInstanceID.Malformed"):
verbose_proxy_logger.debug(
f"terminate: instance {vm.vm_id} already gone ({code})"
)
return
verbose_proxy_logger.exception(
"EC2 terminate failed for vm=%s: %s", vm.vm_id, _safe_aws_error(e)
)
raise ProvisionError(
f"EC2 TerminateInstances failed: {_safe_aws_error(e)}"
) from e
async def status(
self, vm: VMHandle, aws_creds: Optional[AwsCreds] = None
) -> VMStatus:
"""Return current EC2 status. Requires the team's BYOC creds."""
if aws_creds is None:
raise InvalidCredentialsError("EC2 status requires BYOC AWS credentials.")
try:
raw = await asyncio.get_running_loop().run_in_executor(
self._executor,
self._describe_instance_sync,
aws_creds,
vm.region or aws_creds.region,
vm.vm_id,
)
except Exception as e:
code = _aws_error_code(e)
if code in ("InvalidInstanceID.NotFound", "InvalidInstanceID.Malformed"):
return VMStatus(state=VMState.TERMINATED)
raise ProvisionError(
f"EC2 DescribeInstances failed: {_safe_aws_error(e)}"
) from e
if raw is None:
return VMStatus(state=VMState.TERMINATED)
ec2_state = (raw.get("State") or {}).get("Name", "unknown")
return VMStatus(
state=_EC2_STATE_TO_VMSTATE.get(ec2_state, VMState.UNKNOWN),
public_ip=raw.get("PublicIpAddress"),
private_ip=raw.get("PrivateIpAddress"),
raw=raw,
)
# ---------- sync (boto3) helpers, run in the thread pool ----------
def _build_ec2_client(self, creds: AwsCreds, region: str) -> Any:
"""Return a boto3 EC2 client. Only place creds touch boto3."""
# Re-silence each call (defensive: another module may have changed it).
_silence_boto_logs()
try:
import boto3 # noqa: PLC0415 imported here so litellm core works without boto3 installed
from botocore.config import Config # noqa: PLC0415
except ImportError as e:
raise ProvisionError(
"boto3 is required for the EC2 VM provider. "
"Install with: pip install boto3."
) from e
config = Config(
retries={
"max_attempts": _BOTO3_RETRY_MAX_ATTEMPTS,
"mode": _BOTO3_RETRY_MODE,
},
region_name=region,
)
return boto3.client(
"ec2",
region_name=region,
aws_access_key_id=creds.access_key_id,
aws_secret_access_key=creds.secret_access_key,
aws_session_token=creds.session_token,
config=config,
)
def _run_instances_with_fallback(
self,
creds: AwsCreds,
ami_id: str,
instance_type: str,
ec2_config: Ec2Config,
user_data: str,
tag_specs: List[Dict[str, Any]],
use_spot: bool,
) -> Dict[str, Any]:
"""Try spot, fall back to on-demand once if capacity unavailable."""
client = self._build_ec2_client(creds, ec2_config.region)
if use_spot:
try:
instance = self._run_instances_sync(
client,
ami_id=ami_id,
instance_type=instance_type,
ec2_config=ec2_config,
user_data=user_data,
tag_specs=tag_specs,
spot=True,
)
instance["_purchase_mode"] = "spot"
return instance
except Exception as e:
code = _aws_error_code(e)
if code in _INVALID_CRED_CODES:
raise InvalidCredentialsError(
f"AWS rejected the team's credentials: {code}"
) from e
if code in _SPOT_UNAVAILABLE_CODES:
verbose_proxy_logger.warning(
f"Spot capacity unavailable ({code}); falling back to on-demand."
)
# fall through to on-demand
else:
raise
instance = self._run_instances_sync(
client,
ami_id=ami_id,
instance_type=instance_type,
ec2_config=ec2_config,
user_data=user_data,
tag_specs=tag_specs,
spot=False,
)
instance["_purchase_mode"] = "on-demand"
return instance
def _run_instances_sync(
self,
client: Any,
*,
ami_id: str,
instance_type: str,
ec2_config: Ec2Config,
user_data: str,
tag_specs: List[Dict[str, Any]],
spot: bool,
) -> Dict[str, Any]:
kwargs: Dict[str, Any] = {
"ImageId": ami_id,
"InstanceType": instance_type,
"MinCount": 1,
"MaxCount": 1,
"UserData": user_data,
"TagSpecifications": tag_specs,
}
if ec2_config.subnet_id:
kwargs["SubnetId"] = ec2_config.subnet_id
if ec2_config.security_group_id:
kwargs["SecurityGroupIds"] = [ec2_config.security_group_id]
if ec2_config.iam_instance_profile:
kwargs["IamInstanceProfile"] = {"Name": ec2_config.iam_instance_profile}
if spot:
kwargs["InstanceMarketOptions"] = {
"MarketType": "spot",
"SpotOptions": {
"InstanceInterruptionBehavior": _SPOT_INTERRUPTION_BEHAVIOR,
"SpotInstanceType": "one-time",
},
}
try:
response = client.run_instances(**kwargs)
except Exception as e:
code = _aws_error_code(e)
if code in _INVALID_CRED_CODES:
raise InvalidCredentialsError(
f"AWS rejected the team's credentials: {code}"
) from e
raise
instances = response.get("Instances", [])
if not instances:
raise ProvisionError("RunInstances succeeded but returned no instances.")
return instances[0]
def _terminate_instances_sync(
self, creds: AwsCreds, region: str, instance_id: str
) -> None:
client = self._build_ec2_client(creds, region)
client.terminate_instances(InstanceIds=[instance_id])
def _describe_instance_sync(
self, creds: AwsCreds, region: str, instance_id: str
) -> Optional[Dict[str, Any]]:
client = self._build_ec2_client(creds, region)
response = client.describe_instances(InstanceIds=[instance_id])
for res in response.get("Reservations", []):
for inst in res.get("Instances", []):
return inst
return None

View file

@ -0,0 +1,69 @@
"""
Factory for `AgentVMProvider` implementations.
Reads `agent_settings.vm_provider` from the loaded proxy config. Default is
`noop` so the proxy starts cleanly without AWS configured.
Validation criterion #1 (`test_factory`) covers:
- factory returns the right impl for each value
- unknown values raise `ValueError`
- defaults to `noop` when no config is set
This module is the *config-driven* path: at proxy startup we call
``build_vm_provider(agent_settings)`` and register the result via
``register_vm_provider`` (see ``registry.py``). Runtime callers (session
endpoints, sweepers, tests) then look the provider up by name with
``get_vm_provider(name)``.
"""
from __future__ import annotations
from typing import Any, Dict, Optional
from litellm.proxy.agent_session_endpoints.vm_providers.base import AgentVMProvider
from litellm.proxy.agent_session_endpoints.vm_providers.ec2 import EC2Provider
from litellm.proxy.agent_session_endpoints.vm_providers.noop import NoopProvider
# Registry of {name: factory_callable}. The factory takes the raw settings
# dict (e.g. `agent_settings.ec2`) so per-provider config stays inside the
# provider class, not the factory.
_PROVIDER_REGISTRY = {
"noop": lambda settings: NoopProvider(),
"ec2": lambda settings: EC2Provider(settings=settings),
}
SUPPORTED_PROVIDERS = tuple(_PROVIDER_REGISTRY.keys())
def build_vm_provider(
agent_settings: Optional[Dict[str, Any]] = None,
) -> AgentVMProvider:
"""
Build the configured provider from a config block.
`agent_settings` is the `agent_settings` block from `config.yaml`:
```yaml
agent_settings:
vm_provider: ec2
ec2:
default_region: us-west-2
default_ami_id: ami-...
...
```
Defaults to `noop` if `agent_settings` is None or `vm_provider` is missing.
Raises `ValueError` for unknown provider names.
"""
settings = agent_settings or {}
provider_name = settings.get("vm_provider", "noop")
if provider_name not in _PROVIDER_REGISTRY:
raise ValueError(
f"Unknown vm_provider {provider_name!r}. "
f"Supported: {', '.join(sorted(_PROVIDER_REGISTRY.keys()))}."
)
provider_settings = settings.get(provider_name, {}) or {}
return _PROVIDER_REGISTRY[provider_name](provider_settings)

View file

@ -1,79 +1,114 @@
"""
Noop VM provider — for tests.
No-op `AgentVMProvider` implementation.
Records every call so tests can assert provider.provision/terminate were
invoked with expected arguments.
Used in tests and for environments without an EC2 backend. `provision` returns
a fake handle immediately; `status` always reports `running`. The provider is
the default in `factory.py` so the proxy starts up without AWS configured.
Records every `provision`/`terminate` call so legacy A-era tests can assert
the endpoints invoked the provider with the expected arguments. Recording is
thread-safe.
"""
from __future__ import annotations
import threading
import uuid
from typing import Any, Dict, List, Optional
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
AgentVMProvider,
ProvisionResult,
ProvisionContext,
VMHandle,
VMState,
VMStatus,
)
class NoopVMProvider(AgentVMProvider):
"""In-process provider: returns ``noop_<uuid>`` instance ids.
Thread-safe call recording. Tests inspect ``provision_calls`` /
``terminate_calls`` to verify the endpoints invoked the provider
correctly.
"""
class NoopProvider(AgentVMProvider):
"""In-memory provider used by tests and `vm_provider: noop` config."""
name = "noop"
def __init__(self, fail_provision: bool = False) -> None:
self._lock = threading.Lock()
self._terminated: set = set()
# Recording for legacy A-era tests.
self.provision_calls: List[Dict[str, Any]] = []
self.terminate_calls: List[Dict[str, Any]] = []
self.fail_provision = fail_provision
async def provision(
self,
session_id: str,
agent_id: str,
repos: List[Dict[str, Any]],
env_vars: Optional[Dict[str, str]],
daemon_token: str,
proxy_base_url: str,
) -> ProvisionResult:
async def provision(self, ctx: ProvisionContext) -> VMHandle:
with self._lock:
self.provision_calls.append(
{
"session_id": session_id,
"agent_id": agent_id,
"repos": repos,
"env_vars_set": list((env_vars or {}).keys()),
"proxy_base_url": proxy_base_url,
"session_id": ctx.session_id,
"team_id": ctx.team_id,
"agent_id": ctx.agent_id,
"repos": [
{"url": r.url, "ref": r.ref, "path": r.path}
for r in ctx.repos
],
"env_vars_set": list((ctx.env_vars or {}).keys()),
"mode": ctx.mode,
"daemon_base_url": ctx.daemon_base_url,
}
)
if self.fail_provision:
raise RuntimeError("noop provider configured to fail provisioning")
return ProvisionResult(
vm_id=f"noop_{uuid.uuid4().hex[:12]}",
metadata={"provider": "noop"},
vm_id = f"noop-{uuid.uuid4().hex[:12]}"
return VMHandle(
vm_id=vm_id,
provider=self.name,
region=(ctx.ec2_config.region if ctx.ec2_config else None),
metadata={
"session_id": ctx.session_id,
"team_id": ctx.team_id,
"mode": ctx.mode,
},
)
async def terminate(
self,
session_id: str,
vm_id: Optional[str],
metadata: Optional[Dict[str, Any]],
vm: Optional[VMHandle] = None,
*,
session_id: Optional[str] = None,
vm_id: Optional[str] = None,
metadata: Optional[Dict[str, Any]] = None,
) -> None:
"""Idempotent terminate.
Accepts either B's API (``vm: VMHandle``) or the legacy keyword form
(``session_id=..., vm_id=..., metadata=...``) used by A-era callers
that may still construct a partial handle from a session row.
"""
if vm is not None:
recorded_session = (vm.metadata or {}).get("session_id")
recorded_vm_id = vm.vm_id
recorded_meta = vm.metadata
else:
recorded_session = session_id
recorded_vm_id = vm_id
recorded_meta = metadata
with self._lock:
self.terminate_calls.append(
{
"session_id": session_id,
"vm_id": vm_id,
"metadata": metadata,
"session_id": recorded_session,
"vm_id": recorded_vm_id,
"metadata": recorded_meta,
}
)
if recorded_vm_id:
self._terminated.add(recorded_vm_id)
async def status(self, vm: VMHandle) -> VMStatus:
if vm.vm_id in self._terminated:
return VMStatus(state=VMState.TERMINATED)
return VMStatus(state=VMState.RUNNING, public_ip="127.0.0.1")
def reset(self) -> None:
"""Clear recorded calls — for use between tests."""
"""Test helper: clear recorded calls between cases."""
with self._lock:
self.provision_calls.clear()
self.terminate_calls.clear()

View file

@ -1,28 +1,41 @@
"""Process-wide registry of VM providers, keyed by provider name."""
"""Process-wide registry of VM providers, keyed by provider name.
The registry is the runtime lookup path used by ``session_endpoints.py``,
``sweepers.py``, and tests. At proxy startup we typically build the configured
provider via ``factory.build_vm_provider(agent_settings)`` and then call
``register_vm_provider(provider)`` to install it under its ``provider.name``.
For tests / local-dev convenience, ``get_vm_provider("noop")`` lazily
instantiates a default ``NoopProvider`` on first access so callers don't need
a setup hook just to use the noop.
"""
from typing import Dict
from litellm.proxy.agent_session_endpoints.vm_providers.base import AgentVMProvider
from litellm.proxy.agent_session_endpoints.vm_providers.noop import NoopVMProvider
from litellm.proxy.agent_session_endpoints.vm_providers.noop import NoopProvider
_REGISTRY: Dict[str, AgentVMProvider] = {}
def register_vm_provider(provider: AgentVMProvider) -> None:
"""Register a provider by ``provider.name``. Last-write-wins; tests
use this to swap in a fresh ``NoopVMProvider`` between cases."""
use this to swap in a fresh ``NoopProvider`` between cases."""
_REGISTRY[provider.name] = provider
def get_vm_provider(name: str) -> AgentVMProvider:
"""Return the registered provider for ``name``.
Lazily instantiates a default ``NoopVMProvider`` on first access so
tests don't need a setup hook just to use the noop.
Lazily instantiates a default ``NoopProvider`` on first access so tests
don't need a setup hook just to use the noop. For other names (e.g.
``"ec2"``) the caller MUST register an instance first via
``register_vm_provider`` (typically at proxy startup from
``factory.build_vm_provider``).
"""
if name not in _REGISTRY:
if name == "noop":
_REGISTRY[name] = NoopVMProvider()
_REGISTRY[name] = NoopProvider()
else:
raise KeyError(f"No VM provider registered for '{name}'")
return _REGISTRY[name]

View file

@ -0,0 +1,187 @@
"""
Team-scoped BYOC AWS config + creds resolver.
Reads `LiteLLM_AgentVMConfig` (table owned by Epic G / LIT-2891) and decrypts
the per-field-encrypted AWS credentials using the proxy's salt key. Falls back
to env vars `LITELLM_AGENT_AWS_ACCESS_KEY_ID` / `..._SECRET_ACCESS_KEY` /
`..._SESSION_TOKEN` so local dev works without a DB row.
Epic G's schema stores creds as **separate encrypted columns** (not a single
JSON blob). Columns:
provider -- "ec2" | "self_hosted" | "disabled"
aws_auth_method -- "access_keys" | "iam_role" | "instance_metadata"
aws_access_key_id_enc -- encrypt_value_helper(...)
aws_secret_access_key_enc -- encrypt_value_helper(...)
aws_role_arn_enc -- encrypt_value_helper(...) (cross-account mode)
aws_region -- plain
ami_id, instance_type, subnet_id, security_group_id,
iam_instance_profile, use_spot, ...
Reconciliation note: this module was originally written against Epic B's
single-column ``aws_creds_enc`` shape. The integration branch adopted Epic G's
column-per-field shape; this resolver was rewritten to match.
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from typing import Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
AwsCreds,
Ec2Config,
InvalidCredentialsError,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
@dataclass
class TeamVMConfig:
"""Resolved team-scoped VM config (creds + EC2 overrides)."""
aws_creds: AwsCreds
ec2_config: Ec2Config
def encrypt_aws_field(value: Optional[str], field_name: str) -> Optional[str]:
"""Encrypt a single AWS credential field. Empty/None passes through."""
if not value:
return None
return encrypt_value_helper(value)
def decrypt_aws_field(blob: Optional[str], field_name: str) -> Optional[str]:
"""Decrypt a single AWS credential field. Empty/None passes through.
Returns ``None`` (not raises) if the blob is malformed — the caller is
expected to validate that the required fields are present.
"""
if not blob:
return None
try:
return decrypt_value_helper(blob, key=field_name)
except Exception as e:
verbose_proxy_logger.warning(
f"team_config: decrypt failed for {field_name}: {type(e).__name__}"
)
return None
def _ec2_config_from_row(row: Any, default_region: str) -> Ec2Config:
"""Build an `Ec2Config` from a `LiteLLM_AgentVMConfig` Prisma row.
Reads G's column names: ``aws_region`` (not ``region``).
"""
return Ec2Config(
region=getattr(row, "aws_region", None) or default_region,
subnet_id=getattr(row, "subnet_id", None),
security_group_id=getattr(row, "security_group_id", None),
iam_instance_profile=getattr(row, "iam_instance_profile", None),
instance_type=getattr(row, "instance_type", None) or "t3.large",
use_spot=bool(getattr(row, "use_spot", True)),
ami_id=getattr(row, "ami_id", None),
)
def _creds_from_env(default_region: str) -> Optional[AwsCreds]:
"""Build creds from `LITELLM_AGENT_AWS_*` env vars; None if unset."""
access_key_id = os.getenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID")
secret_access_key = os.getenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY")
if not access_key_id or not secret_access_key:
return None
return AwsCreds(
access_key_id=access_key_id,
secret_access_key=secret_access_key,
session_token=os.getenv("LITELLM_AGENT_AWS_SESSION_TOKEN"),
region=os.getenv("LITELLM_AGENT_AWS_REGION", default_region),
)
def _creds_from_row(row: Any, default_region: str) -> Optional[AwsCreds]:
"""Decrypt G's per-field encrypted creds. Returns None if not configured.
Only handles ``aws_auth_method == 'access_keys'``. Other modes
(``iam_role`` / ``instance_metadata``) defer credential lookup to the
boto3 default chain at provision time and don't surface raw keys here.
"""
auth_method = getattr(row, "aws_auth_method", None)
if auth_method and auth_method != "access_keys":
# IAM role / instance metadata — caller should construct AwsCreds via
# boto3's default chain instead of pulling from the DB.
return None
access_key_id = decrypt_aws_field(
getattr(row, "aws_access_key_id_enc", None), "aws_access_key_id"
)
secret_access_key = decrypt_aws_field(
getattr(row, "aws_secret_access_key_enc", None), "aws_secret_access_key"
)
if not access_key_id or not secret_access_key:
return None
return AwsCreds(
access_key_id=access_key_id,
secret_access_key=secret_access_key,
session_token=None,
region=getattr(row, "aws_region", None) or default_region,
)
async def get_team_vm_config(
team_id: str,
prisma_client: Any,
default_region: str = "us-west-2",
) -> TeamVMConfig:
"""
Resolve the team's BYOC AWS creds + EC2 overrides.
Lookup order:
1. `LiteLLM_AgentVMConfig` row keyed by `team_id` (G's per-field columns)
2. `LITELLM_AGENT_AWS_*` env vars (local dev fallback)
Raises `InvalidCredentialsError` if neither path yields creds.
"""
row = None
if prisma_client is not None:
try:
row = await prisma_client.db.litellm_agentvmconfig.find_unique(
where={"team_id": team_id}
)
except Exception as e:
# If the table doesn't exist yet (G1 hasn't run the migration),
# silently fall through to env-var path so local dev is unblocked.
verbose_proxy_logger.debug(
f"LiteLLM_AgentVMConfig lookup failed for team={team_id}: "
f"{type(e).__name__}; falling back to env vars."
)
if row is not None:
ec2_config = _ec2_config_from_row(row, default_region=default_region)
creds = _creds_from_row(row, default_region=default_region)
if creds is None:
raise InvalidCredentialsError(
f"Team {team_id} has no usable AWS credentials configured. "
"Add them under Settings → Cloud Agents."
)
# Override creds region with row.aws_region if explicitly set.
if ec2_config.region:
creds.region = ec2_config.region
return TeamVMConfig(aws_creds=creds, ec2_config=ec2_config)
env_creds = _creds_from_env(default_region=default_region)
if env_creds is None:
raise InvalidCredentialsError(
f"Team {team_id} has no AWS credentials configured "
"(no DB row, no LITELLM_AGENT_AWS_* env vars). "
"Add them under Settings → Cloud Agents."
)
return TeamVMConfig(
aws_creds=env_creds,
ec2_config=Ec2Config(region=env_creds.region),
)

View file

@ -0,0 +1,41 @@
# Example: agent_session VM provider with EC2 (BYOC AWS).
#
# Provisions one EC2 per agent session in the *team's* AWS account.
# Per-team overrides (creds, subnet, IAM profile, AMI, instance type) live in
# `LiteLLM_AgentVMConfig`, populated via Settings UI (Epic G).
#
# See:
# litellm/proxy/agent_session_endpoints/vm_providers/ec2.py
# infra/ami/README.md (how to build the AMI)
model_list:
- model_name: gpt-4
litellm_params:
model: openai/gpt-4
api_key: os.environ/OPENAI_API_KEY
agent_settings:
vm_provider: ec2 # or `noop` for environments without AWS
sweep_interval_seconds: 30 # bootstrap/heartbeat/max-session sweepers tick
ec2:
# Defaults applied when a team's `LiteLLM_AgentVMConfig` row leaves the
# field NULL. Captured from the B0 spike.
default_region: us-west-2
# Baked by `infra/ami/litellm-agent-runtime.pkr.hcl`. The value below is
# the BYOC PoC account's AMI built from this PR; rebuild for your own
# account and replace.
default_ami_id: ami-074a518157fe137b4
default_instance_type: t3.large
use_spot: true
max_session_minutes: 120
bootstrap_timeout_seconds: 180
heartbeat_timeout_seconds: 120
warm_pool:
enabled: false # B2 (LIT-2890) implements the warm-pool path
size: 2
max_idle_minutes: 30
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY

View file

@ -80,6 +80,9 @@ model LiteLLM_AgentsTable {
updated_by String
}
// LiteLLM_AgentVMConfig is defined further down (LIT-2891 / Epic G section).
// Owned by Epic G's Settings UI; consumed by Epic B's EC2 provider.
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String

View file

@ -261,6 +261,7 @@ markers = [
"asyncio: mark test as an asyncio test",
"limit_leaks: mark test with memory limit for leak detection (e.g., '40 MB')",
"no_parallel: mark test to run sequentially (not in parallel) - typically for memory measurement tests",
"slow: mark test as a slow real-cloud test (e.g., real EC2 RunInstances). Skipped by default; opt in with `pytest -m slow`.",
]
filterwarnings = [
# Suppress Pydantic serializer warnings from mock server responses (non-critical for memory tests)

View file

@ -80,6 +80,9 @@ model LiteLLM_AgentsTable {
updated_by String
}
// LiteLLM_AgentVMConfig is defined further down (LIT-2891 / Epic G section).
// Owned by Epic G's Settings UI; consumed by Epic B's EC2 provider.
model LiteLLM_OrganizationTable {
organization_id String @id @default(uuid())
organization_alias String

View file

@ -0,0 +1,270 @@
"""
Tests for the agent-session sweepers.
Covers:
- Validation #7: max_session_minutes — sessions older than ceiling get terminated
- Validation #9: bootstrap_timeout — sessions stuck in `provisioning` get terminated
- Validation #10: heartbeat_loss — `ready` sessions whose daemon went quiet get terminated
- The optimistic re-fetch lock: sweeper does not double-terminate a session
whose status changed underneath it.
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.agent_session_endpoints.sweepers import (
SweeperConfig,
bootstrap_timeout_sweeper,
heartbeat_timeout_sweeper,
max_session_minutes_sweeper,
STATUS_FAILED,
STATUS_PROVISIONING,
STATUS_READY,
STATUS_TERMINATED,
)
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
AgentVMProvider,
ProvisionContext,
VMHandle,
VMState,
VMStatus,
)
class _FakeProvider(AgentVMProvider):
name = "fake"
def __init__(self) -> None:
self.terminate_calls: List[str] = []
async def provision(self, ctx: ProvisionContext) -> VMHandle: # pragma: no cover
raise AssertionError("provision should not be called by sweepers")
async def terminate(self, vm: VMHandle, **_kwargs) -> None:
self.terminate_calls.append(vm.vm_id)
async def status(self, vm: VMHandle, **_kwargs) -> VMStatus: # pragma: no cover
return VMStatus(state=VMState.UNKNOWN)
class _Row:
"""Mutable struct for `LiteLLM_AgentSession` row attribute access."""
def __init__(self, **kwargs: Any) -> None:
for k, v in kwargs.items():
setattr(self, k, v)
def _build_handle(row: Any) -> VMHandle:
return VMHandle(vm_id=row.vm_id, provider="fake", region="us-west-2")
def _build_prisma(rows: Dict[str, Any], find_many_results: List[Any]) -> Any:
"""Wire up a fake Prisma client whose store is a {session_id: row} dict."""
prisma = MagicMock()
prisma.db = MagicMock()
table = MagicMock()
async def find_many(where: Dict[str, Any], **_kwargs):
return list(find_many_results)
async def find_unique(where: Dict[str, Any]):
return rows.get(where["session_id"])
async def update(where: Dict[str, Any], data: Dict[str, Any]):
row = rows.get(where["session_id"])
if row is None:
raise Exception("row not found")
for k, v in data.items():
setattr(row, k, v)
return row
table.find_many = AsyncMock(side_effect=find_many)
table.find_unique = AsyncMock(side_effect=find_unique)
table.update = AsyncMock(side_effect=update)
prisma.db.litellm_agentsession = table
return prisma
@pytest.mark.asyncio
async def test_bootstrap_timeout_sweeper_terminates_stuck_provisioning():
now = datetime.now(timezone.utc)
stuck_row = _Row(
session_id="sess-stuck",
status=STATUS_PROVISIONING,
vm_id="i-stuck",
created_at=now - timedelta(seconds=600),
last_heartbeat_at=None,
)
fresh_row = _Row(
session_id="sess-fresh",
status=STATUS_PROVISIONING,
vm_id="i-fresh",
created_at=now,
last_heartbeat_at=None,
)
rows = {stuck_row.session_id: stuck_row, fresh_row.session_id: fresh_row}
# find_many returns only the stuck one (the WHERE filter is the prisma side).
prisma = _build_prisma(rows, find_many_results=[stuck_row])
provider = _FakeProvider()
config = SweeperConfig(bootstrap_timeout_seconds=180)
swept = await bootstrap_timeout_sweeper(
provider=provider,
prisma_client=prisma,
config=config,
handle_builder=_build_handle,
)
assert swept == 1
assert provider.terminate_calls == ["i-stuck"]
assert stuck_row.status == STATUS_FAILED
assert stuck_row.failure_reason == "bootstrap_timeout"
# Untouched session stays in provisioning.
assert fresh_row.status == STATUS_PROVISIONING
@pytest.mark.asyncio
async def test_heartbeat_timeout_sweeper_terminates_quiet_ready_sessions():
now = datetime.now(timezone.utc)
quiet_row = _Row(
session_id="sess-quiet",
status=STATUS_READY,
vm_id="i-quiet",
created_at=now - timedelta(minutes=10),
last_heartbeat_at=now - timedelta(seconds=600),
)
rows = {quiet_row.session_id: quiet_row}
prisma = _build_prisma(rows, find_many_results=[quiet_row])
provider = _FakeProvider()
config = SweeperConfig(heartbeat_timeout_seconds=120)
swept = await heartbeat_timeout_sweeper(
provider=provider,
prisma_client=prisma,
config=config,
handle_builder=_build_handle,
)
assert swept == 1
assert provider.terminate_calls == ["i-quiet"]
assert quiet_row.status == STATUS_TERMINATED
assert quiet_row.failure_reason == "heartbeat_timeout"
@pytest.mark.asyncio
async def test_max_session_minutes_sweeper_terminates_long_running_sessions():
now = datetime.now(timezone.utc)
old_ready = _Row(
session_id="sess-old",
status=STATUS_READY,
vm_id="i-old",
created_at=now - timedelta(minutes=200),
last_heartbeat_at=now,
)
rows = {old_ready.session_id: old_ready}
prisma = _build_prisma(rows, find_many_results=[old_ready])
provider = _FakeProvider()
config = SweeperConfig(max_session_minutes=120)
swept = await max_session_minutes_sweeper(
provider=provider,
prisma_client=prisma,
config=config,
handle_builder=_build_handle,
)
assert swept == 1
assert provider.terminate_calls == ["i-old"]
assert old_ready.status == STATUS_TERMINATED
assert old_ready.failure_reason == "max_session_minutes"
@pytest.mark.asyncio
async def test_sweeper_skips_session_already_terminated():
"""Optimistic re-fetch lock: if a session moved to TERMINATED between
find_many and our work, the sweeper must NOT re-terminate."""
now = datetime.now(timezone.utc)
candidate = _Row(
session_id="sess-raced",
status=STATUS_TERMINATED, # already moved underneath us
vm_id="i-already-done",
created_at=now - timedelta(minutes=200),
last_heartbeat_at=None,
)
rows = {candidate.session_id: candidate}
prisma = _build_prisma(rows, find_many_results=[candidate])
provider = _FakeProvider()
config = SweeperConfig(max_session_minutes=120)
await max_session_minutes_sweeper(
provider=provider,
prisma_client=prisma,
config=config,
handle_builder=_build_handle,
)
# Provider must not be called.
assert provider.terminate_calls == []
@pytest.mark.asyncio
async def test_sweeper_with_no_vm_id_just_marks_row():
"""If a session never got a VM (e.g. RunInstances failed mid-provision),
the sweeper still moves the row to FAILED."""
def _maybe_handle(row: Any) -> Optional[VMHandle]:
# Build a handle with empty vm_id — sweeper should skip terminate.
return VMHandle(vm_id="", provider="fake", region="us-west-2")
now = datetime.now(timezone.utc)
row = _Row(
session_id="sess-novm",
status=STATUS_PROVISIONING,
vm_id="",
created_at=now - timedelta(seconds=600),
last_heartbeat_at=None,
)
rows = {row.session_id: row}
prisma = _build_prisma(rows, find_many_results=[row])
provider = _FakeProvider()
config = SweeperConfig(bootstrap_timeout_seconds=180)
await bootstrap_timeout_sweeper(
provider=provider,
prisma_client=prisma,
config=config,
handle_builder=_maybe_handle,
)
assert provider.terminate_calls == []
assert row.status == STATUS_FAILED
def test_sweeper_config_from_agent_settings():
cfg = SweeperConfig.from_agent_settings(
{
"sweep_interval_seconds": 17,
"ec2": {
"bootstrap_timeout_seconds": 200,
"heartbeat_timeout_seconds": 99,
"max_session_minutes": 60,
},
}
)
assert cfg.sweep_interval_seconds == 17
assert cfg.bootstrap_timeout_seconds == 200
assert cfg.heartbeat_timeout_seconds == 99
assert cfg.max_session_minutes == 60
def test_sweeper_config_defaults_when_missing():
cfg = SweeperConfig.from_agent_settings(None)
assert cfg.sweep_interval_seconds == 30
assert cfg.bootstrap_timeout_seconds == 180
assert cfg.heartbeat_timeout_seconds == 120
assert cfg.max_session_minutes == 120

View file

@ -0,0 +1,415 @@
"""
Mocked tests for `EC2Provider`.
These tests use a fake boto3 client (no AWS calls). The real-cloud tests
that use BYOC creds are in `test_ec2_provider_real.py` and are gated
behind `pytest --slow` so they don't run in the unit suite.
Coverage:
- Validation #5: spot → on-demand fallback when `InsufficientInstanceCapacity`
- Validation #8: terminate is idempotent (covers cascade-terminate cleanup)
- Validation #11: invalid AWS creds raise InvalidCredentialsError fast (no instance)
- Validation #13: AWS keys never appear in any log record or repr
"""
from __future__ import annotations
import logging
from typing import Any, Dict, List, Optional
from unittest.mock import patch
import pytest
from litellm.proxy.agent_session_endpoints.vm_providers import (
AwsCreds,
EC2Provider,
Ec2Config,
InvalidCredentialsError,
ProvisionContext,
ProvisionError,
Repo,
VMHandle,
VMState,
)
_FAKE_KEY = "AKIAFAKEEXAMPLEKEY12"
_FAKE_SECRET = "FAKEsecretFAKEsecretFAKEsecretFAKEsecre" # 40 chars
_FAKE_KEY_LEAK_CANARY = "AKIATESTLEAKCANARY00"
class _FakeClientError(Exception):
"""Stand-in for botocore.exceptions.ClientError."""
def __init__(self, code: str, message: str = "fake") -> None:
self.response = {"Error": {"Code": code, "Message": message}}
super().__init__(f"{code}: {message}")
class _FakeEc2Client:
"""In-memory boto3 EC2 client double."""
def __init__(
self,
*,
run_responses: Optional[List[Any]] = None,
terminate_should_raise: Optional[Exception] = None,
describe_response: Optional[Dict[str, Any]] = None,
) -> None:
# Each call to run_instances pops the next response. The response can be
# an instance dict OR an Exception subclass to raise.
self.run_responses: List[Any] = list(run_responses or [])
self.terminate_should_raise = terminate_should_raise
self.describe_response = describe_response
self.run_calls: List[Dict[str, Any]] = []
self.terminate_calls: List[Dict[str, Any]] = []
self.describe_calls: List[Dict[str, Any]] = []
def run_instances(self, **kwargs):
self.run_calls.append(kwargs)
if not self.run_responses:
raise AssertionError("run_instances called more times than expected")
nxt = self.run_responses.pop(0)
if isinstance(nxt, Exception):
raise nxt
return {"Instances": [nxt]}
def terminate_instances(self, **kwargs):
self.terminate_calls.append(kwargs)
if self.terminate_should_raise is not None:
raise self.terminate_should_raise
def describe_instances(self, **kwargs):
self.describe_calls.append(kwargs)
if isinstance(self.describe_response, Exception):
raise self.describe_response
if self.describe_response is None:
return {"Reservations": []}
return {"Reservations": [{"Instances": [self.describe_response]}]}
def _patch_provider_client(provider: EC2Provider, fake: _FakeEc2Client):
"""Patch `_build_ec2_client` on this provider instance."""
return patch.object(provider, "_build_ec2_client", lambda creds, region: fake)
def _ctx(team_id: str = "team-1") -> ProvisionContext:
return ProvisionContext(
session_id="sess-1",
team_id=team_id,
agent_id="agent-1",
repos=[Repo(url="https://example.com/x.git")],
env_vars={"FOO": "bar"},
aws_creds=AwsCreds(
access_key_id=_FAKE_KEY,
secret_access_key=_FAKE_SECRET,
region="us-west-2",
),
ec2_config=Ec2Config(
region="us-west-2",
subnet_id="subnet-1",
security_group_id="sg-1",
iam_instance_profile="litellm-ec2-poc",
instance_type="t3.large",
use_spot=True,
ami_id="ami-deadbeef",
),
daemon_jwt="fake.jwt.value",
daemon_base_url="https://proxy.example/",
mode="session",
)
# ---------- Validation #11: invalid creds fail-fast ----------
@pytest.mark.asyncio
async def test_provision_without_creds_raises_invalid_credentials_error():
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
ctx = _ctx()
ctx.aws_creds = None
with pytest.raises(InvalidCredentialsError):
await provider.provision(ctx)
@pytest.mark.asyncio
async def test_provision_invalid_creds_aws_response_raises_invalid_credentials_error():
"""AWS rejects creds with `InvalidClientTokenId` → 400 InvalidCredentialsError, no instance."""
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(
run_responses=[_FakeClientError("InvalidClientTokenId", "bad creds")]
)
with _patch_provider_client(provider, fake):
with pytest.raises(InvalidCredentialsError):
await provider.provision(_ctx())
# No retry / no on-demand fallback for cred errors — we must fail-fast.
assert len(fake.run_calls) == 1
@pytest.mark.asyncio
async def test_provision_signature_does_not_match_raises_invalid_credentials():
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(
run_responses=[_FakeClientError("SignatureDoesNotMatch", "wrong secret")]
)
with _patch_provider_client(provider, fake):
with pytest.raises(InvalidCredentialsError):
await provider.provision(_ctx())
assert len(fake.run_calls) == 1
# ---------- Validation #5: spot → on-demand fallback ----------
@pytest.mark.asyncio
async def test_provision_spot_fallback_to_on_demand():
"""Spot raises `InsufficientInstanceCapacity` → provider retries on-demand once."""
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(
run_responses=[
_FakeClientError("InsufficientInstanceCapacity", "no spot"),
{"InstanceId": "i-on-demand-1"},
]
)
with _patch_provider_client(provider, fake):
handle = await provider.provision(_ctx())
assert handle.vm_id == "i-on-demand-1"
assert handle.metadata["purchase_mode"] == "on-demand"
# First call had spot market options; second did not.
assert "InstanceMarketOptions" in fake.run_calls[0]
assert "InstanceMarketOptions" not in fake.run_calls[1]
@pytest.mark.asyncio
async def test_provision_spot_first_succeeds_no_fallback():
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(run_responses=[{"InstanceId": "i-spot-1"}])
with _patch_provider_client(provider, fake):
handle = await provider.provision(_ctx())
assert handle.metadata["purchase_mode"] == "spot"
assert len(fake.run_calls) == 1
@pytest.mark.asyncio
async def test_provision_no_spot_when_use_spot_false():
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(run_responses=[{"InstanceId": "i-1"}])
ctx = _ctx()
ctx.ec2_config.use_spot = False # type: ignore[union-attr]
with _patch_provider_client(provider, fake):
handle = await provider.provision(ctx)
assert handle.metadata["purchase_mode"] == "on-demand"
assert "InstanceMarketOptions" not in fake.run_calls[0]
# ---------- AMI required ----------
@pytest.mark.asyncio
async def test_provision_no_ami_raises_provision_error():
provider = EC2Provider({}) # no default_ami_id
ctx = _ctx()
ctx.ec2_config.ami_id = None # type: ignore[union-attr]
with pytest.raises(ProvisionError) as exc_info:
await provider.provision(ctx)
assert "AMI" in str(exc_info.value)
# ---------- Tags + IAM passthrough ----------
@pytest.mark.asyncio
async def test_provision_tags_instance_with_session_team_agent_ids():
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(run_responses=[{"InstanceId": "i-tagged"}])
with _patch_provider_client(provider, fake):
await provider.provision(_ctx(team_id="team-tag-test"))
tag_specs = fake.run_calls[0]["TagSpecifications"]
instance_tags = next(
s["Tags"] for s in tag_specs if s["ResourceType"] == "instance"
)
keys_to_values = {t["Key"]: t["Value"] for t in instance_tags}
assert keys_to_values["litellm-session-id"] == "sess-1"
assert keys_to_values["litellm-team-id"] == "team-tag-test"
assert keys_to_values["litellm-agent-id"] == "agent-1"
@pytest.mark.asyncio
async def test_provision_passes_iam_instance_profile():
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(run_responses=[{"InstanceId": "i-iam"}])
with _patch_provider_client(provider, fake):
await provider.provision(_ctx())
assert fake.run_calls[0]["IamInstanceProfile"] == {"Name": "litellm-ec2-poc"}
# ---------- Validation #8 piece: terminate idempotent ----------
@pytest.mark.asyncio
async def test_terminate_calls_aws():
provider = EC2Provider({})
fake = _FakeEc2Client()
handle = VMHandle(vm_id="i-1", provider="ec2", region="us-west-2")
with _patch_provider_client(provider, fake):
await provider.terminate(
handle,
aws_creds=AwsCreds(
access_key_id=_FAKE_KEY,
secret_access_key=_FAKE_SECRET,
region="us-west-2",
),
)
assert fake.terminate_calls == [{"InstanceIds": ["i-1"]}]
@pytest.mark.asyncio
async def test_terminate_already_gone_is_noop():
provider = EC2Provider({})
fake = _FakeEc2Client(
terminate_should_raise=_FakeClientError("InvalidInstanceID.NotFound")
)
handle = VMHandle(vm_id="i-already-gone", provider="ec2", region="us-west-2")
with _patch_provider_client(provider, fake):
# Must not raise.
await provider.terminate(
handle,
aws_creds=AwsCreds(
access_key_id=_FAKE_KEY,
secret_access_key=_FAKE_SECRET,
region="us-west-2",
),
)
@pytest.mark.asyncio
async def test_terminate_without_creds_raises():
provider = EC2Provider({})
handle = VMHandle(vm_id="i-1", provider="ec2", region="us-west-2")
with pytest.raises(InvalidCredentialsError):
await provider.terminate(handle)
# ---------- Status ----------
@pytest.mark.asyncio
async def test_status_running():
provider = EC2Provider({})
fake = _FakeEc2Client(
describe_response={
"InstanceId": "i-1",
"State": {"Name": "running"},
"PublicIpAddress": "1.2.3.4",
"PrivateIpAddress": "10.0.0.1",
}
)
handle = VMHandle(vm_id="i-1", provider="ec2", region="us-west-2")
with _patch_provider_client(provider, fake):
status = await provider.status(
handle,
aws_creds=AwsCreds(
access_key_id=_FAKE_KEY,
secret_access_key=_FAKE_SECRET,
region="us-west-2",
),
)
assert status.state == VMState.RUNNING
assert status.public_ip == "1.2.3.4"
@pytest.mark.asyncio
async def test_status_terminated_when_instance_not_found():
provider = EC2Provider({})
fake = _FakeEc2Client(
describe_response=_FakeClientError("InvalidInstanceID.NotFound")
)
handle = VMHandle(vm_id="i-gone", provider="ec2", region="us-west-2")
with _patch_provider_client(provider, fake):
status = await provider.status(
handle,
aws_creds=AwsCreds(
access_key_id=_FAKE_KEY,
secret_access_key=_FAKE_SECRET,
region="us-west-2",
),
)
assert status.state == VMState.TERMINATED
# ---------- Validation #13: creds never leak ----------
def test_aws_creds_repr_redacts():
creds = AwsCreds(
access_key_id=_FAKE_KEY_LEAK_CANARY,
secret_access_key="topsecret-secret-secret-secret-secret-1",
session_token="some-token",
region="us-west-2",
)
text = repr(creds)
# Neither the access key nor the secret may appear.
assert _FAKE_KEY_LEAK_CANARY not in text
assert "topsecret" not in text
assert "REDACTED" in text
assert str(creds) == repr(creds)
@pytest.mark.asyncio
async def test_aws_creds_never_logged_during_provision(caplog):
"""Even with DEBUG logging, the access key never lands in proxy logs."""
caplog.set_level(logging.DEBUG)
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(run_responses=[{"InstanceId": "i-leakcheck"}])
ctx = _ctx()
ctx.aws_creds = AwsCreds(
access_key_id=_FAKE_KEY_LEAK_CANARY,
secret_access_key="leak-canary-secret",
region="us-west-2",
)
with _patch_provider_client(provider, fake):
await provider.provision(ctx)
full_log = "\n".join(rec.getMessage() for rec in caplog.records)
assert _FAKE_KEY_LEAK_CANARY not in full_log
assert "leak-canary-secret" not in full_log
@pytest.mark.asyncio
async def test_aws_creds_never_in_exception_message():
"""A boto3 error message must not echo the access key."""
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(
run_responses=[_FakeClientError("ValidationError", "bad request")]
)
ctx = _ctx()
ctx.aws_creds = AwsCreds(
access_key_id=_FAKE_KEY_LEAK_CANARY,
secret_access_key="leak-canary-secret",
region="us-west-2",
)
with _patch_provider_client(provider, fake):
with pytest.raises(ProvisionError) as exc_info:
await provider.provision(ctx)
assert _FAKE_KEY_LEAK_CANARY not in str(exc_info.value)
# ---------- User-data shape ----------
@pytest.mark.asyncio
async def test_user_data_includes_session_id_and_jwt():
"""The provider builds user-data with the right env (the daemon reads them)."""
provider = EC2Provider({"default_ami_id": "ami-deadbeef"})
fake = _FakeEc2Client(run_responses=[{"InstanceId": "i-user-data"}])
with _patch_provider_client(provider, fake):
await provider.provision(_ctx())
user_data = fake.run_calls[0]["UserData"]
assert "LITELLM_SESSION_ID=sess-1" in user_data
assert "LITELLM_TEAM_ID=team-1" in user_data
assert "LITELLM_AGENT_ID=agent-1" in user_data
assert "LITELLM_DAEMON_JWT=fake.jwt.value" in user_data
assert "LITELLM_AGENT_MODE=session" in user_data

View file

@ -0,0 +1,278 @@
"""
Real-cloud tests for `EC2Provider`.
Skipped by default. Enable with `pytest -m slow` and the BYOC env vars set:
LITELLM_AGENT_AWS_ACCESS_KEY_ID
LITELLM_AGENT_AWS_SECRET_ACCESS_KEY
LITELLM_AGENT_AWS_REGION (default us-west-2)
LITELLM_TEST_SUBNET_ID
LITELLM_TEST_SECURITY_GROUP_ID
LITELLM_TEST_IAM_INSTANCE_PROFILE
LITELLM_TEST_AMI_ID
These are the resources B0 captured for the BYOC PoC account
(see LIT-2888 deliverables comment).
Each test is wrapped in a try/finally that calls TerminateInstances. A 60-min
process-wide watchdog is also installed so a hung test cannot leak an
instance overnight.
Cost per run: ~$0.03 per session (one t3.large for under a minute).
"""
from __future__ import annotations
import os
import threading
import time
from typing import List, Optional
import pytest
from litellm.proxy.agent_session_endpoints.vm_providers import (
AwsCreds,
EC2Provider,
Ec2Config,
ProvisionContext,
Repo,
VMHandle,
VMState,
)
# Test markers — `slow` is the standard real-cloud gate.
pytestmark = [pytest.mark.slow]
def _have_creds() -> bool:
return bool(
os.getenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID")
and os.getenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY")
and os.getenv("LITELLM_TEST_AMI_ID")
)
_skip_no_creds = pytest.mark.skipif(
not _have_creds(),
reason="real-cloud test requires LITELLM_AGENT_AWS_* + LITELLM_TEST_* env vars",
)
def _ec2_config_from_env() -> Ec2Config:
return Ec2Config(
region=os.getenv("LITELLM_AGENT_AWS_REGION", "us-west-2"),
subnet_id=os.getenv("LITELLM_TEST_SUBNET_ID"),
security_group_id=os.getenv("LITELLM_TEST_SECURITY_GROUP_ID"),
iam_instance_profile=os.getenv("LITELLM_TEST_IAM_INSTANCE_PROFILE"),
instance_type="t3.large",
use_spot=False, # real tests prefer determinism over savings
ami_id=os.getenv("LITELLM_TEST_AMI_ID"),
)
def _aws_creds_from_env() -> AwsCreds:
return AwsCreds(
access_key_id=os.environ["LITELLM_AGENT_AWS_ACCESS_KEY_ID"],
secret_access_key=os.environ["LITELLM_AGENT_AWS_SECRET_ACCESS_KEY"],
session_token=os.getenv("LITELLM_AGENT_AWS_SESSION_TOKEN"),
region=os.getenv("LITELLM_AGENT_AWS_REGION", "us-west-2"),
)
def _install_watchdog(
provider: EC2Provider, handles: List[VMHandle]
) -> threading.Timer:
"""60-min hard watchdog: terminate everything we tracked, no matter what."""
def _kill_all() -> None:
import asyncio # noqa: PLC0415
async def _go():
for h in handles:
try:
await provider.terminate(h, aws_creds=_aws_creds_from_env())
except Exception:
pass
try:
asyncio.run(_go())
except Exception:
pass
timer = threading.Timer(60 * 60, _kill_all)
timer.daemon = True
timer.start()
return timer
@_skip_no_creds
@pytest.mark.asyncio
async def test_real_boot():
"""Validation #3: instance reaches `running` within 90s, then we terminate."""
handles: List[VMHandle] = []
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
watchdog = _install_watchdog(provider, handles)
ctx = ProvisionContext(
session_id=f"sess-test-real-{int(time.time())}",
team_id="team-test-real",
agent_id="agent-test-real",
repos=[Repo(url="https://github.com/octocat/Hello-World")],
env_vars={"FOO": "bar"},
aws_creds=_aws_creds_from_env(),
ec2_config=_ec2_config_from_env(),
daemon_jwt="not-a-real-jwt-test-only",
daemon_base_url="https://example.invalid/",
mode="session",
)
handle: Optional[VMHandle] = None
try:
handle = await provider.provision(ctx)
handles.append(handle)
assert handle.vm_id.startswith("i-")
# Poll for `running` (max 90s).
deadline = time.time() + 90
last: Optional[VMState] = None
while time.time() < deadline:
status = await provider.status(handle, aws_creds=ctx.aws_creds)
last = status.state
if status.state == VMState.RUNNING:
break
time.sleep(2)
assert last == VMState.RUNNING, f"never reached running, last={last}"
finally:
if handle is not None:
await provider.terminate(handle, aws_creds=ctx.aws_creds)
watchdog.cancel()
@_skip_no_creds
@pytest.mark.asyncio
async def test_byoc_invalid_creds_fail_fast():
"""Validation #11: bad creds → InvalidCredentialsError within ~5s, no instance launched."""
from litellm.proxy.agent_session_endpoints.vm_providers import (
InvalidCredentialsError,
)
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
bad_creds = AwsCreds(
access_key_id="AKIAINVALIDFAKEKEY00",
secret_access_key="invalid-secret-do-not-leak",
region=os.getenv("LITELLM_AGENT_AWS_REGION", "us-west-2"),
)
ctx = ProvisionContext(
session_id=f"sess-bad-{int(time.time())}",
team_id="team-bad",
aws_creds=bad_creds,
ec2_config=_ec2_config_from_env(),
)
t0 = time.time()
with pytest.raises(InvalidCredentialsError):
await provider.provision(ctx)
elapsed = time.time() - t0
assert elapsed < 10, f"fail-fast took too long: {elapsed:.1f}s"
@_skip_no_creds
@pytest.mark.asyncio
async def test_terminate_idempotent_real():
"""Validation #8 piece: terminate the same instance twice — second is a no-op."""
handles: List[VMHandle] = []
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
watchdog = _install_watchdog(provider, handles)
ctx = ProvisionContext(
session_id=f"sess-term-{int(time.time())}",
team_id="team-term",
aws_creds=_aws_creds_from_env(),
ec2_config=_ec2_config_from_env(),
daemon_jwt="x",
daemon_base_url="https://example.invalid/",
)
handle: Optional[VMHandle] = None
try:
handle = await provider.provision(ctx)
handles.append(handle)
await provider.terminate(handle, aws_creds=ctx.aws_creds)
# Second call must not raise.
await provider.terminate(handle, aws_creds=ctx.aws_creds)
finally:
if handle is not None:
try:
await provider.terminate(handle, aws_creds=ctx.aws_creds)
except Exception:
pass
watchdog.cancel()
@_skip_no_creds
@pytest.mark.asyncio
async def test_cold_boot_p50_under_30s():
"""Validation #14 — cold-boot P50 (RunInstances → daemon-ready) must be ≤ 30s.
The Packer-baked AMI eliminates the apt-get update step that dominated the B0
spike's 31.8s user-data phase. With pre-installed node/python/git/gh/uv/bun,
cold-boot should land in the 18-25s range. P50 ≤ 30s gives ~5s headroom for
AWS variability.
Daemon-ready is approximated by polling `provider.status()` for `VMState.RUNNING`
— once Epic C (LIT-2879) ships the real daemon callback, swap this for the
session-row `status == 'ready'` check.
"""
import statistics # noqa: PLC0415
handles: List[VMHandle] = []
provider = EC2Provider({"default_ami_id": os.environ["LITELLM_TEST_AMI_ID"]})
watchdog = _install_watchdog(provider, handles)
latencies: List[float] = []
try:
for i in range(5):
ctx = ProvisionContext(
session_id=f"sess-coldboot-{i}-{int(time.time())}",
team_id="team-coldboot",
agent_id="agent-coldboot",
repos=[Repo(url="https://github.com/octocat/Hello-World")],
env_vars={"FOO": "bar"},
aws_creds=_aws_creds_from_env(),
ec2_config=_ec2_config_from_env(),
daemon_jwt="not-a-real-jwt-test-only",
daemon_base_url="https://example.invalid/",
mode="session",
)
handle: Optional[VMHandle] = None
try:
t0 = time.perf_counter()
handle = await provider.provision(ctx)
handles.append(handle)
# Poll for daemon-ready (approximated via VMState.RUNNING).
# 60s timeout gives 2x headroom over the 30s P50 gate.
deadline = time.time() + 60
last: Optional[VMState] = None
while time.time() < deadline:
status = await provider.status(handle, aws_creds=ctx.aws_creds)
last = status.state
if status.state == VMState.RUNNING:
break
time.sleep(1)
assert (
last == VMState.RUNNING
), f"iteration {i}: never reached running, last={last}"
latencies.append(time.perf_counter() - t0)
finally:
if handle is not None:
await provider.terminate(handle, aws_creds=ctx.aws_creds)
p50 = statistics.median(latencies)
# P95 with n=5 → use the max sample as an approximation.
p95 = sorted(latencies)[int(len(latencies) * 0.95)]
print(f"Cold-boot latencies (s): {latencies}")
print(f"P50 = {p50:.1f}s, P95 = {p95:.1f}s")
assert p50 <= 30.0, f"Cold-boot P50 {p50:.1f}s exceeds 30s gate"
finally:
watchdog.cancel()

View file

@ -0,0 +1,85 @@
"""
Mocked tests for the `AgentVMProvider` factory.
Covers Validation #1 from LIT-2878: factory returns the right impl per
config value and rejects unknown values. Validation #6 (provider swap is
config-only) is also exercised here — switching `vm_provider` from `noop` to
`ec2` requires no code path change beyond the factory.
"""
from __future__ import annotations
import pytest
from litellm.proxy.agent_session_endpoints.vm_providers import (
EC2Provider,
NoopProvider,
SUPPORTED_PROVIDERS,
build_vm_provider,
)
def test_factory_default_is_noop():
"""No `agent_settings` block at all → noop provider."""
provider = build_vm_provider(None)
assert isinstance(provider, NoopProvider)
assert provider.name == "noop"
def test_factory_explicit_noop():
provider = build_vm_provider({"vm_provider": "noop"})
assert isinstance(provider, NoopProvider)
def test_factory_ec2_with_settings():
settings = {
"vm_provider": "ec2",
"ec2": {
"default_region": "us-west-2",
"default_ami_id": "ami-deadbeef",
"default_instance_type": "t3.large",
"use_spot": True,
},
}
provider = build_vm_provider(settings)
assert isinstance(provider, EC2Provider)
assert provider.default_region == "us-west-2"
assert provider.default_ami_id == "ami-deadbeef"
assert provider.default_instance_type == "t3.large"
assert provider.default_use_spot is True
def test_factory_ec2_uses_defaults_when_block_missing():
"""`vm_provider: ec2` with no `ec2:` block still works using built-in defaults."""
provider = build_vm_provider({"vm_provider": "ec2"})
assert isinstance(provider, EC2Provider)
assert provider.default_region == "us-west-2" # baked-in default
assert provider.default_instance_type == "t3.large"
def test_factory_unknown_raises():
with pytest.raises(ValueError) as exc_info:
build_vm_provider({"vm_provider": "firecracker-someday"})
msg = str(exc_info.value)
assert "Unknown vm_provider" in msg
# Make sure the error message lists the supported providers so callers
# can fix their config without grep.
for name in SUPPORTED_PROVIDERS:
assert name in msg
def test_factory_supported_providers_set():
assert set(SUPPORTED_PROVIDERS) >= {"noop", "ec2"}
def test_provider_swap_is_config_only():
"""Switching `vm_provider` is the only change required (Validation #6)."""
cfg_noop = {"vm_provider": "noop"}
cfg_ec2 = {**cfg_noop, "vm_provider": "ec2"}
p1 = build_vm_provider(cfg_noop)
p2 = build_vm_provider(cfg_ec2)
# Both implement AgentVMProvider — same call site can use either.
assert hasattr(p1, "provision") and hasattr(p2, "provision")
assert hasattr(p1, "terminate") and hasattr(p2, "terminate")
assert hasattr(p1, "status") and hasattr(p2, "status")
assert p1.name != p2.name

View file

@ -0,0 +1,46 @@
"""Mocked tests for the `NoopProvider`."""
from __future__ import annotations
import pytest
from litellm.proxy.agent_session_endpoints.vm_providers import (
NoopProvider,
ProvisionContext,
VMState,
)
@pytest.mark.asyncio
async def test_noop_provision_returns_handle():
provider = NoopProvider()
ctx = ProvisionContext(session_id="sess-1", team_id="team-1")
handle = await provider.provision(ctx)
assert handle.vm_id.startswith("noop-")
assert handle.provider == "noop"
assert handle.metadata["session_id"] == "sess-1"
assert handle.metadata["team_id"] == "team-1"
@pytest.mark.asyncio
async def test_noop_status_lifecycle():
provider = NoopProvider()
ctx = ProvisionContext(session_id="sess-1", team_id="team-1")
handle = await provider.provision(ctx)
status = await provider.status(handle)
assert status.state == VMState.RUNNING
assert status.public_ip == "127.0.0.1"
await provider.terminate(handle)
status_after = await provider.status(handle)
assert status_after.state == VMState.TERMINATED
@pytest.mark.asyncio
async def test_noop_terminate_idempotent():
provider = NoopProvider()
handle = await provider.provision(ProvisionContext(session_id="s", team_id="t"))
await provider.terminate(handle)
# Second call must not raise.
await provider.terminate(handle)

View file

@ -0,0 +1,267 @@
"""
Tests for `team_config.get_team_vm_config` — BYOC creds resolution against
Epic G's `LiteLLM_AgentVMConfig` schema (per-field encrypted columns).
Covers:
- Validation #11: missing creds → InvalidCredentialsError fast (no instance launched)
- Validation #12: cross-team isolation — two teams resolve to two distinct creds
- Per-field encryption round-trip (encrypt_aws_field + decrypt_aws_field)
- env-var fallback for local dev
- Falls back to env vars when the AgentVMConfig table doesn't exist yet
Reconciliation note: this file was originally written against Epic B's
single-blob `aws_creds_enc` shape. Rewritten on the integration branch to
match Epic G's per-field encrypted columns (`aws_access_key_id_enc`,
`aws_secret_access_key_enc`, `aws_region`, ...).
"""
from __future__ import annotations
from typing import Any, Dict, Optional
from unittest.mock import AsyncMock, MagicMock
import pytest
def _set_master_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "test-salt-do-not-use-in-prod-1234567890")
def test_encrypt_decrypt_aws_field_round_trip(monkeypatch):
_set_master_key(monkeypatch)
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
decrypt_aws_field,
encrypt_aws_field,
)
plain = "AKIAROUNDTRIPTESTKEY"
blob = encrypt_aws_field(plain, "aws_access_key_id")
assert blob is not None
# The encrypted blob must NOT contain the plaintext.
assert plain not in blob
decrypted = decrypt_aws_field(blob, "aws_access_key_id")
assert decrypted == plain
def test_encrypt_aws_field_returns_none_for_empty():
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
encrypt_aws_field,
)
assert encrypt_aws_field(None, "aws_access_key_id") is None
assert encrypt_aws_field("", "aws_access_key_id") is None
def test_decrypt_aws_field_returns_none_for_empty():
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
decrypt_aws_field,
)
assert decrypt_aws_field(None, "aws_access_key_id") is None
assert decrypt_aws_field("", "aws_access_key_id") is None
@pytest.mark.asyncio
async def test_no_db_no_env_raises_invalid_credentials(monkeypatch):
"""Validation #11: missing creds must fail-fast, not launch any instance."""
_set_master_key(monkeypatch)
monkeypatch.delenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID", raising=False)
monkeypatch.delenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY", raising=False)
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
InvalidCredentialsError,
)
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
get_team_vm_config,
)
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_agentvmconfig = MagicMock()
prisma.db.litellm_agentvmconfig.find_unique = AsyncMock(return_value=None)
with pytest.raises(InvalidCredentialsError):
await get_team_vm_config("team-no-creds", prisma_client=prisma)
@pytest.mark.asyncio
async def test_env_fallback_when_no_db_row(monkeypatch):
_set_master_key(monkeypatch)
monkeypatch.setenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID", "AKIAFROMENV")
monkeypatch.setenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY", "secret-from-env")
monkeypatch.setenv("LITELLM_AGENT_AWS_REGION", "us-east-1")
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
get_team_vm_config,
)
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_agentvmconfig = MagicMock()
prisma.db.litellm_agentvmconfig.find_unique = AsyncMock(return_value=None)
cfg = await get_team_vm_config("team-env", prisma_client=prisma)
assert cfg.aws_creds.access_key_id == "AKIAFROMENV"
assert cfg.aws_creds.secret_access_key == "secret-from-env"
assert cfg.aws_creds.region == "us-east-1"
@pytest.mark.asyncio
async def test_falls_back_to_env_when_prisma_table_missing(monkeypatch):
"""If the LiteLLM_AgentVMConfig table doesn't exist yet, the resolver falls
back to env vars instead of crashing."""
_set_master_key(monkeypatch)
monkeypatch.setenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID", "AKIAFROMENVFALLBACK")
monkeypatch.setenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY", "secret-fallback")
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
get_team_vm_config,
)
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_agentvmconfig = MagicMock()
prisma.db.litellm_agentvmconfig.find_unique = AsyncMock(
side_effect=Exception("relation does not exist")
)
cfg = await get_team_vm_config("team-no-table", prisma_client=prisma)
assert cfg.aws_creds.access_key_id == "AKIAFROMENVFALLBACK"
def _row_with_encrypted_creds(
access_key_id_plain: str,
secret_access_key_plain: str,
region: str = "us-west-2",
) -> Any:
"""Build a fake `LiteLLM_AgentVMConfig` Prisma row with G's column shape."""
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
encrypt_aws_field,
)
row = MagicMock()
row.provider = "ec2"
row.aws_auth_method = "access_keys"
row.aws_access_key_id_enc = encrypt_aws_field(
access_key_id_plain, "aws_access_key_id"
)
row.aws_secret_access_key_enc = encrypt_aws_field(
secret_access_key_plain, "aws_secret_access_key"
)
row.aws_role_arn_enc = None
row.aws_region = region
row.subnet_id = "subnet-1"
row.security_group_id = "sg-1"
row.iam_instance_profile = "litellm-ec2-poc"
row.instance_type = "t3.large"
row.use_spot = True
row.ami_id = "ami-deadbeef"
return row
@pytest.mark.asyncio
async def test_byoc_cross_team_isolation(monkeypatch):
"""Validation #12: two teams resolve to two distinct creds objects."""
_set_master_key(monkeypatch)
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
get_team_vm_config,
)
rows: Dict[str, Any] = {
"team-a": _row_with_encrypted_creds(
"AKIATEAMA000000000", "teama-secret", region="us-west-2"
),
"team-b": _row_with_encrypted_creds(
"AKIATEAMB000000000", "teamb-secret", region="us-east-1"
),
}
async def find_unique(where: Dict[str, Any]) -> Optional[Any]:
return rows.get(where["team_id"])
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_agentvmconfig = MagicMock()
prisma.db.litellm_agentvmconfig.find_unique = AsyncMock(side_effect=find_unique)
cfg_a = await get_team_vm_config("team-a", prisma_client=prisma)
cfg_b = await get_team_vm_config("team-b", prisma_client=prisma)
assert cfg_a.aws_creds.access_key_id == "AKIATEAMA000000000"
assert cfg_a.aws_creds.region == "us-west-2"
assert cfg_b.aws_creds.access_key_id == "AKIATEAMB000000000"
assert cfg_b.aws_creds.region == "us-east-1"
assert cfg_a.aws_creds.access_key_id != cfg_b.aws_creds.access_key_id
@pytest.mark.asyncio
async def test_db_row_with_no_creds_raises(monkeypatch):
_set_master_key(monkeypatch)
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
InvalidCredentialsError,
)
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
get_team_vm_config,
)
row = MagicMock()
row.provider = "ec2"
row.aws_auth_method = "access_keys"
row.aws_access_key_id_enc = None
row.aws_secret_access_key_enc = None
row.aws_role_arn_enc = None
row.aws_region = "us-west-2"
row.subnet_id = None
row.security_group_id = None
row.iam_instance_profile = None
row.instance_type = None
row.use_spot = True
row.ami_id = None
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_agentvmconfig = MagicMock()
prisma.db.litellm_agentvmconfig.find_unique = AsyncMock(return_value=row)
with pytest.raises(InvalidCredentialsError):
await get_team_vm_config("team-empty", prisma_client=prisma)
@pytest.mark.asyncio
async def test_iam_role_method_returns_none_creds_falls_back_env(monkeypatch):
"""When `aws_auth_method == 'iam_role'`, the resolver does NOT pull keys
from the DB — it leaves cred resolution to boto3's default chain. With no
env vars set this raises InvalidCredentialsError."""
_set_master_key(monkeypatch)
monkeypatch.delenv("LITELLM_AGENT_AWS_ACCESS_KEY_ID", raising=False)
monkeypatch.delenv("LITELLM_AGENT_AWS_SECRET_ACCESS_KEY", raising=False)
from litellm.proxy.agent_session_endpoints.vm_providers.base import (
InvalidCredentialsError,
)
from litellm.proxy.agent_session_endpoints.vm_providers.team_config import (
get_team_vm_config,
)
row = MagicMock()
row.provider = "ec2"
row.aws_auth_method = "iam_role"
row.aws_access_key_id_enc = "anything"
row.aws_secret_access_key_enc = "anything"
row.aws_role_arn_enc = "anything"
row.aws_region = "us-west-2"
row.subnet_id = None
row.security_group_id = None
row.iam_instance_profile = None
row.instance_type = None
row.use_spot = True
row.ami_id = None
prisma = MagicMock()
prisma.db = MagicMock()
prisma.db.litellm_agentvmconfig = MagicMock()
prisma.db.litellm_agentvmconfig.find_unique = AsyncMock(return_value=row)
with pytest.raises(InvalidCredentialsError):
await get_team_vm_config("team-iam-role", prisma_client=prisma)