mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
commit
6dccf40e0e
26 changed files with 3434 additions and 92 deletions
122
infra/ami/README.md
Normal file
122
infra/ami/README.md
Normal 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
|
||||
```
|
||||
198
infra/ami/files/daemon-stub.py
Normal file
198
infra/ami/files/daemon-stub.py
Normal 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()
|
||||
22
infra/ami/files/litellm-agent-runtime.service
Normal file
22
infra/ami/files/litellm-agent-runtime.service
Normal 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
|
||||
143
infra/ami/litellm-agent-runtime.pkr.hcl
Normal file
143
infra/ami/litellm-agent-runtime.pkr.hcl
Normal 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",
|
||||
]
|
||||
}
|
||||
}
|
||||
98
infra/ami/scripts/install-runtime.sh
Executable file
98
infra/ami/scripts/install-runtime.sh
Executable 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
372
litellm/proxy/agent_session_endpoints/sweepers.py
Normal file
372
litellm/proxy/agent_session_endpoints/sweepers.py
Normal 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
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
488
litellm/proxy/agent_session_endpoints/vm_providers/ec2.py
Normal file
488
litellm/proxy/agent_session_endpoints/vm_providers/ec2.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue