932 lines
32 KiB
Python
932 lines
32 KiB
Python
"""MCP Server for Docker-based DAG Orchestration."""
|
|
|
|
import asyncio
|
|
import json
|
|
import os
|
|
import uuid
|
|
from datetime import datetime
|
|
from typing import Any, Dict, Set
|
|
from mcp.server.fastmcp import FastMCP
|
|
from starlette.applications import Starlette
|
|
from starlette.responses import JSONResponse, StreamingResponse, FileResponse, HTMLResponse
|
|
from starlette.requests import Request
|
|
from starlette.routing import Route, Mount
|
|
from starlette.staticfiles import StaticFiles
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
|
import uvicorn
|
|
import logging
|
|
|
|
from .auth import (
|
|
OAUTH_ENABLED,
|
|
verify_access_token,
|
|
extract_bearer_token,
|
|
get_protected_resource_metadata,
|
|
get_www_authenticate_header,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Server configuration
|
|
MCP_PORT = int(os.environ.get("MCP_PORT", "8000"))
|
|
|
|
# Environment variables for Docker
|
|
DOCKER_HOST = os.environ.get("DOCKER_HOST", "unix:///var/run/docker.sock")
|
|
DOCKER_TLS_VERIFY = os.environ.get("DOCKER_TLS_VERIFY", "")
|
|
DOCKER_CERT_PATH = os.environ.get("DOCKER_CERT_PATH", "")
|
|
|
|
# Docker image for DAG executor
|
|
DAG_EXECUTOR_IMAGE = os.environ.get("DAG_EXECUTOR_IMAGE", "dag-executor:latest")
|
|
|
|
# LLM API keys (passed to container)
|
|
OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY", "")
|
|
ANTHROPIC_API_KEY = os.environ.get("ANTHROPIC_API_KEY", "")
|
|
GOOGLE_API_KEY = os.environ.get("GOOGLE_API_KEY", "")
|
|
|
|
# Create FastMCP server (host="0.0.0.0" allows external access)
|
|
mcp = FastMCP("mcp-docker-orchestrator", host="0.0.0.0", port=MCP_PORT)
|
|
|
|
|
|
async def health_check(request):
|
|
"""Health check endpoint for Docker/load balancer."""
|
|
return JSONResponse({"status": "healthy", "service": "mcp-docker-orchestrator"})
|
|
|
|
|
|
async def oauth_protected_resource(request: Request):
|
|
"""
|
|
OAuth 2.0 Protected Resource Metadata endpoint.
|
|
GPT and other OAuth clients call this first to discover authorization server.
|
|
|
|
RFC 8414: OAuth 2.0 Authorization Server Metadata
|
|
"""
|
|
return JSONResponse(get_protected_resource_metadata())
|
|
|
|
|
|
class OAuthMiddleware(BaseHTTPMiddleware):
|
|
"""
|
|
Middleware for OAuth2 Bearer token authentication.
|
|
|
|
- Skips authentication for health check and metadata endpoints
|
|
- Returns 401 with WWW-Authenticate header for missing/invalid tokens
|
|
- Stores verified user info in request.state.user
|
|
"""
|
|
|
|
# Paths that don't require authentication
|
|
PUBLIC_PATHS = {
|
|
"/health",
|
|
"/.well-known/oauth-protected-resource",
|
|
"/.well-known/oauth-authorization-server",
|
|
"/.well-known/openid-configuration",
|
|
}
|
|
|
|
# Path prefixes that don't require authentication
|
|
PUBLIC_PREFIXES = (
|
|
"/.well-known/",
|
|
"/mcp/.well-known/",
|
|
"/api/", # Task Monitor REST API
|
|
"/sse/", # Task Monitor SSE
|
|
"/monitor", # Task Monitor UI
|
|
"/assets/", # Static assets
|
|
)
|
|
|
|
async def dispatch(self, request: Request, call_next):
|
|
# Skip auth if OAuth is disabled
|
|
if not OAUTH_ENABLED:
|
|
return await call_next(request)
|
|
|
|
# Skip auth for public endpoints
|
|
path = request.url.path
|
|
if path in self.PUBLIC_PATHS or path.startswith(self.PUBLIC_PREFIXES):
|
|
return await call_next(request)
|
|
|
|
# Extract Bearer token
|
|
authorization = request.headers.get("authorization")
|
|
logger.info(f"[AUTH] Request to {request.url.path} - Authorization header present: {authorization is not None}")
|
|
if authorization:
|
|
logger.debug(f"[AUTH] Authorization header value (first 20 chars): {authorization[:20]}...")
|
|
token = extract_bearer_token(authorization)
|
|
|
|
if not token:
|
|
logger.warning(f"Missing or invalid Authorization header for {request.url.path}")
|
|
logger.debug(f"[AUTH] All request headers: {dict(request.headers)}")
|
|
return JSONResponse(
|
|
status_code=401,
|
|
content={"error": "unauthorized", "message": "Missing or invalid Bearer token"},
|
|
headers={"WWW-Authenticate": get_www_authenticate_header()}
|
|
)
|
|
|
|
# Verify token using Keycloak introspection
|
|
try:
|
|
user_info = await verify_access_token(token)
|
|
request.state.user = user_info
|
|
logger.debug(f"Authenticated client: {user_info.get('client_id', 'unknown')}")
|
|
except ValueError as e:
|
|
logger.warning(f"Token verification failed: {e}")
|
|
return JSONResponse(
|
|
status_code=401,
|
|
content={"error": "invalid_token", "message": str(e)},
|
|
headers={"WWW-Authenticate": get_www_authenticate_header()}
|
|
)
|
|
|
|
return await call_next(request)
|
|
|
|
# In-memory storage for tracking executions
|
|
executions: Dict[str, Dict[str, Any]] = {}
|
|
|
|
|
|
def get_docker_client():
|
|
"""Get Docker client with proper configuration."""
|
|
import docker
|
|
|
|
if DOCKER_HOST.startswith("tcp://"):
|
|
if DOCKER_TLS_VERIFY:
|
|
tls_config = docker.tls.TLSConfig(
|
|
client_cert=(
|
|
os.path.join(DOCKER_CERT_PATH, "cert.pem"),
|
|
os.path.join(DOCKER_CERT_PATH, "key.pem")
|
|
),
|
|
ca_cert=os.path.join(DOCKER_CERT_PATH, "ca.pem"),
|
|
verify=True
|
|
)
|
|
return docker.DockerClient(base_url=DOCKER_HOST, tls=tls_config)
|
|
else:
|
|
return docker.DockerClient(base_url=DOCKER_HOST)
|
|
else:
|
|
return docker.from_env()
|
|
|
|
|
|
@mcp.tool()
|
|
async def run_dag(
|
|
execution_name: str,
|
|
task_procedure: dict,
|
|
tasks: dict,
|
|
system_prompt: str,
|
|
tools: dict = None
|
|
) -> str:
|
|
"""
|
|
Execute a DAG (task_procedure) in a Docker container.
|
|
Tasks are executed in parallel based on dependencies using asyncio.
|
|
|
|
Args:
|
|
execution_name: Name for this execution (used for tracking)
|
|
task_procedure: DAG structure defining task execution order (IN, OUT, task nodes with nexts/wait_until)
|
|
tasks: Task definitions mapping task_name to {llm_provider, llm_model, prompts}
|
|
system_prompt: System prompt for all tasks
|
|
tools: Optional MCP tools configuration for tasks
|
|
"""
|
|
execution_id = str(uuid.uuid4())[:8]
|
|
container_name = f"dag-{execution_name}-{execution_id}"
|
|
|
|
# Prepare execution config as JSON
|
|
execution_config = {
|
|
"execution_id": execution_id,
|
|
"execution_name": execution_name,
|
|
"task_procedure": task_procedure,
|
|
"tasks": tasks,
|
|
"system_prompt": system_prompt,
|
|
"tools": tools
|
|
}
|
|
|
|
# Environment variables for the container
|
|
# STATUS_REPORT_URL uses the service name in Docker network
|
|
status_report_url = os.environ.get("STATUS_REPORT_URL", "http://mcp-docker-orchestrator:8000")
|
|
environment = {
|
|
"EXECUTION_CONFIG": json.dumps(execution_config),
|
|
"OPENAI_API_KEY": OPENAI_API_KEY,
|
|
"ANTHROPIC_API_KEY": ANTHROPIC_API_KEY,
|
|
"GOOGLE_API_KEY": GOOGLE_API_KEY,
|
|
"STATUS_REPORT_URL": status_report_url,
|
|
}
|
|
|
|
# Run in thread pool to avoid blocking
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def start_container():
|
|
client = get_docker_client()
|
|
container = client.containers.run(
|
|
image=DAG_EXECUTOR_IMAGE,
|
|
name=container_name,
|
|
environment=environment,
|
|
detach=True,
|
|
remove=False,
|
|
network="agentbackend_agent-network", # Connect to agent-network for MCP server access
|
|
mem_limit="2g", # Set memory limit to 2GB
|
|
memswap_limit="2g", # Disable swap to prevent performance degradation
|
|
)
|
|
return container.id
|
|
|
|
container_id = await loop.run_in_executor(None, start_container)
|
|
|
|
# Store execution info (use UTC with Z suffix for proper timezone handling)
|
|
created_at = datetime.utcnow().isoformat() + "Z"
|
|
executions[execution_id] = {
|
|
"execution_id": execution_id,
|
|
"execution_name": execution_name,
|
|
"container_id": container_id,
|
|
"container_name": container_name,
|
|
"status": "running",
|
|
"task_procedure": task_procedure,
|
|
"tasks": tasks,
|
|
"created_at": created_at
|
|
}
|
|
|
|
# Broadcast new execution to SSE subscribers
|
|
await broadcast_sse_event("execution_update", {
|
|
"execution_id": execution_id,
|
|
"execution_name": execution_name,
|
|
"status": "running",
|
|
"container_name": container_name,
|
|
"created_at": created_at,
|
|
})
|
|
|
|
return json.dumps({
|
|
"success": True,
|
|
"execution_id": execution_id,
|
|
"container_name": container_name,
|
|
"status": "running",
|
|
"message": f"DAG execution '{execution_name}' started in container"
|
|
}, indent=2)
|
|
|
|
|
|
@mcp.tool()
|
|
async def check_execution_status(execution_id: str) -> str:
|
|
"""
|
|
Check the status of a running DAG execution.
|
|
Returns current state (running, completed, failed) and progress.
|
|
|
|
Args:
|
|
execution_id: Execution ID returned from run_dag
|
|
"""
|
|
if execution_id not in executions:
|
|
return json.dumps({"error": f"Execution {execution_id} not found"})
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def get_status():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
return {
|
|
"container_status": container.status,
|
|
"exit_code": container.attrs.get("State", {}).get("ExitCode")
|
|
}
|
|
except Exception as e:
|
|
return {"container_status": "not_found", "error": str(e)}
|
|
|
|
status_info = await loop.run_in_executor(None, get_status)
|
|
|
|
# Update execution status based on container status
|
|
container_status = status_info["container_status"]
|
|
if container_status == "running":
|
|
execution["status"] = "running"
|
|
elif container_status == "exited":
|
|
exit_code = status_info.get("exit_code", -1)
|
|
execution["status"] = "completed" if exit_code == 0 else "failed"
|
|
else:
|
|
execution["status"] = container_status
|
|
|
|
return json.dumps({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
"container_status": container_status,
|
|
"exit_code": status_info.get("exit_code"),
|
|
}, indent=2)
|
|
|
|
|
|
@mcp.tool()
|
|
async def get_execution_logs(execution_id: str, tail: int = 100) -> str:
|
|
"""
|
|
Get logs from a DAG execution container.
|
|
Returns stdout/stderr from the executor.
|
|
|
|
Args:
|
|
execution_id: Execution ID
|
|
tail: Number of lines to return from the end (default: 100)
|
|
"""
|
|
if execution_id not in executions:
|
|
return json.dumps({"error": f"Execution {execution_id} not found"})
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def get_logs():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
logs = container.logs(tail=tail, timestamps=True).decode("utf-8")
|
|
return {"logs": logs}
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
logs_info = await loop.run_in_executor(None, get_logs)
|
|
|
|
return json.dumps({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
**logs_info
|
|
}, indent=2)
|
|
|
|
|
|
@mcp.tool()
|
|
async def get_dag_result(execution_id: str) -> str:
|
|
"""
|
|
Get the final result from a completed DAG execution.
|
|
Returns task outputs and any errors.
|
|
|
|
Args:
|
|
execution_id: Execution ID
|
|
"""
|
|
if execution_id not in executions:
|
|
return json.dumps({"error": f"Execution {execution_id} not found"})
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def get_result():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
|
|
# Check if container is done
|
|
if container.status != "exited":
|
|
return {"status": "still_running", "message": "Container is still running"}
|
|
|
|
# Get logs (result should be in stdout as JSON)
|
|
logs = container.logs().decode("utf-8")
|
|
|
|
# Try to parse result from logs
|
|
result_marker = "=== DAG EXECUTION RESULT ==="
|
|
if result_marker in logs:
|
|
result_start = logs.index(result_marker) + len(result_marker)
|
|
result_json = logs[result_start:].strip()
|
|
try:
|
|
return {"result": json.loads(result_json)}
|
|
except json.JSONDecodeError:
|
|
return {"result": result_json}
|
|
|
|
return {"logs": logs}
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
result_info = await loop.run_in_executor(None, get_result)
|
|
|
|
return json.dumps({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
**result_info
|
|
}, indent=2)
|
|
|
|
|
|
@mcp.tool()
|
|
async def stop_dag_execution(execution_id: str) -> str:
|
|
"""
|
|
Stop a running DAG execution.
|
|
Forcefully terminates the container.
|
|
|
|
Args:
|
|
execution_id: Execution ID
|
|
"""
|
|
if execution_id not in executions:
|
|
return json.dumps({"error": f"Execution {execution_id} not found"})
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def stop_container():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
container.stop(timeout=10)
|
|
return {"stopped": True}
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
stop_info = await loop.run_in_executor(None, stop_container)
|
|
|
|
if stop_info.get("stopped"):
|
|
execution["status"] = "stopped"
|
|
|
|
return json.dumps({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
**stop_info
|
|
}, indent=2)
|
|
|
|
|
|
@mcp.tool()
|
|
async def list_dag_executions(limit: int = 10, status: str = None) -> str:
|
|
"""
|
|
List all DAG executions.
|
|
Shows recent executions with their status.
|
|
|
|
Args:
|
|
limit: Maximum number of executions to return (default: 10)
|
|
status: Filter by status (running, completed, failed)
|
|
"""
|
|
# Filter and sort executions
|
|
filtered = []
|
|
for exec_id, execution in executions.items():
|
|
if status and execution["status"] != status:
|
|
continue
|
|
filtered.append({
|
|
"execution_id": exec_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
"container_name": execution["container_name"],
|
|
})
|
|
|
|
# Sort by created_at (most recent first) and limit
|
|
filtered = filtered[-limit:]
|
|
|
|
return json.dumps({
|
|
"total": len(filtered),
|
|
"executions": filtered
|
|
}, indent=2)
|
|
|
|
|
|
# ============================================================
|
|
# REST API Endpoints for Web UI
|
|
# ============================================================
|
|
|
|
# SSE subscribers for real-time updates
|
|
sse_subscribers: Set[asyncio.Queue] = set()
|
|
|
|
|
|
async def broadcast_sse_event(event_type: str, data: dict):
|
|
"""Broadcast an SSE event to all subscribers."""
|
|
message = f"event: {event_type}\ndata: {json.dumps(data)}\n\n"
|
|
dead_queues = []
|
|
for queue in sse_subscribers:
|
|
try:
|
|
queue.put_nowait(message)
|
|
except asyncio.QueueFull:
|
|
dead_queues.append(queue)
|
|
for q in dead_queues:
|
|
sse_subscribers.discard(q)
|
|
|
|
|
|
async def api_list_executions(request: Request):
|
|
"""GET /api/executions - List all executions."""
|
|
status_filter = request.query_params.get("status")
|
|
limit = int(request.query_params.get("limit", "50"))
|
|
|
|
filtered = []
|
|
for exec_id, execution in executions.items():
|
|
if status_filter and execution["status"] != status_filter:
|
|
continue
|
|
filtered.append({
|
|
"execution_id": exec_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
"container_name": execution["container_name"],
|
|
"created_at": execution.get("created_at"),
|
|
})
|
|
|
|
filtered = filtered[-limit:]
|
|
return JSONResponse({
|
|
"total": len(filtered),
|
|
"executions": filtered
|
|
})
|
|
|
|
|
|
async def api_get_execution(request: Request):
|
|
"""GET /api/executions/{execution_id} - Get execution details."""
|
|
execution_id = request.path_params["execution_id"]
|
|
|
|
if execution_id not in executions:
|
|
return JSONResponse({"error": f"Execution {execution_id} not found"}, status_code=404)
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def get_status():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
return {
|
|
"container_status": container.status,
|
|
"exit_code": container.attrs.get("State", {}).get("ExitCode")
|
|
}
|
|
except Exception as e:
|
|
return {"container_status": "not_found", "error": str(e)}
|
|
|
|
status_info = await loop.run_in_executor(None, get_status)
|
|
|
|
# Update execution status based on container status
|
|
container_status = status_info["container_status"]
|
|
if container_status == "running":
|
|
execution["status"] = "running"
|
|
elif container_status == "exited":
|
|
exit_code = status_info.get("exit_code", -1)
|
|
execution["status"] = "completed" if exit_code == 0 else "failed"
|
|
else:
|
|
execution["status"] = container_status
|
|
|
|
return JSONResponse({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
"container_name": execution["container_name"],
|
|
"container_status": container_status,
|
|
"exit_code": status_info.get("exit_code"),
|
|
"created_at": execution.get("created_at"),
|
|
"task_procedure": execution.get("task_procedure"),
|
|
"task_statuses": execution.get("task_statuses", {}),
|
|
})
|
|
|
|
|
|
async def api_get_execution_logs(request: Request):
|
|
"""GET /api/executions/{execution_id}/logs - Get execution logs."""
|
|
execution_id = request.path_params["execution_id"]
|
|
tail = int(request.query_params.get("tail", "100"))
|
|
|
|
if execution_id not in executions:
|
|
return JSONResponse({"error": f"Execution {execution_id} not found"}, status_code=404)
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def get_logs():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
logs = container.logs(tail=tail, timestamps=True).decode("utf-8")
|
|
return {"logs": logs}
|
|
except Exception as e:
|
|
return {"error": str(e), "logs": ""}
|
|
|
|
logs_info = await loop.run_in_executor(None, get_logs)
|
|
|
|
return JSONResponse({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
**logs_info
|
|
})
|
|
|
|
|
|
async def api_get_execution_result(request: Request):
|
|
"""GET /api/executions/{execution_id}/result - Get execution result."""
|
|
execution_id = request.path_params["execution_id"]
|
|
|
|
if execution_id not in executions:
|
|
return JSONResponse({"error": f"Execution {execution_id} not found"}, status_code=404)
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def get_result():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
|
|
if container.status != "exited":
|
|
return {"status": "still_running", "message": "Container is still running"}
|
|
|
|
logs = container.logs().decode("utf-8")
|
|
|
|
result_marker = "=== DAG EXECUTION RESULT ==="
|
|
if result_marker in logs:
|
|
result_start = logs.index(result_marker) + len(result_marker)
|
|
result_json = logs[result_start:].strip()
|
|
try:
|
|
return {"result": json.loads(result_json)}
|
|
except json.JSONDecodeError:
|
|
return {"result": result_json}
|
|
|
|
return {"logs": logs}
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
result_info = await loop.run_in_executor(None, get_result)
|
|
|
|
return JSONResponse({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
**result_info
|
|
})
|
|
|
|
|
|
async def api_update_task_status(request: Request):
|
|
"""POST /api/executions/{execution_id}/task-status - Update task status from dag_executor."""
|
|
execution_id = request.path_params["execution_id"]
|
|
|
|
if execution_id not in executions:
|
|
return JSONResponse({"error": f"Execution {execution_id} not found"}, status_code=404)
|
|
|
|
try:
|
|
body = await request.json()
|
|
task_name = body.get("task_name")
|
|
status = body.get("status") # "running", "completed", "failed"
|
|
output = body.get("output", "")
|
|
error = body.get("error", "")
|
|
|
|
if not task_name or not status:
|
|
return JSONResponse({"error": "task_name and status are required"}, status_code=400)
|
|
|
|
execution = executions[execution_id]
|
|
|
|
# Initialize task_statuses if not present
|
|
if "task_statuses" not in execution:
|
|
execution["task_statuses"] = {}
|
|
|
|
execution["task_statuses"][task_name] = {
|
|
"status": status,
|
|
"output": output,
|
|
"error": error,
|
|
"updated_at": datetime.utcnow().isoformat() + "Z"
|
|
}
|
|
|
|
# Broadcast task status update
|
|
await broadcast_sse_event("task_status_update", {
|
|
"execution_id": execution_id,
|
|
"task_name": task_name,
|
|
"status": status,
|
|
"output": output,
|
|
"error": error
|
|
})
|
|
|
|
return JSONResponse({"success": True})
|
|
|
|
except json.JSONDecodeError:
|
|
return JSONResponse({"error": "Invalid JSON body"}, status_code=400)
|
|
|
|
|
|
async def api_update_execution_status(request: Request):
|
|
"""POST /api/executions/{execution_id}/status - Update execution status from dag_executor."""
|
|
execution_id = request.path_params["execution_id"]
|
|
|
|
if execution_id not in executions:
|
|
return JSONResponse({"error": f"Execution {execution_id} not found"}, status_code=404)
|
|
|
|
try:
|
|
body = await request.json()
|
|
status = body.get("status") # "completed", "failed", "partial"
|
|
error = body.get("error", "")
|
|
|
|
if not status:
|
|
return JSONResponse({"error": "status is required"}, status_code=400)
|
|
|
|
execution = executions[execution_id]
|
|
execution["status"] = status
|
|
if error:
|
|
execution["error"] = error
|
|
|
|
logger.info(f"[Execution:{execution_id}] Status updated to: {status}")
|
|
|
|
# Broadcast execution status update
|
|
await broadcast_sse_event("execution_update", {
|
|
"execution_id": execution_id,
|
|
"status": status,
|
|
"error": error
|
|
})
|
|
|
|
return JSONResponse({"success": True})
|
|
|
|
except json.JSONDecodeError:
|
|
return JSONResponse({"error": "Invalid JSON body"}, status_code=400)
|
|
|
|
|
|
async def api_stop_execution(request: Request):
|
|
"""POST /api/executions/{execution_id}/stop - Stop execution."""
|
|
execution_id = request.path_params["execution_id"]
|
|
|
|
if execution_id not in executions:
|
|
return JSONResponse({"error": f"Execution {execution_id} not found"}, status_code=404)
|
|
|
|
execution = executions[execution_id]
|
|
container_id = execution["container_id"]
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def stop_container():
|
|
client = get_docker_client()
|
|
try:
|
|
container = client.containers.get(container_id)
|
|
container.stop(timeout=10)
|
|
return {"stopped": True}
|
|
except Exception as e:
|
|
return {"error": str(e)}
|
|
|
|
stop_info = await loop.run_in_executor(None, stop_container)
|
|
|
|
if stop_info.get("stopped"):
|
|
execution["status"] = "stopped"
|
|
# Broadcast status change
|
|
await broadcast_sse_event("execution_update", {
|
|
"execution_id": execution_id,
|
|
"status": "stopped"
|
|
})
|
|
|
|
return JSONResponse({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution["execution_name"],
|
|
**stop_info
|
|
})
|
|
|
|
|
|
async def sse_executions(request: Request):
|
|
"""GET /sse/executions - SSE stream for execution updates."""
|
|
|
|
async def event_generator():
|
|
queue = asyncio.Queue(maxsize=100)
|
|
sse_subscribers.add(queue)
|
|
|
|
try:
|
|
# Send initial execution list
|
|
initial_data = []
|
|
for exec_id, execution in executions.items():
|
|
initial_data.append({
|
|
"execution_id": exec_id,
|
|
"execution_name": execution["execution_name"],
|
|
"status": execution["status"],
|
|
"container_name": execution["container_name"],
|
|
"created_at": execution.get("created_at"),
|
|
})
|
|
|
|
yield f"event: execution_list\ndata: {json.dumps(initial_data)}\n\n"
|
|
|
|
# Keep connection alive and send updates
|
|
while True:
|
|
try:
|
|
# Wait for new messages with timeout for keep-alive
|
|
message = await asyncio.wait_for(queue.get(), timeout=15.0)
|
|
yield message
|
|
except asyncio.TimeoutError:
|
|
# Send keep-alive ping
|
|
yield ": ping\n\n"
|
|
finally:
|
|
sse_subscribers.discard(queue)
|
|
|
|
return StreamingResponse(
|
|
event_generator(),
|
|
media_type="text/event-stream",
|
|
headers={
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
}
|
|
)
|
|
|
|
|
|
async def monitor_container_status():
|
|
"""Background task to monitor container status for running executions."""
|
|
while True:
|
|
try:
|
|
await asyncio.sleep(5) # Check every 5 seconds
|
|
|
|
# Find running executions
|
|
running_executions = [
|
|
(exec_id, exec_data)
|
|
for exec_id, exec_data in executions.items()
|
|
if exec_data["status"] == "running"
|
|
]
|
|
|
|
if not running_executions:
|
|
continue
|
|
|
|
loop = asyncio.get_event_loop()
|
|
|
|
def check_containers():
|
|
client = get_docker_client()
|
|
results = []
|
|
for exec_id, exec_data in running_executions:
|
|
try:
|
|
container = client.containers.get(exec_data["container_id"])
|
|
if container.status == "exited":
|
|
exit_code = container.attrs.get("State", {}).get("ExitCode", -1)
|
|
new_status = "completed" if exit_code == 0 else "failed"
|
|
results.append((exec_id, new_status))
|
|
except Exception:
|
|
# Container not found, mark as failed
|
|
results.append((exec_id, "failed"))
|
|
return results
|
|
|
|
status_changes = await loop.run_in_executor(None, check_containers)
|
|
|
|
for exec_id, new_status in status_changes:
|
|
if executions[exec_id]["status"] == "running":
|
|
executions[exec_id]["status"] = new_status
|
|
logger.info(f"[Monitor] Execution {exec_id} status changed to: {new_status}")
|
|
await broadcast_sse_event("execution_update", {
|
|
"execution_id": exec_id,
|
|
"status": new_status
|
|
})
|
|
|
|
except asyncio.CancelledError:
|
|
logger.info("[Monitor] Container monitoring task cancelled")
|
|
break
|
|
except Exception as e:
|
|
logger.error(f"[Monitor] Error in container monitoring: {e}")
|
|
await asyncio.sleep(5)
|
|
|
|
|
|
async def serve_web_ui(request: Request):
|
|
"""Serve the web UI index.html."""
|
|
web_dir = os.path.join(os.path.dirname(__file__), "..", "..", "web-client", "dist")
|
|
index_path = os.path.join(web_dir, "index.html")
|
|
|
|
if os.path.exists(index_path):
|
|
return FileResponse(index_path)
|
|
else:
|
|
return HTMLResponse(
|
|
"<h1>Web UI Not Built</h1><p>Run 'cd web-client && npm run build' to build the UI.</p>",
|
|
status_code=404
|
|
)
|
|
|
|
|
|
def create_app():
|
|
"""Create the combined ASGI app with health check, OAuth, MCP, and Web UI."""
|
|
from contextlib import asynccontextmanager
|
|
from starlette.middleware import Middleware
|
|
from starlette.middleware.trustedhost import TrustedHostMiddleware
|
|
|
|
@asynccontextmanager
|
|
async def lifespan(app):
|
|
"""Manage MCP session manager lifecycle."""
|
|
async with mcp.session_manager.run():
|
|
logger.info(f"OAuth authentication: {'ENABLED' if OAUTH_ENABLED else 'DISABLED'}")
|
|
logger.info("Web UI available at /monitor")
|
|
|
|
# Start background container monitoring task
|
|
monitor_task = asyncio.create_task(monitor_container_status())
|
|
logger.info("Container status monitoring started")
|
|
|
|
try:
|
|
yield
|
|
finally:
|
|
monitor_task.cancel()
|
|
try:
|
|
await monitor_task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
logger.info("Container status monitoring stopped")
|
|
|
|
# Get the MCP ASGI app
|
|
mcp_app = mcp.streamable_http_app()
|
|
|
|
# Web UI static files directory
|
|
web_dist_dir = os.path.join(os.path.dirname(__file__), "..", "..", "web-client", "dist")
|
|
|
|
# Create Starlette app with health endpoint, OAuth metadata, REST API, and MCP
|
|
routes = [
|
|
# Health check
|
|
Route("/health", health_check, methods=["GET"]),
|
|
# OAuth metadata
|
|
Route("/.well-known/oauth-protected-resource", oauth_protected_resource, methods=["GET"]),
|
|
# REST API for Web UI
|
|
Route("/api/executions", api_list_executions, methods=["GET"]),
|
|
Route("/api/executions/{execution_id}", api_get_execution, methods=["GET"]),
|
|
Route("/api/executions/{execution_id}/logs", api_get_execution_logs, methods=["GET"]),
|
|
Route("/api/executions/{execution_id}/result", api_get_execution_result, methods=["GET"]),
|
|
Route("/api/executions/{execution_id}/stop", api_stop_execution, methods=["POST"]),
|
|
Route("/api/executions/{execution_id}/task-status", api_update_task_status, methods=["POST"]),
|
|
Route("/api/executions/{execution_id}/status", api_update_execution_status, methods=["POST"]),
|
|
# SSE for real-time updates
|
|
Route("/sse/executions", sse_executions, methods=["GET"]),
|
|
# Web UI
|
|
Route("/monitor", serve_web_ui, methods=["GET"]),
|
|
]
|
|
|
|
# Add static files mount if dist directory exists
|
|
if os.path.exists(web_dist_dir):
|
|
routes.append(Mount("/assets", StaticFiles(directory=os.path.join(web_dist_dir, "assets")), name="assets"))
|
|
|
|
app = Starlette(
|
|
routes=routes,
|
|
lifespan=lifespan,
|
|
middleware=[
|
|
Middleware(TrustedHostMiddleware, allowed_hosts=["*"]),
|
|
Middleware(OAuthMiddleware),
|
|
]
|
|
)
|
|
|
|
# Mount MCP app at root (FastMCP provides /mcp endpoint internally)
|
|
app.mount("/", mcp_app)
|
|
|
|
return app
|
|
|
|
|
|
if __name__ == "__main__":
|
|
logging.basicConfig(level=logging.INFO)
|
|
logger.info(f"Starting MCP server on port {MCP_PORT}")
|
|
|
|
app = create_app()
|
|
uvicorn.run(app, host="0.0.0.0", port=MCP_PORT)
|