diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260506220000_add_cloud_agent_settings_tables/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260506220000_add_cloud_agent_settings_tables/migration.sql new file mode 100644 index 00000000000..390ed8ea4c4 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260506220000_add_cloud_agent_settings_tables/migration.sql @@ -0,0 +1,85 @@ +-- Cloud Agents settings (Epic G / LIT-2891). +-- +-- Adds four tables that back the new Settings -> Cloud Agents UI: +-- * LiteLLM_AgentVMConfig -- per-team provider, AWS BYOC, warm pool, network access +-- * LiteLLM_AgentSecret -- per-team encrypted secrets, write-only on read +-- * LiteLLM_AgentWorker -- self-hosted worker registrations +-- * LiteLLM_AgentWorkerPairingToken -- single-use 15-min pairing tokens +-- +-- Encryption uses the existing nacl/SecretBox path in +-- litellm.proxy.common_utils.encrypt_decrypt_utils — no new KMS infra here. +-- Values are stored base64-encoded; raw secrets are NEVER returned from any +-- GET endpoint (LIT-2891 validation #2). + +CREATE TABLE IF NOT EXISTS "LiteLLM_AgentVMConfig" ( + "team_id" TEXT PRIMARY KEY, + "provider" TEXT NOT NULL DEFAULT 'disabled', + "aws_auth_method" TEXT, + "aws_access_key_id_enc" TEXT, + "aws_secret_access_key_enc" TEXT, + "aws_role_arn_enc" TEXT, + "aws_region" TEXT, + "ami_id" TEXT, + "instance_type" TEXT, + "subnet_id" TEXT, + "security_group_id" TEXT, + "iam_instance_profile" TEXT, + "use_spot" BOOLEAN NOT NULL DEFAULT TRUE, + "max_session_minutes" INTEGER NOT NULL DEFAULT 120, + "warm_pool_enabled" BOOLEAN NOT NULL DEFAULT FALSE, + "warm_pool_size" INTEGER NOT NULL DEFAULT 0, + "max_idle_minutes" INTEGER NOT NULL DEFAULT 30, + "hydrate_transport" TEXT NOT NULL DEFAULT 'auto', + "network_access" JSONB NOT NULL DEFAULT '{"mode":"allow_all","allowlist":[]}'::jsonb, + "self_hosted_enabled" BOOLEAN NOT NULL DEFAULT FALSE, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE TABLE IF NOT EXISTS "LiteLLM_AgentSecret" ( + "id" TEXT PRIMARY KEY, + "team_id" TEXT NOT NULL, + "name" TEXT NOT NULL, + "value_enc" TEXT NOT NULL, + "scope" JSONB NOT NULL DEFAULT '"all"'::jsonb, + "type" TEXT NOT NULL DEFAULT 'env', + "file_path" TEXT, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT +); +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_AgentSecret_team_id_name_key" + ON "LiteLLM_AgentSecret" ("team_id", "name"); +CREATE INDEX IF NOT EXISTS "LiteLLM_AgentSecret_team_id_idx" + ON "LiteLLM_AgentSecret" ("team_id"); + +CREATE TABLE IF NOT EXISTS "LiteLLM_AgentWorker" ( + "id" TEXT PRIMARY KEY, + "team_id" TEXT NOT NULL, + "hostname" TEXT NOT NULL, + "status" TEXT NOT NULL DEFAULT 'offline', + "last_seen_at" TIMESTAMP(3), + "cpu_pct" DOUBLE PRECISION, + "mem_gb" DOUBLE PRECISION, + "active_sessions" INTEGER NOT NULL DEFAULT 0, + "worker_jwt_hash" TEXT NOT NULL, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS "LiteLLM_AgentWorker_team_id_status_idx" + ON "LiteLLM_AgentWorker" ("team_id", "status"); +-- Lookup by JWT digest is on the hot path — every worker long-poll +-- heartbeat (LIT-2890 / B2) calls `find_worker_by_jwt`, which filters +-- by this column. Without an index that's a full scan per heartbeat. +CREATE INDEX IF NOT EXISTS "LiteLLM_AgentWorker_worker_jwt_hash_idx" + ON "LiteLLM_AgentWorker" ("worker_jwt_hash"); + +CREATE TABLE IF NOT EXISTS "LiteLLM_AgentWorkerPairingToken" ( + "token_hash" TEXT PRIMARY KEY, + "team_id" TEXT NOT NULL, + "created_by" TEXT NOT NULL, + "expires_at" TIMESTAMP(3) NOT NULL, + "used_at" TIMESTAMP(3), + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP +); +CREATE INDEX IF NOT EXISTS "LiteLLM_AgentWorkerPairingToken_team_id_idx" + ON "LiteLLM_AgentWorkerPairingToken" ("team_id"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 6a2e09e1616..d52083d1a90 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1375,7 +1375,6 @@ model LiteLLM_WorkflowMessage { @@index([run_id]) } - // =========================================================================== // Agent Sessions / Runs (Cursor SDK on LiteLLM) // @@ -1468,3 +1467,88 @@ model LiteLLM_AgentRunEvent { @@unique([run_id, seq]) @@index([run_id, seq]) } + +// =========================================================================== +// Cloud Agent settings (LIT-2891) — per-team VM provider config + secrets + +// self-hosted worker registry +// =========================================================================== + +// Per-team Cloud Agent VM provider configuration. Holds AWS BYOC creds +// (encrypted via the same nacl/SecretBox path as virtual keys), provisioning +// defaults, warm-pool settings, and the network egress allowlist that gets +// pushed to the daemon at hydrate time. +model LiteLLM_AgentVMConfig { + team_id String @id + provider String @default("disabled") // "ec2" | "self_hosted" | "disabled" + aws_auth_method String? // "access_keys" | "iam_role" | "instance_metadata" + aws_access_key_id_enc String? // encrypted + aws_secret_access_key_enc String? // encrypted + aws_role_arn_enc String? // encrypted (cross-account role mode) + aws_region String? + ami_id String? + instance_type String? + subnet_id String? + security_group_id String? + iam_instance_profile String? + use_spot Boolean @default(true) + max_session_minutes Int @default(120) + warm_pool_enabled Boolean @default(false) + warm_pool_size Int @default(0) + max_idle_minutes Int @default(30) + hydrate_transport String @default("auto") // "auto" | "ssm" | "long_poll" + network_access Json @default("{\"mode\":\"allow_all\",\"allowlist\":[]}") + self_hosted_enabled Boolean @default(false) + created_at DateTime @default(now()) + updated_at DateTime @default(now()) @updatedAt +} + +// Per-team encrypted secrets injected into agent VMs at session start. Value +// is ALWAYS write-only: GET endpoints must never return value_enc decrypted. +// Scope is "all" or a list of repo full_name strings; the proxy joins on +// session.repos at hydrate time. +model LiteLLM_AgentSecret { + id String @id @default(uuid()) + team_id String + name String + value_enc String // base64-encoded, nacl.SecretBox encrypted, write-only + scope Json @default("\"all\"") // "all" | string[] + type String @default("env") // "env" | "file" + file_path String? + created_at DateTime @default(now()) + updated_at DateTime @default(now()) @updatedAt + created_by String? + + @@unique([team_id, name]) + @@index([team_id]) +} + +// Self-hosted worker registrations. Each worker holds a long-lived JWT and +// long-polls for hydrate. status is best-effort heartbeat tracking. +model LiteLLM_AgentWorker { + id String @id @default(uuid()) + team_id String + hostname String + status String @default("offline") // "online" | "offline" + last_seen_at DateTime? + cpu_pct Float? + mem_gb Float? + active_sessions Int @default(0) + worker_jwt_hash String // sha256 of issued JWT — never store raw JWT + created_at DateTime @default(now()) + + @@index([team_id, status]) + @@index([worker_jwt_hash]) +} + +// Single-use 15-minute pairing tokens that workers exchange for a long-lived +// worker JWT during the install flow. +model LiteLLM_AgentWorkerPairingToken { + token_hash String @id // sha256 of the raw token (raw token never persisted) + team_id String + created_by String + expires_at DateTime + used_at DateTime? + created_at DateTime @default(now()) + + @@index([team_id]) +} diff --git a/litellm/constants.py b/litellm/constants.py index 1e96f8eac9a..0aba12f9f47 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -144,6 +144,30 @@ MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS = int( os.getenv("MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS", "60") ) +# Cloud Agents (Epic G / LIT-2891) — settings UI defaults. +# Pair tokens are single-use and short-lived: the operator copies the install +# command, runs it on the worker box, and the worker exchanges the token for a +# long-lived JWT. 15 min is enough for a human-in-the-loop install but tight +# enough to limit the blast radius of a leaked token. +CLOUD_AGENT_PAIR_TOKEN_TTL_MINUTES = int( + os.getenv("CLOUD_AGENT_PAIR_TOKEN_TTL_MINUTES", "15") +) +# Hourly costs ($/hr) used by the Warm Pool screen's idle-cost widget. v1 +# hardcodes a small list — a follow-up ticket will pull live AWS pricing via +# the Pricing API. Values are us-east-1 list prices, on-demand. +CLOUD_AGENT_INSTANCE_HOURLY_COST_USD = { + "t3.medium": 0.0416, + "t3.large": 0.0832, + "t3.xlarge": 0.1664, + "t3.2xlarge": 0.3328, + "m5.large": 0.096, + "m5.xlarge": 0.192, + "m5.2xlarge": 0.384, +} +# 730 = average hours per month (24 * 365.25 / 12). Mirrors the AWS billing +# convention so our estimate matches what the customer sees on their bill. +CLOUD_AGENT_HOURS_PER_MONTH = 730 + # MCP timeout defaults (seconds). Override via env vars for slow/custom MCP servers. MCP_CLIENT_TIMEOUT = float(os.getenv("LITELLM_MCP_CLIENT_TIMEOUT", "60.0")) MCP_TOOL_LISTING_TIMEOUT = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIMEOUT", "30.0")) diff --git a/litellm/proxy/agent_settings_endpoints/__init__.py b/litellm/proxy/agent_settings_endpoints/__init__.py new file mode 100644 index 00000000000..aadb47fb9ac --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/__init__.py @@ -0,0 +1,9 @@ +""" +Cloud Agents settings endpoints (Epic G / LIT-2891). + +Backs the Settings -> Cloud Agents UI with four resources: +* AgentVMConfig (provider, AWS BYOC, warm pool, network access) +* AgentSecret (per-team encrypted secrets, write-only on read) +* AgentWorker (self-hosted worker registrations) +* AgentWorkerPairingToken (single-use 15-min tokens for worker install) +""" diff --git a/litellm/proxy/agent_settings_endpoints/encryption.py b/litellm/proxy/agent_settings_endpoints/encryption.py new file mode 100644 index 00000000000..ee547af0436 --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/encryption.py @@ -0,0 +1,44 @@ +""" +Thin encryption wrapper for the Cloud Agents settings endpoints. + +We do NOT roll a new KMS path here — we reuse the existing virtual-key +nacl/SecretBox helpers (`encrypt_value_helper` / `decrypt_value_helper`) so +operators only need to manage one `LITELLM_SALT_KEY`. The wrapper exists only +to give the agent endpoints a single import surface and to centralize the +"None passes through" semantics so callers don't sprinkle `if value is None` +guards everywhere. + +Used for: +* AWS BYOC creds on `LiteLLM_AgentVMConfig` (access key, secret key, role ARN) +* Per-team secret values on `LiteLLM_AgentSecret` +""" + +from typing import Optional + +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) + + +def encrypt_optional(value: Optional[str]) -> Optional[str]: + """Encrypt a value, passing None through. + + Returns base64(urlsafe)-encoded ciphertext, or None if the input was None + or empty. Empty strings are normalized to None so `aws_role_arn=""` from + the wire round-trips cleanly to NULL in the DB. + """ + if value is None or value == "": + return None + return encrypt_value_helper(value) + + +def decrypt_optional(value: Optional[str], *, key: str) -> Optional[str]: + """Decrypt a previously-encrypted value, passing None through. + + `key` is a debug-label only (used by the underlying helper to surface + which field failed to decrypt) — it is NOT a signing key. + """ + if value is None or value == "": + return None + return decrypt_value_helper(value=value, key=key) diff --git a/litellm/proxy/agent_settings_endpoints/pair_tokens.py b/litellm/proxy/agent_settings_endpoints/pair_tokens.py new file mode 100644 index 00000000000..5366de0d4be --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/pair_tokens.py @@ -0,0 +1,104 @@ +""" +Self-hosted worker pairing tokens (LIT-2891 / Screen 3). + +Flow: +1. UI calls `POST /v2/agent-workers/pair-token` → server generates a 32-byte + urlsafe token, stores ONLY its sha256 in + `LiteLLM_AgentWorkerPairingToken`, returns the raw token to the caller + exactly once. TTL = 15 min. +2. The user runs the install one-liner with `--token `. The worker calls + `POST /v2/agent-workers/register`, which calls `consume_pair_token` — that + re-hashes the token, atomically marks the row `used_at=now()`, and returns + the team_id. Single-use is enforced by checking `used_at IS NULL`. +3. The worker is then issued a long-lived JWT (also stored as a sha256 hash + on the `LiteLLM_AgentWorker` row). + +Raw tokens and JWTs are NEVER persisted — only their sha256 digests. This +matches the existing virtual-key hashed-token pattern. +""" + +import hashlib +import secrets +import shlex +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Optional + +from litellm.constants import CLOUD_AGENT_PAIR_TOKEN_TTL_MINUTES + + +@dataclass +class IssuedPairToken: + raw_token: str # returned to the caller ONCE + token_hash: str # what we persist + expires_at: datetime # UTC + + +def hash_pair_token(raw_token: str) -> str: + """sha256 hex digest. Same algorithm used at issue time and consume time.""" + return hashlib.sha256(raw_token.encode("utf-8")).hexdigest() + + +def issue_pair_token( + *, + ttl_minutes: int = CLOUD_AGENT_PAIR_TOKEN_TTL_MINUTES, +) -> IssuedPairToken: + """Generate a fresh pair token. Caller persists `token_hash` + `expires_at`. + + The raw token is 32 bytes of urandom, urlsafe-base64-encoded — that's ~43 + chars of entropy, which is plenty for a single-use 15-min token. + """ + raw_token = secrets.token_urlsafe(32) + token_hash = hash_pair_token(raw_token) + expires_at = datetime.now(timezone.utc) + timedelta(minutes=ttl_minutes) + return IssuedPairToken( + raw_token=raw_token, token_hash=token_hash, expires_at=expires_at + ) + + +def hash_worker_jwt(raw_jwt: str) -> str: + """sha256 of the worker JWT — what we persist on `LiteLLM_AgentWorker`.""" + return hashlib.sha256(raw_jwt.encode("utf-8")).hexdigest() + + +def is_expired( + expires_at: Optional[datetime], *, now: Optional[datetime] = None +) -> bool: + """True iff `expires_at` is in the past. Naive datetimes are treated as UTC.""" + if expires_at is None: + return True + current = now or datetime.now(timezone.utc) + if expires_at.tzinfo is None: + # Match the implicit-UTC convention of the existing proxy DB writes. + expires_at = expires_at.replace(tzinfo=timezone.utc) + if current.tzinfo is None: + current = current.replace(tzinfo=timezone.utc) + return expires_at <= current + + +def build_install_command( + *, + proxy_url: str, + raw_token: str, + install_script_url: str = "https://litellm.ai/install-worker", +) -> str: + """Render the one-liner the UI shows in the Add Machine modal. + + Kept as a helper (not f-string at the callsite) so tests can lock the + exact format and so we can swap the install host without touching + endpoint code. + + All operator-controlled values (`proxy_url`, `raw_token`, and the + install script URL) are run through `shlex.quote` before interpolation + so that spaces, quotes, or other shell metacharacters in any of them + can't break out of the install command. The proxy URL is otherwise + not validated here — the caller is responsible for verifying the + host (see `worker_endpoints._resolve_proxy_url`). + """ + quoted_url = shlex.quote(install_script_url) + quoted_proxy = shlex.quote(proxy_url) + quoted_token = shlex.quote(raw_token) + return ( + f"curl -fsS {quoted_url} | sh -s -- " + f"--proxy {quoted_proxy} --token {quoted_token}" + ) diff --git a/litellm/proxy/agent_settings_endpoints/pool_status_endpoints.py b/litellm/proxy/agent_settings_endpoints/pool_status_endpoints.py new file mode 100644 index 00000000000..b1ea0a1564c --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/pool_status_endpoints.py @@ -0,0 +1,89 @@ +""" +`GET /v2/agent-vm-pool/status` — live VM pool status. + +This endpoint returns the warm-pool / hydrating / attached counts the dashboard +shows under "View live pool status" on the Warm Pool screen. + +The real implementation is owned by Epic B2 (LIT-2890), which actually tracks +warm VMs, hydrate state, and session attachments. To unblock the UI ahead of +B2, we ship a deterministic stub here: + +* When `LITELLM_AGENT_POOL_STATUS_MOCK=1` (the default while B2 is open), the + endpoint returns `{warm: 0, hydrating: 0, attached: 0}`. +* When B2 lands, that flag flips to `0` (or the env var goes away entirely) + and the implementation reads from B2's session/VM tracking tables. + +Keeping the route here means the UI ships against a stable contract today — +no front-end changes required when B2 swaps the body. +""" + +import os +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Request +from pydantic import BaseModel + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +class PoolStatusResponse(BaseModel): + """Match the shape B2 will return so the UI doesn't have to change.""" + + warm: int + hydrating: int + attached: int + # Optional for the v1 stub; B2 will populate. + avg_create_latency_ms: Optional[int] = None + oldest_warm_seconds: Optional[int] = None + + +def _pool_status_mock_enabled() -> bool: + """B2 owns the real status path; default to mock until then.""" + return os.getenv("LITELLM_AGENT_POOL_STATUS_MOCK", "1") == "1" + + +def _resolve_team_id(user_api_key_dict: UserAPIKeyAuth) -> str: + team_id = user_api_key_dict.team_id or (user_api_key_dict.metadata or {}).get( + "team_id" + ) + if not team_id: + raise HTTPException( + status_code=400, + detail="Cloud Agent pool status is scoped to a team.", + ) + return team_id + + +@router.get( + "/v2/agent-vm-pool/status", + dependencies=[Depends(user_api_key_auth)], + response_model=PoolStatusResponse, + tags=["cloud agents"], +) +async def get_agent_vm_pool_status( + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> PoolStatusResponse: + """Return live VM pool counts. Stubbed until B2 (LIT-2890) lands.""" + # Validate team scoping even in the mock so we surface 400s consistently + # with the rest of the cloud-agents endpoints — saves a UI bug later. + _resolve_team_id(user_api_key_dict) + + if _pool_status_mock_enabled(): + return PoolStatusResponse(warm=0, hydrating=0, attached=0) + + # Real path — B2 will fill this in. Until then, refusing loudly is better + # than silently returning zeros and pretending everything is fine. + raise HTTPException( + status_code=501, + detail={ + "error": ( + "Live pool status is owned by LIT-2890 (Epic B2). Set " + "LITELLM_AGENT_POOL_STATUS_MOCK=1 to use the stub." + ), + "code": "not_implemented", + }, + ) diff --git a/litellm/proxy/agent_settings_endpoints/scope_filter.py b/litellm/proxy/agent_settings_endpoints/scope_filter.py new file mode 100644 index 00000000000..08022f9e499 --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/scope_filter.py @@ -0,0 +1,127 @@ +""" +Scope-filtering helper for cloud-agent secrets (LIT-2891 validation #3). + +A secret's `scope` field is either: +* the literal string `"all"` — applies to every session for the team, OR +* a list of repo full-names (e.g. `["BerriAI/litellm", "BerriAI/litellm-docs"]`) + — applies only to sessions whose repo set intersects this list. + +This module is the single source of truth for "should this secret be present in +the hydrate payload for this session?" Both the GET endpoints (for UI display) +and the session-create hydrate path (B2) call into here so the access-control +logic can never drift between UI and the wire. +""" + +from typing import Any, Iterable, List, Optional, Tuple, Union +from urllib.parse import urlparse + +ScopeValue = Union[str, List[str]] + + +def _normalize_repo(repo: Any) -> Optional[str]: + """Reduce any repo reference to its `owner/name` form (lowercase). + + Accepts: + * a plain string (`"BerriAI/litellm"` or `"github.com/BerriAI/litellm"`) + * a `https://github.com/BerriAI/litellm.git` URL + * a dict with `full_name` or `url` + + Returns None for anything we can't parse — caller treats that as "no + match" rather than crashing the hydrate path. + """ + if repo is None: + return None + + if isinstance(repo, dict): + if isinstance(repo.get("full_name"), str): + return _normalize_repo(repo["full_name"]) + if isinstance(repo.get("url"), str): + return _normalize_repo(repo["url"]) + return None + + if not isinstance(repo, str): + return None + + raw = repo.strip() + if not raw: + return None + + # URL form + if "://" in raw: + parsed = urlparse(raw) + path = (parsed.path or "").strip("/") + else: + path = raw + + # Strip leading host fragments (`github.com/owner/name` → `owner/name`) + while path.startswith(("github.com/", "gitlab.com/", "bitbucket.org/")): + path = path.split("/", 1)[1] + + # Strip trailing `.git` + if path.endswith(".git"): + path = path[:-4] + + parts = [p for p in path.split("/") if p] + if len(parts) < 2: + return None + owner, name = parts[0], parts[1] + return f"{owner.lower()}/{name.lower()}" + + +def normalize_repos(repos: Iterable[Any]) -> List[str]: + """Public helper — used by hydrate to canonicalize a session's repo list.""" + out: List[str] = [] + seen = set() + for r in repos or []: + normalized = _normalize_repo(r) + if normalized and normalized not in seen: + seen.add(normalized) + out.append(normalized) + return out + + +def secret_in_scope(scope: ScopeValue, session_repos: Iterable[Any]) -> bool: + """Return True if a secret with this `scope` applies to a session. + + `scope == "all"` always matches. A list scope matches when any of its + entries (normalized) is in the session's normalized repo set. Empty + scope-list matches nothing — that's the safe default if a UI bug ever + writes `scope=[]`. + """ + if scope == "all": + return True + if not isinstance(scope, list): + # Defensive: any unexpected shape is treated as no-match. + return False + if not scope: + return False + + normalized_session = set(normalize_repos(session_repos)) + if not normalized_session: + return False + + for entry in scope: + normalized_entry = _normalize_repo(entry) + if normalized_entry and normalized_entry in normalized_session: + return True + return False + + +def partition_secrets_for_session( + secrets: Iterable[Tuple[str, ScopeValue]], + session_repos: Iterable[Any], +) -> Tuple[List[str], List[str]]: + """Split (name, scope) pairs into (in_scope_names, out_of_scope_names). + + Used by the hydrate-builder to log which secrets it skipped without + leaking values. Order is preserved. + """ + repos = list(session_repos or []) + in_scope: List[str] = [] + out_of_scope: List[str] = [] + for name, scope in secrets: + if secret_in_scope(scope, repos): + in_scope.append(name) + else: + out_of_scope.append(name) + return in_scope, out_of_scope diff --git a/litellm/proxy/agent_settings_endpoints/secrets_endpoints.py b/litellm/proxy/agent_settings_endpoints/secrets_endpoints.py new file mode 100644 index 00000000000..17542d800a7 --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/secrets_endpoints.py @@ -0,0 +1,309 @@ +""" +`/v2/agent-secrets` endpoints (LIT-2891 / Screen 5). + +Per-team encrypted secrets for cloud agent VMs. Two security invariants +worth calling out, since the rest of this file is built around them: + +* **Write-only values.** `value` is accepted on POST/PUT and stored encrypted, + but it is NEVER decrypted onto a GET response. The response schema + (`AgentSecretResponse`) has no `value` field, so even an accidental + `model_dump()` of the row can't leak the plaintext. +* **Per-team isolation.** Every query filters on `team_id` resolved from the + caller's API key. Cross-team reads are not possible at this layer — the + composite unique key `(team_id, name)` makes that explicit. + +The session-create / hydrate path (B2) calls the lower-level +`partition_secrets_for_session` helper from `scope_filter.py` to figure out +which secrets to push into a given session. This module only handles the UI +CRUD surface. +""" + +from typing import Any, Dict, List, Optional + +from fastapi import APIRouter, Depends, HTTPException, Request + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy.agent_settings_endpoints.encryption import encrypt_optional +from litellm.proxy.agent_settings_endpoints.types import ( + AgentSecretCreateRequest, + AgentSecretListResponse, + AgentSecretResponse, + AgentSecretUpdateRequest, +) +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +def _resolve_team_id(user_api_key_dict: UserAPIKeyAuth) -> str: + """Pick the team to scope this request to. Raise 400 if missing. + + Same contract as the VM config endpoints — secrets are per-team and we + refuse to silently fall back to a "default" scope. + """ + team_id = user_api_key_dict.team_id or (user_api_key_dict.metadata or {}).get( + "team_id" + ) + if not team_id: + raise HTTPException( + status_code=400, + detail=( + "Cloud Agent secrets are scoped to a team. Pick a team from the " + "header switcher and try again." + ), + ) + return team_id + + +def _row_to_response(row: Dict[str, Any]) -> AgentSecretResponse: + """Map a Prisma row to the public response shape. + + Note: `value_enc` is intentionally NOT read here. The response model + has no value field, so even if a future caller passed `**row` we + couldn't accidentally surface the ciphertext. + """ + scope = row.get("scope") + if scope is None: + scope = "all" + if isinstance(scope, str) and scope not in ("all",): + # SQLite stores Json columns as raw strings — tolerate both shapes. + import json as _json + + try: + parsed = _json.loads(scope) + except Exception: + parsed = "all" + scope = parsed if parsed in ("all",) or isinstance(parsed, list) else "all" + + created_at = row.get("created_at") + updated_at = row.get("updated_at") + return AgentSecretResponse( + name=row["name"], + scope=scope, + type=row.get("type") or "env", + file_path=row.get("file_path"), + created_at=str(created_at) if created_at is not None else "", + updated_at=str(updated_at) if updated_at is not None else "", + ) + + +def _validate_secret_payload( + *, + type_: Optional[str], + file_path: Optional[str], + is_create: bool, +) -> None: + """Reject `type=file` without a `file_path`. Mirrors the UI form validation + so the rule lives in exactly one place server-side too.""" + if type_ == "file" and not file_path and is_create: + raise HTTPException( + status_code=400, + detail="`file_path` is required when `type` is `file`.", + ) + + +def _build_create_payload( + *, + team_id: str, + body: AgentSecretCreateRequest, + created_by: Optional[str], +) -> Dict[str, Any]: + """Build the dict for prisma.create — encrypts `value`, never stores raw.""" + encrypted = encrypt_optional(body.value) + if encrypted is None: + # Pydantic enforces min_length=1 already, so this is just a belt-and- + # suspenders guard against future signature drift. + raise HTTPException(status_code=400, detail="Secret `value` cannot be empty.") + return { + "team_id": team_id, + "name": body.name, + "value_enc": encrypted, + "scope": body.scope, + "type": body.type, + "file_path": body.file_path, + "created_by": created_by, + } + + +def _build_update_payload(body: AgentSecretUpdateRequest) -> Dict[str, Any]: + """Build the dict for prisma.update. Omits unset fields so partial PATCH- + like updates from the UI don't clobber existing scope/type.""" + payload: Dict[str, Any] = {} + if body.value is not None: + encrypted = encrypt_optional(body.value) + if encrypted is not None: + payload["value_enc"] = encrypted + if body.scope is not None: + payload["scope"] = body.scope + if body.type is not None: + payload["type"] = body.type + if body.file_path is not None: + payload["file_path"] = body.file_path + return payload + + +@router.get( + "/v2/agent-secrets", + dependencies=[Depends(user_api_key_auth)], + response_model=AgentSecretListResponse, + tags=["cloud agents"], +) +async def list_agent_secrets( + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> AgentSecretListResponse: + """List secrets for the team. Returns metadata only — no values.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + rows = await prisma_client.db.litellm_agentsecret.find_many( + where={"team_id": team_id}, + order={"name": "asc"}, + ) + secrets: List[AgentSecretResponse] = [ + _row_to_response(dict(r) if not isinstance(r, dict) else r) for r in rows + ] + return AgentSecretListResponse(secrets=secrets) + + +@router.post( + "/v2/agent-secrets", + dependencies=[Depends(user_api_key_auth)], + response_model=AgentSecretResponse, + tags=["cloud agents"], +) +async def create_agent_secret( + request: Request, + body: AgentSecretCreateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> AgentSecretResponse: + """Create a new secret. Returns metadata only.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + _validate_secret_payload(type_=body.type, file_path=body.file_path, is_create=True) + + # Conflict check — composite unique (team_id, name). + existing = await prisma_client.db.litellm_agentsecret.find_unique( + where={"team_id_name": {"team_id": team_id, "name": body.name}} + ) + if existing is not None: + raise HTTPException( + status_code=409, + detail=f"Secret `{body.name}` already exists for this team. Use PUT to update it.", + ) + + payload = _build_create_payload( + team_id=team_id, + body=body, + created_by=user_api_key_dict.user_id, + ) + + try: + created = await prisma_client.db.litellm_agentsecret.create(data=payload) + except Exception as exc: + verbose_proxy_logger.exception( + "Failed to create agent secret name=%s team=%s: %s", + body.name, + team_id, + exc, + ) + raise HTTPException(status_code=500, detail="Failed to create secret.") + + row = dict(created) if not isinstance(created, dict) else created + return _row_to_response(row) + + +@router.put( + "/v2/agent-secrets/{name}", + dependencies=[Depends(user_api_key_auth)], + response_model=AgentSecretResponse, + tags=["cloud agents"], +) +async def update_agent_secret( + request: Request, + name: str, + body: AgentSecretUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> AgentSecretResponse: + """Update an existing secret by name.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + _validate_secret_payload(type_=body.type, file_path=body.file_path, is_create=False) + + payload = _build_update_payload(body) + if not payload: + raise HTTPException( + status_code=400, + detail="At least one field must be provided to update.", + ) + + existing = await prisma_client.db.litellm_agentsecret.find_unique( + where={"team_id_name": {"team_id": team_id, "name": name}} + ) + if existing is None: + raise HTTPException( + status_code=404, detail=f"Secret `{name}` not found for this team." + ) + + updated = await prisma_client.db.litellm_agentsecret.update( + where={"team_id_name": {"team_id": team_id, "name": name}}, + data=payload, + ) + row = dict(updated) if not isinstance(updated, dict) else updated + return _row_to_response(row) + + +@router.delete( + "/v2/agent-secrets/{name}", + dependencies=[Depends(user_api_key_auth)], + tags=["cloud agents"], +) +async def delete_agent_secret( + request: Request, + name: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> Dict[str, Any]: + """Delete a secret by name. Idempotent — returns 404 if not present.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + + existing = await prisma_client.db.litellm_agentsecret.find_unique( + where={"team_id_name": {"team_id": team_id, "name": name}} + ) + if existing is None: + raise HTTPException( + status_code=404, detail=f"Secret `{name}` not found for this team." + ) + + await prisma_client.db.litellm_agentsecret.delete( + where={"team_id_name": {"team_id": team_id, "name": name}} + ) + return {"deleted": True, "name": name} diff --git a/litellm/proxy/agent_settings_endpoints/types.py b/litellm/proxy/agent_settings_endpoints/types.py new file mode 100644 index 00000000000..ab777e33aa3 --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/types.py @@ -0,0 +1,145 @@ +""" +Pydantic request/response models for the Cloud Agents settings endpoints. + +The split matters: +* `AgentVMConfigUpdateRequest` accepts plaintext AWS keys; the endpoint encrypts + them before they touch the DB. +* `AgentVMConfigResponse` redacts those keys to a "***" sentinel and returns + only the metadata. We never decrypt secrets onto a GET response (LIT-2891 + validation #2). +* `AgentSecretCreateRequest` accepts a plaintext value; `AgentSecretResponse` + has no `value` field at all so the type system itself blocks accidental + exposure. +""" + +from typing import List, Literal, Optional, Union + +from pydantic import BaseModel, Field + +REDACTED_VALUE = "***" + +NetworkAccessMode = Literal["allow_all", "allowlist_only"] +AgentSecretType = Literal["env", "file"] +AgentSecretScope = Union[Literal["all"], List[str]] + + +class NetworkAccessConfig(BaseModel): + mode: NetworkAccessMode = "allow_all" + allowlist: List[str] = Field(default_factory=list) + + +class AgentVMConfigUpdateRequest(BaseModel): + provider: Optional[Literal["ec2", "self_hosted", "disabled"]] = None + aws_auth_method: Optional[ + Literal["access_keys", "iam_role", "instance_metadata"] + ] = None + aws_access_key_id: Optional[str] = None # plaintext on the wire + aws_secret_access_key: Optional[str] = None # plaintext on the wire + aws_role_arn: Optional[str] = None # plaintext on the wire + aws_region: Optional[str] = None + ami_id: Optional[str] = None + instance_type: Optional[str] = None + subnet_id: Optional[str] = None + security_group_id: Optional[str] = None + iam_instance_profile: Optional[str] = None + use_spot: Optional[bool] = None + max_session_minutes: Optional[int] = Field(default=None, ge=1, le=1440) + warm_pool_enabled: Optional[bool] = None + warm_pool_size: Optional[int] = Field(default=None, ge=0, le=100) + max_idle_minutes: Optional[int] = Field(default=None, ge=0, le=1440) + hydrate_transport: Optional[Literal["auto", "ssm", "long_poll"]] = None + network_access: Optional[NetworkAccessConfig] = None + self_hosted_enabled: Optional[bool] = None + + +class AgentVMConfigResponse(BaseModel): + """GET response — AWS creds are redacted to REDACTED_VALUE if set, None if unset.""" + + team_id: str + provider: str + aws_auth_method: Optional[str] + aws_access_key_id: Optional[str] # "***" if set, else None + aws_secret_access_key: Optional[str] # "***" if set, else None + aws_role_arn: Optional[str] # "***" if set, else None + aws_region: Optional[str] + ami_id: Optional[str] + instance_type: Optional[str] + subnet_id: Optional[str] + security_group_id: Optional[str] + iam_instance_profile: Optional[str] + use_spot: bool + max_session_minutes: int + warm_pool_enabled: bool + warm_pool_size: int + max_idle_minutes: int + hydrate_transport: str + network_access: NetworkAccessConfig + self_hosted_enabled: bool + + +class TestConnectionResponse(BaseModel): + ok: bool + account_id: Optional[str] = None + arn: Optional[str] = None + region: Optional[str] = None + error: Optional[str] = None + + +class AgentSecretCreateRequest(BaseModel): + name: str = Field(min_length=1, max_length=128, pattern=r"^[A-Za-z_][A-Za-z0-9_]*$") + value: str = Field(min_length=1) + scope: AgentSecretScope = "all" + type: AgentSecretType = "env" + file_path: Optional[str] = None + + +class AgentSecretUpdateRequest(BaseModel): + value: Optional[str] = Field(default=None, min_length=1) + scope: Optional[AgentSecretScope] = None + type: Optional[AgentSecretType] = None + file_path: Optional[str] = None + + +class AgentSecretResponse(BaseModel): + """Note: NO `value` field — this type itself enforces validation #2.""" + + name: str + scope: AgentSecretScope + type: AgentSecretType + file_path: Optional[str] + created_at: str + updated_at: str + + +class AgentSecretListResponse(BaseModel): + secrets: List[AgentSecretResponse] + + +class PairTokenResponse(BaseModel): + token: str # raw token returned ONCE — never persisted in DB + expires_at: str + install_command: str + + +class AgentWorkerResponse(BaseModel): + id: str + hostname: str + status: str + last_seen_at: Optional[str] + cpu_pct: Optional[float] + mem_gb: Optional[float] + active_sessions: int + + +class AgentWorkerListResponse(BaseModel): + workers: List[AgentWorkerResponse] + + +class AgentWorkerRegisterRequest(BaseModel): + pair_token: str = Field(min_length=1) + hostname: str = Field(min_length=1, max_length=255) + + +class AgentWorkerRegisterResponse(BaseModel): + worker_id: str + worker_jwt: str # long-lived JWT, returned ONCE diff --git a/litellm/proxy/agent_settings_endpoints/vm_config_endpoints.py b/litellm/proxy/agent_settings_endpoints/vm_config_endpoints.py new file mode 100644 index 00000000000..765fd0d5985 --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/vm_config_endpoints.py @@ -0,0 +1,408 @@ +""" +`/v2/agent-vm-config` endpoints (LIT-2891 / Screen 1, 2, 4). + +Backs the Settings -> Cloud Agents -> Provider, Warm Pool, and Network Access +screens. The whole config is one row per team in `LiteLLM_AgentVMConfig`. The +GET response NEVER includes raw AWS creds — fields are returned as +`REDACTED_VALUE` if set, `None` if unset. + +Test Connection currently MOCKS `sts:GetCallerIdentity` until B0 (LIT-2888) +closes — that ticket installs the real boto3 path. The mock is gated on the +`LITELLM_CLOUD_AGENT_MOCK_AWS` env var so we can flip to real once B0 lands +without changing this code. +""" + +import os +from typing import Any, Dict, Optional + +from fastapi import APIRouter, Depends, HTTPException, Request + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy.agent_settings_endpoints.encryption import ( + decrypt_optional, + encrypt_optional, +) +from litellm.proxy.agent_settings_endpoints.types import ( + REDACTED_VALUE, + AgentVMConfigResponse, + AgentVMConfigUpdateRequest, + NetworkAccessConfig, + TestConnectionResponse, +) +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +_DEFAULT_NETWORK_ACCESS: Dict[str, Any] = {"mode": "allow_all", "allowlist": []} + + +def _resolve_team_id(user_api_key_dict: UserAPIKeyAuth) -> str: + """Pick the team to scope this request to. Raise 400 if we can't. + + The settings UI is always opened in the context of a team (the user picks + one from the team-switcher in the header). The dashboard sends the team + ID via the auth context. If neither team_id nor team_alias is set we + refuse — there is no sensible "default" config to return. + """ + team_id = user_api_key_dict.team_id or (user_api_key_dict.metadata or {}).get( + "team_id" + ) + if not team_id: + raise HTTPException( + status_code=400, + detail=( + "Cloud Agent settings are scoped to a team. Pick a team from the " + "header switcher and try again." + ), + ) + return team_id + + +def _redact(value: Optional[str]) -> Optional[str]: + """Return REDACTED_VALUE if a non-empty ciphertext exists, else None.""" + return REDACTED_VALUE if value else None + + +def _row_to_response(row: Dict[str, Any]) -> AgentVMConfigResponse: + """Map a Prisma row dict to the public response shape, redacting secrets.""" + network_access_raw = row.get("network_access") or _DEFAULT_NETWORK_ACCESS + if isinstance(network_access_raw, str): + # Prisma sometimes returns Json fields as already-parsed dicts but + # under SQLite (test) it can hand back the raw string — be tolerant. + import json as _json + + try: + network_access_raw = _json.loads(network_access_raw) + except Exception: + network_access_raw = _DEFAULT_NETWORK_ACCESS + network_access = NetworkAccessConfig(**network_access_raw) + + return AgentVMConfigResponse( + team_id=row["team_id"], + provider=row.get("provider") or "disabled", + aws_auth_method=row.get("aws_auth_method"), + aws_access_key_id=_redact(row.get("aws_access_key_id_enc")), + aws_secret_access_key=_redact(row.get("aws_secret_access_key_enc")), + aws_role_arn=_redact(row.get("aws_role_arn_enc")), + aws_region=row.get("aws_region"), + ami_id=row.get("ami_id"), + instance_type=row.get("instance_type"), + subnet_id=row.get("subnet_id"), + security_group_id=row.get("security_group_id"), + iam_instance_profile=row.get("iam_instance_profile"), + use_spot=bool(row.get("use_spot", True)), + max_session_minutes=int(row.get("max_session_minutes") or 120), + warm_pool_enabled=bool(row.get("warm_pool_enabled", False)), + warm_pool_size=int(row.get("warm_pool_size") or 0), + max_idle_minutes=int(row.get("max_idle_minutes") or 30), + hydrate_transport=row.get("hydrate_transport") or "auto", + network_access=network_access, + self_hosted_enabled=bool(row.get("self_hosted_enabled", False)), + ) + + +def _empty_row(team_id: str) -> Dict[str, Any]: + """Synthesize a default row for teams that have never saved settings.""" + return { + "team_id": team_id, + "provider": "disabled", + "aws_auth_method": None, + "aws_access_key_id_enc": None, + "aws_secret_access_key_enc": None, + "aws_role_arn_enc": None, + "aws_region": None, + "ami_id": None, + "instance_type": None, + "subnet_id": None, + "security_group_id": None, + "iam_instance_profile": None, + "use_spot": True, + "max_session_minutes": 120, + "warm_pool_enabled": False, + "warm_pool_size": 0, + "max_idle_minutes": 30, + "hydrate_transport": "auto", + "network_access": _DEFAULT_NETWORK_ACCESS, + "self_hosted_enabled": False, + } + + +def _build_update_payload( + body: AgentVMConfigUpdateRequest, +) -> Dict[str, Any]: + """Translate the request model into a dict suitable for Prisma upsert. + + AWS fields are encrypted in-place; sentinel `REDACTED_VALUE` from the UI + means "leave the existing value alone" and is filtered out before write. + Network access goes in as raw JSON (Prisma handles the encoding). + """ + payload: Dict[str, Any] = {} + + plain_fields = ( + "provider", + "aws_auth_method", + "aws_region", + "ami_id", + "instance_type", + "subnet_id", + "security_group_id", + "iam_instance_profile", + "use_spot", + "max_session_minutes", + "warm_pool_enabled", + "warm_pool_size", + "max_idle_minutes", + "hydrate_transport", + "self_hosted_enabled", + ) + for field in plain_fields: + value = getattr(body, field, None) + # Use `is not None` (not truthiness): `False`, `0`, and `""` are + # all valid values that callers may legitimately want to write + # (e.g. disabling warm pool with `warm_pool_enabled=False`, + # zeroing `warm_pool_size`). A future contributor adding a field + # here should keep this guard so those updates don't get dropped. + if value is not None: + payload[field] = value + + encrypted_fields = ( + ("aws_access_key_id", "aws_access_key_id_enc"), + ("aws_secret_access_key", "aws_secret_access_key_enc"), + ("aws_role_arn", "aws_role_arn_enc"), + ) + for plain_field, db_field in encrypted_fields: + value = getattr(body, plain_field, None) + if value is None: + # Field omitted entirely — leave existing value. + continue + if value == REDACTED_VALUE: + # UI round-trip: don't overwrite with the redacted sentinel. + continue + payload[db_field] = encrypt_optional(value) + + if body.network_access is not None: + payload["network_access"] = body.network_access.model_dump() + + return payload + + +async def _get_or_create_row(prisma_client: Any, team_id: str) -> Dict[str, Any]: + """Fetch the row, returning a default if none exists.""" + row = await prisma_client.db.litellm_agentvmconfig.find_unique( + where={"team_id": team_id} + ) + if row is None: + return _empty_row(team_id) + return dict(row) if not isinstance(row, dict) else row + + +@router.get( + "/v2/agent-vm-config", + dependencies=[Depends(user_api_key_auth)], + response_model=AgentVMConfigResponse, + tags=["cloud agents"], +) +async def get_agent_vm_config( + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> AgentVMConfigResponse: + """Return the team's VM config with AWS creds redacted.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + row = await _get_or_create_row(prisma_client, team_id) + return _row_to_response(row) + + +@router.put( + "/v2/agent-vm-config", + dependencies=[Depends(user_api_key_auth)], + response_model=AgentVMConfigResponse, + tags=["cloud agents"], +) +async def update_agent_vm_config( + request: Request, + body: AgentVMConfigUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> AgentVMConfigResponse: + """Upsert the team's VM config. AWS creds are encrypted before write.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + payload = _build_update_payload(body) + + create_payload = {**_empty_row(team_id), **payload} + update_payload = payload + + await prisma_client.db.litellm_agentvmconfig.upsert( + where={"team_id": team_id}, + data={ + "create": create_payload, + "update": update_payload, + }, + ) + + refreshed = await _get_or_create_row(prisma_client, team_id) + return _row_to_response(refreshed) + + +async def _resolve_aws_creds( + prisma_client: Any, team_id: str +) -> Dict[str, Optional[str]]: + """Decrypt the team's stored AWS creds. Used by Test Connection + by B2. + + Exposed as a helper so the hydrate path (B2) can reuse it. Returns a dict + with `access_key_id`, `secret_access_key`, `role_arn`, `region` — any of + which may be None. + """ + row = await prisma_client.db.litellm_agentvmconfig.find_unique( + where={"team_id": team_id} + ) + if row is None: + return { + "access_key_id": None, + "secret_access_key": None, + "role_arn": None, + "region": None, + } + row_dict = dict(row) if not isinstance(row, dict) else row + return { + "access_key_id": decrypt_optional( + row_dict.get("aws_access_key_id_enc"), + key="aws_access_key_id", + ), + "secret_access_key": decrypt_optional( + row_dict.get("aws_secret_access_key_enc"), + key="aws_secret_access_key", + ), + "role_arn": decrypt_optional( + row_dict.get("aws_role_arn_enc"), + key="aws_role_arn", + ), + "region": row_dict.get("aws_region"), + } + + +def _mock_caller_identity(creds: Dict[str, Optional[str]]) -> TestConnectionResponse: + """Stand-in for boto3 sts:GetCallerIdentity until B0 ships the real call. + + Returns a deterministic ok/err response shaped like the real STS reply so + the UI can be wired without depending on B0. The mock fails if no access + key is configured — that mirrors the real failure mode and makes the + "no creds yet" UX testable end-to-end. + """ + access_key = creds.get("access_key_id") + if not access_key: + return TestConnectionResponse( + ok=False, + error=( + "No AWS credentials configured for this team. Add an Access Key " + "or IAM Role under Provider Settings and try again." + ), + ) + # The mock account ID is derived from the access key fingerprint so each + # team gets a stable-but-distinct value during development. + suffix = "".join(c for c in access_key if c.isdigit())[-12:].rjust(12, "0") + region = creds.get("region") or "us-west-2" + return TestConnectionResponse( + ok=True, + account_id=suffix, + arn=f"arn:aws:iam::{suffix}:user/litellm-cloud-agents", + region=region, + ) + + +def _aws_mock_enabled() -> bool: + """Whether `test-connection` should return a synthetic mock response + instead of calling the real `sts:GetCallerIdentity`. + + Defaults to OFF — a fresh production proxy must always validate AWS + credentials against STS, never silently return success for invalid + creds. Set `LITELLM_CLOUD_AGENT_MOCK_AWS=1` explicitly to opt into the + mock path during local development / tests. + """ + return os.getenv("LITELLM_CLOUD_AGENT_MOCK_AWS", "0") == "1" + + +@router.post( + "/v2/agent-vm-config/test-connection", + dependencies=[Depends(user_api_key_auth)], + response_model=TestConnectionResponse, + tags=["cloud agents"], +) +async def test_aws_connection( + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> TestConnectionResponse: + """Validate the team's stored AWS creds against `sts:GetCallerIdentity`. + + Phase 1 (current): mocked behind `LITELLM_CLOUD_AGENT_MOCK_AWS=1`. + Phase 2 (post-B0): real boto3 client; same response shape. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + creds = await _resolve_aws_creds(prisma_client, team_id) + + if _aws_mock_enabled(): + return _mock_caller_identity(creds) + + # Real STS path — B0 will land this. Imported lazily to avoid pulling in + # boto3 at module-load time for proxies that don't use cloud agents. + try: + import boto3 # type: ignore + from botocore.exceptions import ClientError # type: ignore + except ImportError: + return TestConnectionResponse( + ok=False, + error="boto3 not installed on this proxy — install litellm[proxy] extras.", + ) + + if not creds.get("access_key_id"): + return TestConnectionResponse( + ok=False, + error="No AWS credentials configured for this team.", + ) + + try: + client = boto3.client( + "sts", + aws_access_key_id=creds["access_key_id"], + aws_secret_access_key=creds["secret_access_key"], + region_name=creds.get("region") or "us-west-2", + ) + identity = client.get_caller_identity() + except ClientError as exc: # pragma: no cover — exercised once B0 lands + verbose_proxy_logger.warning( + "agent-vm-config test-connection failed for team=%s: %s", + team_id, + exc, + ) + return TestConnectionResponse(ok=False, error=str(exc)) + + return TestConnectionResponse( + ok=True, + account_id=identity.get("Account"), + arn=identity.get("Arn"), + region=creds.get("region") or "us-west-2", + ) diff --git a/litellm/proxy/agent_settings_endpoints/worker_endpoints.py b/litellm/proxy/agent_settings_endpoints/worker_endpoints.py new file mode 100644 index 00000000000..780b8fc7f38 --- /dev/null +++ b/litellm/proxy/agent_settings_endpoints/worker_endpoints.py @@ -0,0 +1,352 @@ +""" +`/v2/agent-workers` endpoints (LIT-2891 / Screen 3). + +Self-hosted worker registration flow: + +1. Operator clicks "Add Machine" → UI calls `POST /v2/agent-workers/pair-token`. + We mint a 32-byte urlsafe token, persist only its sha256, return the raw + token + an install one-liner. TTL = 15 min, single use. +2. Operator runs the install command on the worker box. The worker calls + `POST /v2/agent-workers/register` with the raw token + its hostname. We + re-hash, atomically mark the pairing-token row as consumed, mint a long- + lived worker JWT, persist its sha256, and return the raw JWT to the worker. + Both raw values are returned exactly once and never persisted. +3. The dashboard lists workers via `GET /v2/agent-workers` and revokes them + via `DELETE /v2/agent-workers/{id}`. + +The actual long-poll / hydrate transport is owned by Epic B2 — this module +only owns the auth handshake and CRUD surface. +""" + +import os +import secrets as _stdlib_secrets +from datetime import datetime, timezone +from typing import Any, Dict, List, Optional + +from fastapi import APIRouter, Depends, HTTPException, Request + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy.agent_settings_endpoints.pair_tokens import ( + build_install_command, + hash_pair_token, + hash_worker_jwt, + is_expired, + issue_pair_token, +) +from litellm.proxy.agent_settings_endpoints.types import ( + AgentWorkerListResponse, + AgentWorkerRegisterRequest, + AgentWorkerRegisterResponse, + AgentWorkerResponse, + PairTokenResponse, +) +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +def _resolve_team_id(user_api_key_dict: UserAPIKeyAuth) -> str: + """Pick the team to scope this request to. Raise 400 if missing.""" + team_id = user_api_key_dict.team_id or (user_api_key_dict.metadata or {}).get( + "team_id" + ) + if not team_id: + raise HTTPException( + status_code=400, + detail=( + "Cloud Agent workers are scoped to a team. Pick a team from the " + "header switcher and try again." + ), + ) + return team_id + + +def _row_to_worker_response(row: Dict[str, Any]) -> AgentWorkerResponse: + """Map a Prisma worker row to the public response shape.""" + last_seen = row.get("last_seen_at") + return AgentWorkerResponse( + id=row["id"], + hostname=row.get("hostname") or "", + status=row.get("status") or "offline", + last_seen_at=str(last_seen) if last_seen is not None else None, + cpu_pct=row.get("cpu_pct"), + mem_gb=row.get("mem_gb"), + active_sessions=int(row.get("active_sessions") or 0), + ) + + +def _resolve_proxy_url(request: Request) -> str: + """Resolve the public proxy URL embedded in the install one-liner. + + SECURITY: this URL ends up in `--proxy ` of the curl-pipe-sh + install command. If an attacker can influence it, they can redirect + the worker (and the freshly-issued pair token) to a host they control. + Resolution order: + + 1. `LITELLM_CLOUD_AGENT_PROXY_BASE_URL` env var — operator-configured, + fully trusted. This is the recommended path for production. + 2. `X-Forwarded-Host` / `X-Forwarded-Proto` — only honored when the + operator opts in via `LITELLM_TRUST_PROXY_HEADERS=1`. Required for + deployments behind a reverse proxy that terminates TLS. + 3. The request's direct `Host` header + URL scheme. Safe by default + because it reflects the actual TCP-layer destination of the + request, not an attacker-controlled hop hint. + + `localhost:4000` is the last-resort fallback for tests/local dev. + """ + configured = os.getenv("LITELLM_CLOUD_AGENT_PROXY_BASE_URL") + if configured: + return configured.rstrip("/") + + if os.getenv("LITELLM_TRUST_PROXY_HEADERS", "0") == "1": + forwarded_proto = request.headers.get("x-forwarded-proto") + forwarded_host = request.headers.get("x-forwarded-host") + if forwarded_host: + scheme = forwarded_proto or request.url.scheme or "https" + return f"{scheme}://{forwarded_host}" + + host = request.headers.get("host") or "localhost:4000" + scheme = request.url.scheme or "https" + return f"{scheme}://{host}" + + +@router.get( + "/v2/agent-workers", + dependencies=[Depends(user_api_key_auth)], + response_model=AgentWorkerListResponse, + tags=["cloud agents"], +) +async def list_agent_workers( + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> AgentWorkerListResponse: + """List the team's self-hosted workers.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + rows = await prisma_client.db.litellm_agentworker.find_many( + where={"team_id": team_id}, + order={"created_at": "desc"}, + ) + workers: List[AgentWorkerResponse] = [ + _row_to_worker_response(dict(r) if not isinstance(r, dict) else r) for r in rows + ] + return AgentWorkerListResponse(workers=workers) + + +@router.post( + "/v2/agent-workers/pair-token", + dependencies=[Depends(user_api_key_auth)], + response_model=PairTokenResponse, + tags=["cloud agents"], +) +async def create_pair_token( + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> PairTokenResponse: + """Mint a single-use 15-minute pairing token. Raw token returned ONCE.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + issued = issue_pair_token() + + try: + await prisma_client.db.litellm_agentworkerpairingtoken.create( + data={ + "token_hash": issued.token_hash, + "team_id": team_id, + "created_by": user_api_key_dict.user_id or "unknown", + "expires_at": issued.expires_at, + } + ) + except Exception as exc: + verbose_proxy_logger.exception( + "Failed to persist pair token for team=%s: %s", team_id, exc + ) + raise HTTPException(status_code=500, detail="Failed to mint pair token.") + + install_command = build_install_command( + proxy_url=_resolve_proxy_url(request), + raw_token=issued.raw_token, + ) + return PairTokenResponse( + token=issued.raw_token, + expires_at=issued.expires_at.isoformat(), + install_command=install_command, + ) + + +@router.post( + "/v2/agent-workers/register", + response_model=AgentWorkerRegisterResponse, + tags=["cloud agents"], +) +async def register_agent_worker( + request: Request, + body: AgentWorkerRegisterRequest, +) -> AgentWorkerRegisterResponse: + """Worker exchanges its pair token for a long-lived JWT. + + NOTE: this endpoint is NOT behind `user_api_key_auth` — the worker has + no API key yet, the pair token *is* its proof of authorization. We + enforce single-use atomically by checking `used_at IS NULL` on update. + + Defense-in-depth: the pair token is 256 bits of entropy with a 15-min + TTL, so brute-force is not a credible threat. We log each failed + attempt with the source IP so operators running this behind a WAF or + fail2ban can detect and rate-limit abusive callers at the network + layer (the proxy itself doesn't ship a built-in per-IP limiter). + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + client_ip = request.client.host if request.client is not None else "unknown" + + token_hash = hash_pair_token(body.pair_token) + pair_row = await prisma_client.db.litellm_agentworkerpairingtoken.find_unique( + where={"token_hash": token_hash} + ) + if pair_row is None: + verbose_proxy_logger.warning( + "agent-workers/register: invalid pair token from ip=%s hostname=%s", + client_ip, + body.hostname, + ) + raise HTTPException(status_code=401, detail="Invalid pairing token.") + + pair = dict(pair_row) if not isinstance(pair_row, dict) else pair_row + if pair.get("used_at") is not None: + verbose_proxy_logger.warning( + "agent-workers/register: replay of consumed pair token from ip=%s hostname=%s team=%s", + client_ip, + body.hostname, + pair.get("team_id"), + ) + raise HTTPException( + status_code=401, detail="Pairing token has already been used." + ) + if is_expired(pair.get("expires_at")): + verbose_proxy_logger.warning( + "agent-workers/register: expired pair token from ip=%s hostname=%s team=%s", + client_ip, + body.hostname, + pair.get("team_id"), + ) + raise HTTPException(status_code=401, detail="Pairing token has expired.") + + team_id = pair["team_id"] + + # Atomically consume the pair token. `update_many` with the `used_at: None` + # filter is the closest Prisma gives us to a CAS — if a second register + # raced us, the second one matches 0 rows and we 401. + consumed = await prisma_client.db.litellm_agentworkerpairingtoken.update_many( + where={"token_hash": token_hash, "used_at": None}, + data={"used_at": datetime.now(timezone.utc)}, + ) + if not consumed: + raise HTTPException( + status_code=401, detail="Pairing token has already been used." + ) + + # Mint the worker JWT. This is a short opaque urlsafe string; the daemon + # presents it on every long-poll. We persist only its sha256. + raw_jwt = _stdlib_secrets.token_urlsafe(48) + jwt_hash = hash_worker_jwt(raw_jwt) + + try: + worker = await prisma_client.db.litellm_agentworker.create( + data={ + "team_id": team_id, + "hostname": body.hostname, + "status": "online", + "last_seen_at": datetime.now(timezone.utc), + "active_sessions": 0, + "worker_jwt_hash": jwt_hash, + } + ) + except Exception as exc: + verbose_proxy_logger.exception( + "Failed to create agent worker hostname=%s team=%s: %s", + body.hostname, + team_id, + exc, + ) + raise HTTPException(status_code=500, detail="Failed to register worker.") + + worker_dict = dict(worker) if not isinstance(worker, dict) else worker + return AgentWorkerRegisterResponse( + worker_id=worker_dict["id"], + worker_jwt=raw_jwt, + ) + + +@router.delete( + "/v2/agent-workers/{worker_id}", + dependencies=[Depends(user_api_key_auth)], + tags=["cloud agents"], +) +async def delete_agent_worker( + request: Request, + worker_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> Dict[str, Any]: + """Revoke a worker. Idempotent — 404 if not found, no-op if already gone.""" + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + team_id = _resolve_team_id(user_api_key_dict) + existing = await prisma_client.db.litellm_agentworker.find_unique( + where={"id": worker_id} + ) + if existing is None: + raise HTTPException(status_code=404, detail="Worker not found.") + + existing_dict = dict(existing) if not isinstance(existing, dict) else existing + if existing_dict.get("team_id") != team_id: + # Don't leak existence cross-team; same status code as not-found. + raise HTTPException(status_code=404, detail="Worker not found.") + + await prisma_client.db.litellm_agentworker.delete(where={"id": worker_id}) + return {"deleted": True, "id": worker_id} + + +# Re-exported for tests + B2 hydrate path: given a raw worker JWT, look up the +# worker row. Centralized here so the auth scheme lives in exactly one place. +async def find_worker_by_jwt( + prisma_client: Any, raw_jwt: str +) -> Optional[Dict[str, Any]]: + """Look up a worker row by its (raw) JWT. Returns None if not found. + + `worker_jwt_hash` is indexed (see migration), so this is an index + lookup — fine for B2's per-heartbeat call pattern. + """ + jwt_hash = hash_worker_jwt(raw_jwt) + worker = await prisma_client.db.litellm_agentworker.find_first( + where={"worker_jwt_hash": jwt_hash} + ) + if worker is None: + return None + return dict(worker) if not isinstance(worker, dict) else worker diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a962d90a038..b4d891ed16e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -244,6 +244,18 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._types import * from litellm.proxy._lazy_features import attach_lazy_features +from litellm.proxy.agent_settings_endpoints.pool_status_endpoints import ( + router as agent_pool_status_router, +) +from litellm.proxy.agent_settings_endpoints.secrets_endpoints import ( + router as agent_secrets_router, +) +from litellm.proxy.agent_settings_endpoints.vm_config_endpoints import ( + router as agent_vm_config_router, +) +from litellm.proxy.agent_settings_endpoints.worker_endpoints import ( + router as agent_workers_router, +) from litellm.proxy.analytics_endpoints.analytics_endpoints import ( router as analytics_router, ) @@ -14904,6 +14916,11 @@ app.include_router(cache_settings_router) app.include_router(user_agent_analytics_router) app.include_router(enterprise_router) app.include_router(ui_discovery_endpoints_router) +# Cloud Agents settings (LIT-2891) — VM config, secrets, self-hosted workers. +app.include_router(agent_vm_config_router) +app.include_router(agent_secrets_router) +app.include_router(agent_workers_router) +app.include_router(agent_pool_status_router) # Eager: /models/{name}:method overlaps with the OpenAI /models endpoint. app.include_router(google_router) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 6a2e09e1616..d52083d1a90 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1375,7 +1375,6 @@ model LiteLLM_WorkflowMessage { @@index([run_id]) } - // =========================================================================== // Agent Sessions / Runs (Cursor SDK on LiteLLM) // @@ -1468,3 +1467,88 @@ model LiteLLM_AgentRunEvent { @@unique([run_id, seq]) @@index([run_id, seq]) } + +// =========================================================================== +// Cloud Agent settings (LIT-2891) — per-team VM provider config + secrets + +// self-hosted worker registry +// =========================================================================== + +// Per-team Cloud Agent VM provider configuration. Holds AWS BYOC creds +// (encrypted via the same nacl/SecretBox path as virtual keys), provisioning +// defaults, warm-pool settings, and the network egress allowlist that gets +// pushed to the daemon at hydrate time. +model LiteLLM_AgentVMConfig { + team_id String @id + provider String @default("disabled") // "ec2" | "self_hosted" | "disabled" + aws_auth_method String? // "access_keys" | "iam_role" | "instance_metadata" + aws_access_key_id_enc String? // encrypted + aws_secret_access_key_enc String? // encrypted + aws_role_arn_enc String? // encrypted (cross-account role mode) + aws_region String? + ami_id String? + instance_type String? + subnet_id String? + security_group_id String? + iam_instance_profile String? + use_spot Boolean @default(true) + max_session_minutes Int @default(120) + warm_pool_enabled Boolean @default(false) + warm_pool_size Int @default(0) + max_idle_minutes Int @default(30) + hydrate_transport String @default("auto") // "auto" | "ssm" | "long_poll" + network_access Json @default("{\"mode\":\"allow_all\",\"allowlist\":[]}") + self_hosted_enabled Boolean @default(false) + created_at DateTime @default(now()) + updated_at DateTime @default(now()) @updatedAt +} + +// Per-team encrypted secrets injected into agent VMs at session start. Value +// is ALWAYS write-only: GET endpoints must never return value_enc decrypted. +// Scope is "all" or a list of repo full_name strings; the proxy joins on +// session.repos at hydrate time. +model LiteLLM_AgentSecret { + id String @id @default(uuid()) + team_id String + name String + value_enc String // base64-encoded, nacl.SecretBox encrypted, write-only + scope Json @default("\"all\"") // "all" | string[] + type String @default("env") // "env" | "file" + file_path String? + created_at DateTime @default(now()) + updated_at DateTime @default(now()) @updatedAt + created_by String? + + @@unique([team_id, name]) + @@index([team_id]) +} + +// Self-hosted worker registrations. Each worker holds a long-lived JWT and +// long-polls for hydrate. status is best-effort heartbeat tracking. +model LiteLLM_AgentWorker { + id String @id @default(uuid()) + team_id String + hostname String + status String @default("offline") // "online" | "offline" + last_seen_at DateTime? + cpu_pct Float? + mem_gb Float? + active_sessions Int @default(0) + worker_jwt_hash String // sha256 of issued JWT — never store raw JWT + created_at DateTime @default(now()) + + @@index([team_id, status]) + @@index([worker_jwt_hash]) +} + +// Single-use 15-minute pairing tokens that workers exchange for a long-lived +// worker JWT during the install flow. +model LiteLLM_AgentWorkerPairingToken { + token_hash String @id // sha256 of the raw token (raw token never persisted) + team_id String + created_by String + expires_at DateTime + used_at DateTime? + created_at DateTime @default(now()) + + @@index([team_id]) +} diff --git a/schema.prisma b/schema.prisma index 6a2e09e1616..d52083d1a90 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1375,7 +1375,6 @@ model LiteLLM_WorkflowMessage { @@index([run_id]) } - // =========================================================================== // Agent Sessions / Runs (Cursor SDK on LiteLLM) // @@ -1468,3 +1467,88 @@ model LiteLLM_AgentRunEvent { @@unique([run_id, seq]) @@index([run_id, seq]) } + +// =========================================================================== +// Cloud Agent settings (LIT-2891) — per-team VM provider config + secrets + +// self-hosted worker registry +// =========================================================================== + +// Per-team Cloud Agent VM provider configuration. Holds AWS BYOC creds +// (encrypted via the same nacl/SecretBox path as virtual keys), provisioning +// defaults, warm-pool settings, and the network egress allowlist that gets +// pushed to the daemon at hydrate time. +model LiteLLM_AgentVMConfig { + team_id String @id + provider String @default("disabled") // "ec2" | "self_hosted" | "disabled" + aws_auth_method String? // "access_keys" | "iam_role" | "instance_metadata" + aws_access_key_id_enc String? // encrypted + aws_secret_access_key_enc String? // encrypted + aws_role_arn_enc String? // encrypted (cross-account role mode) + aws_region String? + ami_id String? + instance_type String? + subnet_id String? + security_group_id String? + iam_instance_profile String? + use_spot Boolean @default(true) + max_session_minutes Int @default(120) + warm_pool_enabled Boolean @default(false) + warm_pool_size Int @default(0) + max_idle_minutes Int @default(30) + hydrate_transport String @default("auto") // "auto" | "ssm" | "long_poll" + network_access Json @default("{\"mode\":\"allow_all\",\"allowlist\":[]}") + self_hosted_enabled Boolean @default(false) + created_at DateTime @default(now()) + updated_at DateTime @default(now()) @updatedAt +} + +// Per-team encrypted secrets injected into agent VMs at session start. Value +// is ALWAYS write-only: GET endpoints must never return value_enc decrypted. +// Scope is "all" or a list of repo full_name strings; the proxy joins on +// session.repos at hydrate time. +model LiteLLM_AgentSecret { + id String @id @default(uuid()) + team_id String + name String + value_enc String // base64-encoded, nacl.SecretBox encrypted, write-only + scope Json @default("\"all\"") // "all" | string[] + type String @default("env") // "env" | "file" + file_path String? + created_at DateTime @default(now()) + updated_at DateTime @default(now()) @updatedAt + created_by String? + + @@unique([team_id, name]) + @@index([team_id]) +} + +// Self-hosted worker registrations. Each worker holds a long-lived JWT and +// long-polls for hydrate. status is best-effort heartbeat tracking. +model LiteLLM_AgentWorker { + id String @id @default(uuid()) + team_id String + hostname String + status String @default("offline") // "online" | "offline" + last_seen_at DateTime? + cpu_pct Float? + mem_gb Float? + active_sessions Int @default(0) + worker_jwt_hash String // sha256 of issued JWT — never store raw JWT + created_at DateTime @default(now()) + + @@index([team_id, status]) + @@index([worker_jwt_hash]) +} + +// Single-use 15-minute pairing tokens that workers exchange for a long-lived +// worker JWT during the install flow. +model LiteLLM_AgentWorkerPairingToken { + token_hash String @id // sha256 of the raw token (raw token never persisted) + team_id String + created_by String + expires_at DateTime + used_at DateTime? + created_at DateTime @default(now()) + + @@index([team_id]) +} diff --git a/tests/test_litellm/proxy/agent_settings_endpoints/__init__.py b/tests/test_litellm/proxy/agent_settings_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/agent_settings_endpoints/test_aws_mock_default.py b/tests/test_litellm/proxy/agent_settings_endpoints/test_aws_mock_default.py new file mode 100644 index 00000000000..82c17e563cc --- /dev/null +++ b/tests/test_litellm/proxy/agent_settings_endpoints/test_aws_mock_default.py @@ -0,0 +1,45 @@ +""" +The `_aws_mock_enabled` flag controls whether `POST /v2/agent-vm-config/ +test-connection` returns a synthetic success response or actually calls +`sts:GetCallerIdentity`. The default MUST be off, so a freshly deployed +production proxy does not silently green-light invalid AWS credentials. + +This is the fix for the P1 Greptile flagged on PR #27332 — previously the +default was `"1"` (mock-on), so an operator who saved bad credentials +would see a fake success and only discover the failure when VMs failed +to launch. +""" + +import pytest + +from litellm.proxy.agent_settings_endpoints.vm_config_endpoints import ( + _aws_mock_enabled, +) + + +class TestAwsMockDefault: + def test_mock_disabled_when_env_unset(self, monkeypatch): + monkeypatch.delenv("LITELLM_CLOUD_AGENT_MOCK_AWS", raising=False) + assert _aws_mock_enabled() is False + + def test_mock_disabled_when_env_zero(self, monkeypatch): + monkeypatch.setenv("LITELLM_CLOUD_AGENT_MOCK_AWS", "0") + assert _aws_mock_enabled() is False + + def test_mock_disabled_for_arbitrary_truthy_strings(self, monkeypatch): + # We require an EXPLICIT "1" — no fuzzy-truthy parsing. A typo + # like "true" or "yes" must not silently flip on the mock. + for bad_value in ("true", "yes", "on", "TRUE", " 1 "): + monkeypatch.setenv("LITELLM_CLOUD_AGENT_MOCK_AWS", bad_value) + assert ( + _aws_mock_enabled() is False + ), f"value {bad_value!r} unexpectedly enabled mock mode" + + def test_mock_enabled_only_with_explicit_one(self, monkeypatch): + monkeypatch.setenv("LITELLM_CLOUD_AGENT_MOCK_AWS", "1") + assert _aws_mock_enabled() is True + + +@pytest.fixture(autouse=True) +def _no_op_fixture(): + yield diff --git a/tests/test_litellm/proxy/agent_settings_endpoints/test_pair_tokens.py b/tests/test_litellm/proxy/agent_settings_endpoints/test_pair_tokens.py new file mode 100644 index 00000000000..40e1f3ebffa --- /dev/null +++ b/tests/test_litellm/proxy/agent_settings_endpoints/test_pair_tokens.py @@ -0,0 +1,157 @@ +""" +Unit tests for the pair-token + worker-JWT helpers (LIT-2891 validation #5). + +These helpers underpin the "Add Machine" install flow. The invariants we +care about are: + +* The raw token is high-entropy and unique per call. +* Only the sha256 is intended to be persisted; the helper returns both so + the caller can do the right thing, but the digest is what lands in the DB. +* Hashing is deterministic — same raw token always hashes to the same + digest, so the consume-side lookup works. +* Expiry is honored, with naive datetimes treated as UTC (matches the + proxy's existing convention for DB-stored timestamps). +* The install command renders exactly the format the docs/UX promise + (`curl ... | sh -s -- --proxy ... --token ...`). Tests pin this so + changing the install host doesn't silently break the on-box install. +""" + +from datetime import datetime, timedelta, timezone + +import pytest + +from litellm.proxy.agent_settings_endpoints.pair_tokens import ( + build_install_command, + hash_pair_token, + hash_worker_jwt, + is_expired, + issue_pair_token, +) + + +class TestIssuePairToken: + def test_returns_distinct_raw_tokens_per_call(self): + a = issue_pair_token() + b = issue_pair_token() + assert a.raw_token != b.raw_token + assert a.token_hash != b.token_hash + + def test_token_hash_matches_raw(self): + issued = issue_pair_token() + assert issued.token_hash == hash_pair_token(issued.raw_token) + + def test_default_expiry_is_in_the_future(self): + issued = issue_pair_token() + now = datetime.now(timezone.utc) + assert issued.expires_at > now + # Should be ~15 minutes; allow generous slack so test isn't flaky. + assert issued.expires_at - now <= timedelta(minutes=20) + + def test_custom_ttl_respected(self): + issued = issue_pair_token(ttl_minutes=1) + delta = issued.expires_at - datetime.now(timezone.utc) + assert delta <= timedelta(minutes=2) + assert delta >= timedelta(seconds=30) + + def test_raw_token_has_meaningful_entropy(self): + # 32 bytes urlsafe-base64 → 43 chars unpadded; we just need to + # confirm the helper isn't accidentally returning a short string. + issued = issue_pair_token() + assert len(issued.raw_token) >= 32 + + +class TestHashHelpers: + def test_pair_token_hash_is_deterministic(self): + token = "abc123" + assert hash_pair_token(token) == hash_pair_token(token) + + def test_pair_token_hash_distinguishes_inputs(self): + assert hash_pair_token("a") != hash_pair_token("b") + + def test_worker_jwt_hash_uses_same_algo(self): + # Both helpers should be sha256 hex; identical input → identical + # output. (Not strictly required by the spec but a useful sanity + # check that we didn't accidentally swap the algorithm in one.) + assert hash_pair_token("zzz") == hash_worker_jwt("zzz") + + def test_hashes_are_64_hex_chars(self): + digest = hash_pair_token("anything") + assert len(digest) == 64 + int(digest, 16) # valid hex; raises if not + + +class TestIsExpired: + def test_future_is_not_expired(self): + future = datetime.now(timezone.utc) + timedelta(minutes=5) + assert is_expired(future) is False + + def test_past_is_expired(self): + past = datetime.now(timezone.utc) - timedelta(minutes=5) + assert is_expired(past) is True + + def test_none_is_treated_as_expired(self): + # Defensive default — missing expiry must never be treated as + # "valid forever". + assert is_expired(None) is True + + def test_naive_datetime_treated_as_utc(self): + # Matches the proxy DB convention; naive timestamps come back from + # SQLite without tzinfo. + future_naive = datetime.utcnow() + timedelta( + minutes=5 + ) # noqa: DTZ003 — intentionally naive for the test + assert is_expired(future_naive) is False + + def test_explicit_now_argument_used(self): + expires = datetime.now(timezone.utc) + timedelta(minutes=1) + far_future = expires + timedelta(hours=1) + assert is_expired(expires, now=far_future) is True + + +class TestBuildInstallCommand: + def test_default_renders_expected_one_liner(self): + cmd = build_install_command( + proxy_url="https://proxy.example.com", raw_token="TOK" + ) + # All values pass through shlex.quote — simple alnum/`:/.-` strings + # come back unquoted, matching the documented one-liner. + assert cmd == ( + "curl -fsS https://litellm.ai/install-worker | sh -s -- " + "--proxy https://proxy.example.com --token TOK" + ) + + def test_install_url_override(self): + cmd = build_install_command( + proxy_url="https://p", + raw_token="T", + install_script_url="https://internal.example/install.sh", + ) + assert "https://internal.example/install.sh" in cmd + assert "--proxy https://p" in cmd + assert "--token T" in cmd + + def test_proxy_url_with_metacharacters_is_shell_quoted(self): + # Defense against header-injection / config bugs that might land a + # space, semicolon, or quote in the proxy URL. shlex.quote wraps the + # value in single quotes so it can't break out of the install line. + cmd = build_install_command(proxy_url="https://h; rm -rf /", raw_token="TOK") + assert "'https://h; rm -rf /'" in cmd + # Sanity: the dangerous payload must NOT appear unquoted. + assert "--proxy https://h; rm -rf /" not in cmd + + def test_raw_token_with_metacharacters_is_shell_quoted(self): + cmd = build_install_command(proxy_url="https://p", raw_token="abc def$(whoami)") + assert "'abc def$(whoami)'" in cmd + + def test_install_script_url_is_shell_quoted(self): + cmd = build_install_command( + proxy_url="https://p", + raw_token="T", + install_script_url="https://hosts space.example/install.sh", + ) + assert "'https://hosts space.example/install.sh'" in cmd + + +@pytest.fixture(autouse=True) +def _no_op_fixture(): + yield diff --git a/tests/test_litellm/proxy/agent_settings_endpoints/test_proxy_url_resolution.py b/tests/test_litellm/proxy/agent_settings_endpoints/test_proxy_url_resolution.py new file mode 100644 index 00000000000..00976fdb283 --- /dev/null +++ b/tests/test_litellm/proxy/agent_settings_endpoints/test_proxy_url_resolution.py @@ -0,0 +1,140 @@ +""" +Tests for `_resolve_proxy_url` — the helper that picks the URL embedded in +the install one-liner returned to the operator. + +This is security-sensitive: if an attacker can influence the URL, they can +redirect the worker (and the freshly-issued pair token) to a host they +control. The contract we test: + +* `LITELLM_CLOUD_AGENT_PROXY_BASE_URL` (when set) wins, regardless of what + the request claims. Operator-configured value is the only fully trusted + source. +* `X-Forwarded-Host` / `X-Forwarded-Proto` are IGNORED unless the operator + explicitly opts in via `LITELLM_TRUST_PROXY_HEADERS=1`. This is the fix + for the header-injection path Greptile flagged on PR #27332. +* When neither override is present, we fall back to the request's direct + `Host` header — that reflects the actual TCP destination of the request, + not an attacker-controlled hop hint. +""" + +from types import SimpleNamespace + +import pytest + +from litellm.proxy.agent_settings_endpoints.worker_endpoints import ( + _resolve_proxy_url, +) + + +def _fake_request( + *, + host: str = "proxy.example.com", + scheme: str = "https", + forwarded_host: str = "", + forwarded_proto: str = "", +): + """Build a minimal stand-in for fastapi.Request with just the bits + `_resolve_proxy_url` actually reads. Avoids spinning up the full + Starlette request object.""" + headers = {"host": host} + if forwarded_host: + headers["x-forwarded-host"] = forwarded_host + if forwarded_proto: + headers["x-forwarded-proto"] = forwarded_proto + return SimpleNamespace( + headers=headers, + url=SimpleNamespace(scheme=scheme), + client=SimpleNamespace(host="127.0.0.1"), + ) + + +class TestEnvOverrideWins: + def test_explicit_proxy_base_url_used_verbatim(self, monkeypatch): + monkeypatch.setenv( + "LITELLM_CLOUD_AGENT_PROXY_BASE_URL", "https://configured.example" + ) + url = _resolve_proxy_url( + _fake_request(forwarded_host="attacker.example", forwarded_proto="http") + ) + assert url == "https://configured.example" + + def test_proxy_base_url_strips_trailing_slash(self, monkeypatch): + monkeypatch.setenv( + "LITELLM_CLOUD_AGENT_PROXY_BASE_URL", "https://configured.example/" + ) + url = _resolve_proxy_url(_fake_request()) + assert url == "https://configured.example" + + +class TestForwardedHeadersIgnoredByDefault: + def test_xff_host_ignored_without_trust_flag(self, monkeypatch): + monkeypatch.delenv("LITELLM_CLOUD_AGENT_PROXY_BASE_URL", raising=False) + monkeypatch.delenv("LITELLM_TRUST_PROXY_HEADERS", raising=False) + url = _resolve_proxy_url( + _fake_request( + host="trusted.example.com", + forwarded_host="attacker.example", + forwarded_proto="http", + ) + ) + # The forged X-Forwarded-Host MUST NOT show up. + assert "attacker.example" not in url + assert url == "https://trusted.example.com" + + def test_xff_host_explicit_off(self, monkeypatch): + monkeypatch.delenv("LITELLM_CLOUD_AGENT_PROXY_BASE_URL", raising=False) + monkeypatch.setenv("LITELLM_TRUST_PROXY_HEADERS", "0") + url = _resolve_proxy_url( + _fake_request( + host="trusted.example.com", + forwarded_host="attacker.example", + ) + ) + assert "attacker.example" not in url + + +class TestForwardedHeadersOptIn: + def test_xff_honored_when_trust_flag_set(self, monkeypatch): + monkeypatch.delenv("LITELLM_CLOUD_AGENT_PROXY_BASE_URL", raising=False) + monkeypatch.setenv("LITELLM_TRUST_PROXY_HEADERS", "1") + url = _resolve_proxy_url( + _fake_request( + host="origin.internal", + forwarded_host="public.example.com", + forwarded_proto="https", + ) + ) + assert url == "https://public.example.com" + + def test_xff_falls_back_to_host_when_only_proto_forwarded(self, monkeypatch): + # If only x-forwarded-proto is present (no host), we still use the + # direct Host header — we never silently mix forwarded scheme with + # request-direct host or vice versa. + monkeypatch.delenv("LITELLM_CLOUD_AGENT_PROXY_BASE_URL", raising=False) + monkeypatch.setenv("LITELLM_TRUST_PROXY_HEADERS", "1") + url = _resolve_proxy_url( + _fake_request( + host="proxy.example.com", + scheme="https", + forwarded_proto="http", + ) + ) + assert url == "https://proxy.example.com" + + +class TestFallbacks: + def test_fallback_to_localhost_when_no_host(self, monkeypatch): + monkeypatch.delenv("LITELLM_CLOUD_AGENT_PROXY_BASE_URL", raising=False) + monkeypatch.delenv("LITELLM_TRUST_PROXY_HEADERS", raising=False) + request = SimpleNamespace( + headers={}, + url=SimpleNamespace(scheme="https"), + client=SimpleNamespace(host="127.0.0.1"), + ) + url = _resolve_proxy_url(request) + assert url == "https://localhost:4000" + + +@pytest.fixture(autouse=True) +def _no_op_fixture(): + yield diff --git a/tests/test_litellm/proxy/agent_settings_endpoints/test_scope_filter.py b/tests/test_litellm/proxy/agent_settings_endpoints/test_scope_filter.py new file mode 100644 index 00000000000..4b70635f6c2 --- /dev/null +++ b/tests/test_litellm/proxy/agent_settings_endpoints/test_scope_filter.py @@ -0,0 +1,148 @@ +""" +Unit tests for `partition_secrets_for_session` (LIT-2891 validation #3). + +The scope filter is the single source of truth for "which secrets get pushed +into a hydrated agent session". Both the UI display path and the B2 hydrate +path call into here, so the access-control behavior must NEVER drift. + +Critical paths covered: +1. `scope == "all"` matches every session. +2. List scope matches when session.repos intersects the list. +3. List scope does NOT match when there's no intersection — including the + particularly subtle case where the session repo is from the same org + but a different repo (e.g. `BerriAI/other` vs scoped `BerriAI/litellm`). +4. Empty session repos with a list scope = no match (don't accidentally + leak when the session forgot to declare repos). +5. Repo URL normalization: `https://github.com/BerriAI/litellm.git` and + `BerriAI/litellm` and the dict shapes all match. +""" + +import pytest + +from litellm.proxy.agent_settings_endpoints.scope_filter import ( + normalize_repos, + partition_secrets_for_session, + secret_in_scope, +) + + +class TestSecretInScopeAll: + """`scope='all'` is the easy case but worth pinning.""" + + def test_all_scope_matches_with_empty_repos(self): + assert secret_in_scope("all", []) is True + + def test_all_scope_matches_with_repos(self): + assert secret_in_scope("all", ["BerriAI/litellm"]) is True + + def test_all_scope_matches_with_dict_repos(self): + assert secret_in_scope("all", [{"full_name": "BerriAI/litellm"}]) is True + + +class TestSecretInScopeList: + """The validation #3 core: per-repo scope must isolate.""" + + def test_list_scope_matches_intersecting_repo(self): + assert secret_in_scope(["BerriAI/litellm"], ["BerriAI/litellm"]) is True + + def test_list_scope_rejects_different_repo_same_org(self): + # This is the LIT-2891 validation #3 case: a secret scoped to + # BerriAI/litellm must NOT appear in a hydrate payload for + # BerriAI/other-repo, even though they're the same org. + assert secret_in_scope(["BerriAI/litellm"], ["BerriAI/other"]) is False + + def test_list_scope_rejects_different_org(self): + assert secret_in_scope(["BerriAI/litellm"], ["openai/openai-python"]) is False + + def test_list_scope_with_empty_session_repos_does_not_match(self): + # Defense-in-depth: a session that forgot to declare repos must + # NOT inherit list-scoped secrets. + assert secret_in_scope(["BerriAI/litellm"], []) is False + + def test_empty_list_scope_never_matches(self): + # `scope=[]` is treated as "no repos" — explicit safe default. + assert secret_in_scope([], ["BerriAI/litellm"]) is False + + def test_multi_repo_scope_matches_any(self): + scope = ["BerriAI/litellm", "BerriAI/litellm-docs"] + assert secret_in_scope(scope, ["BerriAI/litellm-docs"]) is True + assert secret_in_scope(scope, ["BerriAI/other"]) is False + + +class TestRepoNormalization: + """Different repo reference shapes must canonicalize to the same form.""" + + def test_plain_owner_name(self): + assert normalize_repos(["BerriAI/litellm"]) == ["berriai/litellm"] + + def test_url_with_https_scheme(self): + assert normalize_repos(["https://github.com/BerriAI/litellm"]) == [ + "berriai/litellm" + ] + + def test_url_with_dot_git_suffix(self): + assert normalize_repos(["https://github.com/BerriAI/litellm.git"]) == [ + "berriai/litellm" + ] + + def test_dict_with_full_name(self): + assert normalize_repos([{"full_name": "BerriAI/litellm"}]) == [ + "berriai/litellm" + ] + + def test_dict_with_url(self): + assert normalize_repos([{"url": "https://github.com/BerriAI/litellm.git"}]) == [ + "berriai/litellm" + ] + + def test_case_insensitive(self): + # Same repo with different casing — should dedupe to one. + result = normalize_repos( + ["BerriAI/LiteLLM", "berriai/litellm", "BERRIAI/LITELLM"] + ) + assert result == ["berriai/litellm"] + + def test_unparseable_returns_none(self): + # Garbage in, no entries out — must never crash the hydrate path. + assert normalize_repos(["not-a-repo"]) == [] + assert normalize_repos([None]) == [] + assert normalize_repos([42]) == [] + + +class TestPartitionSecretsForSession: + """The public entrypoint used by B2 hydrate. Output ordering matters + for the audit log so we pin it explicitly.""" + + def test_partitions_in_and_out_of_scope(self): + secrets = [ + ("DATABASE_URL", ["BerriAI/litellm"]), + ("OPENAI_API_KEY", "all"), + ("INTERNAL_TOKEN", ["BerriAI/internal"]), + ] + in_scope, out_of_scope = partition_secrets_for_session( + secrets, ["BerriAI/litellm"] + ) + assert in_scope == ["DATABASE_URL", "OPENAI_API_KEY"] + assert out_of_scope == ["INTERNAL_TOKEN"] + + def test_preserves_input_order(self): + # Audit logs read in order, so the partition has to keep order. + secrets = [ + ("Z_SECRET", "all"), + ("A_SECRET", "all"), + ("M_SECRET", ["BerriAI/litellm"]), + ] + in_scope, _ = partition_secrets_for_session(secrets, ["BerriAI/litellm"]) + assert in_scope == ["Z_SECRET", "A_SECRET", "M_SECRET"] + + def test_empty_secrets_returns_empty(self): + assert partition_secrets_for_session([], ["BerriAI/litellm"]) == ( + [], + [], + ) + + +# Make pytest pick the file up without an explicit `pytest_plugins`. +@pytest.fixture(autouse=True) +def _no_op_fixture(): + yield diff --git a/tests/test_litellm/proxy/agent_settings_endpoints/test_secrets_write_only.py b/tests/test_litellm/proxy/agent_settings_endpoints/test_secrets_write_only.py new file mode 100644 index 00000000000..1e895a765d8 --- /dev/null +++ b/tests/test_litellm/proxy/agent_settings_endpoints/test_secrets_write_only.py @@ -0,0 +1,211 @@ +""" +Write-only secret invariants (LIT-2891 validation #2). + +The single most important property of `/v2/agent-secrets`: a stored secret +value MUST never reappear on any GET response, ever. The defense is layered: + +* **Type-level**: `AgentSecretResponse` has no `value` field. Even an + accidental `**row` splat can't surface the ciphertext through the model. +* **Endpoint-level**: every `response_model` annotation references one of the + value-free models, and `_row_to_response` reads named columns explicitly + rather than splatting the whole Prisma row. +* **Schema-level**: `AgentSecretListResponse` only nests + `AgentSecretResponse`, so a list response can't smuggle values either. + +These tests pin all three so a future refactor can't quietly weaken the +contract. They deliberately avoid importing the endpoints module (which +would pull in the full proxy auth stack and its `orjson` dependency) — we +stress the type contract directly and grep the endpoint source as text. +""" + +import pathlib + +import pytest +from pydantic import ValidationError + +from litellm.proxy.agent_settings_endpoints.types import ( + AgentSecretCreateRequest, + AgentSecretListResponse, + AgentSecretResponse, + AgentSecretUpdateRequest, +) + +_SECRETS_ENDPOINTS_PATH = ( + pathlib.Path(__file__).resolve().parents[4] + / "litellm" + / "proxy" + / "agent_settings_endpoints" + / "secrets_endpoints.py" +) + + +def _read_endpoints_source() -> str: + # Read as text instead of importing the module — keeps these tests + # collectable without the full proxy dependency tree (orjson, prisma, + # etc.) installed. + assert ( + _SECRETS_ENDPOINTS_PATH.exists() + ), f"secrets_endpoints.py missing at {_SECRETS_ENDPOINTS_PATH}" + return _SECRETS_ENDPOINTS_PATH.read_text() + + +def _strip_docstrings_and_comments(source: str) -> str: + """Remove triple-quoted strings and `# ...` line comments. + + Used by the source-grep tests to avoid tripping on legitimate mentions + of `value_enc` inside docstrings (e.g. "intentionally NOT read here") + and inline explanatory comments. Crude regex is fine for our use — + the source file is small and we're checking a denylist. + """ + import re + + # Triple-quoted strings (greedy across lines). + no_docstrings = re.sub(r'"""[\s\S]*?"""', "", source) + no_docstrings = re.sub(r"'''[\s\S]*?'''", "", no_docstrings) + # `# ...` line comments. + no_comments = re.sub(r"#[^\n]*", "", no_docstrings) + return no_comments + + +class TestResponseModelHasNoValueField: + """The type system itself enforces validation #2.""" + + def test_agent_secret_response_has_no_value_field(self): + assert "value" not in AgentSecretResponse.model_fields + assert "value_enc" not in AgentSecretResponse.model_fields + + def test_response_model_dump_never_contains_value(self): + # Even if a future refactor added `value` to the row dict, the + # response schema would silently drop it (Pydantic ignores extras + # by default for BaseModel). + resp = AgentSecretResponse( + name="X", + scope="all", + type="env", + file_path=None, + created_at="now", + updated_at="now", + ) + dumped = resp.model_dump() + assert "value" not in dumped + assert "value_enc" not in dumped + + def test_list_response_only_nests_value_free_models(self): + # A list response is a wrapper around AgentSecretResponse — assert + # the nested type has the right shape. + resp = AgentSecretListResponse(secrets=[]) + dumped = resp.model_dump() + assert dumped == {"secrets": []} + # The model_fields annotation must be List[AgentSecretResponse]. + secrets_field = AgentSecretListResponse.model_fields["secrets"] + annotation = str(secrets_field.annotation) + assert "AgentSecretResponse" in annotation + + def test_response_model_drops_value_extra_via_validation(self): + # Round-trip through validate to confirm Pydantic strips an extra + # `value` even if a caller tries to slip it past the type. + raw = { + "name": "X", + "scope": "all", + "type": "env", + "file_path": None, + "created_at": "t", + "updated_at": "t", + "value": "DO-NOT-LEAK", + } + resp = AgentSecretResponse.model_validate(raw) + assert "value" not in resp.model_dump() + assert "DO-NOT-LEAK" not in str(resp.model_dump()) + + +class TestRequestModelsAcceptValue: + """The flip side: write-only means values DO go in on POST/PUT.""" + + def test_create_request_requires_value(self): + with pytest.raises(ValidationError): + AgentSecretCreateRequest(name="X", value="") # too short + with pytest.raises(ValidationError): + AgentSecretCreateRequest(name="X") # missing entirely + + def test_create_request_validates_name_charset(self): + # We use names as env-var keys, so they must follow shell-safe + # identifier rules. Reject names with dashes or starting with a digit. + with pytest.raises(ValidationError): + AgentSecretCreateRequest(name="bad-name", value="x") + with pytest.raises(ValidationError): + AgentSecretCreateRequest(name="9starting_digit", value="x") + # Valid identifier passes. + AgentSecretCreateRequest(name="OPENAI_API_KEY", value="x") + + def test_update_request_value_is_optional(self): + # PUT is partial — UI may send only a scope change. + body = AgentSecretUpdateRequest(scope="all") + assert body.value is None + + def test_update_request_value_min_length_1(self): + with pytest.raises(ValidationError): + AgentSecretUpdateRequest(value="") + + +class TestEndpointSourceContainsNoValueLeak: + """Belt-and-suspenders: scan the secrets endpoint source for accidental + references that could leak `value_enc` onto a response. If a future PR + introduces something like `response['value'] = decrypt(...)`, these + tests catch it before it ships. + """ + + def test_no_decrypt_call_in_endpoints_module(self): + source = _read_endpoints_source() + # The secrets module must not import the decrypt helper — only the + # VM config endpoints module needs it (for Test Connection). + assert "decrypt_optional" not in source + assert "decrypt_value_helper" not in source + + def test_endpoints_use_value_free_response_models(self): + source = _read_endpoints_source() + # Both response_model occurrences should reference the value-free + # types — never some bare dict that could drift. + assert "response_model=AgentSecretResponse" in source + assert "response_model=AgentSecretListResponse" in source + + def test_row_to_response_does_not_read_value_enc(self): + source = _read_endpoints_source() + # Find the helper that maps DB rows to API responses. The chokepoint + # is the GET path: it MUST NOT read value_enc. Write-side payload + # builders (create/update) are allowed to mention it because they + # write the ciphertext into the DB. + helper_marker = "def _row_to_response(row:" + helper_start = source.find(helper_marker) + assert helper_start != -1, ( + "_row_to_response helper missing — refactor must not lose this " + "chokepoint." + ) + next_def = source.find("\ndef ", helper_start + len(helper_marker)) + helper_body = ( + source[helper_start:] if next_def == -1 else source[helper_start:next_def] + ) + # Strip docstrings and `# ...` comments before checking. The helper's + # docstring legitimately calls out "value_enc is intentionally NOT + # read here", and we don't want that mention to trip the test. + stripped = _strip_docstrings_and_comments(helper_body) + assert "value_enc" not in stripped, ( + "_row_to_response references value_enc — the GET path must " + "never touch the ciphertext column." + ) + + def test_no_response_dict_assigns_value(self): + # Make sure no code path builds a response dict and adds `value` to + # it. (`response_model` enforces this on FastAPI's side, but we + # double-check the module doesn't construct such dicts directly.) + source = _read_endpoints_source() + # Disallow patterns like `"value": ...,` in response-shaped dicts. + # Allow `body.value` (a request field). The simplest portable rule + # is: the literal `"value":` must never appear in the source. + assert ( + '"value":' not in source + ), "secrets endpoint constructs a dict with a `value` key — possible leak." + + +@pytest.fixture(autouse=True) +def _no_op_fixture(): + yield