446 lines
15 KiB
Python
446 lines
15 KiB
Python
#!/usr/bin/env python3
|
|
"""
|
|
DAG Executor - Runs inside Docker container.
|
|
Executes task_procedure (DAG) with asyncio parallelism.
|
|
Uses llm_bridge for LLM calls and MCP tool support.
|
|
"""
|
|
|
|
import asyncio
|
|
import json
|
|
import os
|
|
import sys
|
|
import urllib.request
|
|
import urllib.error
|
|
from typing import Any, Dict, Set
|
|
from dataclasses import dataclass, field
|
|
from datetime import datetime
|
|
import logging
|
|
from concurrent.futures import ThreadPoolExecutor, as_completed
|
|
|
|
from llm_bridge.bridge import Bridge
|
|
from llm_bridge.client import MCPClientManager
|
|
|
|
logging.basicConfig(level=logging.INFO)
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Server URL for status reporting (inside Docker network)
|
|
STATUS_REPORT_URL = os.environ.get("STATUS_REPORT_URL", "http://mcp-docker-orchestrator:8000")
|
|
|
|
|
|
@dataclass
|
|
class TaskResult:
|
|
"""Result of a single task execution."""
|
|
task_name: str
|
|
status: str # "completed", "failed"
|
|
output: str = ""
|
|
error: str = ""
|
|
iterations: int = 0
|
|
started_at: str = ""
|
|
completed_at: str = ""
|
|
|
|
|
|
@dataclass
|
|
class DAGExecutionResult:
|
|
"""Result of entire DAG execution."""
|
|
execution_id: str
|
|
execution_name: str
|
|
status: str # "completed", "failed", "partial"
|
|
task_results: Dict[str, TaskResult] = field(default_factory=dict)
|
|
started_at: str = ""
|
|
completed_at: str = ""
|
|
error: str = ""
|
|
|
|
|
|
def report_task_status(execution_id: str, task_name: str, status: str, output: str = "", error: str = ""):
|
|
"""Report task status to the server."""
|
|
try:
|
|
url = f"{STATUS_REPORT_URL}/api/executions/{execution_id}/task-status"
|
|
data = json.dumps({
|
|
"task_name": task_name,
|
|
"status": status,
|
|
"output": output,
|
|
"error": error
|
|
}).encode('utf-8')
|
|
|
|
req = urllib.request.Request(
|
|
url,
|
|
data=data,
|
|
headers={"Content-Type": "application/json"},
|
|
method="POST"
|
|
)
|
|
|
|
with urllib.request.urlopen(req, timeout=5):
|
|
logger.debug(f"[Task:{task_name}] Status reported: {status}")
|
|
except Exception as e:
|
|
logger.warning(f"[Task:{task_name}] Failed to report status: {e}")
|
|
|
|
|
|
def report_execution_status(execution_id: str, status: str, error: str = ""):
|
|
"""Report execution completion status to the server."""
|
|
try:
|
|
url = f"{STATUS_REPORT_URL}/api/executions/{execution_id}/status"
|
|
data = json.dumps({
|
|
"status": status,
|
|
"error": error
|
|
}).encode('utf-8')
|
|
|
|
req = urllib.request.Request(
|
|
url,
|
|
data=data,
|
|
headers={"Content-Type": "application/json"},
|
|
method="POST"
|
|
)
|
|
|
|
with urllib.request.urlopen(req, timeout=5):
|
|
logger.info(f"[DAG] Execution status reported: {status}")
|
|
except Exception as e:
|
|
logger.warning(f"[DAG] Failed to report execution status: {e}")
|
|
|
|
|
|
def get_api_key(provider: str) -> str:
|
|
"""Get API key for the given provider."""
|
|
if provider == "openai":
|
|
return os.environ.get("OPENAI_API_KEY", "")
|
|
elif provider == "anthropic":
|
|
return os.environ.get("ANTHROPIC_API_KEY", "")
|
|
elif provider == "google":
|
|
return os.environ.get("GOOGLE_API_KEY", "")
|
|
else:
|
|
return ""
|
|
|
|
|
|
def create_llm_bridge(provider: str, model: str, api_key: str, mcp_manager: MCPClientManager = None) -> Bridge:
|
|
"""Create LLM Bridge based on provider."""
|
|
if provider == "openai":
|
|
return Bridge(model=model, api_key=api_key, mcp_client_manager=mcp_manager)
|
|
elif provider == "anthropic":
|
|
base_url = "https://api.anthropic.com/v1/"
|
|
return Bridge(model=model, api_key=api_key, mcp_client_manager=mcp_manager, base_url=base_url)
|
|
elif provider == "google":
|
|
base_url = "https://generativelanguage.googleapis.com/v1beta/openai/"
|
|
return Bridge(model=model, api_key=api_key, mcp_client_manager=mcp_manager, base_url=base_url)
|
|
else:
|
|
raise ValueError(f"Unknown LLM provider: {provider}")
|
|
|
|
|
|
async def execute_task(
|
|
task_name: str,
|
|
task_config: Dict[str, Any],
|
|
system_prompt: str,
|
|
mcp_manager: MCPClientManager = None,
|
|
execution_id: str = None
|
|
) -> TaskResult:
|
|
"""Execute a single task using llm_bridge."""
|
|
started_at = datetime.utcnow().isoformat()
|
|
logger.info(f"[Task:{task_name}] Starting...")
|
|
|
|
# Report task as running
|
|
if execution_id:
|
|
report_task_status(execution_id, task_name, "running")
|
|
|
|
try:
|
|
provider = task_config.get("llm_provider", "openai")
|
|
model = task_config.get("llm_model", "gpt-4o")
|
|
prompts = task_config.get("prompts", [])
|
|
|
|
api_key = get_api_key(provider)
|
|
if not api_key:
|
|
raise ValueError(f"No API key found for provider: {provider}")
|
|
|
|
logger.info(f"[Task:{task_name}] Using {provider}/{model}")
|
|
|
|
# Create LLM Bridge
|
|
llm = create_llm_bridge(provider, model, api_key, mcp_manager)
|
|
|
|
# Build messages
|
|
messages = [{"role": "system", "content": system_prompt}]
|
|
messages.extend(prompts)
|
|
|
|
# Execute task loop (similar to Task.run in agent.py)
|
|
iterations = 0
|
|
final_output = ""
|
|
|
|
for i in range(8):
|
|
iterations = i + 1
|
|
logger.info(f"[Task:{task_name}] Iteration {iterations}/8")
|
|
|
|
# Call LLM via Bridge
|
|
result = await llm.process_messages(messages)
|
|
|
|
# Convert to dict format
|
|
messages = [msg.model_dump() if hasattr(msg, 'model_dump') else msg for msg in result]
|
|
|
|
# Find latest assistant content
|
|
content = ""
|
|
for msg in reversed(messages):
|
|
msg_role = msg.get('role') if isinstance(msg, dict) else getattr(msg, 'role', None)
|
|
msg_content = msg.get('content') if isinstance(msg, dict) else getattr(msg, 'content', None)
|
|
if msg_role == "assistant" and msg_content:
|
|
content = msg_content
|
|
break
|
|
|
|
final_output = content
|
|
|
|
# Check for termination
|
|
if content and (content.strip().endswith("**terminate**") or content.strip().endswith("terminate")):
|
|
logger.info(f"[Task:{task_name}] Termination signal received")
|
|
break
|
|
|
|
# Continue prompt
|
|
messages.append({"role": "user", "content": "Continue the next step."})
|
|
|
|
completed_at = datetime.utcnow().isoformat()
|
|
logger.info(f"[Task:{task_name}] Completed in {iterations} iterations")
|
|
|
|
# Report task as completed
|
|
if execution_id:
|
|
report_task_status(execution_id, task_name, "completed", output=final_output[:500] if final_output else "")
|
|
|
|
return TaskResult(
|
|
task_name=task_name,
|
|
status="completed",
|
|
output=final_output,
|
|
iterations=iterations,
|
|
started_at=started_at,
|
|
completed_at=completed_at
|
|
)
|
|
|
|
except Exception as e:
|
|
completed_at = datetime.utcnow().isoformat()
|
|
logger.error(f"[Task:{task_name}] Failed: {e}")
|
|
import traceback
|
|
traceback.print_exc()
|
|
|
|
# Report task as failed
|
|
if execution_id:
|
|
report_task_status(execution_id, task_name, "failed", error=str(e))
|
|
|
|
return TaskResult(
|
|
task_name=task_name,
|
|
status="failed",
|
|
error=str(e),
|
|
started_at=started_at,
|
|
completed_at=completed_at
|
|
)
|
|
|
|
|
|
async def run_dag(
|
|
execution_id: str,
|
|
execution_name: str,
|
|
task_procedure: Dict[str, Dict],
|
|
tasks: Dict[str, Dict],
|
|
system_prompt: str,
|
|
tools_config: Dict = None
|
|
) -> DAGExecutionResult:
|
|
"""
|
|
Execute DAG with proper dependency handling.
|
|
Tasks without dependencies run in parallel.
|
|
"""
|
|
started_at = datetime.utcnow().isoformat()
|
|
logger.info(f"[DAG] Starting execution: {execution_name} ({execution_id})")
|
|
logger.info(f"[DAG] Tasks to execute: {list(task_procedure.keys())}")
|
|
|
|
# Create shared MCP client manager if tools are configured
|
|
mcp_manager = None
|
|
if tools_config:
|
|
try:
|
|
# Use 600 second (10 minute) timeout for MCP connections
|
|
mcp_manager = MCPClientManager.from_dict(tools_config, timeout=600.0)
|
|
logger.info("[DAG] MCP client manager created with 600s timeout")
|
|
except Exception as e:
|
|
logger.warning(f"[DAG] Failed to create MCP manager: {e}")
|
|
|
|
completed: Set[str] = set()
|
|
task_results: Dict[str, TaskResult] = {}
|
|
|
|
# Always mark IN as completed
|
|
completed.add("IN")
|
|
|
|
while "OUT" not in completed:
|
|
# Find tasks ready to execute (all dependencies satisfied)
|
|
ready_tasks = []
|
|
for task_name, config in task_procedure.items():
|
|
if task_name in ("IN", "OUT"):
|
|
continue
|
|
if task_name in completed:
|
|
continue
|
|
|
|
wait_until = config.get("wait_until", [])
|
|
# Remove IN from dependencies (it's always done)
|
|
wait_until = [dep for dep in wait_until if dep != "IN"]
|
|
|
|
if all(dep in completed for dep in wait_until):
|
|
ready_tasks.append(task_name)
|
|
|
|
if not ready_tasks:
|
|
# Check if OUT can be completed
|
|
out_config = task_procedure.get("OUT", {})
|
|
out_deps = out_config.get("wait_until", [])
|
|
out_deps = [dep for dep in out_deps if dep not in ("IN",)]
|
|
|
|
if all(dep in completed for dep in out_deps):
|
|
completed.add("OUT")
|
|
break
|
|
else:
|
|
# Deadlock or no more tasks
|
|
logger.warning("[DAG] No ready tasks but OUT not completed")
|
|
logger.warning(f"[DAG] Completed: {completed}")
|
|
logger.warning(f"[DAG] OUT dependencies: {out_deps}")
|
|
break
|
|
|
|
logger.info(f"[DAG] Executing tasks in parallel with threads: {ready_tasks}")
|
|
|
|
# Execute ready tasks in parallel using ThreadPoolExecutor
|
|
def run_task_sync(name: str) -> TaskResult:
|
|
"""Wrapper to run async task in thread."""
|
|
task_config = tasks.get(name, {})
|
|
# Create new event loop for this thread
|
|
loop = asyncio.new_event_loop()
|
|
asyncio.set_event_loop(loop)
|
|
try:
|
|
return loop.run_until_complete(
|
|
execute_task(name, task_config, system_prompt, mcp_manager, execution_id)
|
|
)
|
|
finally:
|
|
loop.close()
|
|
|
|
# Use ThreadPoolExecutor for true parallel execution
|
|
with ThreadPoolExecutor(max_workers=len(ready_tasks)) as executor:
|
|
futures = {executor.submit(run_task_sync, name): name for name in ready_tasks}
|
|
for future in as_completed(futures):
|
|
task_name = futures[future]
|
|
try:
|
|
result = future.result()
|
|
task_results[result.task_name] = result
|
|
completed.add(result.task_name)
|
|
if result.status == "failed":
|
|
logger.warning(f"[DAG] Task {result.task_name} failed, continuing with other tasks")
|
|
except Exception as e:
|
|
logger.error(f"[DAG] Task {task_name} raised exception: {e}")
|
|
task_results[task_name] = TaskResult(
|
|
task_name=task_name,
|
|
status="failed",
|
|
error=str(e),
|
|
started_at=datetime.utcnow().isoformat(),
|
|
completed_at=datetime.utcnow().isoformat()
|
|
)
|
|
completed.add(task_name)
|
|
|
|
completed_at = datetime.utcnow().isoformat()
|
|
|
|
# Determine overall status
|
|
failed_tasks = [name for name, result in task_results.items() if result.status == "failed"]
|
|
if failed_tasks:
|
|
status = "partial" if len(failed_tasks) < len(task_results) else "failed"
|
|
else:
|
|
status = "completed"
|
|
|
|
logger.info(f"[DAG] Execution {status}. Tasks: {len(task_results)}, Failed: {len(failed_tasks)}")
|
|
|
|
return DAGExecutionResult(
|
|
execution_id=execution_id,
|
|
execution_name=execution_name,
|
|
status=status,
|
|
task_results=task_results,
|
|
started_at=started_at,
|
|
completed_at=completed_at
|
|
)
|
|
|
|
|
|
def main():
|
|
"""Main entry point for DAG executor."""
|
|
print("=" * 60)
|
|
print("DAG Executor Starting")
|
|
print("=" * 60)
|
|
|
|
# Get execution config from environment
|
|
execution_config_str = os.environ.get("EXECUTION_CONFIG", "")
|
|
|
|
if not execution_config_str:
|
|
print("Error: EXECUTION_CONFIG environment variable not set")
|
|
sys.exit(1)
|
|
|
|
try:
|
|
config = json.loads(execution_config_str)
|
|
except json.JSONDecodeError as e:
|
|
print(f"Error: Invalid EXECUTION_CONFIG JSON: {e}")
|
|
sys.exit(1)
|
|
|
|
execution_id = config.get("execution_id", "unknown")
|
|
execution_name = config.get("execution_name", "unknown")
|
|
task_procedure = config.get("task_procedure", {})
|
|
tasks = config.get("tasks", {})
|
|
system_prompt = config.get("system_prompt", "You are a helpful assistant.")
|
|
tools_config = config.get("tools", None)
|
|
|
|
print(f"Execution ID: {execution_id}")
|
|
print(f"Execution Name: {execution_name}")
|
|
print(f"Task Procedure: {json.dumps(task_procedure, indent=2)}")
|
|
print(f"Tasks: {list(tasks.keys())}")
|
|
print(f"Tools configured: {tools_config is not None}")
|
|
print("=" * 60)
|
|
|
|
# Run DAG
|
|
try:
|
|
result = asyncio.run(run_dag(
|
|
execution_id=execution_id,
|
|
execution_name=execution_name,
|
|
task_procedure=task_procedure,
|
|
tasks=tasks,
|
|
system_prompt=system_prompt,
|
|
tools_config=tools_config
|
|
))
|
|
|
|
# Output result as JSON
|
|
print("\n=== DAG EXECUTION RESULT ===")
|
|
|
|
# Convert dataclass to dict for JSON serialization
|
|
result_dict = {
|
|
"execution_id": result.execution_id,
|
|
"execution_name": result.execution_name,
|
|
"status": result.status,
|
|
"started_at": result.started_at,
|
|
"completed_at": result.completed_at,
|
|
"task_results": {
|
|
name: {
|
|
"task_name": tr.task_name,
|
|
"status": tr.status,
|
|
"output": tr.output,
|
|
"error": tr.error,
|
|
"iterations": tr.iterations,
|
|
"started_at": tr.started_at,
|
|
"completed_at": tr.completed_at
|
|
}
|
|
for name, tr in result.task_results.items()
|
|
}
|
|
}
|
|
|
|
print(json.dumps(result_dict, indent=2))
|
|
|
|
# Report execution completion to server
|
|
report_execution_status(execution_id, result.status)
|
|
|
|
# Exit with appropriate code
|
|
sys.exit(0 if result.status == "completed" else 1)
|
|
|
|
except Exception as e:
|
|
print("\n=== DAG EXECUTION RESULT ===")
|
|
error_msg = str(e)
|
|
print(json.dumps({
|
|
"execution_id": execution_id,
|
|
"execution_name": execution_name,
|
|
"status": "failed",
|
|
"error": error_msg
|
|
}, indent=2))
|
|
import traceback
|
|
traceback.print_exc()
|
|
|
|
# Report failure to server
|
|
report_execution_status(execution_id, "failed", error_msg)
|
|
|
|
sys.exit(1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|