From 977ab6682c002248db4710083994b09474862dfa Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 7 May 2026 15:17:21 -0700 Subject: [PATCH] feat(managed agents): add v1 Fargate-backed managed-agent endpoints MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- .../migration.sql | 104 ++++ .../migration.sql | 28 ++ .../litellm_proxy_extras/schema.prisma | 87 +++- .../managed_agents_endpoints/__init__.py | 10 + .../managed_agents_endpoints/config_loader.py | 163 +++++++ .../managed_agents_endpoints/endpoints.py | 291 +++++++++++ .../endpoints_agents.py | 89 ++++ .../endpoints_passthrough.py | 205 ++++++++ .../endpoints_sessions.py | 324 +++++++++++++ .../fargate/__init__.py | 0 .../fargate/bootstrap.py | 326 +++++++++++++ .../managed_agents_endpoints/fargate/build.py | 167 +++++++ .../fargate/registry.py | 267 ++++++++++ .../managed_agents_endpoints/fargate/tasks.py | 171 +++++++ .../git_validation.py | 133 +++++ .../harness_client.py | 65 +++ .../managed_agents_endpoints/lifecycle.py | 205 ++++++++ .../sample_harnesses/opencode/Dockerfile | 35 ++ .../sample_harnesses/opencode/entrypoint.sh | 73 +++ .../proxy/managed_agents_endpoints/types.py | 115 +++++ litellm/proxy/proxy_server.py | 42 ++ litellm/proxy/schema.prisma | 87 +++- .../managed_agents_endpoints/__init__.py | 0 .../test_bootstrap.py | 338 +++++++++++++ .../test_endpoints_sessions.py | 458 ++++++++++++++++++ .../test_endpoints_templates.py | 231 +++++++++ .../test_git_validation.py | 253 ++++++++++ .../test_lifecycle.py | 233 +++++++++ .../managed_agents_endpoints/test_registry.py | 203 ++++++++ 29 files changed, 4701 insertions(+), 2 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260507120000_add_managed_agents_tables/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260507130000_managed_agents_template_v2/migration.sql create mode 100644 litellm/proxy/managed_agents_endpoints/__init__.py create mode 100644 litellm/proxy/managed_agents_endpoints/config_loader.py create mode 100644 litellm/proxy/managed_agents_endpoints/endpoints.py create mode 100644 litellm/proxy/managed_agents_endpoints/endpoints_agents.py create mode 100644 litellm/proxy/managed_agents_endpoints/endpoints_passthrough.py create mode 100644 litellm/proxy/managed_agents_endpoints/endpoints_sessions.py create mode 100644 litellm/proxy/managed_agents_endpoints/fargate/__init__.py create mode 100644 litellm/proxy/managed_agents_endpoints/fargate/bootstrap.py create mode 100644 litellm/proxy/managed_agents_endpoints/fargate/build.py create mode 100644 litellm/proxy/managed_agents_endpoints/fargate/registry.py create mode 100644 litellm/proxy/managed_agents_endpoints/fargate/tasks.py create mode 100644 litellm/proxy/managed_agents_endpoints/git_validation.py create mode 100644 litellm/proxy/managed_agents_endpoints/harness_client.py create mode 100644 litellm/proxy/managed_agents_endpoints/lifecycle.py create mode 100644 litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/Dockerfile create mode 100755 litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/entrypoint.sh create mode 100644 litellm/proxy/managed_agents_endpoints/types.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/test_bootstrap.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_sessions.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_templates.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/test_git_validation.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/test_lifecycle.py create mode 100644 tests/test_litellm/proxy/managed_agents_endpoints/test_registry.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260507120000_add_managed_agents_tables/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260507120000_add_managed_agents_tables/migration.sql new file mode 100644 index 00000000000..601a16fc076 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260507120000_add_managed_agents_tables/migration.sql @@ -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; diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260507130000_managed_agents_template_v2/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260507130000_managed_agents_template_v2/migration.sql new file mode 100644 index 00000000000..9d1aaff2748 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260507130000_managed_agents_template_v2/migration.sql @@ -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"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 84ce99557e3..47297e8b226 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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]) +} diff --git a/litellm/proxy/managed_agents_endpoints/__init__.py b/litellm/proxy/managed_agents_endpoints/__init__.py new file mode 100644 index 00000000000..17497ad12ac --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/__init__.py @@ -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 diff --git a/litellm/proxy/managed_agents_endpoints/config_loader.py b/litellm/proxy/managed_agents_endpoints/config_loader.py new file mode 100644 index 00000000000..7875e5cdbbe --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/config_loader.py @@ -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())] diff --git a/litellm/proxy/managed_agents_endpoints/endpoints.py b/litellm/proxy/managed_agents_endpoints/endpoints.py new file mode 100644 index 00000000000..dac1f325519 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/endpoints.py @@ -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"} diff --git a/litellm/proxy/managed_agents_endpoints/endpoints_agents.py b/litellm/proxy/managed_agents_endpoints/endpoints_agents.py new file mode 100644 index 00000000000..19fffc6cca7 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/endpoints_agents.py @@ -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) diff --git a/litellm/proxy/managed_agents_endpoints/endpoints_passthrough.py b/litellm/proxy/managed_agents_endpoints/endpoints_passthrough.py new file mode 100644 index 00000000000..2ef85a8a85a --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/endpoints_passthrough.py @@ -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), + ) diff --git a/litellm/proxy/managed_agents_endpoints/endpoints_sessions.py b/litellm/proxy/managed_agents_endpoints/endpoints_sessions.py new file mode 100644 index 00000000000..d2b33d021b1 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/endpoints_sessions.py @@ -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"} diff --git a/litellm/proxy/managed_agents_endpoints/fargate/__init__.py b/litellm/proxy/managed_agents_endpoints/fargate/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/proxy/managed_agents_endpoints/fargate/bootstrap.py b/litellm/proxy/managed_agents_endpoints/fargate/bootstrap.py new file mode 100644 index 00000000000..c77f905dd48 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/fargate/bootstrap.py @@ -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, + ) diff --git a/litellm/proxy/managed_agents_endpoints/fargate/build.py b/litellm/proxy/managed_agents_endpoints/fargate/build.py new file mode 100644 index 00000000000..2a95f4aedbe --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/fargate/build.py @@ -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}" + ) diff --git a/litellm/proxy/managed_agents_endpoints/fargate/registry.py b/litellm/proxy/managed_agents_endpoints/fargate/registry.py new file mode 100644 index 00000000000..95104a28d02 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/fargate/registry.py @@ -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 diff --git a/litellm/proxy/managed_agents_endpoints/fargate/tasks.py b/litellm/proxy/managed_agents_endpoints/fargate/tasks.py new file mode 100644 index 00000000000..5dace2f6697 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/fargate/tasks.py @@ -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 diff --git a/litellm/proxy/managed_agents_endpoints/git_validation.py b/litellm/proxy/managed_agents_endpoints/git_validation.py new file mode 100644 index 00000000000..2cecc431b33 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/git_validation.py @@ -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 diff --git a/litellm/proxy/managed_agents_endpoints/harness_client.py b/litellm/proxy/managed_agents_endpoints/harness_client.py new file mode 100644 index 00000000000..2715ea32cb9 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/harness_client.py @@ -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'") diff --git a/litellm/proxy/managed_agents_endpoints/lifecycle.py b/litellm/proxy/managed_agents_endpoints/lifecycle.py new file mode 100644 index 00000000000..c6762683376 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/lifecycle.py @@ -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) diff --git a/litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/Dockerfile b/litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/Dockerfile new file mode 100644 index 00000000000..ae8065a140c --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/Dockerfile @@ -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"] diff --git a/litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/entrypoint.sh b/litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/entrypoint.sh new file mode 100755 index 00000000000..8309140e474 --- /dev/null +++ b/litellm/proxy/managed_agents_endpoints/sample_harnesses/opencode/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 < .opencode/agent/default.md < 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" diff --git a/tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_sessions.py b/tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_sessions.py new file mode 100644 index 00000000000..be217fa8542 --- /dev/null +++ b/tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_sessions.py @@ -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() diff --git a/tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_templates.py b/tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_templates.py new file mode 100644 index 00000000000..e7b9b1a5777 --- /dev/null +++ b/tests/test_litellm/proxy/managed_agents_endpoints/test_endpoints_templates.py @@ -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 diff --git a/tests/test_litellm/proxy/managed_agents_endpoints/test_git_validation.py b/tests/test_litellm/proxy/managed_agents_endpoints/test_git_validation.py new file mode 100644 index 00000000000..00770556b28 --- /dev/null +++ b/tests/test_litellm/proxy/managed_agents_endpoints/test_git_validation.py @@ -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- + 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" diff --git a/tests/test_litellm/proxy/managed_agents_endpoints/test_lifecycle.py b/tests/test_litellm/proxy/managed_agents_endpoints/test_lifecycle.py new file mode 100644 index 00000000000..2eeb5e20c39 --- /dev/null +++ b/tests/test_litellm/proxy/managed_agents_endpoints/test_lifecycle.py @@ -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() diff --git a/tests/test_litellm/proxy/managed_agents_endpoints/test_registry.py b/tests/test_litellm/proxy/managed_agents_endpoints/test_registry.py new file mode 100644 index 00000000000..c334c4cd0d7 --- /dev/null +++ b/tests/test_litellm/proxy/managed_agents_endpoints/test_registry.py @@ -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()