""" Ray-based distributed execution engine for hyperscale network automation. Enables parallel execution of: - Config generation across thousands of devices - Batfish analysis on device groups - Concurrent GNS3 deployments - Validation and remediation at scale Works locally (single machine) or on Ray clusters with zero code changes. """ import logging from typing import List, Dict, Any, Optional, Callable, Tuple from dataclasses import dataclass, field from enum import Enum import asyncio # Optional ray import - gracefully degrade if not available try: import ray from ray.util.queue import Queue as RayQueue RAY_AVAILABLE = True except ImportError: RAY_AVAILABLE = False logging.warning("Ray not installed - distributed execution disabled. Install with: pip install ray") import time from datetime import datetime logger = logging.getLogger(__name__) class TaskStatus(Enum): """Task execution status""" PENDING = "pending" RUNNING = "running" SUCCESS = "success" FAILED = "failed" RETRYING = "retrying" @dataclass class TaskResult: """Result from a distributed task execution""" device_id: str status: TaskStatus result: Any = None error: Optional[str] = None duration_seconds: float = 0.0 retry_count: int = 0 timestamp: datetime = field(default_factory=datetime.now) @dataclass class ExecutionProgress: """Real-time progress tracking for fleet operations""" total_devices: int completed: int = 0 failed: int = 0 running: int = 0 pending: int = 0 start_time: datetime = field(default_factory=datetime.now) @property def completion_percentage(self) -> float: """Calculate completion percentage""" if self.total_devices == 0: return 0.0 return (self.completed / self.total_devices) * 100 @property def success_rate(self) -> float: """Calculate success rate of completed tasks""" total_finished = self.completed + self.failed if total_finished == 0: return 0.0 return (self.completed / total_finished) * 100 @property def elapsed_seconds(self) -> float: """Time elapsed since start""" return (datetime.now() - self.start_time).total_seconds() @property def estimated_time_remaining(self) -> Optional[float]: """Estimate time remaining based on current progress""" if self.completed == 0: return None rate = self.completed / self.elapsed_seconds remaining = self.total_devices - (self.completed + self.failed) return remaining / rate if rate > 0 else None def to_dict(self) -> Dict[str, Any]: """Convert to dictionary for serialization""" return { "total_devices": self.total_devices, "completed": self.completed, "failed": self.failed, "running": self.running, "pending": self.pending, "completion_percentage": self.completion_percentage, "success_rate": self.success_rate, "elapsed_seconds": self.elapsed_seconds, "estimated_time_remaining": self.estimated_time_remaining } # Only define ray-dependent classes if ray is available if RAY_AVAILABLE: @ray.remote class ProgressTracker: """Actor for tracking execution progress across distributed workers""" def __init__(self, total_devices: int): self.progress = ExecutionProgress(total_devices=total_devices) self.results: List[TaskResult] = [] def update_status(self, device_id: str, status: TaskStatus): """Update device status""" if status == TaskStatus.RUNNING: self.progress.running += 1 self.progress.pending -= 1 elif status == TaskStatus.SUCCESS: self.progress.running -= 1 self.progress.completed += 1 elif status == TaskStatus.FAILED: self.progress.running -= 1 self.progress.failed += 1 def add_result(self, result: TaskResult): """Add task result""" self.results.append(result) def get_progress(self) -> Dict[str, Any]: """Get current progress""" return self.progress.to_dict() def get_results(self) -> List[TaskResult]: """Get all results""" return self.results def get_failed_devices(self) -> List[str]: """Get list of failed device IDs""" return [r.device_id for r in self.results if r.status == TaskStatus.FAILED] @ray.remote def generate_device_config(device_id: str, device_data: Dict[str, Any], template_fn: Callable, progress_tracker: Any) -> TaskResult: """ Ray remote function for parallel config generation. Args: device_id: Unique device identifier device_data: Device parameters (hostname, ip, role, etc.) template_fn: Function to generate config from device data progress_tracker: Progress tracking actor Returns: TaskResult with generated config or error """ start_time = time.time() try: # Update status to running ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.RUNNING)) # Generate config config = template_fn(device_data) duration = time.time() - start_time result = TaskResult( device_id=device_id, status=TaskStatus.SUCCESS, result=config, duration_seconds=duration ) # Update status to success ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.SUCCESS)) ray.get(progress_tracker.add_result.remote(result)) return result except Exception as e: duration = time.time() - start_time result = TaskResult( device_id=device_id, status=TaskStatus.FAILED, error=str(e), duration_seconds=duration ) ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.FAILED)) ray.get(progress_tracker.add_result.remote(result)) return result @ray.remote def analyze_device_config(device_id: str, config: str, batfish_client: Any, progress_tracker: Any) -> TaskResult: """ Ray remote function for parallel Batfish analysis. Args: device_id: Unique device identifier config: Device configuration to analyze batfish_client: Batfish client instance progress_tracker: Progress tracking actor Returns: TaskResult with analysis results or error """ start_time = time.time() try: ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.RUNNING)) # Run Batfish analysis analysis = batfish_client.analyze_configs({device_id: config}) duration = time.time() - start_time result = TaskResult( device_id=device_id, status=TaskStatus.SUCCESS, result=analysis, duration_seconds=duration ) ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.SUCCESS)) ray.get(progress_tracker.add_result.remote(result)) return result except Exception as e: duration = time.time() - start_time result = TaskResult( device_id=device_id, status=TaskStatus.FAILED, error=str(e), duration_seconds=duration ) ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.FAILED)) ray.get(progress_tracker.add_result.remote(result)) return result @ray.remote def deploy_to_device(device_id: str, config: str, gns3_client: Any, progress_tracker: Any, max_retries: int = 3) -> TaskResult: """ Ray remote function for parallel device deployment. Args: device_id: Unique device identifier config: Configuration to deploy gns3_client: GNS3 client instance progress_tracker: Progress tracking actor max_retries: Maximum retry attempts on failure Returns: TaskResult with deployment status or error """ start_time = time.time() retry_count = 0 while retry_count <= max_retries: try: if retry_count > 0: ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.RETRYING)) else: ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.RUNNING)) # Deploy config to device deployment_result = gns3_client.apply_config(device_id, config) duration = time.time() - start_time result = TaskResult( device_id=device_id, status=TaskStatus.SUCCESS, result=deployment_result, duration_seconds=duration, retry_count=retry_count ) ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.SUCCESS)) ray.get(progress_tracker.add_result.remote(result)) return result except Exception as e: retry_count += 1 if retry_count > max_retries: duration = time.time() - start_time result = TaskResult( device_id=device_id, status=TaskStatus.FAILED, error=f"Failed after {retry_count} retries: {str(e)}", duration_seconds=duration, retry_count=retry_count - 1 ) ray.get(progress_tracker.update_status.remote(device_id, TaskStatus.FAILED)) ray.get(progress_tracker.add_result.remote(result)) return result # Exponential backoff time.sleep(2 ** retry_count) class RayExecutor: """ Distributed execution engine for hyperscale network automation. Provides parallel execution of config generation, analysis, and deployment across thousands of devices using Ray's distributed computing framework. """ def __init__(self, ray_address: Optional[str] = None, num_cpus: Optional[int] = None): """ Initialize Ray executor. Args: ray_address: Ray cluster address (None for local mode) num_cpus: Number of CPUs to use (None for auto-detect) """ self.ray_address = ray_address self.num_cpus = num_cpus self.initialized = False self._progress_tracker = NoneNone def initialize(self): """Initialize Ray runtime""" if self.initialized: return try: # Check if Ray is already initialized if ray.is_initialized(): logger.info("Ray already initialized") else: # Initialize Ray if self.ray_address: # Connect to existing cluster ray.init(address=self.ray_address) logger.info(f"Connected to Ray cluster at {self.ray_address}") else: # Start local Ray instance init_kwargs = {} if self.num_cpus: init_kwargs['num_cpus'] = self.num_cpus ray.init(**init_kwargs) logger.info(f"Started local Ray instance with {ray.available_resources().get('CPU', 0)} CPUs") self.initialized = True except Exception as e: logger.error(f"Failed to initialize Ray: {e}") raise def shutdown(self): """Shutdown Ray runtime""" if self.initialized and ray.is_initialized(): ray.shutdown() self.initialized = False logger.info("Ray shutdown complete") def parallel_config_generation(self, devices: List[Dict[str, Any]], template_fn: Callable, batch_size: int = 100) -> Tuple[List[TaskResult], ExecutionProgress]: """ Generate configs for multiple devices in parallel. Args: devices: List of device data dicts template_fn: Function to generate config from device data batch_size: Number of devices to process in each batch Returns: Tuple of (results, final_progress) """ self.initialize() # Create progress tracker progress_tracker = ProgressTracker.remote(total_devices=len(devices)) # Initialize pending count ray.get(progress_tracker.update_status.remote("_init_", TaskStatus.PENDING)) for _ in range(len(devices) - 1): ray.get(progress_tracker.update_status.remote("_init_", TaskStatus.PENDING)) # Launch parallel tasks futures = [] for device in devices: future = generate_device_config.remote( device_id=device['device_id'], device_data=device, template_fn=template_fn, progress_tracker=progress_tracker ) futures.append(future) # Process in batches to avoid overwhelming the cluster if len(futures) >= batch_size: ray.get(futures) futures = [] # Wait for remaining tasks if futures: ray.get(futures) # Get final results results = ray.get(progress_tracker.get_results.remote()) final_progress = ray.get(progress_tracker.get_progress.remote()) return results, final_progress def parallel_batfish_analysis(self, configs: Dict[str, str], batfish_client: Any, batch_size: int = 50) -> Tuple[List[TaskResult], ExecutionProgress]: """ Analyze configs in parallel using Batfish. Args: configs: Dict mapping device_id to config string batfish_client: Batfish client instance batch_size: Number of configs to analyze in each batch Returns: Tuple of (results, final_progress) """ self.initialize() progress_tracker = ProgressTracker.remote(total_devices=len(configs)) # Initialize pending count for _ in range(len(configs)): ray.get(progress_tracker.update_status.remote("_init_", TaskStatus.PENDING)) # Launch parallel analysis tasks futures = [] for device_id, config in configs.items(): future = analyze_device_config.remote( device_id=device_id, config=config, batfish_client=batfish_client, progress_tracker=progress_tracker ) futures.append(future) if len(futures) >= batch_size: ray.get(futures) futures = [] if futures: ray.get(futures) results = ray.get(progress_tracker.get_results.remote()) final_progress = ray.get(progress_tracker.get_progress.remote()) return results, final_progressress def parallel_deployment(self, deployments: Dict[str, str], gns3_client: Any, batch_size: int = 20, max_retries: int = 3) -> Tuple[List[TaskResult], ExecutionProgress]: """ Deploy configs to multiple devices in parallel. Args: deployments: Dict mapping device_id to config string gns3_client: GNS3 client instance batch_size: Number of devices to deploy to simultaneously max_retries: Maximum retry attempts per device Returns: Tuple of (results, final_progress) """ self.initialize() progress_tracker = ProgressTracker.remote(total_devices=len(deployments)) # Initialize pending count for _ in range(len(deployments)): ray.get(progress_tracker.update_status.remote("_init_", TaskStatus.PENDING)) # Launch parallel deployment tasks futures = [] for device_id, config in deployments.items(): future = deploy_to_device.remote( device_id=device_id, config=config, gns3_client=gns3_client, progress_tracker=progress_tracker, max_retries=max_retries ) futures.append(future) # Deploy in smaller batches to avoid overwhelming network if len(futures) >= batch_size: ray.get(futures) futures = [] if futures: ray.get(futures) results = ray.get(progress_tracker.get_results.remote()) final_progress = ray.get(progress_tracker.get_progress.remote()) return results, final_progress def get_cluster_resources(self) -> Dict[str, Any]: """Get available cluster resources""" self.initialize() return { 'available': ray.available_resources(), 'total': ray.cluster_resources() } def staggered_rollout(self, deployments: Dict[str, str], gns3_client: Any, stages: List[float] = [0.01, 0.1, 0.5, 1.0], validation_fn: Optional[Callable] = None) -> Tuple[List[TaskResult], ExecutionProgress]: """ Deploy to devices in stages with validation between stages. Implements canary deployment pattern: - Stage 1: 1% of fleet - Stage 2: 10% of fleet - Stage 3: 50% of fleet - Stage 4: 100% of fleet Args: deployments: Dict mapping device_id to config gns3_client: GNS3 client instance stages: List of percentages for each stage (0.0 to 1.0) validation_fn: Optional function to validate stage success Returns: Tuple of (results, final_progress) """ self.initialize() device_ids = list(deployments.keys()) total_devices = len(device_ids) all_results = [] current_index = 0 for stage_pct in stages: stage_count = int(total_devices * stage_pct) - current_index if stage_count <= 0: continue stage_devices = device_ids[current_index:current_index + stage_count] stage_deployments = {did: deployments[did] for did in stage_devices} logger.info(f"Starting stage {stage_pct*100}%: deploying to {len(stage_devices)} devices") # Deploy this stage results, progress = self.parallel_deployment( deployments=stage_deployments, gns3_client=gns3_client, batch_size=min(20, len(stage_devices)) ) all_results.extend(results) # Check for failures failed_count = sum(1 for r in results if r.status == TaskStatus.FAILED) failure_rate = failed_count / len(results) if results else 0 if failure_rate > 0.1: # More than 10% failure rate logger.error(f"Stage failed with {failure_rate*100}% failure rate. Stopping rollout.") # Return partial results final_progress = ExecutionProgress( total_devices=total_devices, completed=sum(1 for r in all_results if r.status == TaskStatus.SUCCESS), failed=sum(1 for r in all_results if r.status == TaskStatus.FAILED) ) return all_results, final_progress.to_dict() # Run validation if provided if validation_fn: try: if not validation_fn(stage_devices, results): logger.error("Stage validation failed. Stopping rollout.") final_progress = ExecutionProgress( total_devices=total_devices, completed=sum(1 for r in all_results if r.status == TaskStatus.SUCCESS), failed=sum(1 for r in all_results if r.status == TaskStatus.FAILED) ) return all_results, final_progress.to_dict() except Exception as e: logger.error(f"Stage validation error: {e}. Stopping rollout.") final_progress = ExecutionProgress( total_devices=total_devices, completed=sum(1 for r in all_results if r.status == TaskStatus.SUCCESS), failed=sum(1 for r in all_results if r.status == TaskStatus.FAILED) ) return all_results, final_progress.to_dict() logger.info(f"Stage {stage_pct*100}% completed successfully") current_index += stage_count # Create final progress final_progress = ExecutionProgress( total_devices=total_devices, completed=sum(1 for r in all_results if r.status == TaskStatus.SUCCESS), failed=sum(1 for r in all_results if r.status == TaskStatus.FAILED) ) return all_results, final_progress.to_dict() else: # Fallback executor when ray is not available class RayExecutor: """Fallback executor without distributed capabilities""" def __init__(self, ray_address: Optional[str] = None, num_cpus: Optional[int] = None): logging.warning("Ray not available - using fallback sequential executor") self.ray_address = ray_address self.num_cpus = num_cpus self.initialized = False def initialize(self): """No-op initialization for fallback""" self.initialized = True def shutdown(self): """No-op shutdown for fallback""" self.initialized = False def parallel_config_generation(self, devices: List[Dict[str, Any]], template_fn: Callable, batch_size: int = 100) -> Tuple[List[TaskResult], ExecutionProgress]: """Sequential fallback for config generation""" raise NotImplementedError( "Distributed execution requires ray. Install with: pip install ray" ) def parallel_batfish_analysis(self, configs: Dict[str, str], batfish_client: Any, batch_size: int = 50) -> Tuple[List[TaskResult], ExecutionProgress]: """Sequential fallback for analysis""" raise NotImplementedError( "Distributed execution requires ray. Install with: pip install ray" ) def parallel_deployment(self, deployments: Dict[str, str], gns3_client: Any, batch_size: int = 20, max_retries: int = 3) -> Tuple[List[TaskResult], ExecutionProgress]: """Sequential fallback for deployment""" raise NotImplementedError( "Distributed execution requires ray. Install with: pip install ray" ) def staged_rollout(self, deployments: Dict[str, str], gns3_client: Any, stages: List[float] = None, validation_fn: Optional[Callable] = None) -> Tuple[List[TaskResult], Dict[str, Any]]: """Sequential fallback for staged rollout""" raise NotImplementedError( "Distributed execution requires ray. Install with: pip install ray" )