feat(managed agents): add v1 Fargate-backed managed-agent endpoints

Adds the full v1 managed_agents stack under
`/v1/managed_agents/*`:

- Sandbox templates: `POST/GET/DELETE /v1/managed_agents/sandbox-templates`
  Builds a Docker harness, pushes to ECR, registers a Fargate task definition.
- Agents: `POST/GET /v1/managed_agents/agents` with model + prompt + branch
  override.
- Sessions: `POST /v1/managed_agents/agents/{id}/session`,
  `GET/DELETE /v1/managed_agents/sessions/{id}`. Spawns a Fargate task,
  waits for ready, seeds the harness chat session (and optional first
  prompt) in one call.
- Passthrough: `POST /v1/managed_agents/sessions/{id}/message`,
  SSE `/events`, raw `/raw/{path}` for direct opencode access.

Includes:
  * Prisma schema for `LiteLLM_ManagedAgent{SandboxTemplate,,Session}Table`
    plus migrations under litellm-proxy-extras.
  * Background reconciler that stops orphaned Fargate tasks every 60s,
    wired into proxy_server lifespan.
  * Sample opencode Dockerfile + entrypoint under
    `managed_agents_endpoints/sample_harnesses/opencode/`.
  * Unit tests for templates, agents, sessions, lifecycle, bootstrap,
    git validation, and dockerfile registry.

Schema-validated end-to-end against a real Postgres instance (round-trip
create→include→update→delete). All 60 unit tests pass.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Krrish Dholakia 2026-05-07 15:17:21 -07:00
parent 8c9830eef9
commit 977ab6682c
29 changed files with 4701 additions and 2 deletions

View file

@ -0,0 +1,104 @@
-- CreateTable
CREATE TABLE "LiteLLM_ManagedAgentSandboxTemplateTable" (
"template_id" TEXT NOT NULL,
"template_name" TEXT,
"harness" TEXT NOT NULL,
"image_uri" TEXT NOT NULL,
"container_port" INTEGER NOT NULL DEFAULT 4096,
"image_env" JSONB DEFAULT '{}',
"repo_url" TEXT NOT NULL,
"default_branch" TEXT NOT NULL DEFAULT 'main',
"visibility" TEXT NOT NULL DEFAULT 'public',
"git_credential_id" TEXT,
"description" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"created_by" TEXT,
"updated_at" TIMESTAMP(3) NOT NULL,
"updated_by" TEXT,
CONSTRAINT "LiteLLM_ManagedAgentSandboxTemplateTable_pkey" PRIMARY KEY ("template_id")
);
-- CreateTable
CREATE TABLE "LiteLLM_ManagedAgentTable" (
"agent_id" TEXT NOT NULL,
"agent_name" TEXT,
"model" TEXT NOT NULL,
"prompt" TEXT,
"tools" JSONB NOT NULL DEFAULT '[]',
"template_id" TEXT NOT NULL,
"branch" TEXT,
"metadata" JSONB DEFAULT '{}',
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"created_by" TEXT,
"team_id" TEXT,
"organization_id" TEXT,
"updated_at" TIMESTAMP(3) NOT NULL,
"updated_by" TEXT,
CONSTRAINT "LiteLLM_ManagedAgentTable_pkey" PRIMARY KEY ("agent_id")
);
-- CreateTable
CREATE TABLE "LiteLLM_ManagedAgentSessionTable" (
"session_id" TEXT NOT NULL,
"agent_id" TEXT NOT NULL,
"status" TEXT NOT NULL DEFAULT 'creating',
"task_arn" TEXT,
"sandbox_url" TEXT,
"harness_session_id" TEXT,
"fargate_cluster" TEXT,
"fargate_task_def_arn" TEXT,
"virtual_key_hash" TEXT,
"failure_reason" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"created_by" TEXT,
"team_id" TEXT,
"last_seen_at" TIMESTAMP(3),
"expires_at" TIMESTAMP(3),
"stopped_at" TIMESTAMP(3),
CONSTRAINT "LiteLLM_ManagedAgentSessionTable_pkey" PRIMARY KEY ("session_id")
);
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSandboxTemplateTable_harness_idx" ON "LiteLLM_ManagedAgentSandboxTemplateTable"("harness");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSandboxTemplateTable_created_by_idx" ON "LiteLLM_ManagedAgentSandboxTemplateTable"("created_by");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSandboxTemplateTable_git_credential_id_idx" ON "LiteLLM_ManagedAgentSandboxTemplateTable"("git_credential_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentTable_template_id_idx" ON "LiteLLM_ManagedAgentTable"("template_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentTable_created_by_idx" ON "LiteLLM_ManagedAgentTable"("created_by");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentTable_team_id_idx" ON "LiteLLM_ManagedAgentTable"("team_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentTable_organization_id_idx" ON "LiteLLM_ManagedAgentTable"("organization_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSessionTable_agent_id_idx" ON "LiteLLM_ManagedAgentSessionTable"("agent_id");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSessionTable_status_idx" ON "LiteLLM_ManagedAgentSessionTable"("status");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSessionTable_created_by_idx" ON "LiteLLM_ManagedAgentSessionTable"("created_by");
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSessionTable_last_seen_at_idx" ON "LiteLLM_ManagedAgentSessionTable"("last_seen_at");
-- AddForeignKey
ALTER TABLE "LiteLLM_ManagedAgentSandboxTemplateTable" ADD CONSTRAINT "LiteLLM_ManagedAgentSandboxTemplateTable_git_credential_id_fkey" FOREIGN KEY ("git_credential_id") REFERENCES "LiteLLM_CredentialsTable"("credential_id") ON DELETE SET NULL ON UPDATE CASCADE;
-- AddForeignKey
ALTER TABLE "LiteLLM_ManagedAgentTable" ADD CONSTRAINT "LiteLLM_ManagedAgentTable_template_id_fkey" FOREIGN KEY ("template_id") REFERENCES "LiteLLM_ManagedAgentSandboxTemplateTable"("template_id") ON DELETE CASCADE ON UPDATE CASCADE;
-- AddForeignKey
ALTER TABLE "LiteLLM_ManagedAgentSessionTable" ADD CONSTRAINT "LiteLLM_ManagedAgentSessionTable_agent_id_fkey" FOREIGN KEY ("agent_id") REFERENCES "LiteLLM_ManagedAgentTable"("agent_id") ON DELETE CASCADE ON UPDATE CASCADE;

View file

@ -0,0 +1,28 @@
-- Schema v2 for managed-agents sandbox template:
-- replace `harness` with `dockerfile_id`, drop `image_env`, make `image_uri` nullable,
-- and add build-pipeline fields (task_def_arn, image_hash, build_status, build_error).
-- DropIndex
DROP INDEX "LiteLLM_ManagedAgentSandboxTemplateTable_harness_idx";
-- AlterTable
ALTER TABLE "LiteLLM_ManagedAgentSandboxTemplateTable"
DROP COLUMN "harness",
DROP COLUMN "image_env",
ALTER COLUMN "image_uri" DROP NOT NULL,
ADD COLUMN "task_def_arn" TEXT,
ADD COLUMN "image_hash" TEXT,
ADD COLUMN "build_status" TEXT NOT NULL DEFAULT 'pending',
ADD COLUMN "build_error" TEXT;
-- Add NOT NULL `dockerfile_id`. We use a transient default of '' so the ADD COLUMN
-- succeeds even on a populated table, then drop the default so future inserts must
-- supply a real value. There is no production data on this table yet (v1 just shipped),
-- so no backfill statement is required.
ALTER TABLE "LiteLLM_ManagedAgentSandboxTemplateTable"
ADD COLUMN "dockerfile_id" TEXT NOT NULL DEFAULT '';
ALTER TABLE "LiteLLM_ManagedAgentSandboxTemplateTable"
ALTER COLUMN "dockerfile_id" DROP DEFAULT;
-- CreateIndex
CREATE INDEX "LiteLLM_ManagedAgentSandboxTemplateTable_dockerfile_id_idx" ON "LiteLLM_ManagedAgentSandboxTemplateTable"("dockerfile_id");

View file

@ -38,11 +38,13 @@ model LiteLLM_CredentialsTable {
credential_id String @id @default(uuid())
credential_name String @unique
credential_values Json
credential_info Json?
credential_info Json?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
updated_by String
managed_agent_sandbox_templates LiteLLM_ManagedAgentSandboxTemplateTable[]
}
// Models on proxy
@ -1374,3 +1376,86 @@ model LiteLLM_WorkflowMessage {
@@unique([run_id, sequence_number])
@@index([run_id])
}
// Managed Agents: sandbox harness template (admin-defined recipe).
// Maps harness type → ECR image. e.g. opencode/claude-code/aider.
model LiteLLM_ManagedAgentSandboxTemplateTable {
template_id String @id @default(uuid())
template_name String?
dockerfile_id String
image_uri String?
task_def_arn String?
image_hash String?
build_status String @default("pending")
build_error String?
container_port Int @default(4096)
repo_url String
default_branch String @default("main")
visibility String @default("public")
git_credential_id String?
description String?
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
git_credential LiteLLM_CredentialsTable? @relation(fields: [git_credential_id], references: [credential_id], onDelete: SetNull)
agents LiteLLM_ManagedAgentTable[]
@@index([dockerfile_id])
@@index([created_by])
@@index([git_credential_id])
}
// Managed Agents: agent config — model + prompt + tools + template ref.
model LiteLLM_ManagedAgentTable {
agent_id String @id @default(uuid())
agent_name String?
model String
prompt String?
tools Json @default("[]")
template_id String
branch String?
metadata Json? @default("{}")
created_at DateTime @default(now())
created_by String?
team_id String?
organization_id String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
template LiteLLM_ManagedAgentSandboxTemplateTable @relation(fields: [template_id], references: [template_id], onDelete: Cascade)
sessions LiteLLM_ManagedAgentSessionTable[]
@@index([template_id])
@@index([created_by])
@@index([team_id])
@@index([organization_id])
}
// Managed Agents: live Fargate task instance bound to one agent.
model LiteLLM_ManagedAgentSessionTable {
session_id String @id @default(uuid())
agent_id String
status String @default("creating")
task_arn String?
sandbox_url String?
harness_session_id String?
fargate_cluster String?
fargate_task_def_arn String?
virtual_key_hash String?
failure_reason String?
created_at DateTime @default(now())
created_by String?
team_id String?
last_seen_at DateTime?
expires_at DateTime?
stopped_at DateTime?
agent LiteLLM_ManagedAgentTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
@@index([agent_id])
@@index([status])
@@index([created_by])
@@index([last_seen_at])
}

View file

@ -0,0 +1,10 @@
from litellm.proxy.managed_agents_endpoints.endpoints import router
from litellm.proxy.managed_agents_endpoints import (
endpoints_agents,
) # noqa: F401 registers /agents routes
from litellm.proxy.managed_agents_endpoints import (
endpoints_passthrough,
) # noqa: F401 registers passthrough routes
from litellm.proxy.managed_agents_endpoints import (
endpoints_sessions,
) # noqa: F401 registers /sessions routes

View file

@ -0,0 +1,163 @@
"""Config loader for managed_agents.
Reads the `general_settings.managed_agents` yaml block at startup, validates
that each declared dockerfile path exists and is readable, hashes the
dockerfile + context contents, and builds an in-memory ``DOCKERFILE_REGISTRY``.
The template-create endpoint validates incoming ``dockerfile_id`` values
against this registry.
"""
import os
import re
from dataclasses import dataclass
from typing import Dict, List, Optional
from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy.managed_agents_endpoints.fargate.registry import (
compute_dockerfile_hash,
)
from litellm.proxy.managed_agents_endpoints.types import ManagedAgentsConfig
_DOCKERFILE_ID_RE = re.compile(r"^[a-zA-Z0-9_-]+$")
@dataclass(frozen=True)
class DockerfileEntry:
dockerfile_id: str
path: str # absolute path to Dockerfile
context_dir: str # absolute path to context dir (default: dockerfile dir)
container_port: int
content_hash: str # sha256 of dockerfile + context
# Module-level state. Populated by ``initialize`` at proxy startup.
DOCKERFILE_REGISTRY: Dict[str, DockerfileEntry] = {}
MANAGED_AGENTS_CONFIG: Optional[ManagedAgentsConfig] = None
def _validate_dockerfile_id(dockerfile_id: str) -> None:
"""Reject ids with chars that aren't safe for ECR repo / task family names."""
if not _DOCKERFILE_ID_RE.match(dockerfile_id):
raise ValueError(
f"dockerfile_id '{dockerfile_id}' invalid: only [a-zA-Z0-9_-] allowed"
)
def _resolve_path(path: str) -> str:
return os.path.abspath(os.path.expanduser(path))
def load_managed_agents_config(
general_settings: dict,
) -> Optional[ManagedAgentsConfig]:
"""Parse the ``managed_agents`` block.
Returns ``None`` if the block is absent or ``enabled`` is false. Raises
``ValueError`` on invalid shape (forbidden keys, wrong types, etc.).
"""
raw = general_settings.get("managed_agents")
if not raw:
return None
if not isinstance(raw, dict):
raise ValueError(
f"managed_agents config must be a mapping, got {type(raw).__name__}"
)
try:
config = ManagedAgentsConfig(**raw)
except ValidationError as e:
raise ValueError(f"managed_agents config invalid: {e}") from e
if not config.enabled:
return None
return config
def build_dockerfile_registry(
config: ManagedAgentsConfig,
) -> Dict[str, DockerfileEntry]:
"""Validate each dockerfile path exists + is readable, compute hashes.
Raises ``ValueError`` if a ``dockerfile_id`` contains disallowed chars and
``FileNotFoundError`` if any declared path is missing or unreadable.
"""
registry: Dict[str, DockerfileEntry] = {}
for dockerfile_id, dockerfile_cfg in config.dockerfiles.items():
_validate_dockerfile_id(dockerfile_id)
abs_path = _resolve_path(dockerfile_cfg.path)
if not os.path.isfile(abs_path):
raise FileNotFoundError(
f"managed_agents: dockerfile path for id '{dockerfile_id}' "
f"does not exist: {abs_path}"
)
if not os.access(abs_path, os.R_OK):
raise FileNotFoundError(
f"managed_agents: dockerfile path for id '{dockerfile_id}' "
f"is not readable: {abs_path}"
)
context_dir = os.path.dirname(abs_path)
content_hash = compute_dockerfile_hash(abs_path, context_dir)
entry = DockerfileEntry(
dockerfile_id=dockerfile_id,
path=abs_path,
context_dir=context_dir,
container_port=dockerfile_cfg.container_port,
content_hash=content_hash,
)
registry[dockerfile_id] = entry
verbose_proxy_logger.info(
"managed_agents: loaded dockerfile id=%s path=%s hash=%s",
dockerfile_id,
abs_path,
content_hash[:12],
)
return registry
def initialize(general_settings: dict) -> None:
"""Populate module-level ``DOCKERFILE_REGISTRY`` + ``MANAGED_AGENTS_CONFIG``.
Called from ``proxy_server`` startup. No-op if managed_agents is absent or
disabled. Refuses to start (raises) on any path missing — fail-fast at
boot so bad config never makes it into a running proxy.
"""
global DOCKERFILE_REGISTRY, MANAGED_AGENTS_CONFIG
config = load_managed_agents_config(general_settings)
if config is None:
DOCKERFILE_REGISTRY = {}
MANAGED_AGENTS_CONFIG = None
return
registry = build_dockerfile_registry(config)
DOCKERFILE_REGISTRY = registry
MANAGED_AGENTS_CONFIG = config
verbose_proxy_logger.info(
"managed_agents: registry initialized with %d dockerfile(s)",
len(registry),
)
def get_dockerfile(dockerfile_id: str) -> DockerfileEntry:
"""Lookup helper for endpoints. Raises ``KeyError`` if id not in registry."""
if dockerfile_id not in DOCKERFILE_REGISTRY:
raise KeyError(dockerfile_id)
return DOCKERFILE_REGISTRY[dockerfile_id]
def list_dockerfiles() -> List[DockerfileEntry]:
"""Snapshot of registry, sorted by id."""
return [DOCKERFILE_REGISTRY[k] for k in sorted(DOCKERFILE_REGISTRY.keys())]

View file

@ -0,0 +1,291 @@
from typing import List
import boto3
from fastapi import APIRouter, Depends, HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.managed_agents_endpoints import config_loader as _config_loader
from litellm.proxy.managed_agents_endpoints.config_loader import (
get_dockerfile,
list_dockerfiles,
)
from litellm.proxy.managed_agents_endpoints.fargate.build import provision_template
from litellm.proxy.managed_agents_endpoints.git_validation import (
encrypt_and_store_git_token,
validate_repo_branch,
)
from litellm.proxy.managed_agents_endpoints.lifecycle import stop_sessions_for_template
from litellm.proxy.managed_agents_endpoints.types import (
AwsOverrides,
DockerfileOut,
TemplateCreate,
TemplateOut,
)
from litellm.proxy.utils import jsonify_object
router = APIRouter(prefix="/v1/managed_agents", tags=["managed_agents"])
def _template_row_to_out(row) -> TemplateOut:
return TemplateOut(
id=row.template_id,
name=row.template_name,
dockerfile_id=row.dockerfile_id,
container_port=row.container_port,
repo_url=row.repo_url,
default_branch=row.default_branch,
visibility=row.visibility,
image_uri=row.image_uri,
task_def_arn=row.task_def_arn,
build_status=row.build_status,
build_error=row.build_error,
)
def _resolve_region() -> str:
cfg = _config_loader.MANAGED_AGENTS_CONFIG
if cfg is not None and cfg.aws_region:
return cfg.aws_region
return "us-east-1"
def _resolve_aws_overrides() -> AwsOverrides:
cfg = _config_loader.MANAGED_AGENTS_CONFIG
if cfg is not None:
return cfg.aws
return AwsOverrides()
def _require_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(status_code=403, detail="admin role required")
@router.get("/dockerfiles", response_model=List[DockerfileOut])
async def list_dockerfiles_endpoint(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> List[DockerfileOut]:
return [
DockerfileOut(id=entry.dockerfile_id, container_port=entry.container_port)
for entry in list_dockerfiles()
]
@router.post("/sandbox-templates", response_model=TemplateOut)
async def create_sandbox_template(
body: TemplateCreate,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> TemplateOut:
from litellm.proxy.proxy_server import prisma_client
_require_admin(user_api_key_dict)
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
try:
dockerfile_entry = get_dockerfile(body.dockerfile_id)
except KeyError:
available = [entry.dockerfile_id for entry in list_dockerfiles()]
raise HTTPException(
status_code=400,
detail={
"error": f"unknown dockerfile_id '{body.dockerfile_id}'",
"available_dockerfile_ids": available,
},
)
if body.visibility == "private" and not body.git_token:
raise HTTPException(
status_code=400,
detail="visibility=private requires git_token",
)
if body.visibility == "public" and body.git_token:
raise HTTPException(
status_code=400,
detail="visibility=public must not include git_token",
)
validate_repo_branch(body.repo_url, body.default_branch, body.git_token)
git_credential_id = None
if body.git_token:
git_credential_id = await encrypt_and_store_git_token(
prisma_client,
raw_token=body.git_token,
created_by=user_api_key_dict.user_id or "",
)
create_data = jsonify_object(
{
"template_name": body.name,
"dockerfile_id": body.dockerfile_id,
"container_port": dockerfile_entry.container_port,
"repo_url": body.repo_url,
"default_branch": body.default_branch,
"visibility": body.visibility,
"git_credential_id": git_credential_id,
"description": None,
"build_status": "pending",
"created_by": user_api_key_dict.user_id,
"updated_by": user_api_key_dict.user_id,
}
)
row = await prisma_client.db.litellm_managedagentsandboxtemplatetable.create(
data=create_data
)
region = _resolve_region()
aws_overrides = _resolve_aws_overrides()
try:
provisioned = await provision_template(
dockerfile_id=dockerfile_entry.dockerfile_id,
dockerfile_path=dockerfile_entry.path,
context_dir=dockerfile_entry.context_dir,
container_port=dockerfile_entry.container_port,
region=region,
aws_overrides=aws_overrides,
)
except Exception as e:
verbose_proxy_logger.exception(
"managed_agents: provision_template failed for template_id=%s: %s",
row.template_id,
e,
)
failed_row = (
await prisma_client.db.litellm_managedagentsandboxtemplatetable.update(
where={"template_id": row.template_id},
data={
"build_status": "failed",
"build_error": str(e),
"updated_by": user_api_key_dict.user_id,
},
)
)
raise HTTPException(
status_code=500,
detail={
"error": "provision_template failed",
"template": _template_row_to_out(failed_row).model_dump(),
},
)
updated_row = (
await prisma_client.db.litellm_managedagentsandboxtemplatetable.update(
where={"template_id": row.template_id},
data={
"image_uri": provisioned.image_uri,
"task_def_arn": provisioned.task_def_arn,
"image_hash": provisioned.image_hash,
"build_status": "ready",
"updated_by": user_api_key_dict.user_id,
},
)
)
return _template_row_to_out(updated_row)
@router.get("/sandbox-templates", response_model=List[TemplateOut])
async def list_sandbox_templates(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> List[TemplateOut]:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
rows = await prisma_client.db.litellm_managedagentsandboxtemplatetable.find_many()
return [_template_row_to_out(row) for row in rows]
@router.get("/sandbox-templates/{template_id}", response_model=TemplateOut)
async def get_sandbox_template(
template_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> TemplateOut:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsandboxtemplatetable.find_unique(
where={"template_id": template_id}
)
if row is None:
raise HTTPException(
status_code=404, detail=f"template '{template_id}' not found"
)
return _template_row_to_out(row)
@router.delete("/sandbox-templates/{template_id}")
async def delete_sandbox_template(
template_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.proxy_server import prisma_client
_require_admin(user_api_key_dict)
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsandboxtemplatetable.find_unique(
where={"template_id": template_id}
)
if row is None:
raise HTTPException(
status_code=404, detail=f"template '{template_id}' not found"
)
agent_count = await prisma_client.db.litellm_managedagenttable.count(
where={"template_id": template_id}
)
if agent_count > 0:
raise HTTPException(
status_code=409,
detail=(
f"cannot delete template '{template_id}': "
f"{agent_count} agent(s) still reference it"
),
)
region = _resolve_region()
aws_overrides = _resolve_aws_overrides()
cluster = aws_overrides.cluster or "litellm-managed-agents"
try:
await stop_sessions_for_template(
prisma_client=prisma_client,
region=region,
cluster=cluster,
template_id=template_id,
)
except Exception as e:
verbose_proxy_logger.warning(
"managed_agents: stop_sessions_for_template failed for template_id=%s: %s",
template_id,
e,
)
if row.task_def_arn:
try:
ecs = boto3.client("ecs", region_name=region)
ecs.deregister_task_definition(taskDefinition=row.task_def_arn)
except Exception as e:
verbose_proxy_logger.warning(
"managed_agents: deregister_task_definition failed for arn=%s: %s",
row.task_def_arn,
e,
)
await prisma_client.db.litellm_managedagentsandboxtemplatetable.delete(
where={"template_id": template_id}
)
return {"id": template_id, "status": "deleted"}

View file

@ -0,0 +1,89 @@
import asyncio
import json
from fastapi import Depends, HTTPException
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.managed_agents_endpoints.endpoints import router
from litellm.proxy.managed_agents_endpoints.git_validation import (
decrypt_git_token,
validate_repo_branch,
)
from litellm.proxy.managed_agents_endpoints.types import AgentCreate, AgentOut
from litellm.proxy.utils import jsonify_object
def _agent_row_to_out(row) -> AgentOut:
return AgentOut(
id=row.agent_id,
name=row.agent_name,
model=row.model,
template_id=row.template_id,
branch=row.branch,
)
@router.post("/agents", response_model=AgentOut)
async def create_agent(
body: AgentCreate,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> AgentOut:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
template = (
await prisma_client.db.litellm_managedagentsandboxtemplatetable.find_unique(
where={"template_id": body.template_id}
)
)
if template is None:
raise HTTPException(
status_code=404, detail=f"template '{body.template_id}' not found"
)
branch = body.branch or template.default_branch
git_token = await decrypt_git_token(prisma_client, template.git_credential_id)
await asyncio.to_thread(validate_repo_branch, template.repo_url, branch, git_token)
create_data = jsonify_object(
{
"agent_name": body.name,
"model": body.model,
"prompt": body.prompt,
"tools": json.dumps(body.tools),
"branch": branch,
"metadata": {
"litellm_api_key": body.litellm_api_key,
"litellm_api_base": body.litellm_api_base,
},
"created_by": user_api_key_dict.user_id,
"updated_by": user_api_key_dict.user_id,
}
)
create_data["template"] = {"connect": {"template_id": body.template_id}}
row = await prisma_client.db.litellm_managedagenttable.create(data=create_data)
return _agent_row_to_out(row)
@router.get("/agents/{agent_id}", response_model=AgentOut)
async def get_agent(
agent_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> AgentOut:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagenttable.find_unique(
where={"agent_id": agent_id}
)
if row is None:
raise HTTPException(status_code=404, detail=f"agent '{agent_id}' not found")
return _agent_row_to_out(row)

View file

@ -0,0 +1,205 @@
import asyncio
from datetime import datetime, timezone
from typing import Any, Dict
import httpx
from fastapi import Depends, HTTPException, Request
from fastapi.responses import StreamingResponse
from starlette.background import BackgroundTask
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.managed_agents_endpoints.endpoints import router
from litellm.proxy.managed_agents_endpoints.harness_client import (
expand_message,
harness_send_message,
)
from litellm.proxy.managed_agents_endpoints.types import MessageIn
HOP_BY_HOP = {
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailers",
"transfer-encoding",
"upgrade",
}
async def _touch_session(prisma_client, session_id: str) -> None:
try:
await prisma_client.db.litellm_managedagentsessiontable.update(
where={"session_id": session_id},
data={"last_seen_at": datetime.now(timezone.utc)},
)
except Exception:
pass
def _build_http_client() -> httpx.AsyncClient:
return httpx.AsyncClient(
timeout=httpx.Timeout(connect=10, read=None, write=None, pool=10)
)
@router.post("/sessions/{session_id}/message")
async def session_message(
session_id: str,
body: MessageIn,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Dict[str, Any]:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsessiontable.find_unique(
where={"session_id": session_id},
include={"agent": True},
)
if row is None:
raise HTTPException(status_code=404, detail="session not found")
if (
row.status != "ready"
or row.sandbox_url is None
or row.harness_session_id is None
):
raise HTTPException(
status_code=409,
detail=f"session not ready (status={row.status})",
)
parts = expand_message(body.text, body.parts)
client = _build_http_client()
try:
try:
result = await harness_send_message(
row.sandbox_url,
row.harness_session_id,
client,
model=row.agent.model,
parts=parts,
)
except httpx.HTTPStatusError as e:
raise HTTPException(e.response.status_code, e.response.text)
except httpx.HTTPError as e:
raise HTTPException(502, f"upstream error: {e}")
finally:
asyncio.create_task(client.aclose())
asyncio.create_task(_touch_session(prisma_client, session_id))
return result
@router.get("/sessions/{session_id}/events")
async def session_events(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> StreamingResponse:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsessiontable.find_unique(
where={"session_id": session_id}
)
if row is None:
raise HTTPException(status_code=404, detail="session not found")
if row.status != "ready" or row.sandbox_url is None:
raise HTTPException(
status_code=409,
detail=f"session not ready (status={row.status})",
)
client = _build_http_client()
req = client.build_request("GET", f"{row.sandbox_url}/event", timeout=None)
try:
upstream = await client.send(req, stream=True)
except httpx.HTTPError as e:
await client.aclose()
raise HTTPException(502, f"upstream error: {e}")
resp_headers = {
k: v for k, v in upstream.headers.items() if k.lower() not in HOP_BY_HOP
}
async def _close() -> None:
await upstream.aclose()
await client.aclose()
return StreamingResponse(
upstream.aiter_raw(),
status_code=upstream.status_code,
headers=resp_headers,
media_type=upstream.headers.get("content-type", "text/event-stream"),
background=BackgroundTask(_close),
)
@router.api_route(
"/sessions/{session_id}/raw/{path:path}",
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"],
)
async def session_raw_proxy(
session_id: str,
path: str,
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> StreamingResponse:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsessiontable.find_unique(
where={"session_id": session_id}
)
if row is None:
raise HTTPException(status_code=404, detail="session not found")
if row.status != "ready" or row.sandbox_url is None:
raise HTTPException(
status_code=409,
detail=f"session not ready (status={row.status})",
)
target = f"{row.sandbox_url}/{path}"
fwd_headers = {
k: v
for k, v in request.headers.items()
if k.lower() not in HOP_BY_HOP and k.lower() != "host"
}
body = await request.body()
client = _build_http_client()
req = client.build_request(
method=request.method,
url=target,
params=request.query_params,
headers=fwd_headers,
content=body if body else None,
)
try:
upstream = await client.send(req, stream=True)
except httpx.HTTPError as e:
await client.aclose()
raise HTTPException(502, f"upstream error: {e}")
asyncio.create_task(_touch_session(prisma_client, session_id))
resp_headers = {
k: v for k, v in upstream.headers.items() if k.lower() not in HOP_BY_HOP
}
async def _close() -> None:
await upstream.aclose()
await client.aclose()
return StreamingResponse(
upstream.aiter_raw(),
status_code=upstream.status_code,
headers=resp_headers,
background=BackgroundTask(_close),
)

View file

@ -0,0 +1,324 @@
import asyncio
import json
from datetime import datetime, timezone
from typing import Any, Dict, Optional
import httpx
from fastapi import Depends, HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.managed_agents_endpoints import config_loader as _config_loader
from litellm.proxy.managed_agents_endpoints.endpoints import router
from litellm.proxy.managed_agents_endpoints.fargate.bootstrap import (
bootstrap_shared_infra,
)
from litellm.proxy.managed_agents_endpoints.fargate.tasks import (
run_task_sync,
stop_task_sync,
wait_http_ready,
wait_running_get_ip_sync,
)
from litellm.proxy.managed_agents_endpoints.git_validation import decrypt_git_token
from litellm.proxy.managed_agents_endpoints.harness_client import (
harness_create_session,
harness_send_message,
)
from litellm.proxy.managed_agents_endpoints.lifecycle import stop_session_task
from litellm.proxy.managed_agents_endpoints.types import (
AwsOverrides,
SessionCreateIn,
SessionOut,
)
from litellm.proxy.utils import jsonify_object
def _session_row_to_out(row, response: Optional[Dict[str, Any]] = None) -> SessionOut:
return SessionOut(
id=row.session_id,
agent_id=row.agent_id,
sandbox_url=row.sandbox_url,
status=row.status,
task_arn=row.task_arn,
response=response,
)
def _resolve_region() -> str:
cfg = _config_loader.MANAGED_AGENTS_CONFIG
if cfg is not None and cfg.aws_region:
return cfg.aws_region
return "us-east-1"
def _resolve_aws_overrides() -> AwsOverrides:
cfg = _config_loader.MANAGED_AGENTS_CONFIG
if cfg is not None:
return cfg.aws
return AwsOverrides()
def _resolve_cluster(aws_overrides: AwsOverrides) -> str:
return aws_overrides.cluster or "litellm-agents"
def _coerce_metadata(raw: Any) -> Dict[str, Any]:
if raw is None:
return {}
if isinstance(raw, dict):
return raw
if isinstance(raw, str):
try:
parsed = json.loads(raw)
except (TypeError, ValueError):
return {}
if isinstance(parsed, dict):
return parsed
return {}
def _now_utc() -> datetime:
return datetime.now(timezone.utc)
async def _mark_session_failed(
prisma_client: Any, session_id: str, failure_reason: str
) -> None:
try:
await prisma_client.db.litellm_managedagentsessiontable.update(
where={"session_id": session_id},
data={
"status": "failed",
"failure_reason": failure_reason,
"stopped_at": _now_utc(),
},
)
except Exception as e:
verbose_proxy_logger.warning(
f"managed_agents: failed to mark session {session_id} failed: {e}"
)
@router.post("/agents/{agent_id}/session", response_model=SessionOut)
async def create_session(
agent_id: str,
body: Optional[SessionCreateIn] = None,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> SessionOut:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
body = body or SessionCreateIn()
agent = await prisma_client.db.litellm_managedagenttable.find_unique(
where={"agent_id": agent_id}, include={"template": True}
)
if agent is None:
raise HTTPException(status_code=404, detail=f"agent '{agent_id}' not found")
template = getattr(agent, "template", None)
if template is None:
raise HTTPException(
status_code=404,
detail=f"template for agent '{agent_id}' not found",
)
if template.build_status != "ready":
raise HTTPException(
status_code=409,
detail=(
f"template '{template.template_id}' is not ready "
f"(build_status={template.build_status})"
),
)
if not template.task_def_arn:
raise HTTPException(
status_code=409,
detail=f"template '{template.template_id}' has no task_def_arn",
)
region = _resolve_region()
aws_overrides = _resolve_aws_overrides()
cluster = _resolve_cluster(aws_overrides)
metadata = _coerce_metadata(getattr(agent, "metadata", None))
session_create_data = jsonify_object(
{
"agent_id": agent_id,
"status": "creating",
"fargate_cluster": cluster,
"fargate_task_def_arn": template.task_def_arn,
"created_by": user_api_key_dict.user_id,
"team_id": user_api_key_dict.team_id,
}
)
row = await prisma_client.db.litellm_managedagentsessiontable.create(
data=session_create_data
)
session_id = row.session_id
env: Dict[str, str] = {
"LITELLM_API_KEY": metadata.get("litellm_api_key", "") or "",
"LITELLM_API_BASE": metadata.get("litellm_api_base", "") or "",
"LITELLM_DEFAULT_MODEL": agent.model,
"REPO_URL": template.repo_url,
"BRANCH": agent.branch or template.default_branch,
}
git_token = await decrypt_git_token(prisma_client, template.git_credential_id)
if git_token:
env["GIT_TOKEN"] = git_token
if agent.prompt:
env["AGENT_PROMPT"] = agent.prompt
client = httpx.AsyncClient(
timeout=httpx.Timeout(connect=10, read=None, write=None, pool=10)
)
task_arn: Optional[str] = None
try:
infra = await asyncio.to_thread(
bootstrap_shared_infra, region, aws_overrides, template.container_port
)
if not infra.subnet_ids:
raise RuntimeError("bootstrap_shared_infra returned no subnets")
subnet = infra.subnet_ids[0]
security_group = infra.security_group_id
task_arn = await asyncio.to_thread(
run_task_sync,
region=region,
cluster=cluster,
task_def_arn=template.task_def_arn,
container_name="harness",
subnet=subnet,
security_group=security_group,
env=env,
session_id=session_id,
agent_id=agent_id,
)
await prisma_client.db.litellm_managedagentsessiontable.update(
where={"session_id": session_id},
data={"task_arn": task_arn},
)
public_ip = await asyncio.to_thread(
wait_running_get_ip_sync, region, cluster, task_arn, 300
)
sandbox_url = f"http://{public_ip}:{template.container_port}"
await wait_http_ready(sandbox_url, client, timeout=600)
title = (body.title if body else None) or agent.agent_name or "default"
harness_session_id = await harness_create_session(
sandbox_url, client, title=title
)
await prisma_client.db.litellm_managedagentsessiontable.update(
where={"session_id": session_id},
data={
"sandbox_url": sandbox_url,
"harness_session_id": harness_session_id,
"status": "ready",
"last_seen_at": _now_utc(),
},
)
response_body: Optional[Dict[str, Any]] = None
if body and body.initial_prompt:
response_body = await harness_send_message(
sandbox_url,
harness_session_id,
client,
model=agent.model,
parts=[{"type": "text", "text": body.initial_prompt}],
)
return SessionOut(
id=session_id,
agent_id=agent_id,
sandbox_url=sandbox_url,
status="ready",
task_arn=task_arn,
response=response_body,
)
except Exception as e:
verbose_proxy_logger.exception(
"managed_agents: create_session failed for agent=%s session=%s: %s",
agent_id,
session_id,
e,
)
await _mark_session_failed(prisma_client, session_id, str(e))
if task_arn:
try:
await asyncio.to_thread(
stop_task_sync, region, cluster, task_arn, "session create failed"
)
except Exception as stop_err:
verbose_proxy_logger.warning(
f"managed_agents: stop_task after failure raised: {stop_err}"
)
if isinstance(e, HTTPException):
raise
raise HTTPException(status_code=500, detail=f"session create failed: {e}")
finally:
await client.aclose()
@router.get("/sessions/{session_id}", response_model=SessionOut)
async def get_session(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> SessionOut:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsessiontable.find_unique(
where={"session_id": session_id}
)
if row is None:
raise HTTPException(status_code=404, detail=f"session '{session_id}' not found")
return _session_row_to_out(row)
@router.delete("/sessions/{session_id}")
async def delete_session(
session_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
) -> Dict[str, str]:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
row = await prisma_client.db.litellm_managedagentsessiontable.find_unique(
where={"session_id": session_id}
)
if row is None:
raise HTTPException(status_code=404, detail=f"session '{session_id}' not found")
region = _resolve_region()
aws_overrides = _resolve_aws_overrides()
cluster = row.fargate_cluster or _resolve_cluster(aws_overrides)
if row.task_arn:
await stop_session_task(
region=region,
cluster=cluster,
task_arn=row.task_arn,
session_id=session_id,
)
await prisma_client.db.litellm_managedagentsessiontable.update(
where={"session_id": session_id},
data={"status": "dead", "stopped_at": _now_utc()},
)
return {"id": session_id, "status": "dead"}

View file

@ -0,0 +1,326 @@
"""Idempotent bootstrap for shared Fargate infrastructure used by managed agents."""
import json
import time
from dataclasses import dataclass
from typing import Any, Dict, List, Tuple
import boto3
from botocore.exceptions import ClientError
from litellm._logging import verbose_proxy_logger
from litellm.proxy.managed_agents_endpoints.fargate.tasks import _ec2, _ecs
from litellm.proxy.managed_agents_endpoints.types import AwsOverrides
DEFAULT_CLUSTER_NAME = "litellm-agents"
DEFAULT_TASK_EXEC_ROLE_NAME = "litellm-agents-task-exec"
DEFAULT_SECURITY_GROUP_NAME = "litellm-agents-sg"
DEFAULT_LOG_GROUP_NAME = "/ecs/litellm-agents"
TASK_EXEC_ROLE_POLICY_ARN = (
"arn:aws:iam::aws:policy/service-role/AmazonECSTaskExecutionRolePolicy"
)
_TASK_EXEC_TRUST_POLICY = {
"Version": "2012-10-17",
"Statement": [
{
"Effect": "Allow",
"Principal": {"Service": "ecs-tasks.amazonaws.com"},
"Action": "sts:AssumeRole",
}
],
}
_bootstrap_clients: Dict[str, Any] = {}
def _iam():
if "iam" not in _bootstrap_clients:
_bootstrap_clients["iam"] = boto3.client("iam")
return _bootstrap_clients["iam"]
def _logs(region: str):
key = f"logs:{region}"
if key not in _bootstrap_clients:
_bootstrap_clients[key] = boto3.client("logs", region_name=region)
return _bootstrap_clients[key]
@dataclass(frozen=True)
class SharedInfra:
cluster_arn: str
task_exec_role_arn: str
security_group_id: str
log_group_name: str
vpc_id: str
subnet_ids: List[str]
def ensure_cluster(region: str, cluster_name: str) -> str:
ecs = _ecs(region)
r = ecs.describe_clusters(clusters=[cluster_name])
clusters = r.get("clusters", [])
if clusters and clusters[0].get("status") == "ACTIVE":
verbose_proxy_logger.debug(
f"ECS cluster {cluster_name} already ACTIVE in {region}"
)
return clusters[0]["clusterArn"]
verbose_proxy_logger.info(f"Creating ECS cluster {cluster_name} in {region}")
created = ecs.create_cluster(clusterName=cluster_name)
arn = created.get("cluster", {}).get("clusterArn")
if arn:
return arn
r2 = ecs.describe_clusters(clusters=[cluster_name])
return r2["clusters"][0]["clusterArn"]
def ensure_task_exec_role(role_name: str) -> str:
iam = _iam()
try:
arn = iam.get_role(RoleName=role_name)["Role"]["Arn"]
verbose_proxy_logger.debug(f"IAM role {role_name} already exists")
return arn
except ClientError as e:
if e.response.get("Error", {}).get("Code") != "NoSuchEntity":
raise
verbose_proxy_logger.info(f"Creating IAM role {role_name}")
arn = iam.create_role(
RoleName=role_name,
AssumeRolePolicyDocument=json.dumps(_TASK_EXEC_TRUST_POLICY),
Description="ECS task execution role for litellm managed agents",
)["Role"]["Arn"]
iam.attach_role_policy(
RoleName=role_name,
PolicyArn=TASK_EXEC_ROLE_POLICY_ARN,
)
time.sleep(10)
return arn
def discover_vpc_subnet(region: str) -> Tuple[str, str, str]:
ec2 = _ec2(region)
r = ec2.describe_vpcs(Filters=[{"Name": "is-default", "Values": ["true"]}])
vpcs = r.get("Vpcs", [])
if not vpcs:
raise RuntimeError(f"no default VPC in region {region}")
vpc_id = vpcs[0]["VpcId"]
vpc_cidr = vpcs[0]["CidrBlock"]
s = ec2.describe_subnets(
Filters=[
{"Name": "vpc-id", "Values": [vpc_id]},
{"Name": "map-public-ip-on-launch", "Values": ["true"]},
]
)
subnets = s.get("Subnets", [])
if not subnets:
raise RuntimeError(f"no public subnet in VPC {vpc_id} ({region})")
verbose_proxy_logger.debug(
f"Discovered default VPC {vpc_id} cidr={vpc_cidr} in {region}"
)
return vpc_id, vpc_cidr, subnets[0]["SubnetId"]
def ensure_security_group(
region: str,
sg_name: str,
vpc_id: str,
vpc_cidr: str,
container_port: int,
) -> str:
ec2 = _ec2(region)
r = ec2.describe_security_groups(
Filters=[
{"Name": "vpc-id", "Values": [vpc_id]},
{"Name": "group-name", "Values": [sg_name]},
]
)
existing = r.get("SecurityGroups", [])
if existing:
verbose_proxy_logger.debug(
f"Security group {sg_name} already exists in VPC {vpc_id}"
)
return existing[0]["GroupId"]
verbose_proxy_logger.info(f"Creating security group {sg_name} in VPC {vpc_id}")
sg_id = ec2.create_security_group(
GroupName=sg_name,
Description="litellm managed agents sandbox",
VpcId=vpc_id,
)["GroupId"]
ec2.authorize_security_group_ingress(
GroupId=sg_id,
IpPermissions=[
{
"IpProtocol": "tcp",
"FromPort": container_port,
"ToPort": container_port,
"IpRanges": [{"CidrIp": "0.0.0.0/0"}],
}
],
)
# Revoke default allow-all egress so we can install a restricted set.
try:
ec2.revoke_security_group_egress(
GroupId=sg_id,
IpPermissions=[{"IpProtocol": "-1", "IpRanges": [{"CidrIp": "0.0.0.0/0"}]}],
)
except ClientError as e:
verbose_proxy_logger.debug(f"revoke default egress on {sg_id} no-op: {e}")
ec2.authorize_security_group_egress(
GroupId=sg_id,
IpPermissions=[
{
"IpProtocol": "tcp",
"FromPort": 443,
"ToPort": 443,
"IpRanges": [{"CidrIp": "0.0.0.0/0", "Description": "HTTPS"}],
},
{
"IpProtocol": "udp",
"FromPort": 53,
"ToPort": 53,
"IpRanges": [
{"CidrIp": vpc_cidr, "Description": "DNS to VPC resolver"}
],
},
{
"IpProtocol": "tcp",
"FromPort": 53,
"ToPort": 53,
"IpRanges": [{"CidrIp": vpc_cidr, "Description": "DNS TCP fallback"}],
},
],
)
return sg_id
def ensure_log_group(region: str, log_group_name: str) -> None:
logs = _logs(region)
try:
logs.create_log_group(logGroupName=log_group_name)
verbose_proxy_logger.info(
f"Created CloudWatch log group {log_group_name} in {region}"
)
except ClientError as e:
if e.response.get("Error", {}).get("Code") == "ResourceAlreadyExistsException":
verbose_proxy_logger.debug(
f"Log group {log_group_name} already exists in {region}"
)
return
raise
def _validate_existing_cluster(region: str, cluster_name_or_arn: str) -> str:
ecs = _ecs(region)
r = ecs.describe_clusters(clusters=[cluster_name_or_arn])
clusters = r.get("clusters", [])
if not clusters or clusters[0].get("status") != "ACTIVE":
raise RuntimeError(
f"override cluster {cluster_name_or_arn} not ACTIVE in {region}"
)
return clusters[0]["clusterArn"]
def _validate_existing_role(role_arn_or_name: str) -> str:
iam = _iam()
if role_arn_or_name.startswith("arn:"):
role_name = role_arn_or_name.rsplit("/", 1)[-1]
else:
role_name = role_arn_or_name
return iam.get_role(RoleName=role_name)["Role"]["Arn"]
def _validate_existing_security_group(region: str, sg_id: str, vpc_id: str) -> str:
ec2 = _ec2(region)
r = ec2.describe_security_groups(GroupIds=[sg_id])
groups = r.get("SecurityGroups", [])
if not groups:
raise RuntimeError(f"override security_group {sg_id} not found in {region}")
if groups[0].get("VpcId") != vpc_id:
raise RuntimeError(
f"override security_group {sg_id} is in VPC {groups[0].get('VpcId')}, "
f"expected {vpc_id}"
)
return groups[0]["GroupId"]
def _validate_existing_subnets(
region: str, subnet_ids: List[str]
) -> Tuple[str, List[str]]:
ec2 = _ec2(region)
r = ec2.describe_subnets(SubnetIds=subnet_ids)
subnets = r.get("Subnets", [])
if len(subnets) != len(subnet_ids):
found = {s["SubnetId"] for s in subnets}
missing = [s for s in subnet_ids if s not in found]
raise RuntimeError(f"override subnets not found: {missing}")
vpc_ids = {s["VpcId"] for s in subnets}
if len(vpc_ids) != 1:
raise RuntimeError(f"override subnets span multiple VPCs: {vpc_ids}")
return vpc_ids.pop(), [s["SubnetId"] for s in subnets]
def _validate_existing_log_group(region: str, log_group_name: str) -> None:
logs = _logs(region)
r = logs.describe_log_groups(logGroupNamePrefix=log_group_name)
for g in r.get("logGroups", []):
if g.get("logGroupName") == log_group_name:
return
raise RuntimeError(f"override log_group {log_group_name} not found in {region}")
def bootstrap_shared_infra(
region: str,
overrides: AwsOverrides,
container_port: int = 4096,
) -> SharedInfra:
if overrides.cluster:
cluster_arn = _validate_existing_cluster(region, overrides.cluster)
else:
cluster_arn = ensure_cluster(region, DEFAULT_CLUSTER_NAME)
if overrides.task_exec_role_arn:
task_exec_role_arn = _validate_existing_role(overrides.task_exec_role_arn)
else:
task_exec_role_arn = ensure_task_exec_role(DEFAULT_TASK_EXEC_ROLE_NAME)
if overrides.subnets:
vpc_id, subnet_ids = _validate_existing_subnets(region, overrides.subnets)
vpc_cidr_resp = _ec2(region).describe_vpcs(VpcIds=[vpc_id])
vpc_cidr = vpc_cidr_resp["Vpcs"][0]["CidrBlock"]
else:
vpc_id, vpc_cidr, subnet = discover_vpc_subnet(region)
subnet_ids = [subnet]
if overrides.security_group:
security_group_id = _validate_existing_security_group(
region, overrides.security_group, vpc_id
)
else:
security_group_id = ensure_security_group(
region, DEFAULT_SECURITY_GROUP_NAME, vpc_id, vpc_cidr, container_port
)
if overrides.log_group:
_validate_existing_log_group(region, overrides.log_group)
log_group_name = overrides.log_group
else:
ensure_log_group(region, DEFAULT_LOG_GROUP_NAME)
log_group_name = DEFAULT_LOG_GROUP_NAME
return SharedInfra(
cluster_arn=cluster_arn,
task_exec_role_arn=task_exec_role_arn,
security_group_id=security_group_id,
log_group_name=log_group_name,
vpc_id=vpc_id,
subnet_ids=subnet_ids,
)

View file

@ -0,0 +1,167 @@
"""Top-level orchestrator: bootstrap shared infra, build/push image, register task def."""
import asyncio
import os
import re
from dataclasses import dataclass
from typing import Callable, Dict, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.managed_agents_endpoints.fargate.bootstrap import (
SharedInfra,
bootstrap_shared_infra,
)
from litellm.proxy.managed_agents_endpoints.fargate.registry import (
build_and_push,
compute_dockerfile_hash,
)
from litellm.proxy.managed_agents_endpoints.fargate.tasks import _ecs as _ecs_client
from litellm.proxy.managed_agents_endpoints.types import AwsOverrides
_DOCKERFILE_ID_SANITIZE_RE = re.compile(r"[^a-z0-9_-]")
@dataclass(frozen=True)
class ProvisionedTemplate:
image_uri: str
task_def_arn: str
image_hash: str
container_port: int
cluster_arn: str
security_group_id: str
subnet_ids: List[str]
_locks: Dict[str, asyncio.Lock] = {}
_locks_guard = asyncio.Lock()
async def _get_lock(dockerfile_id: str) -> asyncio.Lock:
async with _locks_guard:
if dockerfile_id not in _locks:
_locks[dockerfile_id] = asyncio.Lock()
return _locks[dockerfile_id]
def _sanitize_dockerfile_id(dockerfile_id: str) -> str:
return _DOCKERFILE_ID_SANITIZE_RE.sub("-", dockerfile_id.lower())
def _register_task_definition(
*,
region: str,
family: str,
image_uri: str,
container_port: int,
shared_infra: SharedInfra,
) -> str:
ecs = _ecs_client(region)
r = ecs.register_task_definition(
family=family,
networkMode="awsvpc",
requiresCompatibilities=["FARGATE"],
cpu="512",
memory="1024",
executionRoleArn=shared_infra.task_exec_role_arn,
runtimePlatform={
"cpuArchitecture": "X86_64",
"operatingSystemFamily": "LINUX",
},
containerDefinitions=[
{
"name": "harness",
"image": image_uri,
"essential": True,
"portMappings": [{"containerPort": container_port, "protocol": "tcp"}],
"logConfiguration": {
"logDriver": "awslogs",
"options": {
"awslogs-group": shared_infra.log_group_name,
"awslogs-region": region,
"awslogs-stream-prefix": "harness",
},
},
}
],
)
return r["taskDefinition"]["taskDefinitionArn"]
async def provision_template(
*,
dockerfile_id: str,
dockerfile_path: str,
context_dir: Optional[str],
container_port: int,
region: str,
aws_overrides: AwsOverrides,
log_callback: Optional[Callable[[str], None]] = None,
) -> ProvisionedTemplate:
"""Bootstrap shared infra → build/push image → register task def. Idempotent."""
lock = await _get_lock(dockerfile_id)
verbose_proxy_logger.debug(
f"provision_template waiting for lock dockerfile_id={dockerfile_id}"
)
async with lock:
verbose_proxy_logger.debug(
f"provision_template acquired lock dockerfile_id={dockerfile_id}"
)
try:
ctx_dir = (
context_dir
if context_dir
else os.path.dirname(os.path.abspath(dockerfile_path))
)
image_hash = await asyncio.to_thread(
compute_dockerfile_hash, dockerfile_path, ctx_dir
)
shared_infra = await asyncio.to_thread(
bootstrap_shared_infra, region, aws_overrides, container_port
)
sanitized = _sanitize_dockerfile_id(dockerfile_id)
repo_name = f"litellm-agents-{sanitized}"
family = f"litellm-agents-{sanitized}"
verbose_proxy_logger.info(
f"provision_template build start dockerfile_id={dockerfile_id} "
f"hash={image_hash[:12]} repo={repo_name}"
)
image_uri = await asyncio.to_thread(
build_and_push,
region=region,
repo_name=repo_name,
dockerfile_path=dockerfile_path,
context_dir=ctx_dir,
content_hash=image_hash,
log_callback=log_callback,
)
task_def_arn = await asyncio.to_thread(
_register_task_definition,
region=region,
family=family,
image_uri=image_uri,
container_port=container_port,
shared_infra=shared_infra,
)
verbose_proxy_logger.info(
f"provision_template register-task-def complete family={family} "
f"task_def_arn={task_def_arn}"
)
return ProvisionedTemplate(
image_uri=image_uri,
task_def_arn=task_def_arn,
image_hash=image_hash,
container_port=container_port,
cluster_arn=shared_infra.cluster_arn,
security_group_id=shared_infra.security_group_id,
subnet_ids=shared_infra.subnet_ids,
)
finally:
verbose_proxy_logger.debug(
f"provision_template released lock dockerfile_id={dockerfile_id}"
)

View file

@ -0,0 +1,267 @@
"""ECR repo lifecycle + docker build/push mechanics for managed agent images."""
import base64
import hashlib
import os
import subprocess
from typing import Any, Callable, Dict, Optional
import boto3
from botocore.exceptions import ClientError
from litellm._logging import verbose_proxy_logger
_BUILD_TIMEOUT_SECONDS = 30 * 60
_PUSH_TIMEOUT_SECONDS = 30 * 60
_clients: Dict[str, Any] = {}
def _ecr(region: str):
key = f"ecr:{region}"
if key not in _clients:
_clients[key] = boto3.client("ecr", region_name=region)
return _clients[key]
def ensure_ecr_repo(region: str, repo_name: str) -> str:
client = _ecr(region)
try:
r = client.describe_repositories(repositoryNames=[repo_name])
return r["repositories"][0]["repositoryUri"]
except ClientError as e:
if e.response.get("Error", {}).get("Code") != "RepositoryNotFoundException":
raise
try:
r = client.create_repository(
repositoryName=repo_name, imageScanningConfiguration={"scanOnPush": True}
)
return r["repository"]["repositoryUri"]
except ClientError as e:
if (
e.response.get("Error", {}).get("Code")
!= "RepositoryAlreadyExistsException"
):
raise
r = client.describe_repositories(repositoryNames=[repo_name])
return r["repositories"][0]["repositoryUri"]
def image_exists(region: str, repo_name: str, tag: str) -> bool:
client = _ecr(region)
try:
client.describe_images(repositoryName=repo_name, imageIds=[{"imageTag": tag}])
return True
except ClientError as e:
if e.response.get("Error", {}).get("Code") == "ImageNotFoundException":
return False
raise
def compute_dockerfile_hash(dockerfile_path: str, context_dir: Optional[str]) -> str:
h = hashlib.sha256()
with open(dockerfile_path, "rb") as f:
h.update(b"dockerfile:")
h.update(f.read())
ctx = (
context_dir
if context_dir is not None
else os.path.dirname(os.path.abspath(dockerfile_path))
)
if ctx and os.path.isdir(ctx):
entries = []
for root, dirs, files in os.walk(ctx):
dirs.sort()
for fname in files:
full = os.path.join(root, fname)
rel = os.path.relpath(full, ctx)
entries.append((rel, full))
entries.sort(key=lambda e: e[0])
for rel, full in entries:
h.update(b"\x00path:")
h.update(rel.encode("utf-8"))
try:
with open(full, "rb") as f:
while True:
chunk = f.read(1024 * 1024)
if not chunk:
break
h.update(chunk)
except OSError:
continue
return h.hexdigest()
def docker_login(region: str) -> None:
client = _ecr(region)
r = client.get_authorization_token()
auth = r["authorizationData"][0]
token = auth["authorizationToken"]
registry = auth["proxyEndpoint"]
decoded = base64.b64decode(token).decode("utf-8")
user, _, password = decoded.partition(":")
if not user or not password:
raise RuntimeError("ECR get_authorization_token returned malformed credentials")
try:
result = subprocess.run(
["docker", "login", "-u", user, "--password-stdin", registry],
input=password,
text=True,
capture_output=True,
check=False,
timeout=120,
)
except FileNotFoundError:
raise RuntimeError("docker daemon required on proxy host")
if result.returncode != 0:
tail = (result.stderr or result.stdout or "").strip()[-2000:]
raise RuntimeError(f"docker login failed: {tail}")
def _stream_subprocess(
cmd: list,
*,
timeout: int,
log_callback: Optional[Callable[[str], None]],
op_name: str,
) -> None:
try:
proc = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
bufsize=1,
)
except FileNotFoundError:
raise RuntimeError("docker daemon required on proxy host")
stderr_tail: list = []
stdout_lines: list = []
try:
if proc.stdout is not None:
for line in proc.stdout:
line = line.rstrip("\n")
stdout_lines.append(line)
if log_callback is not None:
try:
log_callback(line)
except Exception:
pass
proc.wait(timeout=timeout)
except subprocess.TimeoutExpired:
proc.kill()
try:
proc.wait(timeout=10)
except subprocess.TimeoutExpired:
pass
raise RuntimeError(f"{op_name} timed out after {timeout}s")
if proc.stderr is not None:
try:
stderr_tail.append(proc.stderr.read())
except Exception:
pass
if proc.returncode != 0:
tail = "".join(stderr_tail).strip()
if not tail:
tail = "\n".join(stdout_lines[-50:])
tail = tail[-4000:]
raise RuntimeError(f"{op_name} failed (exit {proc.returncode}): {tail}")
def docker_build(
dockerfile_path: str,
context_dir: str,
image_uri: str,
*,
log_callback: Optional[Callable[[str], None]] = None,
) -> None:
if not os.path.isfile(dockerfile_path):
raise RuntimeError(f"dockerfile not found: {dockerfile_path}")
ctx = (
context_dir
if context_dir
else os.path.dirname(os.path.abspath(dockerfile_path))
)
if not os.path.isdir(ctx):
raise RuntimeError(f"context dir not found: {ctx}")
cmd = [
"docker",
"build",
"--platform",
"linux/amd64",
"-f",
dockerfile_path,
"-t",
image_uri,
ctx,
]
_stream_subprocess(
cmd,
timeout=_BUILD_TIMEOUT_SECONDS,
log_callback=log_callback,
op_name="docker build",
)
def docker_push(
image_uri: str,
*,
log_callback: Optional[Callable[[str], None]] = None,
) -> None:
cmd = ["docker", "push", image_uri]
_stream_subprocess(
cmd,
timeout=_PUSH_TIMEOUT_SECONDS,
log_callback=log_callback,
op_name="docker push",
)
def build_and_push(
*,
region: str,
repo_name: str,
dockerfile_path: str,
context_dir: str,
content_hash: str,
log_callback: Optional[Callable[[str], None]] = None,
) -> str:
if not os.path.isfile(dockerfile_path):
raise RuntimeError(f"dockerfile not found: {dockerfile_path}")
ctx = (
context_dir
if context_dir
else os.path.dirname(os.path.abspath(dockerfile_path))
)
if not os.path.isdir(ctx):
raise RuntimeError(f"context dir not found: {ctx}")
repo_uri = ensure_ecr_repo(region, repo_name)
tag = content_hash
image_uri = f"{repo_uri}:{tag}"
if image_exists(region, repo_name, tag):
verbose_proxy_logger.info(
f"ECR cache hit for {repo_name}:{tag} — skipping build"
)
return image_uri
verbose_proxy_logger.info(
f"Building image {repo_name}:{tag} from {dockerfile_path} (context={ctx})"
)
docker_login(region)
docker_build(dockerfile_path, ctx, image_uri, log_callback=log_callback)
docker_push(image_uri, log_callback=log_callback)
verbose_proxy_logger.info(f"Pushed image {image_uri}")
return image_uri

View file

@ -0,0 +1,171 @@
"""ECS Fargate task lifecycle for managed agent sandboxes."""
import asyncio
import time
from typing import Any, Dict, List, Optional
import boto3
import httpx
from botocore.exceptions import ClientError
from litellm._logging import verbose_proxy_logger
TAG_SESSION_ID = "litellm:managed_agent_session_id"
TAG_AGENT_ID = "litellm:managed_agent_id"
_clients: Dict[str, Any] = {}
def _ecs(region: str):
key = f"ecs:{region}"
if key not in _clients:
_clients[key] = boto3.client("ecs", region_name=region)
return _clients[key]
def _ec2(region: str):
key = f"ec2:{region}"
if key not in _clients:
_clients[key] = boto3.client("ec2", region_name=region)
return _clients[key]
def run_task_sync(
*,
region: str,
cluster: str,
task_def_arn: str,
container_name: str,
subnet: str,
security_group: str,
env: Dict[str, str],
session_id: str,
agent_id: str,
) -> str:
"""Launch Fargate task. Tags task w/ session_id + agent_id for orphan reconciliation."""
overrides = {
"containerOverrides": [
{
"name": container_name,
"environment": [{"name": k, "value": v} for k, v in env.items()],
}
]
}
r = _ecs(region).run_task(
cluster=cluster,
taskDefinition=task_def_arn,
launchType="FARGATE",
networkConfiguration={
"awsvpcConfiguration": {
"subnets": [subnet],
"securityGroups": [security_group],
"assignPublicIp": "ENABLED",
}
},
overrides=overrides,
tags=[
{"key": TAG_SESSION_ID, "value": session_id},
{"key": TAG_AGENT_ID, "value": agent_id},
],
propagateTags="TASK_DEFINITION",
count=1,
)
if r.get("failures"):
raise RuntimeError(f"run_task failures: {r['failures']}")
return r["tasks"][0]["taskArn"]
def stop_task_sync(
region: str, cluster: str, task_arn: str, reason: str = "litellm cleanup"
) -> None:
"""Best-effort stop. Idempotent — swallows ClientError on already-stopped tasks."""
try:
_ecs(region).stop_task(cluster=cluster, task=task_arn, reason=reason)
except ClientError as e:
verbose_proxy_logger.warning(f"stop_task failed for {task_arn}: {e}")
def wait_running_get_ip_sync(
region: str, cluster: str, task_arn: str, timeout: int = 300
) -> str:
deadline = time.time() + timeout
ecs_client = _ecs(region)
ec2_client = _ec2(region)
while time.time() < deadline:
d = ecs_client.describe_tasks(cluster=cluster, tasks=[task_arn])["tasks"][0]
st = d["lastStatus"]
if st == "STOPPED":
reasons = [c.get("reason") for c in d.get("containers", [])]
raise RuntimeError(
f"task stopped: {d.get('stoppedReason')} containers={reasons}"
)
if st == "RUNNING":
for att in d.get("attachments", []):
eni = next(
(
kv["value"]
for kv in att.get("details", [])
if kv["name"] == "networkInterfaceId"
),
None,
)
if eni:
ni = ec2_client.describe_network_interfaces(
NetworkInterfaceIds=[eni]
)["NetworkInterfaces"][0]
ip = ni.get("Association", {}).get("PublicIp")
if ip:
return ip
time.sleep(3)
raise TimeoutError("task never reached RUNNING with public IP")
async def wait_http_ready(
url: str, client: httpx.AsyncClient, timeout: int = 240
) -> None:
deadline = time.time() + timeout
last: Optional[Exception] = None
while time.time() < deadline:
try:
r = await client.get(url, timeout=3)
if r.status_code < 500:
return
except (httpx.HTTPError, OSError) as e:
last = e
await asyncio.sleep(2)
raise TimeoutError(f"sandbox never ready at {url}: {last}")
def list_tagged_task_arns(region: str, cluster: str) -> List[str]:
"""All task ARNs in cluster regardless of status. Reconciler filters by tag."""
ecs_client = _ecs(region)
arns: List[str] = []
for status in ("RUNNING", "PENDING"):
paginator = ecs_client.get_paginator("list_tasks")
for page in paginator.paginate(cluster=cluster, desiredStatus=status):
arns.extend(page.get("taskArns", []))
return arns
def describe_tasks_with_tags(
region: str, cluster: str, task_arns: List[str]
) -> List[Dict[str, Any]]:
"""Returns list of {taskArn, tags: {key: value}} for arns. Batches in 100s (ECS API limit)."""
if not task_arns:
return []
ecs_client = _ecs(region)
out: List[Dict[str, Any]] = []
for i in range(0, len(task_arns), 100):
batch = task_arns[i : i + 100]
r = ecs_client.describe_tasks(cluster=cluster, tasks=batch, include=["TAGS"])
for t in r.get("tasks", []):
tags = {tag["key"]: tag["value"] for tag in t.get("tags", [])}
out.append(
{
"taskArn": t["taskArn"],
"tags": tags,
"lastStatus": t.get("lastStatus"),
}
)
return out

View file

@ -0,0 +1,133 @@
"""Git repo + branch validation utilities for managed-agent template create."""
import os
import subprocess
import uuid
from typing import Any, Optional
from urllib.parse import urlparse, urlunparse
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
def authed_repo_url(repo_url: str, git_token: Optional[str]) -> str:
if not git_token:
return repo_url
parsed = urlparse(repo_url)
if parsed.scheme != "https":
return repo_url
if not parsed.hostname:
return repo_url
netloc = f"x-access-token:{git_token}@{parsed.hostname}"
if parsed.port:
netloc += f":{parsed.port}"
return urlunparse(
(
parsed.scheme,
netloc,
parsed.path,
parsed.params,
parsed.query,
parsed.fragment,
)
)
def validate_repo_branch(
repo_url: str, branch: str, git_token: Optional[str] = None
) -> None:
env = {**os.environ, "GIT_TERMINAL_PROMPT": "0"}
url = authed_repo_url(repo_url, git_token)
try:
result = subprocess.run(
["git", "ls-remote", "--heads", "--tags", url, branch],
capture_output=True,
text=True,
timeout=15,
check=False,
env=env,
)
except FileNotFoundError as e:
raise HTTPException(
status_code=500, detail=f"git not installed on proxy host: {e}"
)
except subprocess.TimeoutExpired:
raise HTTPException(status_code=400, detail=f"timed out reaching {repo_url}")
if result.returncode != 0:
msg_lines = (result.stderr or result.stdout).strip().splitlines()
tail = msg_lines[-1] if msg_lines else "unknown error"
# Scrub authed URL so token is not echoed back to clients/logs.
tail = tail.replace(url, repo_url)
raise HTTPException(
status_code=400,
detail=f"git ls-remote failed for {repo_url}: {tail}",
)
if not result.stdout.strip():
raise HTTPException(
status_code=400,
detail=f"branch or tag '{branch}' not found in {repo_url}",
)
async def decrypt_git_token(
prisma_client: Any, credential_id: Optional[str]
) -> Optional[str]:
if credential_id is None:
return None
if prisma_client is None:
return None
row = await prisma_client.db.litellm_credentialstable.find_unique(
where={"credential_id": credential_id}
)
if row is None:
return None
credential_values = getattr(row, "credential_values", None)
if credential_values is None and isinstance(row, dict):
credential_values = row.get("credential_values")
if not isinstance(credential_values, dict):
return None
encrypted = credential_values.get("git_token")
if not encrypted:
return None
return decrypt_value_helper(encrypted, key="git_token")
async def encrypt_and_store_git_token(
prisma_client: Any, *, raw_token: str, created_by: str
) -> str:
if prisma_client is None:
raise HTTPException(status_code=500, detail="prisma client not available")
credential_name = f"managed-agent-git-token-{uuid.uuid4()}"
encrypted_token = encrypt_value_helper(raw_token)
created = await prisma_client.db.litellm_credentialstable.create(
data={
"credential_name": credential_name,
"credential_values": {"git_token": encrypted_token},
"credential_info": {"source": "managed_agent_template"},
"created_by": created_by,
"updated_by": created_by,
}
)
credential_id = getattr(created, "credential_id", None)
if credential_id is None and isinstance(created, dict):
credential_id = created.get("credential_id")
if not credential_id:
raise HTTPException(
status_code=500, detail="failed to persist git token credential"
)
verbose_proxy_logger.debug(
"Stored managed-agent git token credential credential_id=%s credential_name=%s",
credential_id,
credential_name,
)
return credential_id

View file

@ -0,0 +1,65 @@
from typing import Any, Dict, List, Optional
import httpx
from fastapi import HTTPException
async def harness_create_session(
sandbox_url: str,
client: httpx.AsyncClient,
*,
title: str = "default",
timeout: int = 30,
) -> str:
"""POST {sandbox_url}/session w/ {"title": ...}. Returns session id.
Handle response shape: bare object OR single-element array (proto comment)."""
r = await client.post(
f"{sandbox_url}/session",
json={"title": title},
timeout=timeout,
)
r.raise_for_status()
data = r.json()
if isinstance(data, list):
if not data:
raise RuntimeError(f"unexpected harness session response: {data}")
data = data[0]
if not isinstance(data, dict) or "id" not in data:
raise RuntimeError(f"unexpected harness session response: {data}")
return data["id"]
async def harness_send_message(
sandbox_url: str,
harness_session_id: str,
client: httpx.AsyncClient,
*,
model: str,
parts: List[Dict[str, Any]],
timeout: int = 240,
) -> Dict[str, Any]:
"""POST {sandbox_url}/session/{id}/message w/ {model:{providerID:'litellm',modelID:model}, parts:...}.
Returns full response JSON."""
body = {
"model": {"providerID": "litellm", "modelID": model},
"parts": parts,
}
r = await client.post(
f"{sandbox_url}/session/{harness_session_id}/message",
json=body,
timeout=timeout,
)
r.raise_for_status()
return r.json()
def expand_message(
text: Optional[str],
parts: Optional[List[Dict[str, Any]]],
) -> List[Dict[str, Any]]:
"""Coerce {text} or {parts} into harness parts list. Raises HTTPException(400) if both missing."""
if parts is not None:
return parts
if text is not None:
return [{"type": "text", "text": text}]
raise HTTPException(400, "message body must include 'text' or 'parts'")

View file

@ -0,0 +1,205 @@
"""
Session lifecycle for managed agents.
Two cleanup paths:
1. Pre-delete (handler-driven): before deleting agent or template, stop all live
Fargate tasks for affected sessions, then mark session rows 'dead'. DB cascade
removes session rows after agent/template delete completes.
2. Reconciler (background sweep): every 60s, list all tagged Fargate tasks in the
configured cluster. For each running task, look up its session_id tag in DB.
Orphan = task running but session row missing OR status in {'dead', 'failed',
'stopped'} OR creating > 10min. Stop orphan task. Inverse direction handled
by pre-delete path.
Reconciler runs at proxy startup (catch crashes mid-spawn) and on a 60s interval.
"""
import asyncio
import time
from typing import Any, Dict, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.managed_agents_endpoints.fargate.tasks import (
TAG_SESSION_ID,
describe_tasks_with_tags,
list_tagged_task_arns,
stop_task_sync,
)
ALIVE_STATUSES = ("creating", "ready")
DEAD_STATUSES = ("dead", "failed", "stopped")
CREATING_TIMEOUT_SECONDS = 600
RECONCILE_INTERVAL_SECONDS = 60
async def stop_session_task(
*, region: str, cluster: str, task_arn: Optional[str], session_id: str
) -> None:
"""Stop a single session's Fargate task. Idempotent."""
if not task_arn:
return
try:
await asyncio.to_thread(
stop_task_sync, region, cluster, task_arn, f"session {session_id} delete"
)
except Exception as e:
verbose_proxy_logger.warning(
f"stop_session_task failed (session={session_id}, arn={task_arn}): {e}"
)
async def stop_sessions_for_agent(
*, prisma_client: Any, region: str, cluster: str, agent_id: str
) -> int:
"""Stop all live Fargate tasks for an agent's sessions. Returns count stopped.
Caller deletes agent row after this. DB CASCADE removes session rows.
"""
sessions = await prisma_client.db.litellm_managedagentsessiontable.find_many(
where={"agent_id": agent_id, "status": {"in": list(ALIVE_STATUSES)}}
)
if not sessions:
return 0
await asyncio.gather(
*[
stop_session_task(
region=region,
cluster=cluster,
task_arn=s.task_arn,
session_id=s.session_id,
)
for s in sessions
],
return_exceptions=True,
)
now = time.time()
await prisma_client.db.litellm_managedagentsessiontable.update_many(
where={"agent_id": agent_id, "status": {"in": list(ALIVE_STATUSES)}},
data={"status": "stopped", "stopped_at": _datetime_from_ts(now)},
)
return len(sessions)
async def stop_sessions_for_template(
*, prisma_client: Any, region: str, cluster: str, template_id: str
) -> int:
"""Stop all live Fargate tasks for sessions of all agents under a template."""
agents = await prisma_client.db.litellm_managedagenttable.find_many(
where={"template_id": template_id}
)
if not agents:
return 0
counts = await asyncio.gather(
*[
stop_sessions_for_agent(
prisma_client=prisma_client,
region=region,
cluster=cluster,
agent_id=a.agent_id,
)
for a in agents
],
return_exceptions=True,
)
return sum(c for c in counts if isinstance(c, int))
async def reconcile_orphans(
*, prisma_client: Any, region: str, cluster: str
) -> Dict[str, int]:
"""One-shot orphan sweep. Stops Fargate tasks whose session is missing/dead.
Returns counts: {scanned, orphaned_stopped, stale_creating_stopped}.
"""
arns = await asyncio.to_thread(list_tagged_task_arns, region, cluster)
tasks = await asyncio.to_thread(describe_tasks_with_tags, region, cluster, arns)
# Filter to only managed-agent tasks (have our session tag).
managed_tasks = [t for t in tasks if TAG_SESSION_ID in t["tags"]]
if not managed_tasks:
return {"scanned": 0, "orphaned_stopped": 0, "stale_creating_stopped": 0}
session_ids = [t["tags"][TAG_SESSION_ID] for t in managed_tasks]
rows = await prisma_client.db.litellm_managedagentsessiontable.find_many(
where={"session_id": {"in": session_ids}}
)
by_id = {r.session_id: r for r in rows}
orphaned = 0
stale = 0
now_ts = time.time()
for task in managed_tasks:
sid = task["tags"][TAG_SESSION_ID]
arn = task["taskArn"]
row = by_id.get(sid)
if row is None:
await _stop_and_log(region, cluster, arn, sid, "missing_db_row")
orphaned += 1
continue
if row.status in DEAD_STATUSES:
await _stop_and_log(region, cluster, arn, sid, f"db_status={row.status}")
orphaned += 1
continue
if row.status == "creating":
row_created_ts = row.created_at.timestamp() if row.created_at else now_ts
if now_ts - row_created_ts > CREATING_TIMEOUT_SECONDS:
await _stop_and_log(region, cluster, arn, sid, "creating_timeout")
await prisma_client.db.litellm_managedagentsessiontable.update(
where={"session_id": sid},
data={
"status": "failed",
"failure_reason": "spawn timeout — task killed by reconciler",
"stopped_at": _datetime_from_ts(now_ts),
},
)
stale += 1
return {
"scanned": len(managed_tasks),
"orphaned_stopped": orphaned,
"stale_creating_stopped": stale,
}
async def reconcile_loop(
*,
prisma_client: Any,
region: str,
cluster: str,
interval_seconds: int = RECONCILE_INTERVAL_SECONDS,
) -> None:
"""Long-running reconciler. Run as fire-and-forget asyncio task."""
while True:
try:
stats = await reconcile_orphans(
prisma_client=prisma_client, region=region, cluster=cluster
)
if stats["orphaned_stopped"] or stats["stale_creating_stopped"]:
verbose_proxy_logger.info(f"managed_agents reconciler: {stats}")
except asyncio.CancelledError:
raise
except Exception as e:
verbose_proxy_logger.exception(f"managed_agents reconciler error: {e}")
await asyncio.sleep(interval_seconds)
async def _stop_and_log(
region: str, cluster: str, task_arn: str, session_id: str, reason: str
) -> None:
verbose_proxy_logger.info(
f"managed_agents: stopping orphan task arn={task_arn} session={session_id} reason={reason}"
)
await asyncio.to_thread(
stop_task_sync, region, cluster, task_arn, f"orphan: {reason}"
)
def _datetime_from_ts(ts: float):
from datetime import datetime, timezone
return datetime.fromtimestamp(ts, tz=timezone.utc)

View file

@ -0,0 +1,35 @@
# ------------------------------------------------------------------ installer
FROM debian:bookworm-slim AS installer
RUN apt-get update && apt-get install -y --no-install-recommends \
curl ca-certificates unzip \
&& rm -rf /var/lib/apt/lists/*
RUN curl -fsSL https://opencode.ai/install | bash
# Drop install bundles / caches if any
RUN find /root/.opencode -maxdepth 2 -type d -name 'cache' -exec rm -rf {} + 2>/dev/null || true \
&& find /root/.opencode -maxdepth 2 -name '*.zip' -delete 2>/dev/null || true \
&& find /root/.opencode -maxdepth 2 -name '*.tar*' -delete 2>/dev/null || true
# ------------------------------------------------------------------ runtime
FROM debian:bookworm-slim
RUN apt-get update && apt-get install -y --no-install-recommends \
git ca-certificates bash \
&& rm -rf /var/lib/apt/lists/* \
&& useradd -m -u 1000 -s /bin/bash sandbox \
&& mkdir -p /work \
&& chown sandbox:sandbox /work
COPY --from=installer --chown=sandbox:sandbox /root/.opencode /home/sandbox/.opencode
ENV PATH="/home/sandbox/.opencode/bin:${PATH}"
WORKDIR /work
COPY --chown=sandbox:sandbox entrypoint.sh /entrypoint.sh
RUN chmod +x /entrypoint.sh
USER sandbox
EXPOSE 4096
ENTRYPOINT ["/entrypoint.sh"]

View file

@ -0,0 +1,73 @@
#!/usr/bin/env bash
set -euo pipefail
: "${REPO_URL:?REPO_URL required}"
: "${LITELLM_API_KEY:?LITELLM_API_KEY required}"
: "${LITELLM_API_BASE:?LITELLM_API_BASE required}"
: "${LITELLM_DEFAULT_MODEL:?LITELLM_DEFAULT_MODEL required}"
: "${BRANCH:=main}"
: "${PORT:=4096}"
: "${REPO_DIR:=/work/repo}"
# Normalize base URL: strip trailing slash, ensure /v1 suffix
BASE="${LITELLM_API_BASE%/}"
case "$BASE" in
*/v1) ;;
*) BASE="${BASE}/v1" ;;
esac
# Clone. Token (if present) is fed via stdin to credential helper so it never
# lands in argv, env of child processes, .git/config, or shell history.
if [ ! -d "$REPO_DIR/.git" ]; then
if [ -n "${GIT_TOKEN:-}" ]; then
git -c credential.helper= \
-c "credential.helper=!f() { echo username=x-access-token; echo password=$GIT_TOKEN; }; f" \
clone --depth 1 --branch "$BRANCH" "$REPO_URL" "$REPO_DIR"
else
git clone --depth 1 --branch "$BRANCH" "$REPO_URL" "$REPO_DIR"
fi
fi
# Wipe token from env so opencode shell tool can't `printenv GIT_TOKEN`.
unset GIT_TOKEN
cd "$REPO_DIR"
# Belt-and-suspenders: ensure .git/config has clean remote (no embedded creds).
git remote set-url origin "$REPO_URL" 2>/dev/null || true
# Wire LiteLLM as OpenAI-compatible provider
cat > opencode.json <<EOF
{
"\$schema": "https://opencode.ai/config.json",
"provider": {
"litellm": {
"npm": "@ai-sdk/openai-compatible",
"options": {
"baseURL": "${BASE}",
"apiKey": "${LITELLM_API_KEY}"
},
"models": {
"${LITELLM_DEFAULT_MODEL}": {}
}
}
},
"model": "litellm/${LITELLM_DEFAULT_MODEL}"
}
EOF
if [ -n "${AGENT_PROMPT:-}" ]; then
mkdir -p .opencode/agent
cat > .opencode/agent/default.md <<EOF
---
description: sandbox agent
---
${AGENT_PROMPT}
EOF
fi
echo "[entrypoint] booting opencode serve on 0.0.0.0:${PORT}"
echo "[entrypoint] base=${BASE} model=${LITELLM_DEFAULT_MODEL} repo=${REPO_DIR}"
exec opencode serve --hostname 0.0.0.0 --port "$PORT"

View file

@ -0,0 +1,115 @@
"""Pydantic v2 type definitions for the managed_agents proxy feature."""
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, ConfigDict, Field
class DockerfileConfig(BaseModel):
model_config = ConfigDict(extra="forbid")
path: str
container_port: int = 4096
class AwsOverrides(BaseModel):
model_config = ConfigDict(extra="forbid")
cluster: Optional[str] = None
subnets: Optional[List[str]] = None
security_group: Optional[str] = None
task_role_arn: Optional[str] = None
task_exec_role_arn: Optional[str] = None
log_group: Optional[str] = None
class ManagedAgentsConfig(BaseModel):
model_config = ConfigDict(extra="forbid")
enabled: bool = False
aws_region: Optional[str] = None
dockerfiles: Dict[str, DockerfileConfig] = Field(default_factory=dict)
aws: AwsOverrides = Field(default_factory=AwsOverrides)
reconcile_interval_seconds: int = 60
class DockerfileOut(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
container_port: int
class TemplateCreate(BaseModel):
model_config = ConfigDict(extra="forbid")
name: Optional[str] = None
dockerfile_id: str
repo_url: str
default_branch: str
visibility: str = Field(pattern="^(public|private)$")
git_token: Optional[str] = None
class TemplateOut(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
name: Optional[str] = None
dockerfile_id: str
container_port: int
repo_url: str
default_branch: str
visibility: str
image_uri: Optional[str] = None
task_def_arn: Optional[str] = None
build_status: str
build_error: Optional[str] = None
class AgentCreate(BaseModel):
model_config = ConfigDict(extra="forbid")
name: Optional[str] = None
model: str
prompt: Optional[str] = None
tools: List[Any] = Field(default_factory=list)
template_id: str
branch: Optional[str] = None
litellm_api_key: Optional[str] = None
litellm_api_base: Optional[str] = None
class AgentOut(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
name: Optional[str] = None
model: str
template_id: str
branch: str
class SessionCreateIn(BaseModel):
model_config = ConfigDict(extra="forbid")
initial_prompt: Optional[str] = None
title: Optional[str] = None
class SessionOut(BaseModel):
model_config = ConfigDict(extra="forbid")
id: str
agent_id: str
sandbox_url: Optional[str] = None
status: str
task_arn: Optional[str] = None
response: Optional[Dict[str, Any]] = None
class MessageIn(BaseModel):
model_config = ConfigDict(extra="forbid")
text: Optional[str] = None
parts: Optional[List[Dict[str, Any]]] = None

View file

@ -309,6 +309,7 @@ from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router
from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router
from litellm.proxy.fine_tuning_endpoints.endpoints import set_fine_tuning_config
from litellm.proxy.google_endpoints.endpoints import router as google_router
from litellm.proxy.managed_agents_endpoints import router as managed_agents_router
from litellm.proxy.guardrails.init_guardrails import (
init_guardrails_v2,
initialize_guardrails,
@ -924,9 +925,49 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
## Initialize shared aiohttp session for connection reuse
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
## [Optional] Managed agents — Fargate orphan task reconciler
managed_agents_reconciler_task: Optional[asyncio.Task] = None
from litellm.proxy.managed_agents_endpoints import (
config_loader as managed_agents_config_loader,
)
managed_agents_config_loader.initialize(general_settings)
managed_agents_config = managed_agents_config_loader.MANAGED_AGENTS_CONFIG
if (
managed_agents_config is not None
and managed_agents_config.enabled
and prisma_client is not None
):
from litellm.proxy.managed_agents_endpoints.lifecycle import reconcile_loop
cluster_name = managed_agents_config.aws.cluster or "litellm-agents"
managed_agents_reconciler_task = asyncio.create_task(
reconcile_loop(
prisma_client=prisma_client,
region=managed_agents_config.aws_region or "us-east-1",
cluster=cluster_name,
interval_seconds=managed_agents_config.reconcile_interval_seconds,
)
)
verbose_proxy_logger.info(
f"managed_agents: started Fargate orphan reconciler (cluster={cluster_name})"
)
# End of startup event
yield
# Shutdown event - cancel managed_agents reconciler
if managed_agents_reconciler_task is not None:
managed_agents_reconciler_task.cancel()
try:
await managed_agents_reconciler_task
except asyncio.CancelledError:
pass
except Exception as e:
verbose_proxy_logger.error(
f"Error stopping managed_agents reconciler: {e}"
)
# Shutdown event - close shared aiohttp session
if shared_aiohttp_session is not None:
try:
@ -14879,6 +14920,7 @@ app.include_router(enterprise_router)
app.include_router(ui_discovery_endpoints_router)
# Eager: /models/{name}:method overlaps with the OpenAI /models endpoint.
app.include_router(google_router)
app.include_router(managed_agents_router)
attach_lazy_features(app)

View file

@ -38,11 +38,13 @@ model LiteLLM_CredentialsTable {
credential_id String @id @default(uuid())
credential_name String @unique
credential_values Json
credential_info Json?
credential_info Json?
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
updated_by String
managed_agent_sandbox_templates LiteLLM_ManagedAgentSandboxTemplateTable[]
}
// Models on proxy
@ -1374,3 +1376,86 @@ model LiteLLM_WorkflowMessage {
@@unique([run_id, sequence_number])
@@index([run_id])
}
// Managed Agents: sandbox harness template (admin-defined recipe).
// Maps harness type → ECR image. e.g. opencode/claude-code/aider.
model LiteLLM_ManagedAgentSandboxTemplateTable {
template_id String @id @default(uuid())
template_name String?
dockerfile_id String
image_uri String?
task_def_arn String?
image_hash String?
build_status String @default("pending")
build_error String?
container_port Int @default(4096)
repo_url String
default_branch String @default("main")
visibility String @default("public")
git_credential_id String?
description String?
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
git_credential LiteLLM_CredentialsTable? @relation(fields: [git_credential_id], references: [credential_id], onDelete: SetNull)
agents LiteLLM_ManagedAgentTable[]
@@index([dockerfile_id])
@@index([created_by])
@@index([git_credential_id])
}
// Managed Agents: agent config — model + prompt + tools + template ref.
model LiteLLM_ManagedAgentTable {
agent_id String @id @default(uuid())
agent_name String?
model String
prompt String?
tools Json @default("[]")
template_id String
branch String?
metadata Json? @default("{}")
created_at DateTime @default(now())
created_by String?
team_id String?
organization_id String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
template LiteLLM_ManagedAgentSandboxTemplateTable @relation(fields: [template_id], references: [template_id], onDelete: Cascade)
sessions LiteLLM_ManagedAgentSessionTable[]
@@index([template_id])
@@index([created_by])
@@index([team_id])
@@index([organization_id])
}
// Managed Agents: live Fargate task instance bound to one agent.
model LiteLLM_ManagedAgentSessionTable {
session_id String @id @default(uuid())
agent_id String
status String @default("creating")
task_arn String?
sandbox_url String?
harness_session_id String?
fargate_cluster String?
fargate_task_def_arn String?
virtual_key_hash String?
failure_reason String?
created_at DateTime @default(now())
created_by String?
team_id String?
last_seen_at DateTime?
expires_at DateTime?
stopped_at DateTime?
agent LiteLLM_ManagedAgentTable @relation(fields: [agent_id], references: [agent_id], onDelete: Cascade)
@@index([agent_id])
@@index([status])
@@index([created_by])
@@index([last_seen_at])
}

View file

@ -0,0 +1,338 @@
"""Unit tests for fargate bootstrap module.
Mocks boto3 clients used by `bootstrap.py` (`_ecs`, `_ec2`, `_iam`, `_logs`).
No real AWS calls.
"""
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from botocore.exceptions import ClientError
sys.path.insert(0, str(Path(__file__).resolve().parents[3]))
from litellm.proxy.managed_agents_endpoints.fargate import bootstrap
from litellm.proxy.managed_agents_endpoints.types import AwsOverrides
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _client_error(code: str) -> ClientError:
return ClientError(
error_response={"Error": {"Code": code, "Message": code}},
operation_name="op",
)
def _no_such_entity_exception_class():
return type("NoSuchEntityException", (Exception,), {})
def _resource_already_exists_class():
return type("ResourceAlreadyExistsException", (Exception,), {})
@pytest.fixture
def mock_ecs():
return MagicMock()
@pytest.fixture
def mock_ec2():
return MagicMock()
@pytest.fixture
def mock_iam():
return MagicMock()
@pytest.fixture
def mock_logs():
m = MagicMock()
m.exceptions.ResourceAlreadyExistsException = _resource_already_exists_class()
return m
# ---------------------------------------------------------------------------
# ensure_cluster
# ---------------------------------------------------------------------------
def test_ensure_cluster_existing_active_returns_arn(mock_ecs):
mock_ecs.describe_clusters.return_value = {
"clusters": [
{"status": "ACTIVE", "clusterArn": "arn:aws:ecs:us-west-2:1:cluster/foo"}
]
}
with patch.object(bootstrap, "_ecs", return_value=mock_ecs):
arn = bootstrap.ensure_cluster("us-west-2", "foo")
assert arn == "arn:aws:ecs:us-west-2:1:cluster/foo"
mock_ecs.describe_clusters.assert_called_once_with(clusters=["foo"])
mock_ecs.create_cluster.assert_not_called()
def test_ensure_cluster_missing_creates(mock_ecs):
mock_ecs.describe_clusters.return_value = {"clusters": []}
mock_ecs.create_cluster.return_value = {
"cluster": {"clusterArn": "arn:aws:ecs:us-west-2:1:cluster/new"}
}
with patch.object(bootstrap, "_ecs", return_value=mock_ecs):
arn = bootstrap.ensure_cluster("us-west-2", "new")
assert arn == "arn:aws:ecs:us-west-2:1:cluster/new"
mock_ecs.create_cluster.assert_called_once_with(clusterName="new")
# ---------------------------------------------------------------------------
# ensure_task_exec_role
# ---------------------------------------------------------------------------
def test_ensure_task_exec_role_existing_returns_arn(mock_iam):
mock_iam.get_role.return_value = {
"Role": {"Arn": "arn:aws:iam::1:role/litellm-agents-task-exec"}
}
with patch.object(bootstrap, "_iam", return_value=mock_iam):
arn = bootstrap.ensure_task_exec_role("litellm-agents-task-exec")
assert arn == "arn:aws:iam::1:role/litellm-agents-task-exec"
mock_iam.create_role.assert_not_called()
mock_iam.attach_role_policy.assert_not_called()
def test_ensure_task_exec_role_missing_creates_attaches_and_sleeps(mock_iam):
mock_iam.get_role.side_effect = _client_error("NoSuchEntity")
mock_iam.create_role.return_value = {
"Role": {"Arn": "arn:aws:iam::1:role/new-role"}
}
with (
patch.object(bootstrap, "_iam", return_value=mock_iam),
patch.object(bootstrap.time, "sleep") as sleep_mock,
):
arn = bootstrap.ensure_task_exec_role("new-role")
assert arn == "arn:aws:iam::1:role/new-role"
mock_iam.create_role.assert_called_once()
mock_iam.attach_role_policy.assert_called_once_with(
RoleName="new-role",
PolicyArn=bootstrap.TASK_EXEC_ROLE_POLICY_ARN,
)
sleep_mock.assert_called_once_with(10)
# ---------------------------------------------------------------------------
# discover_vpc_subnet
# ---------------------------------------------------------------------------
def test_discover_vpc_subnet_happy_path(mock_ec2):
mock_ec2.describe_vpcs.return_value = {
"Vpcs": [{"VpcId": "vpc-1", "CidrBlock": "10.0.0.0/16"}]
}
mock_ec2.describe_subnets.return_value = {"Subnets": [{"SubnetId": "subnet-1"}]}
with patch.object(bootstrap, "_ec2", return_value=mock_ec2):
vpc_id, vpc_cidr, subnet_id = bootstrap.discover_vpc_subnet("us-west-2")
assert vpc_id == "vpc-1"
assert vpc_cidr == "10.0.0.0/16"
assert subnet_id == "subnet-1"
def test_discover_vpc_subnet_no_default_vpc_raises(mock_ec2):
mock_ec2.describe_vpcs.return_value = {"Vpcs": []}
with patch.object(bootstrap, "_ec2", return_value=mock_ec2):
with pytest.raises(RuntimeError, match="no default VPC"):
bootstrap.discover_vpc_subnet("us-west-2")
def test_discover_vpc_subnet_no_public_subnet_raises(mock_ec2):
mock_ec2.describe_vpcs.return_value = {
"Vpcs": [{"VpcId": "vpc-1", "CidrBlock": "10.0.0.0/16"}]
}
mock_ec2.describe_subnets.return_value = {"Subnets": []}
with patch.object(bootstrap, "_ec2", return_value=mock_ec2):
with pytest.raises(RuntimeError, match="no public subnet"):
bootstrap.discover_vpc_subnet("us-west-2")
# ---------------------------------------------------------------------------
# ensure_security_group
# ---------------------------------------------------------------------------
def test_ensure_security_group_existing_returns_id(mock_ec2):
mock_ec2.describe_security_groups.return_value = {
"SecurityGroups": [{"GroupId": "sg-existing"}]
}
with patch.object(bootstrap, "_ec2", return_value=mock_ec2):
sg_id = bootstrap.ensure_security_group(
"us-west-2", "litellm-sg", "vpc-1", "10.0.0.0/16", 4096
)
assert sg_id == "sg-existing"
mock_ec2.create_security_group.assert_not_called()
mock_ec2.authorize_security_group_ingress.assert_not_called()
mock_ec2.revoke_security_group_egress.assert_not_called()
mock_ec2.authorize_security_group_egress.assert_not_called()
def test_ensure_security_group_missing_creates_ingress_revoke_and_authorize(mock_ec2):
mock_ec2.describe_security_groups.return_value = {"SecurityGroups": []}
mock_ec2.create_security_group.return_value = {"GroupId": "sg-new"}
with patch.object(bootstrap, "_ec2", return_value=mock_ec2):
sg_id = bootstrap.ensure_security_group(
"us-west-2", "litellm-sg", "vpc-1", "10.0.0.0/16", 4096
)
assert sg_id == "sg-new"
mock_ec2.create_security_group.assert_called_once()
mock_ec2.authorize_security_group_ingress.assert_called_once()
ingress_kwargs = mock_ec2.authorize_security_group_ingress.call_args.kwargs
assert ingress_kwargs["GroupId"] == "sg-new"
assert ingress_kwargs["IpPermissions"][0]["FromPort"] == 4096
assert ingress_kwargs["IpPermissions"][0]["ToPort"] == 4096
mock_ec2.revoke_security_group_egress.assert_called_once()
mock_ec2.authorize_security_group_egress.assert_called_once()
egress_kwargs = mock_ec2.authorize_security_group_egress.call_args.kwargs
perms = egress_kwargs["IpPermissions"]
ports = sorted({(p["IpProtocol"], p["FromPort"]) for p in perms})
assert ("tcp", 443) in ports
assert ("tcp", 53) in ports
assert ("udp", 53) in ports
# ---------------------------------------------------------------------------
# ensure_log_group
# ---------------------------------------------------------------------------
def test_ensure_log_group_missing_creates(mock_logs):
with patch.object(bootstrap, "_logs", return_value=mock_logs):
bootstrap.ensure_log_group("us-west-2", "/ecs/foo")
mock_logs.create_log_group.assert_called_once_with(logGroupName="/ecs/foo")
def test_ensure_log_group_already_exists_swallowed(mock_logs):
mock_logs.create_log_group.side_effect = _client_error(
"ResourceAlreadyExistsException"
)
with patch.object(bootstrap, "_logs", return_value=mock_logs):
# Should NOT raise
bootstrap.ensure_log_group("us-west-2", "/ecs/foo")
mock_logs.create_log_group.assert_called_once()
def test_ensure_log_group_other_error_reraised(mock_logs):
mock_logs.create_log_group.side_effect = _client_error("AccessDenied")
with patch.object(bootstrap, "_logs", return_value=mock_logs):
with pytest.raises(ClientError):
bootstrap.ensure_log_group("us-west-2", "/ecs/foo")
# ---------------------------------------------------------------------------
# bootstrap_shared_infra
# ---------------------------------------------------------------------------
def test_bootstrap_shared_infra_no_overrides_calls_every_ensure():
overrides = AwsOverrides()
with (
patch.object(
bootstrap, "ensure_cluster", return_value="cluster-arn"
) as ensure_cluster,
patch.object(
bootstrap, "ensure_task_exec_role", return_value="role-arn"
) as ensure_role,
patch.object(
bootstrap,
"discover_vpc_subnet",
return_value=("vpc-1", "10.0.0.0/16", "subnet-1"),
) as discover,
patch.object(
bootstrap, "ensure_security_group", return_value="sg-1"
) as ensure_sg,
patch.object(bootstrap, "ensure_log_group") as ensure_lg,
):
infra = bootstrap.bootstrap_shared_infra("us-west-2", overrides)
ensure_cluster.assert_called_once()
ensure_role.assert_called_once()
discover.assert_called_once()
ensure_sg.assert_called_once()
ensure_lg.assert_called_once()
assert infra.cluster_arn == "cluster-arn"
assert infra.task_exec_role_arn == "role-arn"
assert infra.security_group_id == "sg-1"
assert infra.log_group_name == bootstrap.DEFAULT_LOG_GROUP_NAME
assert infra.vpc_id == "vpc-1"
assert infra.subnet_ids == ["subnet-1"]
def test_bootstrap_shared_infra_with_overrides_uses_given_skips_ensures():
overrides = AwsOverrides(
cluster="my-cluster",
security_group="sg-given",
task_exec_role_arn="arn:aws:iam::1:role/given-role",
)
with (
patch.object(
bootstrap, "_validate_existing_cluster", return_value="cluster-arn-given"
) as v_cluster,
patch.object(
bootstrap, "_validate_existing_role", return_value="role-arn-given"
) as v_role,
patch.object(
bootstrap, "_validate_existing_security_group", return_value="sg-given"
) as v_sg,
patch.object(
bootstrap,
"discover_vpc_subnet",
return_value=("vpc-1", "10.0.0.0/16", "subnet-1"),
) as discover,
patch.object(bootstrap, "ensure_log_group") as ensure_lg,
patch.object(bootstrap, "ensure_cluster") as ensure_cluster,
patch.object(bootstrap, "ensure_task_exec_role") as ensure_role,
patch.object(bootstrap, "ensure_security_group") as ensure_sg,
):
infra = bootstrap.bootstrap_shared_infra("us-west-2", overrides)
# Validators called for overrides
v_cluster.assert_called_once()
v_role.assert_called_once()
v_sg.assert_called_once()
# ensure_* skipped for overridden resources
ensure_cluster.assert_not_called()
ensure_role.assert_not_called()
ensure_sg.assert_not_called()
# Subnet not overridden → discover_vpc_subnet called
discover.assert_called_once()
# log_group not overridden → ensure_log_group called
ensure_lg.assert_called_once()
assert infra.cluster_arn == "cluster-arn-given"
assert infra.task_exec_role_arn == "role-arn-given"
assert infra.security_group_id == "sg-given"

View file

@ -0,0 +1,458 @@
"""Tests for managed_agents_endpoints/endpoints_sessions.py.
Covers POST /v1/managed_agents/agents/{agent_id}/session,
GET /v1/managed_agents/sessions/{session_id},
DELETE /v1/managed_agents/sessions/{session_id}.
AWS lifecycle (bootstrap_shared_infra, run_task_sync, wait_running_get_ip_sync,
wait_http_ready, stop_task_sync, stop_session_task) is fully mocked.
Harness HTTP (harness_create_session, harness_send_message) is also mocked —
this file is a unit test, not the smoke test that hits a live container.
"""
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.managed_agents_endpoints.endpoints import router
# Importing the module registers /agents/{agent_id}/session, /sessions/{id}
# routes onto `router`. Without this, the routes don't exist on the test app.
import litellm.proxy.managed_agents_endpoints.endpoints_sessions # noqa: F401
@pytest.fixture
def user():
return UserAPIKeyAuth(
api_key="sk-user", user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
)
@pytest.fixture
def app_factory():
def make(auth_user):
app = FastAPI()
app.include_router(router)
app.dependency_overrides[user_api_key_auth] = lambda: auth_user
return TestClient(app)
return make
def _make_template(build_status="ready"):
return SimpleNamespace(
template_id="tmpl-1",
template_name="t",
dockerfile_id="opencode",
container_port=4096,
repo_url="https://github.com/x/y",
default_branch="main",
visibility="public",
git_credential_id=None,
image_uri="img:abc",
task_def_arn="arn:td",
image_hash="abc",
build_status=build_status,
build_error=None,
)
def _make_agent(template):
a = SimpleNamespace(
agent_id="agt-1",
agent_name="a",
model="anthropic/claude-sonnet-4-6",
prompt="be concise",
tools=[],
template_id=template.template_id,
branch="main",
metadata={"litellm_api_key": "sk-x", "litellm_api_base": "http://x"},
)
a.template = template
return a
def _make_session(session_id="sess-1", **kw):
base = dict(
session_id=session_id,
agent_id="agt-1",
status="creating",
task_arn=None,
sandbox_url=None,
harness_session_id=None,
fargate_cluster="litellm-agents",
fargate_task_def_arn="arn:td",
failure_reason=None,
stopped_at=None,
last_seen_at=None,
created_by="u1",
team_id=None,
)
base.update(kw)
return SimpleNamespace(**base)
def _make_prisma(agent=None, session=None):
p = MagicMock()
agent_t = MagicMock()
agent_t.find_unique = AsyncMock(return_value=agent)
p.db.litellm_managedagenttable = agent_t
sess_t = MagicMock()
if session is not None:
sess_t.create = AsyncMock(return_value=session)
sess_t.find_unique = AsyncMock(return_value=session)
else:
sess_t.create = AsyncMock()
sess_t.find_unique = AsyncMock(return_value=None)
sess_t.update = AsyncMock()
p.db.litellm_managedagentsessiontable = sess_t
return p
# ---------------------------------------------------------------------------
# create_session
# ---------------------------------------------------------------------------
def test_create_session_404_when_agent_missing(app_factory, user):
client = app_factory(user)
prisma = _make_prisma(agent=None)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = client.post("/v1/managed_agents/agents/missing/session", json={})
assert resp.status_code == 404
assert "missing" in resp.json()["detail"]
def test_create_session_409_when_template_not_ready(app_factory, user):
client = app_factory(user)
template = _make_template(build_status="pending")
agent = _make_agent(template)
prisma = _make_prisma(agent=agent)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = client.post("/v1/managed_agents/agents/agt-1/session", json={})
assert resp.status_code == 409
assert "not ready" in resp.json()["detail"]
def test_create_session_happy_path_with_initial_prompt(app_factory, user):
client = app_factory(user)
template = _make_template()
agent = _make_agent(template)
created_session = _make_session()
prisma = _make_prisma(agent=agent, session=created_session)
infra = SimpleNamespace(
cluster_arn="arn:cluster",
task_exec_role_arn="arn:role",
security_group_id="sg-1",
log_group_name="/ecs/x",
vpc_id="vpc-1",
subnet_ids=["subnet-1"],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.bootstrap_shared_infra",
return_value=infra,
) as mock_boot,
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.run_task_sync",
return_value="arn:task/abc",
) as mock_run,
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.wait_running_get_ip_sync",
return_value="1.2.3.4",
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.wait_http_ready",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.harness_create_session",
new=AsyncMock(return_value="harness-sess-1"),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.harness_send_message",
new=AsyncMock(return_value={"parts": [{"type": "text", "text": "hi"}]}),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.decrypt_git_token",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.managed_agents_endpoints.config_loader.MANAGED_AGENTS_CONFIG",
SimpleNamespace(
aws_region="us-west-2",
aws=SimpleNamespace(cluster=None),
),
),
):
resp = client.post(
"/v1/managed_agents/agents/agt-1/session",
json={"title": "smoke", "initial_prompt": "what is this repo?"},
)
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["status"] == "ready"
assert body["sandbox_url"] == "http://1.2.3.4:4096"
assert body["task_arn"] == "arn:task/abc"
assert body["response"]["parts"][0]["text"] == "hi"
# Bootstrap was called with config region + container port
mock_boot.assert_called_once()
args, _ = mock_boot.call_args
assert args[0] == "us-west-2"
assert args[2] == 4096
# run_task_sync got the env vars from agent.metadata + template
_, run_kwargs = mock_run.call_args
env = run_kwargs["env"]
assert env["LITELLM_API_KEY"] == "sk-x"
assert env["LITELLM_API_BASE"] == "http://x"
assert env["LITELLM_DEFAULT_MODEL"] == "anthropic/claude-sonnet-4-6"
assert env["REPO_URL"] == "https://github.com/x/y"
assert env["BRANCH"] == "main"
assert env["AGENT_PROMPT"] == "be concise"
assert "GIT_TOKEN" not in env # template has no credential
def test_create_session_happy_path_no_initial_prompt(app_factory, user):
client = app_factory(user)
template = _make_template()
agent = _make_agent(template)
prisma = _make_prisma(agent=agent, session=_make_session())
infra = SimpleNamespace(
cluster_arn="arn:cluster",
task_exec_role_arn="arn:role",
security_group_id="sg-1",
log_group_name="/ecs/x",
vpc_id="vpc-1",
subnet_ids=["subnet-1"],
)
send_mock = AsyncMock(return_value={"parts": []})
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.bootstrap_shared_infra",
return_value=infra,
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.run_task_sync",
return_value="arn:task/abc",
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.wait_running_get_ip_sync",
return_value="1.2.3.4",
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.wait_http_ready",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.harness_create_session",
new=AsyncMock(return_value="harness-sess-1"),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.harness_send_message",
new=send_mock,
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.decrypt_git_token",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.managed_agents_endpoints.config_loader.MANAGED_AGENTS_CONFIG",
SimpleNamespace(
aws_region="us-west-2",
aws=SimpleNamespace(cluster=None),
),
),
):
resp = client.post("/v1/managed_agents/agents/agt-1/session", json={})
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["status"] == "ready"
# No initial_prompt → harness_send_message NOT called
send_mock.assert_not_called()
assert body["response"] is None
def test_create_session_marks_failed_on_exception(app_factory, user):
client = app_factory(user)
template = _make_template()
agent = _make_agent(template)
prisma = _make_prisma(agent=agent, session=_make_session())
infra = SimpleNamespace(
cluster_arn="arn:cluster",
task_exec_role_arn="arn:role",
security_group_id="sg-1",
log_group_name="/ecs/x",
vpc_id="vpc-1",
subnet_ids=["subnet-1"],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.bootstrap_shared_infra",
return_value=infra,
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.run_task_sync",
return_value="arn:task/abc",
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.wait_running_get_ip_sync",
side_effect=TimeoutError("never reached RUNNING"),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.stop_task_sync",
return_value=None,
) as mock_stop,
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.decrypt_git_token",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy.managed_agents_endpoints.config_loader.MANAGED_AGENTS_CONFIG",
SimpleNamespace(
aws_region="us-west-2",
aws=SimpleNamespace(cluster=None),
),
),
):
resp = client.post("/v1/managed_agents/agents/agt-1/session", json={})
assert resp.status_code == 500
assert "session create failed" in resp.json()["detail"]
# Session row was marked failed
update_calls = prisma.db.litellm_managedagentsessiontable.update.call_args_list
failed_calls = [
c for c in update_calls if c.kwargs.get("data", {}).get("status") == "failed"
]
assert failed_calls, f"expected status=failed update, got {update_calls}"
# Best-effort stop_task_sync was attempted
mock_stop.assert_called_once()
# ---------------------------------------------------------------------------
# get_session
# ---------------------------------------------------------------------------
def test_get_session_happy(app_factory, user):
client = app_factory(user)
sess = _make_session(
session_id="sess-9",
status="ready",
task_arn="arn:task/9",
sandbox_url="http://1.2.3.4:4096",
)
prisma = _make_prisma(session=sess)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = client.get("/v1/managed_agents/sessions/sess-9")
assert resp.status_code == 200
body = resp.json()
assert body["id"] == "sess-9"
assert body["status"] == "ready"
assert body["sandbox_url"] == "http://1.2.3.4:4096"
assert body["task_arn"] == "arn:task/9"
def test_get_session_404(app_factory, user):
client = app_factory(user)
prisma = _make_prisma()
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = client.get("/v1/managed_agents/sessions/nope")
assert resp.status_code == 404
# ---------------------------------------------------------------------------
# delete_session
# ---------------------------------------------------------------------------
def test_delete_session_404(app_factory, user):
client = app_factory(user)
prisma = _make_prisma()
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = client.delete("/v1/managed_agents/sessions/nope")
assert resp.status_code == 404
def test_delete_session_stops_task_and_marks_dead(app_factory, user):
client = app_factory(user)
sess = _make_session(
session_id="sess-9",
status="ready",
task_arn="arn:task/9",
fargate_cluster="my-cluster",
)
prisma = _make_prisma(session=sess)
stop_mock = AsyncMock(return_value=None)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.stop_session_task",
new=stop_mock,
),
patch(
"litellm.proxy.managed_agents_endpoints.config_loader.MANAGED_AGENTS_CONFIG",
SimpleNamespace(
aws_region="us-west-2",
aws=SimpleNamespace(cluster=None),
),
),
):
resp = client.delete("/v1/managed_agents/sessions/sess-9")
assert resp.status_code == 200
assert resp.json() == {"id": "sess-9", "status": "dead"}
stop_mock.assert_awaited_once()
_, kwargs = stop_mock.call_args
# Cluster comes from row.fargate_cluster, not config default
assert kwargs["cluster"] == "my-cluster"
assert kwargs["task_arn"] == "arn:task/9"
assert kwargs["session_id"] == "sess-9"
# Row was updated to dead
update_calls = prisma.db.litellm_managedagentsessiontable.update.call_args_list
assert any(
c.kwargs.get("data", {}).get("status") == "dead" for c in update_calls
), f"expected status=dead update, got {update_calls}"
def test_delete_session_no_task_arn_skips_stop(app_factory, user):
client = app_factory(user)
sess = _make_session(session_id="sess-9", status="creating", task_arn=None)
prisma = _make_prisma(session=sess)
stop_mock = AsyncMock(return_value=None)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints_sessions.stop_session_task",
new=stop_mock,
),
patch(
"litellm.proxy.managed_agents_endpoints.config_loader.MANAGED_AGENTS_CONFIG",
SimpleNamespace(
aws_region="us-west-2",
aws=SimpleNamespace(cluster=None),
),
),
):
resp = client.delete("/v1/managed_agents/sessions/sess-9")
assert resp.status_code == 200
stop_mock.assert_not_called()

View file

@ -0,0 +1,231 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.managed_agents_endpoints.endpoints import router
@pytest.fixture
def admin():
return UserAPIKeyAuth(
api_key="sk-admin", user_id="a1", user_role=LitellmUserRoles.PROXY_ADMIN
)
@pytest.fixture
def user():
return UserAPIKeyAuth(
api_key="sk-user", user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
)
@pytest.fixture
def app_factory():
def make(auth_user):
app = FastAPI()
app.include_router(router)
app.dependency_overrides[user_api_key_auth] = lambda: auth_user
return TestClient(app)
return make
@pytest.fixture
def fake_prisma():
p = MagicMock()
table = MagicMock()
table.create = AsyncMock()
table.find_unique = AsyncMock()
table.find_many = AsyncMock(return_value=[])
table.update = AsyncMock()
table.delete = AsyncMock()
p.db.litellm_managedagentsandboxtemplatetable = table
agents_t = MagicMock()
agents_t.count = AsyncMock(return_value=0)
p.db.litellm_managedagenttable = agents_t
cred_t = MagicMock()
cred_t.create = AsyncMock(return_value=SimpleNamespace(credential_id="cred-1"))
p.db.litellm_credentialstable = cred_t
return p
def _public_body():
return {
"name": "tpl-1",
"dockerfile_id": "opencode",
"repo_url": "https://github.com/x/y",
"default_branch": "main",
"visibility": "public",
"git_token": None,
}
def test_dockerfiles_lists(app_factory, user):
client = app_factory(user)
with patch(
"litellm.proxy.managed_agents_endpoints.endpoints.list_dockerfiles",
return_value=[SimpleNamespace(dockerfile_id="opencode", container_port=4096)],
):
resp = client.get("/v1/managed_agents/dockerfiles")
assert resp.status_code == 200
data = resp.json()
assert len(data) == 1
assert data[0]["id"] == "opencode"
assert data[0]["container_port"] == 4096
def test_template_create_non_admin_403(app_factory, user, fake_prisma):
client = app_factory(user)
with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma):
resp = client.post("/v1/managed_agents/sandbox-templates", json=_public_body())
assert resp.status_code == 403
def test_template_create_unknown_dockerfile_id_400(app_factory, admin, fake_prisma):
client = app_factory(admin)
body = _public_body()
body["dockerfile_id"] = "nope"
with (
patch("litellm.proxy.proxy_server.prisma_client", fake_prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints.get_dockerfile",
side_effect=KeyError("nope"),
),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints.list_dockerfiles",
return_value=[
SimpleNamespace(dockerfile_id="opencode", container_port=4096)
],
),
):
resp = client.post("/v1/managed_agents/sandbox-templates", json=body)
assert resp.status_code == 400
assert "available" in str(resp.json()["detail"]).lower()
def test_template_create_happy_path(app_factory, admin, fake_prisma):
client = app_factory(admin)
created_row = SimpleNamespace(
template_id="t1",
template_name=None,
dockerfile_id="opencode",
container_port=4096,
repo_url="https://github.com/x/y",
default_branch="main",
visibility="public",
image_uri=None,
task_def_arn=None,
build_status="pending",
build_error=None,
)
updated_row = SimpleNamespace(
template_id="t1",
template_name=None,
dockerfile_id="opencode",
container_port=4096,
repo_url="https://github.com/x/y",
default_branch="main",
visibility="public",
image_uri="img:abc",
task_def_arn="arn:td",
build_status="ready",
build_error=None,
)
fake_prisma.db.litellm_managedagentsandboxtemplatetable.create = AsyncMock(
return_value=created_row
)
fake_prisma.db.litellm_managedagentsandboxtemplatetable.update = AsyncMock(
return_value=updated_row
)
provisioned = SimpleNamespace(
image_uri="img:abc",
task_def_arn="arn:td",
image_hash="abc",
container_port=4096,
cluster_arn="arn:cluster",
security_group_id="sg-1",
subnet_ids=["sub-1"],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", fake_prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints.get_dockerfile",
return_value=SimpleNamespace(
dockerfile_id="opencode",
container_port=4096,
path="/x",
context_dir="/x",
),
),
patch("litellm.proxy.managed_agents_endpoints.endpoints.validate_repo_branch"),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints.provision_template",
AsyncMock(return_value=provisioned),
),
patch(
"litellm.proxy.managed_agents_endpoints.config_loader.MANAGED_AGENTS_CONFIG",
SimpleNamespace(aws_region="us-west-2", aws=SimpleNamespace(cluster=None)),
),
patch("boto3.client", return_value=MagicMock()),
):
resp = client.post("/v1/managed_agents/sandbox-templates", json=_public_body())
assert resp.status_code == 200, resp.text
payload = resp.json()
assert payload["image_uri"] == "img:abc"
assert payload["build_status"] == "ready"
fake_prisma.db.litellm_managedagentsandboxtemplatetable.update.assert_called()
def test_template_create_private_requires_token_400(app_factory, admin, fake_prisma):
client = app_factory(admin)
body = _public_body()
body["visibility"] = "private"
body["git_token"] = None
with (
patch("litellm.proxy.proxy_server.prisma_client", fake_prisma),
patch(
"litellm.proxy.managed_agents_endpoints.endpoints.get_dockerfile",
return_value=SimpleNamespace(
dockerfile_id="opencode",
container_port=4096,
path="/x",
context_dir="/x",
),
),
):
resp = client.post("/v1/managed_agents/sandbox-templates", json=body)
assert resp.status_code == 400
def test_template_delete_with_agents_409(app_factory, admin, fake_prisma):
client = app_factory(admin)
fake_prisma.db.litellm_managedagentsandboxtemplatetable.find_unique = AsyncMock(
return_value=SimpleNamespace(
template_id="t1",
template_name=None,
dockerfile_id="opencode",
container_port=4096,
repo_url="https://github.com/x/y",
default_branch="main",
visibility="public",
image_uri=None,
task_def_arn=None,
build_status="ready",
build_error=None,
)
)
fake_prisma.db.litellm_managedagenttable.count = AsyncMock(return_value=1)
with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma):
resp = client.delete("/v1/managed_agents/sandbox-templates/t1")
assert resp.status_code == 409

View file

@ -0,0 +1,253 @@
"""Unit tests for git_validation module.
Mocks subprocess + prisma + encrypt/decrypt helpers. Verifies:
- authed_repo_url: HTTPS w/ token, HTTPS no token, ssh, malformed
- validate_repo_branch: success, empty stdout, nonzero exit, FileNotFoundError, TimeoutExpired
- decrypt_git_token: None id, found w/ token, not found, missing key
- encrypt_and_store_git_token: happy path, credential_name shape, created_by/updated_by
"""
import re
import subprocess
import sys
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
sys.path.insert(0, str(Path(__file__).resolve().parents[3]))
from litellm.proxy.managed_agents_endpoints.git_validation import (
authed_repo_url,
decrypt_git_token,
encrypt_and_store_git_token,
validate_repo_branch,
)
# ---------------------------------------------------------------------------
# authed_repo_url
# ---------------------------------------------------------------------------
def test_authed_repo_url_https_with_token_injects_netloc():
url = authed_repo_url("https://github.com/org/repo.git", "secret-token")
assert "x-access-token:secret-token@github.com" in url
assert url.endswith("/org/repo.git")
def test_authed_repo_url_https_no_token_returns_input_unchanged():
original = "https://github.com/org/repo.git"
assert authed_repo_url(original, None) == original
assert authed_repo_url(original, "") == original
def test_authed_repo_url_ssh_returns_input_unchanged():
original = "git@github.com:org/repo.git"
assert authed_repo_url(original, "secret-token") == original
def test_authed_repo_url_malformed_returns_input_unchanged():
# No scheme/host → not https, returned untouched.
original = "not-a-real-url"
assert authed_repo_url(original, "secret-token") == original
# ---------------------------------------------------------------------------
# validate_repo_branch (sync, raises HTTPException on failure)
# ---------------------------------------------------------------------------
def _completed(returncode: int, stdout: str = "", stderr: str = ""):
return SimpleNamespace(returncode=returncode, stdout=stdout, stderr=stderr)
def test_validate_repo_branch_success_no_exception():
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.subprocess.run",
return_value=_completed(0, stdout="abc123\trefs/heads/main\n", stderr=""),
) as run_mock:
# Should not raise.
validate_repo_branch("https://github.com/org/repo.git", "main", git_token=None)
run_mock.assert_called_once()
def test_validate_repo_branch_empty_stdout_raises_400_branch_not_found():
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.subprocess.run",
return_value=_completed(0, stdout="", stderr=""),
):
with pytest.raises(HTTPException) as exc_info:
validate_repo_branch(
"https://github.com/org/repo.git", "missing-branch", git_token=None
)
assert exc_info.value.status_code == 400
assert "missing-branch" in exc_info.value.detail
assert "not found" in exc_info.value.detail
def test_validate_repo_branch_nonzero_exit_raises_400_and_scrubs_token():
repo_url = "https://github.com/org/repo.git"
token = "super-secret-token"
stderr = (
"remote: Repository not found.\n"
"fatal: repository 'https://x-access-token:super-secret-token@github.com/org/repo.git/' not found"
)
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.subprocess.run",
return_value=_completed(128, stdout="", stderr=stderr),
):
with pytest.raises(HTTPException) as exc_info:
validate_repo_branch(repo_url, "main", git_token=token)
assert exc_info.value.status_code == 400
detail = exc_info.value.detail
# Plain repo URL should appear in the error.
assert repo_url in detail
# Authed URL (with token) must NOT leak.
assert token not in detail
assert "x-access-token" not in detail
def test_validate_repo_branch_file_not_found_raises_500():
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.subprocess.run",
side_effect=FileNotFoundError("git not on PATH"),
):
with pytest.raises(HTTPException) as exc_info:
validate_repo_branch(
"https://github.com/org/repo.git", "main", git_token=None
)
assert exc_info.value.status_code == 500
assert "git not installed" in exc_info.value.detail
def test_validate_repo_branch_timeout_raises_400():
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.subprocess.run",
side_effect=subprocess.TimeoutExpired(cmd="git", timeout=15),
):
with pytest.raises(HTTPException) as exc_info:
validate_repo_branch(
"https://github.com/org/repo.git", "main", git_token=None
)
assert exc_info.value.status_code == 400
assert "timed out" in exc_info.value.detail
# ---------------------------------------------------------------------------
# decrypt_git_token (async)
# ---------------------------------------------------------------------------
def _fake_prisma_credentials(row):
"""Build a fake prisma client whose litellm_credentialstable.find_unique returns row."""
table = MagicMock()
table.find_unique = AsyncMock(return_value=row)
db = MagicMock()
db.litellm_credentialstable = table
client = MagicMock()
client.db = db
return client, table
@pytest.mark.asyncio
async def test_decrypt_git_token_none_id_returns_none_no_db_call():
client, table = _fake_prisma_credentials(row=None)
result = await decrypt_git_token(client, credential_id=None)
assert result is None
table.find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_decrypt_git_token_found_returns_plaintext():
row = SimpleNamespace(
credential_id="cred-1",
credential_values={"git_token": "ENCRYPTED-BLOB"},
)
client, table = _fake_prisma_credentials(row=row)
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.decrypt_value_helper",
return_value="plaintext-token",
) as decrypt_mock:
result = await decrypt_git_token(client, credential_id="cred-1")
assert result == "plaintext-token"
decrypt_mock.assert_called_once_with("ENCRYPTED-BLOB", key="git_token")
table.find_unique.assert_awaited_once_with(where={"credential_id": "cred-1"})
@pytest.mark.asyncio
async def test_decrypt_git_token_not_found_returns_none():
client, table = _fake_prisma_credentials(row=None)
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.decrypt_value_helper"
) as decrypt_mock:
result = await decrypt_git_token(client, credential_id="missing")
assert result is None
decrypt_mock.assert_not_called()
@pytest.mark.asyncio
async def test_decrypt_git_token_missing_git_token_key_returns_none():
row = SimpleNamespace(
credential_id="cred-1",
credential_values={"some_other_key": "blob"},
)
client, _ = _fake_prisma_credentials(row=row)
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.decrypt_value_helper"
) as decrypt_mock:
result = await decrypt_git_token(client, credential_id="cred-1")
assert result is None
decrypt_mock.assert_not_called()
# ---------------------------------------------------------------------------
# encrypt_and_store_git_token (async)
# ---------------------------------------------------------------------------
def _fake_prisma_credentials_create(returned):
table = MagicMock()
table.create = AsyncMock(return_value=returned)
db = MagicMock()
db.litellm_credentialstable = table
client = MagicMock()
client.db = db
return client, table
@pytest.mark.asyncio
async def test_encrypt_and_store_git_token_happy_path():
created = SimpleNamespace(credential_id="cred-new-123")
client, table = _fake_prisma_credentials_create(returned=created)
with patch(
"litellm.proxy.managed_agents_endpoints.git_validation.encrypt_value_helper",
return_value="ENCRYPTED-OUTPUT",
) as encrypt_mock:
new_id = await encrypt_and_store_git_token(
client, raw_token="raw-secret", created_by="user-alice"
)
assert new_id == "cred-new-123"
encrypt_mock.assert_called_once_with("raw-secret")
table.create.assert_awaited_once()
create_kwargs = table.create.call_args.kwargs
data = create_kwargs["data"]
# Encrypted value is what the create was called with.
assert data["credential_values"] == {"git_token": "ENCRYPTED-OUTPUT"}
# credential_name format: managed-agent-git-token-<uuid4>
assert re.match(r"^managed-agent-git-token-[0-9a-f-]+$", data["credential_name"])
# created_by + updated_by both set to caller-supplied value.
assert data["created_by"] == "user-alice"
assert data["updated_by"] == "user-alice"

View file

@ -0,0 +1,233 @@
"""Smoke tests for managed_agents reconciler.
Mocks Prisma + boto3. Verifies orphan stop_task_sync called for each scenario:
- task w/ no DB row → stopped
- task w/ row.status = 'dead' → stopped
- task w/ row.status = 'stopped' → stopped
- task w/ row.status = 'failed' → stopped
- task w/ row.status = 'creating' young → skipped
- task w/ row.status = 'creating' stale → stopped + row marked failed
- task w/ row.status = 'ready' → skipped
"""
import sys
from datetime import datetime, timedelta, timezone
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, str(Path(__file__).resolve().parents[3]))
from litellm.proxy.managed_agents_endpoints.fargate.tasks import TAG_SESSION_ID
from litellm.proxy.managed_agents_endpoints.lifecycle import reconcile_orphans
def _make_session_row(session_id: str, status: str, created_at: datetime):
return SimpleNamespace(
session_id=session_id,
status=status,
created_at=created_at,
task_arn=f"arn:aws:ecs:us-west-2:123:task/{session_id}",
)
def _fake_prisma(rows):
by_id = {r.session_id: r for r in rows}
async def find_many(where):
ids = where["session_id"]["in"]
return [by_id[i] for i in ids if i in by_id]
update = AsyncMock()
table = MagicMock()
table.find_many = AsyncMock(side_effect=find_many)
table.update = update
db = MagicMock()
db.litellm_managedagentsessiontable = table
client = MagicMock()
client.db = db
return client, update
def _fake_tasks(*session_ids):
return [
{
"taskArn": f"arn:aws:ecs:us-west-2:123:task/{sid}",
"tags": {TAG_SESSION_ID: sid},
"lastStatus": "RUNNING",
}
for sid in session_ids
]
@pytest.mark.asyncio
async def test_orphan_no_db_row_stopped():
tasks = _fake_tasks("s_orphan")
arns = [t["taskArn"] for t in tasks]
prisma, update = _fake_prisma([])
with patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.list_tagged_task_arns",
return_value=arns,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.describe_tasks_with_tags",
return_value=tasks,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.stop_task_sync"
) as stop_mock:
stats = await reconcile_orphans(
prisma_client=prisma, region="us-west-2", cluster="test"
)
assert stats == {"scanned": 1, "orphaned_stopped": 1, "stale_creating_stopped": 0}
stop_mock.assert_called_once()
assert "missing_db_row" in stop_mock.call_args.args[3]
update.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["dead", "failed", "stopped"])
async def test_orphan_dead_row_stopped(status):
rows = [_make_session_row("s1", status, datetime.now(timezone.utc))]
tasks = _fake_tasks("s1")
arns = [t["taskArn"] for t in tasks]
prisma, update = _fake_prisma(rows)
with patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.list_tagged_task_arns",
return_value=arns,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.describe_tasks_with_tags",
return_value=tasks,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.stop_task_sync"
) as stop_mock:
stats = await reconcile_orphans(
prisma_client=prisma, region="us-west-2", cluster="test"
)
assert stats["orphaned_stopped"] == 1
assert stats["stale_creating_stopped"] == 0
stop_mock.assert_called_once()
update.assert_not_called()
@pytest.mark.asyncio
async def test_creating_young_skipped():
rows = [
_make_session_row(
"s_young", "creating", datetime.now(timezone.utc) - timedelta(seconds=30)
)
]
tasks = _fake_tasks("s_young")
arns = [t["taskArn"] for t in tasks]
prisma, update = _fake_prisma(rows)
with patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.list_tagged_task_arns",
return_value=arns,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.describe_tasks_with_tags",
return_value=tasks,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.stop_task_sync"
) as stop_mock:
stats = await reconcile_orphans(
prisma_client=prisma, region="us-west-2", cluster="test"
)
assert stats == {"scanned": 1, "orphaned_stopped": 0, "stale_creating_stopped": 0}
stop_mock.assert_not_called()
update.assert_not_called()
@pytest.mark.asyncio
async def test_creating_stale_stopped_and_marked_failed():
rows = [
_make_session_row(
"s_stale", "creating", datetime.now(timezone.utc) - timedelta(hours=1)
)
]
tasks = _fake_tasks("s_stale")
arns = [t["taskArn"] for t in tasks]
prisma, update = _fake_prisma(rows)
with patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.list_tagged_task_arns",
return_value=arns,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.describe_tasks_with_tags",
return_value=tasks,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.stop_task_sync"
) as stop_mock:
stats = await reconcile_orphans(
prisma_client=prisma, region="us-west-2", cluster="test"
)
assert stats["stale_creating_stopped"] == 1
assert stats["orphaned_stopped"] == 0
stop_mock.assert_called_once()
update.assert_called_once()
update_kwargs = update.call_args.kwargs
assert update_kwargs["where"] == {"session_id": "s_stale"}
assert update_kwargs["data"]["status"] == "failed"
assert "spawn timeout" in update_kwargs["data"]["failure_reason"]
@pytest.mark.asyncio
async def test_ready_row_skipped():
rows = [_make_session_row("s_live", "ready", datetime.now(timezone.utc))]
tasks = _fake_tasks("s_live")
arns = [t["taskArn"] for t in tasks]
prisma, update = _fake_prisma(rows)
with patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.list_tagged_task_arns",
return_value=arns,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.describe_tasks_with_tags",
return_value=tasks,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.stop_task_sync"
) as stop_mock:
stats = await reconcile_orphans(
prisma_client=prisma, region="us-west-2", cluster="test"
)
assert stats == {"scanned": 1, "orphaned_stopped": 0, "stale_creating_stopped": 0}
stop_mock.assert_not_called()
update.assert_not_called()
@pytest.mark.asyncio
async def test_no_managed_tasks_returns_zero():
untagged = [
{"taskArn": "arn:aws:ecs:us-west-2:123:task/other", "tags": {}, "lastStatus": "RUNNING"}
]
arns = [t["taskArn"] for t in untagged]
prisma, _ = _fake_prisma([])
with patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.list_tagged_task_arns",
return_value=arns,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.describe_tasks_with_tags",
return_value=untagged,
), patch(
"litellm.proxy.managed_agents_endpoints.lifecycle.stop_task_sync"
) as stop_mock:
stats = await reconcile_orphans(
prisma_client=prisma, region="us-west-2", cluster="test"
)
assert stats == {"scanned": 0, "orphaned_stopped": 0, "stale_creating_stopped": 0}
stop_mock.assert_not_called()

View file

@ -0,0 +1,203 @@
"""Unit tests for fargate registry module.
Mocks boto3 ECR client + subprocess. No real AWS / docker calls.
"""
import sys
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from botocore.exceptions import ClientError
sys.path.insert(0, str(Path(__file__).resolve().parents[3]))
from litellm.proxy.managed_agents_endpoints.fargate import registry
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _client_error(code: str) -> ClientError:
return ClientError(
error_response={"Error": {"Code": code, "Message": code}},
operation_name="op",
)
@pytest.fixture
def mock_ecr():
return MagicMock()
# ---------------------------------------------------------------------------
# compute_dockerfile_hash
# ---------------------------------------------------------------------------
def test_compute_dockerfile_hash_deterministic(tmp_path):
df = tmp_path / "Dockerfile"
df.write_text("FROM python:3.11\nRUN pip install foo\n")
ctx = tmp_path / "ctx"
ctx.mkdir()
(ctx / "app.py").write_text("print('hi')\n")
h1 = registry.compute_dockerfile_hash(str(df), str(ctx))
h2 = registry.compute_dockerfile_hash(str(df), str(ctx))
assert h1 == h2
assert len(h1) == 64 # sha256 hex
def test_compute_dockerfile_hash_changes_on_dockerfile_change(tmp_path):
df = tmp_path / "Dockerfile"
ctx = tmp_path / "ctx"
ctx.mkdir()
(ctx / "app.py").write_text("print('hi')\n")
df.write_text("FROM python:3.11\n")
h1 = registry.compute_dockerfile_hash(str(df), str(ctx))
df.write_text("FROM python:3.12\n")
h2 = registry.compute_dockerfile_hash(str(df), str(ctx))
assert h1 != h2
def test_compute_dockerfile_hash_changes_on_context_change(tmp_path):
df = tmp_path / "Dockerfile"
df.write_text("FROM python:3.11\n")
ctx = tmp_path / "ctx"
ctx.mkdir()
(ctx / "app.py").write_text("print('hi')\n")
h1 = registry.compute_dockerfile_hash(str(df), str(ctx))
# Add a new file → hash should differ
(ctx / "extra.py").write_text("print('extra')\n")
h2 = registry.compute_dockerfile_hash(str(df), str(ctx))
assert h1 != h2
# ---------------------------------------------------------------------------
# image_exists
# ---------------------------------------------------------------------------
def test_image_exists_returns_true_on_found(mock_ecr):
mock_ecr.describe_images.return_value = {"imageDetails": [{"imageTags": ["abc"]}]}
with patch.object(registry, "_ecr", return_value=mock_ecr):
assert registry.image_exists("us-west-2", "repo", "abc") is True
def test_image_exists_returns_false_on_not_found(mock_ecr):
mock_ecr.describe_images.side_effect = _client_error("ImageNotFoundException")
with patch.object(registry, "_ecr", return_value=mock_ecr):
assert registry.image_exists("us-west-2", "repo", "abc") is False
# ---------------------------------------------------------------------------
# ensure_ecr_repo
# ---------------------------------------------------------------------------
def test_ensure_ecr_repo_creates_when_missing(mock_ecr):
uri = "123456789012.dkr.ecr.us-west-2.amazonaws.com/newrepo"
mock_ecr.describe_repositories.side_effect = _client_error(
"RepositoryNotFoundException"
)
mock_ecr.create_repository.return_value = {"repository": {"repositoryUri": uri}}
with patch.object(registry, "_ecr", return_value=mock_ecr):
result = registry.ensure_ecr_repo("us-west-2", "newrepo")
assert result == uri
assert ".dkr.ecr.us-west-2.amazonaws.com/newrepo" in result
mock_ecr.create_repository.assert_called_once()
# ---------------------------------------------------------------------------
# build_and_push orchestration
# ---------------------------------------------------------------------------
def test_build_and_push_cache_hit_skips_build(tmp_path):
df = tmp_path / "Dockerfile"
df.write_text("FROM scratch\n")
ctx = tmp_path / "ctx"
ctx.mkdir()
repo_uri = "123.dkr.ecr.us-west-2.amazonaws.com/repo"
with (
patch.object(registry, "ensure_ecr_repo", return_value=repo_uri),
patch.object(registry, "image_exists", return_value=True),
patch.object(registry, "docker_login") as login_mock,
patch.object(registry, "docker_build") as build_mock,
patch.object(registry, "docker_push") as push_mock,
patch.object(registry.subprocess, "Popen") as popen_mock,
patch.object(registry.subprocess, "run") as run_mock,
):
result = registry.build_and_push(
region="us-west-2",
repo_name="repo",
dockerfile_path=str(df),
context_dir=str(ctx),
content_hash="abc123",
)
assert result == f"{repo_uri}:abc123"
login_mock.assert_not_called()
build_mock.assert_not_called()
push_mock.assert_not_called()
popen_mock.assert_not_called()
run_mock.assert_not_called()
def test_build_and_push_cache_miss_builds_and_pushes(tmp_path):
df = tmp_path / "Dockerfile"
df.write_text("FROM scratch\n")
ctx = tmp_path / "ctx"
ctx.mkdir()
repo_uri = "123.dkr.ecr.us-west-2.amazonaws.com/repo"
call_order: list = []
with (
patch.object(registry, "ensure_ecr_repo", return_value=repo_uri),
patch.object(registry, "image_exists", return_value=False),
patch.object(
registry,
"docker_login",
side_effect=lambda *a, **kw: call_order.append("login"),
) as login_mock,
patch.object(
registry,
"docker_build",
side_effect=lambda *a, **kw: call_order.append("build"),
) as build_mock,
patch.object(
registry,
"docker_push",
side_effect=lambda *a, **kw: call_order.append("push"),
) as push_mock,
):
result = registry.build_and_push(
region="us-west-2",
repo_name="repo",
dockerfile_path=str(df),
context_dir=str(ctx),
content_hash="abc123",
)
assert result == f"{repo_uri}:abc123"
assert call_order == ["login", "build", "push"]
login_mock.assert_called_once()
build_mock.assert_called_once()
push_mock.assert_called_once()