Files
Liti-agent-Development/dag_executor.py
T

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()