Team Ai
Apppublic

Veer15/openenv-distributed-systems-debugging

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
env.py468 linesDownload Raw Back to server
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