mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge branch 'litellm_lit-2891-settings-ui-cloud-agents' into litellm_lit-2890-b2-warm-pool
This commit is contained in:
commit
76b56b18e5
21 changed files with 2654 additions and 0 deletions
|
|
@ -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");
|
||||
|
|
@ -1397,3 +1397,83 @@ model LiteLLM_WorkflowMessage {
|
|||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
||||
// 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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
9
litellm/proxy/agent_settings_endpoints/__init__.py
Normal file
9
litellm/proxy/agent_settings_endpoints/__init__.py
Normal file
|
|
@ -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)
|
||||
"""
|
||||
44
litellm/proxy/agent_settings_endpoints/encryption.py
Normal file
44
litellm/proxy/agent_settings_endpoints/encryption.py
Normal file
|
|
@ -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)
|
||||
104
litellm/proxy/agent_settings_endpoints/pair_tokens.py
Normal file
104
litellm/proxy/agent_settings_endpoints/pair_tokens.py
Normal file
|
|
@ -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 <raw>`. 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}"
|
||||
)
|
||||
|
|
@ -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",
|
||||
},
|
||||
)
|
||||
127
litellm/proxy/agent_settings_endpoints/scope_filter.py
Normal file
127
litellm/proxy/agent_settings_endpoints/scope_filter.py
Normal file
|
|
@ -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
|
||||
309
litellm/proxy/agent_settings_endpoints/secrets_endpoints.py
Normal file
309
litellm/proxy/agent_settings_endpoints/secrets_endpoints.py
Normal file
|
|
@ -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}
|
||||
145
litellm/proxy/agent_settings_endpoints/types.py
Normal file
145
litellm/proxy/agent_settings_endpoints/types.py
Normal file
|
|
@ -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
|
||||
408
litellm/proxy/agent_settings_endpoints/vm_config_endpoints.py
Normal file
408
litellm/proxy/agent_settings_endpoints/vm_config_endpoints.py
Normal file
|
|
@ -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",
|
||||
)
|
||||
352
litellm/proxy/agent_settings_endpoints/worker_endpoints.py
Normal file
352
litellm/proxy/agent_settings_endpoints/worker_endpoints.py
Normal file
|
|
@ -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 <url>` 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
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -14880,6 +14892,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)
|
||||
|
||||
|
|
|
|||
|
|
@ -1397,3 +1397,83 @@ model LiteLLM_WorkflowMessage {
|
|||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
||||
// 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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1397,3 +1397,83 @@ model LiteLLM_WorkflowMessage {
|
|||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
||||
// 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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue