#!/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()