Download tests/test_infrastructure_plugins.py from windows2t2/cad-bench: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/windows2t2/cad-bench/resolve/main/tests/test_infrastructure_plugins.py
- Command line
-
hf download hf://windows2t2/cad-bench/tests/test_infrastructure_plugins.py
-
curl -L -o test_infrastructure_plugins.py https://huggingface.co/windows2t2/cad-bench/resolve/main/tests/test_infrastructure_plugins.py
12.8 kB
| from __future__ import annotations | |
| import asyncio | |
| import json | |
| import subprocess | |
| import sys | |
| import threading | |
| from collections.abc import Mapping, Sequence | |
| from pathlib import Path | |
| from typing import Any, BinaryIO | |
| import pytest | |
| from autocad_bench.infrastructure import ( | |
| InfrastructureError, | |
| InfrastructurePreflightRequest, | |
| InfrastructurePreflightResult, | |
| InfrastructureSpec, | |
| RecoveryLaunchRequest, | |
| RolloutLaunchRequest, | |
| WorkerStatus, | |
| available_infrastructure_plugins, | |
| get_infrastructure_plugin, | |
| get_infrastructure_recovery_plugin, | |
| register_infrastructure_plugin, | |
| unregister_infrastructure_plugin, | |
| ) | |
| from autocad_bench.orchestration.batch import ( | |
| BatchConfig, | |
| build_rollout_command, | |
| execute_batch, | |
| preflight, | |
| ) | |
| from autocad_bench.orchestration import resume as resume_batch | |
| class _FakeInfrastructurePlugin: | |
| name = "fake" | |
| def __init__(self) -> None: | |
| self.preflight_requests: list[InfrastructurePreflightRequest] = [] | |
| self.launch_requests: list[RolloutLaunchRequest] = [] | |
| self.shared_commands: list[tuple[str, ...]] = [] | |
| self.reaped_batches: list[str] = [] | |
| self.preflight_report: dict[str, Any] = { | |
| "worker_pool": "test-workers" | |
| } | |
| def validate_spec(self, spec: InfrastructureSpec) -> None: | |
| if spec.settings != {"pool": "test-workers"}: | |
| raise InfrastructureError("fake plugin requires its test worker pool") | |
| def create_session(self, spec: InfrastructureSpec) -> object: | |
| self.validate_spec(spec) | |
| return object() | |
| async def preflight( | |
| self, | |
| spec: InfrastructureSpec, | |
| request: InfrastructurePreflightRequest, | |
| *, | |
| session: Any | None = None, | |
| ) -> InfrastructurePreflightResult: | |
| self.validate_spec(spec) | |
| self.preflight_requests.append(request) | |
| return InfrastructurePreflightResult(report=dict(self.preflight_report)) | |
| def build_rollout_command( | |
| self, | |
| spec: InfrastructureSpec, | |
| request: RolloutLaunchRequest, | |
| ) -> list[str]: | |
| self.validate_spec(spec) | |
| self.launch_requests.append(request) | |
| return ["fake-worker", *request.child_command] | |
| def run_shared( | |
| self, | |
| spec: InfrastructureSpec, | |
| command: Sequence[str], | |
| *, | |
| environment: Mapping[str, str], | |
| output: BinaryIO, | |
| stop_event: threading.Event, | |
| ) -> int: | |
| self.shared_commands.append(tuple(command)) | |
| output_dir = Path(command[command.index("--output-dir") + 1]) | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| (output_dir / "result.json").write_text( | |
| json.dumps( | |
| { | |
| "completed": True, | |
| "artifact_bytes": 1, | |
| "evaluation": {"status": "completed"}, | |
| } | |
| ), | |
| encoding="utf-8", | |
| ) | |
| return 0 | |
| def create_trace_publisher( | |
| self, | |
| spec: InfrastructureSpec, | |
| *, | |
| root: Path, | |
| destination: str, | |
| prefix: str, | |
| session: Any | None, | |
| interval_s: float, | |
| max_concurrency: int, | |
| ) -> Any: | |
| raise InfrastructureError("fake plugin has no trace publisher") | |
| def reap_batch( | |
| self, | |
| spec: InfrastructureSpec, | |
| *, | |
| batch_id: str, | |
| session: Any | None = None, | |
| ) -> list[str]: | |
| self.reaped_batches.append(batch_id) | |
| return ["fake-worker-1"] | |
| class _RecoverableFakeInfrastructurePlugin(_FakeInfrastructurePlugin): | |
| name = "recoverable-fake" | |
| def __init__(self) -> None: | |
| super().__init__() | |
| self.active_workers = {"custom-one": "worker-custom-1"} | |
| self.recovery_commands: list[RecoveryLaunchRequest] = [] | |
| self.terminated_workers: list[str] = [] | |
| def validate_recovery_spec(self, spec: InfrastructureSpec) -> None: | |
| self.validate_spec(spec) | |
| def create_recovery_session(self, spec: InfrastructureSpec) -> object: | |
| self.validate_recovery_spec(spec) | |
| return object() | |
| def list_active_workers( | |
| self, | |
| spec: InfrastructureSpec, | |
| *, | |
| batch_id: str, | |
| session: Any | None = None, | |
| ) -> dict[str, str]: | |
| self.validate_recovery_spec(spec) | |
| return dict(self.active_workers) | |
| def worker_id_from_log( | |
| self, | |
| spec: InfrastructureSpec, | |
| log_text: str, | |
| ) -> str | None: | |
| self.validate_recovery_spec(spec) | |
| marker = "worker_id=" | |
| return log_text.split(marker, 1)[1].split()[0] if marker in log_text else None | |
| def worker_status( | |
| self, | |
| spec: InfrastructureSpec, | |
| *, | |
| worker_id: str, | |
| session: Any | None = None, | |
| ) -> WorkerStatus: | |
| self.validate_recovery_spec(spec) | |
| return WorkerStatus(state="ready", available=True) | |
| def build_recovery_command( | |
| self, | |
| spec: InfrastructureSpec, | |
| request: RecoveryLaunchRequest, | |
| ) -> list[str]: | |
| self.validate_recovery_spec(spec) | |
| self.recovery_commands.append(request) | |
| output_dir = request.child_command[ | |
| request.child_command.index("--output-dir") + 1 | |
| ] | |
| script = ( | |
| "import json, pathlib; " | |
| f"p=pathlib.Path({output_dir!r}); " | |
| "p.mkdir(parents=True, exist_ok=True); " | |
| "(p/'result.json').write_text(json.dumps({" | |
| "'completed': True, 'artifact_bytes': 1, " | |
| "'evaluation': {'status': 'completed'}}))" | |
| ) | |
| return [sys.executable, "-c", script] | |
| def terminate_worker( | |
| self, | |
| spec: InfrastructureSpec, | |
| *, | |
| worker_id: str, | |
| session: Any | None = None, | |
| ) -> None: | |
| self.validate_recovery_spec(spec) | |
| self.terminated_workers.append(worker_id) | |
| def fake_plugin() -> _FakeInfrastructurePlugin: | |
| plugin = _FakeInfrastructurePlugin() | |
| register_infrastructure_plugin("fake", lambda: plugin) | |
| try: | |
| yield plugin | |
| finally: | |
| unregister_infrastructure_plugin("fake") | |
| def recoverable_fake_plugin() -> _RecoverableFakeInfrastructurePlugin: | |
| plugin = _RecoverableFakeInfrastructurePlugin() | |
| register_infrastructure_plugin(plugin.name, lambda: plugin) | |
| try: | |
| yield plugin | |
| finally: | |
| unregister_infrastructure_plugin(plugin.name) | |
| def _config() -> BatchConfig: | |
| return BatchConfig.model_validate( | |
| { | |
| "expected_rollouts": 1, | |
| "max_concurrency": 1, | |
| "evaluation": {"enabled": False}, | |
| "infrastructure": { | |
| "backend": "fake", | |
| "broker_version": "custom-broker-v1", | |
| "settings": {"pool": "test-workers"}, | |
| }, | |
| "rollouts": [ | |
| { | |
| "rollout_id": "custom-one", | |
| "display_name": "Custom worker", | |
| "provider": "bedrock", | |
| "model_id": "example.model", | |
| } | |
| ], | |
| } | |
| ) | |
| def _recoverable_config() -> BatchConfig: | |
| value = _config().model_dump(mode="json") | |
| value["infrastructure"]["backend"] = "recoverable-fake" | |
| return BatchConfig.model_validate(value) | |
| def test_custom_plugin_owns_preflight_and_launch_command( | |
| fake_plugin: _FakeInfrastructurePlugin, | |
| tmp_path: Path, | |
| ) -> None: | |
| config = _config() | |
| report = asyncio.run( | |
| preflight( | |
| config, | |
| environment={}, | |
| allow_partial=False, | |
| check_direct_models=False, | |
| ) | |
| ) | |
| command = build_rollout_command( | |
| config, | |
| config.enabled_rollouts[0], | |
| output_root=tmp_path, | |
| batch_id="batch-custom", | |
| rollout_index=0, | |
| ) | |
| assert report["ready"] is True | |
| assert report["infrastructure_backend"] == "fake" | |
| assert report["worker_pool"] == "test-workers" | |
| assert fake_plugin.preflight_requests[0].requested_workers == 1 | |
| assert command[0] == "fake-worker" | |
| assert "autocad_bench.harness.run" in command | |
| assert "--aws-profile" not in command | |
| assert fake_plugin.launch_requests[0].evaluation_arguments == ( | |
| "--no-auto-evaluate", | |
| ) | |
| def test_registry_reports_builtin_and_rejects_unknown( | |
| fake_plugin: _FakeInfrastructurePlugin, | |
| ) -> None: | |
| assert {"aws", "fake"}.issubset(available_infrastructure_plugins()) | |
| assert get_infrastructure_plugin("fake") is fake_plugin | |
| with pytest.raises(InfrastructureError, match="unknown infrastructure plugin"): | |
| get_infrastructure_plugin("does-not-exist") | |
| with pytest.raises(InfrastructureError, match="does not support.*recovery"): | |
| get_infrastructure_recovery_plugin("fake") | |
| def test_core_controller_imports_without_aws_sdk_modules() -> None: | |
| script = """ | |
| import sys | |
| sys.modules["boto3"] = None | |
| sys.modules["botocore"] = None | |
| sys.modules["botocore.exceptions"] = None | |
| import autocad_bench.harness.run | |
| import autocad_bench.orchestration.batch | |
| import autocad_bench.orchestration.resume | |
| """ | |
| completed = subprocess.run( | |
| [sys.executable, "-c", script], | |
| check=False, | |
| capture_output=True, | |
| text=True, | |
| ) | |
| assert completed.returncode == 0, completed.stderr | |
| def test_plugin_cannot_overwrite_controller_preflight_fields( | |
| fake_plugin: _FakeInfrastructurePlugin, | |
| ) -> None: | |
| fake_plugin.preflight_report = {"ready": True} | |
| report = asyncio.run( | |
| preflight( | |
| _config(), | |
| environment={}, | |
| allow_partial=False, | |
| check_direct_models=False, | |
| ) | |
| ) | |
| assert report["ready"] is False | |
| assert report["issues"] == [ | |
| "fake infrastructure preflight failed: infrastructure preflight report " | |
| "uses controller-owned keys: ready" | |
| ] | |
| def test_execute_routes_shared_lifecycle_and_cleanup_through_plugin( | |
| fake_plugin: _FakeInfrastructurePlugin, | |
| tmp_path: Path, | |
| ) -> None: | |
| config = _config() | |
| report = asyncio.run( | |
| preflight( | |
| config, | |
| environment={}, | |
| allow_partial=False, | |
| check_direct_models=False, | |
| ) | |
| ) | |
| output_root = tmp_path / "custom-run" | |
| exit_code = asyncio.run( | |
| execute_batch( | |
| config, | |
| output_root=output_root, | |
| preflight_report=report, | |
| environment={}, | |
| ) | |
| ) | |
| state = json.loads( | |
| (output_root / "batch-state.json").read_text(encoding="utf-8") | |
| ) | |
| assert exit_code == 0 | |
| assert len(fake_plugin.shared_commands) == 1 | |
| assert fake_plugin.reaped_batches == [state["batch_id"]] | |
| assert state["infrastructure_backend"] == "fake" | |
| assert state["reaped_workers"] == ["fake-worker-1"] | |
| assert state["reaped_instances"] == ["fake-worker-1"] | |
| def test_custom_plugin_discovers_reconnects_and_releases_worker( | |
| recoverable_fake_plugin: _RecoverableFakeInfrastructurePlugin, | |
| tmp_path: Path, | |
| ) -> None: | |
| config = _recoverable_config() | |
| batch_root = tmp_path / "benchmarks" / "custom-run" | |
| rollout_dir = batch_root / "rollouts" / "custom-one" | |
| rollout_dir.mkdir(parents=True) | |
| (batch_root / "logs").mkdir() | |
| (rollout_dir / "run-state.json").write_text( | |
| json.dumps( | |
| { | |
| "task_id": "task-001", | |
| "provider": "bedrock", | |
| "model_id": "example.model", | |
| "broker_version": config.infrastructure.broker_version, | |
| } | |
| ), | |
| encoding="utf-8", | |
| ) | |
| (batch_root / "batch-state.json").write_text( | |
| json.dumps( | |
| { | |
| "batch_id": "batch-custom", | |
| "infrastructure_backend": config.infrastructure.backend, | |
| "infrastructure": config.infrastructure.model_dump(mode="json"), | |
| } | |
| ), | |
| encoding="utf-8", | |
| ) | |
| plan = resume_batch._discover_plan( | |
| [batch_root], | |
| sessions={"recoverable-fake": object()}, | |
| ) | |
| plan_path = tmp_path / "custom-resume-plan.json" | |
| exit_code = asyncio.run(resume_batch.execute_plan(plan, plan_path=plan_path)) | |
| state = json.loads( | |
| Path(plan["state_path"]).read_text(encoding="utf-8") | |
| ) | |
| row = state["entries"]["custom-run/custom-one"] | |
| assert exit_code == 0 | |
| assert plan["infrastructure_backends"] == ["recoverable-fake"] | |
| assert plan["entries"][0]["worker_id"] == "worker-custom-1" | |
| assert "instance_id" not in plan["entries"][0] | |
| assert len(recoverable_fake_plugin.recovery_commands) == 1 | |
| assert recoverable_fake_plugin.terminated_workers == ["worker-custom-1"] | |
| assert row["worker_id"] == "worker-custom-1" | |
| assert row["infrastructure_backend"] == "recoverable-fake" | |
| assert "instance_id" not in row | |