Veer15/openenv-distributed-systems-debugging
0
1import json2import os3import subprocess4import time5from pathlib import Path6from typing import Any7 8from .constants import (9 DEFAULT_CONFIGS,10 NO_COMMAND_PROVIDED_SENTINEL,11 TASK_MAX_STEPS,12 TaskName,13)14from .fault_injector import inject_fault15from .graders import grade_task16from .metrics_poller import MetricsPoller17from .models import Action, Observation, StepResult18from .process_manager import ProcessManager19 20 21class DistributedDebugEnv:22 """OpenEnv-compatible distributed systems debugging environment."""23 24 def __init__(25 self, project_root: Path | None = None, mesh_root: Path | None = None26 ) -> None:27 self.project_root = (28 project_root or Path(__file__).resolve().parent.parent29 ).resolve()30 self.mesh_root = (31 mesh_root or Path(os.getenv("MESH_ROOT", self.project_root / "mesh"))32 ).resolve()33 34 self._process_manager = ProcessManager(35 project_root=self.project_root, mesh_root=self.mesh_root36 )37 self._metrics_poller = MetricsPoller(poll_interval_s=2.0)38 39 self.current_task: TaskName | None = None40 self.max_steps: int = 041 self.step_count: int = 042 self.last_exit_code: int = 043 self.prev_observation: Observation | None = None44 self._baselines: dict[str, int] = {45 "baseline_worker_restart_count": 0,46 "baseline_consumer_stall_count": 0,47 }48 self._seen_diagnostic_signatures: set[str] = set()49 self._command_counts: dict[str, int] = {}50 self._last_grader_score: float = 0.051 52 def start(self) -> None:53 if not self._metrics_poller.is_alive():54 self._metrics_poller.start()55 56 def close(self) -> None:57 self._metrics_poller.stop()58 59 def _write_json(self, path: Path, payload: dict[str, Any]) -> None:60 path.parent.mkdir(parents=True, exist_ok=True)61 path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")62 63 def _restore_defaults(self) -> None:64 self._write_json(65 self.mesh_root / "registry.json",66 {67 "services": {68 "auth": {"host": "localhost", "port": 3001, "protocol": "http"},69 "redis": {"host": "localhost", "port": 6379, "protocol": "tcp"},70 "worker": {71 "host": "localhost",72 "port": None,73 "protocol": "internal",74 },75 }76 },77 )78 self._write_json(79 self.mesh_root / "auth" / "config.json", DEFAULT_CONFIGS["auth"]80 )81 self._write_json(82 self.mesh_root / "gateway" / "config.json", DEFAULT_CONFIGS["gateway"]83 )84 self._write_json(85 self.mesh_root / "gateway" / "blocked_routes.json",86 DEFAULT_CONFIGS["blocked_routes"],87 )88 self._write_json(89 self.mesh_root / "worker" / "config.json", DEFAULT_CONFIGS["worker"]90 )91 self._write_json(92 self.mesh_root / "worker" / "job_generator_config.json",93 DEFAULT_CONFIGS["job_generator"],94 )95 96 def _truncate_logs(self) -> None:97 for service in ["gateway", "auth", "worker", "job_gen"]:98 Path(f"/tmp/{service}.log").write_text("", encoding="utf-8")99 100 def _reset_runtime_counters(self) -> None:101 Path("/tmp/worker_restart_count").write_text("0", encoding="utf-8")102 Path("/tmp/consumer_stall_count").write_text("0", encoding="utf-8")103 104 def _redis_flush(self) -> None:105 subprocess.run(106 ["redis-cli", "FLUSHDB"], check=True, capture_output=True, text=True107 )108 109 def _read_float(self, value: str, default: float = 0.0) -> float:110 try:111 return float(value)112 except (TypeError, ValueError):113 return default114 115 def _is_route_blocked(self) -> bool:116 blocked_file = self.mesh_root / "gateway" / "blocked_routes.json"117 try:118 payload = json.loads(blocked_file.read_text(encoding="utf-8"))119 blocked = payload.get("blocked", [])120 return "gateway->redis" in blocked121 except Exception:122 return False123 124 def _is_lock_present(self) -> bool:125 result = subprocess.run(126 ["redis-cli", "EXISTS", "LOCK:job_processor"],127 capture_output=True,128 text=True,129 timeout=2,130 check=False,131 )132 return result.stdout.strip() == "1"133 134 def _is_cascading_timeout_resolved(self) -> bool:135 auth_config_file = self.mesh_root / "auth" / "config.json"136 gateway_config_file = self.mesh_root / "gateway" / "config.json"137 try:138 auth_payload = json.loads(auth_config_file.read_text(encoding="utf-8"))139 gateway_payload = json.loads(140 gateway_config_file.read_text(encoding="utf-8")141 )142 except Exception:143 return False144 145 auth_delay_ms = self._read_float(auth_payload.get("delay_ms"), default=0.0)146 auth_timeout_ms = self._read_float(147 gateway_payload.get("auth_timeout_ms"), default=0.0148 )149 if auth_timeout_ms <= 0:150 return False151 return auth_delay_ms <= auth_timeout_ms152 153 def _is_registry_auth_default(self) -> bool:154 registry_file = self.mesh_root / "registry.json"155 try:156 payload = json.loads(registry_file.read_text(encoding="utf-8"))157 auth_service = payload["services"]["auth"]158 except Exception:159 return False160 161 return (162 auth_service.get("host") == "localhost"163 and int(auth_service.get("port", 0)) == 3001164 and auth_service.get("protocol") == "http"165 )166 167 def _job_generator_interval_ms(self) -> int:168 config_file = self.mesh_root / "worker" / "job_generator_config.json"169 try:170 payload = json.loads(config_file.read_text(encoding="utf-8"))171 except Exception:172 return 0173 174 try:175 return int(payload.get("interval_ms", 0))176 except (TypeError, ValueError):177 return 0178 179 def _is_job_generator_rate_resolved(self) -> bool:180 return self._job_generator_interval_ms() >= int(181 DEFAULT_CONFIGS["job_generator"]["interval_ms"]182 )183 184 def _build_grader_context(self) -> dict[str, Any]:185 return {186 **self._baselines,187 "route_blocked": self._is_route_blocked(),188 "lock_exists": self._is_lock_present(),189 "cascading_timeout_resolved": self._is_cascading_timeout_resolved(),190 "registry_auth_matches_default": self._is_registry_auth_default(),191 "job_generator_interval_ms": self._job_generator_interval_ms(),192 "job_generator_rate_resolved": self._is_job_generator_rate_resolved(),193 }194 195 def _blocked_command(self, command: str) -> bool:196 dangerous_patterns = [197 "rm -rf /",198 "kill -9 1",199 "pkill -f uvicorn",200 "> /tmp/gateway.log",201 "> /tmp/auth.log",202 "> /tmp/worker.log",203 ]204 normalized = command.strip().lower()205 return any(pattern in normalized for pattern in dangerous_patterns)206 207 def _run_command(self, command: str) -> tuple[str, str | None]:208 if command.strip() == NO_COMMAND_PROVIDED_SENTINEL:209 self.last_exit_code = 2210 return (211 "No command provided by model. Expected JSON with a command field.",212 "no_command_provided",213 )214 215 if self._blocked_command(command):216 self.last_exit_code = 1217 return (218 "BLOCKED: This command would damage the environment infrastructure.",219 "blocked_command",220 )221 222 try:223 result = subprocess.run(224 command,225 shell=True,226 capture_output=True,227 text=True,228 timeout=10,229 cwd="/",230 env={231 **os.environ,232 "PATH": "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin",233 },234 check=False,235 )236 self.last_exit_code = result.returncode237 output = (result.stdout + result.stderr).strip() or "(no output)"238 return output, None239 except subprocess.TimeoutExpired:240 self.last_exit_code = 124241 return "Command timed out after 10 seconds.", "timeout"242 except Exception as exc:243 self.last_exit_code = 1244 return f"Command execution error: {exc}", str(exc)245 246 def _command_signature(self, command: str) -> str:247 return " ".join(command.strip().lower().split())248 249 def _is_diagnostic_command(self, command: str) -> bool:250 diagnostic_keywords = [251 "cat",252 "curl",253 "redis-cli",254 "ps",255 "ls",256 "grep",257 "tail",258 "jq",259 "lrange",260 "llen",261 "keys",262 "ttl",263 "get",264 ]265 normalized = command.lower()266 return any(keyword in normalized for keyword in diagnostic_keywords)267 268 def _is_state_change_command(self, command: str) -> bool:269 normalized = command.lower()270 state_change_patterns = [271 "kill -hup",272 "redis-cli del",273 "redis-cli lrem",274 "redis-cli set",275 "redis-cli flushdb",276 "echo '{",277 "> /mesh/",278 "tee /mesh/",279 ]280 return any(pattern in normalized for pattern in state_change_patterns)281 282 def _compute_reward(283 self,284 command: str,285 current: Observation,286 previous: Observation,287 grader_score: float,288 previous_grader_score: float,289 command_error: str | None,290 ) -> float:291 if command_error == "no_command_provided":292 return 0.01293 294 if grader_score >= 0.95:295 return 0.99296 297 reward = grader_score * 0.75298 signature = self._command_signature(command)299 signature_count = self._command_counts.get(signature, 0) + 1300 self._command_counts[signature] = signature_count301 302 if (303 self._is_diagnostic_command(command)304 and signature not in self._seen_diagnostic_signatures305 ):306 reward += 0.02307 self._seen_diagnostic_signatures.add(signature)308 309 if self._is_state_change_command(command):310 reward += 0.03311 312 if grader_score > previous_grader_score + 1e-4:313 reward += 0.15314 else:315 reward -= 0.05316 317 if (318 current.metrics.gateway_success_rate319 > previous.metrics.gateway_success_rate + 1e-3320 ):321 reward += 0.05322 323 if current.metrics.queue_depth < previous.metrics.queue_depth:324 reward += 0.05325 326 if current.metrics.worker_restart_count < previous.metrics.worker_restart_count:327 reward += 0.03328 329 if current.metrics.consumer_stall_count < previous.metrics.consumer_stall_count:330 reward += 0.03331 332 if signature_count > 1:333 reward -= min(0.12, 0.04 * (signature_count - 1))334 335 if command.strip().lower() in {336 "echo",337 "pwd",338 "whoami",339 "date",340 "true",341 "false",342 }:343 reward -= 0.08344 345 if self.last_exit_code != 0 and command_error not in {346 "blocked_command",347 "no_command_provided",348 }:349 reward -= 0.08350 351 if command_error == "blocked_command":352 reward -= 0.25353 354 return max(0.01, min(0.99, reward))355 356 def _status_block(self, metrics: Any) -> str:357 return (358 "=== pipeline status after reset ===\n"359 "gateway: running\n"360 "auth: running\n"361 "worker: running\n"362 f"queue_depth: {metrics.queue_depth}\n"363 f"gateway_success_rate: {metrics.gateway_success_rate:.2f}"364 )365 366 def reset(self, task_name: TaskName | str) -> Observation:367 task = TaskName.parse(task_name) if isinstance(task_name, str) else task_name368 369 self.current_task = task370 self.max_steps = TASK_MAX_STEPS[task]371 self.step_count = 0372 self._seen_diagnostic_signatures = set()373 self._command_counts = {}374 self._last_grader_score = 0.0375 376 self._truncate_logs()377 self._restore_defaults()378 self._redis_flush()379 self._reset_runtime_counters()380 381 Path("/tmp/current_task").write_text(task.value, encoding="utf-8")382 383 self._process_manager.restart_all()384 if not self._process_manager.wait_healthy(timeout_s=30):385 raise RuntimeError("Services failed health checks after reset")386 387 inject_fault(task, self._process_manager)388 time.sleep(1.0)389 390 self._metrics_poller.poll_once()391 metrics = self._metrics_poller.get_current_metrics()392 393 self._baselines = {394 "baseline_worker_restart_count": metrics.worker_restart_count,395 "baseline_consumer_stall_count": metrics.consumer_stall_count,396 }397 self._last_grader_score = grade_task(398 task, metrics, self._build_grader_context()399 )400 401 observation = Observation(402 command_output=self._status_block(metrics),403 metrics=metrics,404 process_status=self._process_manager.get_status(),405 )406 self.prev_observation = observation407 return observation408 409 def step(self, action: Action) -> StepResult:410 if not self.current_task:411 raise RuntimeError(412 "Environment not initialized. Call reset(task_name) first."413 )414 415 self.step_count += 1416 command_output, command_error = self._run_command(action.command)417 418 self._metrics_poller.poll_once()419 metrics = self._metrics_poller.get_current_metrics()420 421 observation = Observation(422 command_output=command_output,423 metrics=metrics,424 process_status=self._process_manager.get_status(),425 )426 427 previous = self.prev_observation or observation428 previous_grader_score = self._last_grader_score429 grader_score = grade_task(430 self.current_task, metrics, self._build_grader_context()431 )432 reward = self._compute_reward(433 action.command,434 observation,435 previous,436 grader_score,437 previous_grader_score,438 command_error,439 )440 if command_error == "no_command_provided":441 done = self.step_count >= self.max_steps442 else:443 done = grader_score >= 0.95 or self.step_count >= self.max_steps444 445 self._last_grader_score = grader_score446 self.prev_observation = observation447 448 info: dict[str, Any] = {449 "grader_score": round(grader_score, 4),450 "error": command_error,451 "exit_code": self.last_exit_code,452 "task": self.current_task.value if self.current_task else None,453 }454 455 return StepResult(observation=observation, reward=reward, done=done, info=info)456 457 def state(self) -> dict[str, Any]:458 self._metrics_poller.poll_once()459 metrics = self._metrics_poller.get_current_metrics()460 return {461 "task": self.current_task.value if self.current_task else None,462 "step_count": self.step_count,463 "max_steps": self.max_steps,464 "metrics": metrics.model_dump(),465 "process_status": self._process_manager.get_status(),466 "baselines": dict(self._baselines),467 }468 