mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
8c9830eef9
commit
977ab6682c
29 changed files with 4701 additions and 2 deletions
|
|
@ -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;
|
||||
|
|
@ -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");
|
||||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
10
litellm/proxy/managed_agents_endpoints/__init__.py
Normal file
10
litellm/proxy/managed_agents_endpoints/__init__.py
Normal 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
|
||||
163
litellm/proxy/managed_agents_endpoints/config_loader.py
Normal file
163
litellm/proxy/managed_agents_endpoints/config_loader.py
Normal 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())]
|
||||
291
litellm/proxy/managed_agents_endpoints/endpoints.py
Normal file
291
litellm/proxy/managed_agents_endpoints/endpoints.py
Normal 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"}
|
||||
89
litellm/proxy/managed_agents_endpoints/endpoints_agents.py
Normal file
89
litellm/proxy/managed_agents_endpoints/endpoints_agents.py
Normal 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)
|
||||
205
litellm/proxy/managed_agents_endpoints/endpoints_passthrough.py
Normal file
205
litellm/proxy/managed_agents_endpoints/endpoints_passthrough.py
Normal 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),
|
||||
)
|
||||
324
litellm/proxy/managed_agents_endpoints/endpoints_sessions.py
Normal file
324
litellm/proxy/managed_agents_endpoints/endpoints_sessions.py
Normal 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"}
|
||||
326
litellm/proxy/managed_agents_endpoints/fargate/bootstrap.py
Normal file
326
litellm/proxy/managed_agents_endpoints/fargate/bootstrap.py
Normal 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,
|
||||
)
|
||||
167
litellm/proxy/managed_agents_endpoints/fargate/build.py
Normal file
167
litellm/proxy/managed_agents_endpoints/fargate/build.py
Normal 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}"
|
||||
)
|
||||
267
litellm/proxy/managed_agents_endpoints/fargate/registry.py
Normal file
267
litellm/proxy/managed_agents_endpoints/fargate/registry.py
Normal 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
|
||||
171
litellm/proxy/managed_agents_endpoints/fargate/tasks.py
Normal file
171
litellm/proxy/managed_agents_endpoints/fargate/tasks.py
Normal 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
|
||||
133
litellm/proxy/managed_agents_endpoints/git_validation.py
Normal file
133
litellm/proxy/managed_agents_endpoints/git_validation.py
Normal 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
|
||||
65
litellm/proxy/managed_agents_endpoints/harness_client.py
Normal file
65
litellm/proxy/managed_agents_endpoints/harness_client.py
Normal 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'")
|
||||
205
litellm/proxy/managed_agents_endpoints/lifecycle.py
Normal file
205
litellm/proxy/managed_agents_endpoints/lifecycle.py
Normal 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)
|
||||
|
|
@ -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"]
|
||||
|
|
@ -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"
|
||||
115
litellm/proxy/managed_agents_endpoints/types.py
Normal file
115
litellm/proxy/managed_agents_endpoints/types.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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()
|
||||
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue