from __future__ import annotations import asyncio import json import signal from pathlib import Path from typing import Any import pytest from botocore.exceptions import ClientError import autocad_bench.infrastructure.aws as aws_module import autocad_bench.orchestration.batch as batch_module from autocad_bench.orchestration.batch import ( BatchState, BatchConfig, build_rollout_command, load_config, preflight, quota_snapshot, reap_batch_instances, rollout_environment, validate_infrastructure, ) def _sample_config() -> BatchConfig: return BatchConfig.model_validate( { "task_id": "task-001", "expected_rollouts": 2, "max_concurrency": 2, "no_wall_timeout": True, "infrastructure": { "backend": "aws", "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a", "subnet-b"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "aws_region": "us-east-1", "session_manager_plugin": "/plugin", }, "evaluation": {"bucket": "test-evaluation-bucket"}, "rollouts": [ { "rollout_id": "openai-one", "display_name": "OpenAI One", "provider": "openai", "model_id": "gpt-one", "reasoning_effort": "xhigh", }, { "rollout_id": "anthropic-two", "display_name": "Anthropic Two", "provider": "anthropic", "model_id": "claude-two", }, ], } ) def test_batch_finish_reconciles_terminal_result_over_stale_row( tmp_path: Path, ) -> None: config = _sample_config() state = BatchState( output_root=tmp_path, batch_id="batch-test", config=config, preflight_report={}, ) state.update( "openai-one", state="infrastructure_failed", execution_status="infrastructure_failed", evaluation_status="failed", exit_code=-1, failure="stale wrapper classification", ) rollout = tmp_path / "rollouts" / "openai-one" rollout.mkdir(parents=True) (rollout / "result.json").write_text( json.dumps( { "execution_status": "completed", "completed": True, "artifact_bytes": 128, "evaluation": {"status": "completed"}, } ) ) state.update( "anthropic-two", state="completed", execution_status="completed", evaluation_status="completed", exit_code=0, ) state.finish(interrupted=False, reaped_instances=[]) value = json.loads((tmp_path / "batch-state.json").read_text()) row = value["rollouts"]["openai-one"] assert row["state"] == "completed" assert row["execution_status"] == "completed" assert row["evaluation_status"] == "completed" assert row["exit_code"] == 0 assert "failure" not in row assert value["state"] == "completed" def test_public_aws_example_is_valid_and_requires_operator_resources() -> None: config_path = Path(__file__).parents[1] / "configs" / "aws.example.json" config = load_config(config_path) assert config.expected_rollouts == 1 assert config.infrastructure.backend == "aws" assert config.infrastructure.aws_profile is None assert config.infrastructure.image_id == "ami-REPLACE_ME" assert config.evaluation.bucket == "your-autocad-bench-bucket" def test_aws_read_retries_transient_signature_failure() -> None: attempts = 0 sleeps: list[float] = [] def call() -> dict[str, bool]: nonlocal attempts attempts += 1 if attempts < 3: raise ClientError( { "Error": { "Code": "InvalidSignatureException", "Message": "clock skew or transient signing failure", } }, "DescribeImages", ) return {"ok": True} assert batch_module._aws_read("test operation", call, sleep=sleeps.append) == { "ok": True } assert attempts == 3 assert sleeps == [1.0, 2.0] def test_infrastructure_preflight_requires_no_ingress_and_one_vpc( monkeypatch: pytest.MonkeyPatch, ) -> None: class Ec2: def describe_images(self, **_: Any) -> dict[str, Any]: return {"Images": [{"ImageId": "ami-test", "State": "available"}]} def describe_subnets(self, **_: Any) -> dict[str, Any]: return { "Subnets": [ {"SubnetId": "subnet-a", "State": "available", "VpcId": "vpc-1"}, {"SubnetId": "subnet-b", "State": "available", "VpcId": "vpc-1"}, ] } def describe_security_groups(self, **_: Any) -> dict[str, Any]: return { "SecurityGroups": [ { "GroupId": "sg-test", "VpcId": "vpc-1", "IpPermissions": [{"IpProtocol": "-1"}], } ] } class Session: def client(self, *_: Any, **__: Any) -> Ec2: return Ec2() monkeypatch.setattr(Path, "is_file", lambda _self: True) with pytest.raises(batch_module.PreflightError, match="no inbound rules"): validate_infrastructure(Session(), _sample_config()) def test_infrastructure_preflight_rejects_stale_broker_image_before_aws() -> None: sample = _sample_config() config = sample.model_copy( update={ "infrastructure": sample.infrastructure.model_copy( update={"broker_version": "windows-autocad-2019-v9"} ) } ) class NoAwsCalls: def client(self, *_: Any, **__: Any) -> Any: raise AssertionError("stale image label must fail before AWS") with pytest.raises(batch_module.PreflightError, match="does not support"): validate_infrastructure(NoAwsCalls(), config) def test_rollout_command_is_isolated_tagged_and_enables_automatic_evaluation() -> None: config = _sample_config() command = build_rollout_command( config, config.enabled_rollouts[1], output_root=Path("/tmp/batch-output"), batch_id="batch-test", rollout_index=1, ) serialized = json.dumps(command) assert "subnet-b" in command assert ["--batch-id", "batch-test"] == command[ command.index("--batch-id") : command.index("--batch-id") + 2 ] assert ["--rollout-id", "anthropic-two"] == command[ command.index("--rollout-id") : command.index("--rollout-id") + 2 ] assert "/tmp/batch-output/rollouts/anthropic-two" in command assert "--cleanup" in command and "terminate" in command admission_index = command.index("--admission-attempts") assert command[admission_index : admission_index + 2] == [ "--admission-attempts", "2", ] assert command[command.index("--handoff-attempts") + 1] == "2" assert command[command.index("--tunnel-startup-attempts") + 1] == "3" assert command[command.index("--tunnel-startup-timeout-s") + 1] == "45.0" assert "--no-wall-timeout" in command assert "--auto-evaluate" in command assert "--vision-judge" in command assert "--evaluator-version" in command assert "OPENAI_API_KEY" in serialized assert "ANTHROPIC_API_KEY" not in serialized assert "openai-secret" not in serialized assert "anthropic-secret" not in serialized assert "Bearer" not in serialized def test_same_model_can_run_distinct_tasks_and_uses_rollout_task_id() -> None: config = BatchConfig.model_validate( { "task_id": "task-001", "expected_rollouts": 2, "max_concurrency": 1, "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "aws_region": "us-east-1", "session_manager_plugin": "/plugin", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "task-001", "display_name": "Basic 001", "task_id": "task-001", "provider": "openai", "model_id": "gpt-same", }, { "rollout_id": "task-002", "display_name": "Basic 002", "task_id": "task-002", "provider": "openai", "model_id": "gpt-same", }, ], } ) command = build_rollout_command( config, config.enabled_rollouts[1], output_root=Path("/tmp/batch-output"), batch_id="batch-test", rollout_index=1, ) task_index = command.index("--task-id") assert command[task_index : task_index + 2] == ["--task-id", "task-002"] def test_rollout_environment_contains_model_and_evaluation_keys_only() -> None: config = _sample_config() environment = { "PATH": "/bin", "OPENAI_API_KEY": "openai-secret", "ANTHROPIC_API_KEY": "anthropic-secret", "AWS_PROFILE": "test-profile", } openai = rollout_environment(config, config.enabled_rollouts[0], environment) anthropic = rollout_environment(config, config.enabled_rollouts[1], environment) assert openai["OPENAI_API_KEY"] == "openai-secret" assert "ANTHROPIC_API_KEY" not in openai assert anthropic["ANTHROPIC_API_KEY"] == "anthropic-secret" assert anthropic["OPENAI_API_KEY"] == "openai-secret" assert openai["AWS_PROFILE"] == anthropic["AWS_PROFILE"] == "test-profile" def test_rollout_command_can_explicitly_disable_automatic_evaluation() -> None: sample = _sample_config() config = sample.model_copy( update={"evaluation": sample.evaluation.model_copy(update={"enabled": False})} ) command = build_rollout_command( config, config.enabled_rollouts[0], output_root=Path("/tmp/batch-output"), batch_id="batch-test", rollout_index=0, ) assert "--no-auto-evaluate" in command assert "--auto-evaluate" not in command assert "--vision-judge" not in command def test_mantle_kimi_profile_pins_chat_model_and_tool_choice() -> None: config = BatchConfig.model_validate( { "task_id": "task-001", "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "mantle-kimi-k2.5-chat", "tool_choice_mode": "specified", "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "session_manager_plugin": "/plugin", "aws_profile": "test-profile", "aws_region": "us-east-1", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "kimi", "display_name": "Kimi K2.5 via AWS Mantle", "provider": "mantle", "model_id": "moonshotai.kimi-k2.5", } ], } ) command = build_rollout_command( config, config.enabled_rollouts[0], output_root=Path("/tmp/mantle-batch"), batch_id="batch-mantle", rollout_index=0, ) assert command[command.index("--provider") + 1] == "mantle" assert command[command.index("--model-id") + 1] == "moonshotai.kimi-k2.5" assert command[command.index("--tool-choice-mode") + 1] == "specified" assert command[command.index("--aws-profile") + 1] == "test-profile" assert config.enabled_rollouts[0].required_key_env is None def test_mantle_kimi_profile_rejects_another_model() -> None: payload = { "task_id": "task-001", "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "mantle-kimi-k2.5-chat", "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "wrong", "display_name": "Wrong model", "provider": "mantle", "model_id": "another-model", } ], } with pytest.raises(ValueError, match="moonshotai.kimi-k2.5"): BatchConfig.model_validate(payload) def test_mantle_grok_profile_pins_exact_model() -> None: payload = { "task_id": "task-001", "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "mantle-grok-4.3-responses", "tool_choice_mode": "specified", "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "grok", "display_name": "Grok 4.3 via AWS Mantle", "provider": "mantle", "model_id": "xai.grok-4.3", } ], } config = BatchConfig.model_validate(payload) assert config.harness_profile == "mantle-grok-4.3-responses" assert config.enabled_rollouts[0].required_key_env is None payload["rollouts"][0]["model_id"] = "another-model" with pytest.raises(ValueError, match="xai.grok-4.3"): BatchConfig.model_validate(payload) def test_fireworks_kimi_fast_profile_pins_router_and_secret() -> None: config = BatchConfig.model_validate( { "task_id": "task-001", "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "fireworks-kimi-k2p6-fast-chat", "tool_choice_mode": "specified", "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "session_manager_plugin": "/plugin", "aws_profile": "test-profile", "aws_region": "us-east-1", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "kimi-fast", "display_name": "Kimi K2.6 Fast via Fireworks", "provider": "fireworks", "model_id": "accounts/fireworks/routers/kimi-k2p6-fast", } ], } ) rollout = config.enabled_rollouts[0] command = build_rollout_command( config, rollout, output_root=Path("/tmp/fireworks-batch"), batch_id="batch-fireworks", rollout_index=0, ) assert command[command.index("--provider") + 1] == "fireworks" assert command[command.index("--model-id") + 1] == ( "accounts/fireworks/routers/kimi-k2p6-fast" ) assert command[command.index("--tool-choice-mode") + 1] == "specified" assert rollout.required_key_env == "FIREWORKS_API_KEY" def test_fireworks_kimi_fast_profile_rejects_another_model() -> None: payload = { "task_id": "task-001", "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "fireworks-kimi-k2p6-fast-chat", "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "wrong", "display_name": "Wrong model", "provider": "fireworks", "model_id": "accounts/fireworks/models/another-model", } ], } with pytest.raises( ValueError, match="accounts/fireworks/routers/kimi-k2p6-fast", ): BatchConfig.model_validate(payload) def test_fireworks_qwen_profile_pins_model_and_secret() -> None: config = BatchConfig.model_validate( { "task_id": "task-001", "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "fireworks-qwen3p7-plus-chat", "tool_choice_mode": "specified", "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-a"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "aws_region": "us-east-1", }, "evaluation": {"enabled": False}, "rollouts": [ { "rollout_id": "qwen", "display_name": "Qwen 3.7 Plus via Fireworks", "provider": "fireworks", "model_id": "accounts/fireworks/models/qwen3p7-plus", } ], } ) rollout = config.enabled_rollouts[0] command = build_rollout_command( config, rollout, output_root=Path("/tmp/fireworks-qwen-batch"), batch_id="batch-fireworks-qwen", rollout_index=0, ) assert command[command.index("--provider") + 1] == "fireworks" assert command[command.index("--model-id") + 1] == ( "accounts/fireworks/models/qwen3p7-plus" ) assert rollout.required_key_env == "FIREWORKS_API_KEY" def test_preflight_checks_full_expected_capacity_before_any_launch( monkeypatch, ) -> None: sample = _sample_config() config = sample.model_copy( update={"evaluation": sample.evaluation.model_copy(update={"enabled": False})} ) requested: list[int] = [] monkeypatch.setattr( aws_module, "validate_infrastructure", lambda _session, _config: {"ami": "ami-test"}, ) def fake_quota(_session: Any, **kwargs: Any) -> dict[str, Any]: requested.append(kwargs["requested_instances"]) return { "enough": True, "requested_vcpus": 8, "remaining_vcpus": 40, } monkeypatch.setattr(aws_module, "quota_snapshot", fake_quota) report = asyncio.run( preflight( config, environment={}, allow_partial=False, check_direct_models=False, session=object(), ) ) assert requested == [2] assert report["ready"] is False assert report["issues"] == [ "missing provider credentials: ANTHROPIC_API_KEY, OPENAI_API_KEY" ] def test_preflight_fails_before_launch_when_automatic_gold_cache_is_missing( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: sample = _sample_config() config = sample.model_copy( update={ "evaluation": sample.evaluation.model_copy( update={ "gold_cache_root": str(tmp_path), "vision_enabled": False, } ) } ) monkeypatch.setattr( aws_module, "validate_infrastructure", lambda _session, _config: {"ami": "ami-test"}, ) monkeypatch.setattr( aws_module, "quota_snapshot", lambda *_args, **_kwargs: { "enough": True, "requested_vcpus": 8, "remaining_vcpus": 40, }, ) report = asyncio.run( preflight( config, environment={ "OPENAI_API_KEY": "openai-secret", "ANTHROPIC_API_KEY": "anthropic-secret", }, allow_partial=False, check_direct_models=False, session=object(), ) ) assert report["ready"] is False assert len(report["issues"]) == 1 assert report["issues"][0].startswith( "automatic evaluation gold cache is not ready for " ) assert "task-001" in report["issues"][0] def test_preflight_rejects_aws_model_without_advertised_image_input( monkeypatch, ) -> None: config = BatchConfig.model_validate( { "expected_rollouts": 1, "max_concurrency": 1, "evaluation": {"enabled": False}, "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-test"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "aws_region": "us-east-1", "session_manager_plugin": "/plugin", }, "rollouts": [ { "rollout_id": "glm-five", "display_name": "GLM 5", "provider": "mantle", "model_id": "zai.glm-5", } ], } ) monkeypatch.setattr( aws_module, "validate_infrastructure", lambda *_args: {"ami": "ami-test"}, ) monkeypatch.setattr( aws_module, "quota_snapshot", lambda *_args, **_kwargs: { "enough": True, "requested_vcpus": 4, "remaining_vcpus": 40, }, ) monkeypatch.setattr( aws_module, "list_mantle_models", lambda *_args, **_kwargs: {"zai.glm-5"}, ) monkeypatch.setattr( aws_module, "list_bedrock_model_capabilities", lambda *_args, **_kwargs: {"zai.glm-5": {"TEXT"}}, ) report = asyncio.run( preflight( config, environment={}, allow_partial=False, check_direct_models=False, session=object(), ) ) assert report["ready"] is False assert report["issues"] == ["AWS models without advertised IMAGE input: zai.glm-5"] def test_preflight_accepts_verified_multimodal_mantle_kimi(monkeypatch) -> None: config = BatchConfig.model_validate( { "expected_rollouts": 1, "max_concurrency": 1, "harness_profile": "mantle-kimi-k2.5-chat", "evaluation": {"enabled": False}, "infrastructure": { "image_id": "ami-test", "broker_version": "windows-autocad-2019-v10", "subnet_ids": ["subnet-test"], "security_group_id": "sg-test", "instance_profile_name": "worker-profile", "instance_type": "g4dn.xlarge", "aws_region": "us-east-1", "session_manager_plugin": "/plugin", }, "rollouts": [ { "rollout_id": "kimi", "display_name": "Kimi K2.5", "provider": "mantle", "model_id": "moonshotai.kimi-k2.5", } ], } ) monkeypatch.setattr( aws_module, "validate_infrastructure", lambda *_args: {"ami": "ami-test"}, ) monkeypatch.setattr( aws_module, "quota_snapshot", lambda *_args, **_kwargs: { "enough": True, "requested_vcpus": 4, "remaining_vcpus": 40, }, ) monkeypatch.setattr( aws_module, "list_mantle_models", lambda *_args, **_kwargs: {"moonshotai.kimi-k2.5"}, ) monkeypatch.setattr( aws_module, "list_bedrock_model_capabilities", lambda *_args, **_kwargs: {}, ) report = asyncio.run( preflight( config, environment={}, allow_partial=False, check_direct_models=False, session=object(), ) ) assert report["ready"] is True assert report["issues"] == [] def test_quota_snapshot_counts_only_on_demand_g_vt_instances() -> None: class Paginator: def paginate(self, **_: Any): return [ { "Reservations": [ { "Instances": [ { "InstanceType": "g5.12xlarge", "CpuOptions": { "CoreCount": 24, "ThreadsPerCore": 2, }, }, { "InstanceType": "g4dn.xlarge", "CpuOptions": { "CoreCount": 2, "ThreadsPerCore": 2, }, }, { "InstanceType": "g4dn.xlarge", "InstanceLifecycle": "spot", "CpuOptions": { "CoreCount": 2, "ThreadsPerCore": 2, }, }, { "InstanceType": "c7i.xlarge", "CpuOptions": { "CoreCount": 2, "ThreadsPerCore": 2, }, }, ] } ] } ] class Ec2: def describe_instance_types(self, **_: Any): return {"InstanceTypes": [{"VCpuInfo": {"DefaultVCpus": 4}}]} def get_paginator(self, name: str): assert name == "describe_instances" return Paginator() class Quotas: def get_service_quota(self, **_: Any): return {"Quota": {"Value": 128.0}} class Session: def client(self, service: str, **_: Any): return Ec2() if service == "ec2" else Quotas() snapshot = quota_snapshot( Session(), region_name="us-east-1", instance_type="g4dn.xlarge", requested_instances=10, ) assert snapshot["used_vcpus"] == 52 assert snapshot["requested_vcpus"] == 40 assert snapshot["remaining_vcpus"] == 76 assert snapshot["enough"] is True def test_quota_snapshot_uses_standard_quota_for_m7_and_excludes_g() -> None: class Paginator: def paginate(self, **_: Any): return [ { "Reservations": [ { "Instances": [ { "InstanceType": "m7i.xlarge", "CpuOptions": {"CoreCount": 2, "ThreadsPerCore": 2}, }, { "InstanceType": "c5.4xlarge", "CpuOptions": {"CoreCount": 8, "ThreadsPerCore": 2}, }, { "InstanceType": "g5.12xlarge", "CpuOptions": { "CoreCount": 24, "ThreadsPerCore": 2, }, }, ] } ] } ] class Ec2: def describe_instance_types(self, **_: Any): return {"InstanceTypes": [{"VCpuInfo": {"DefaultVCpus": 4}}]} def get_paginator(self, name: str): assert name == "describe_instances" return Paginator() class Quotas: def __init__(self) -> None: self.code = "" def get_service_quota(self, **kwargs: Any): self.code = kwargs["QuotaCode"] return {"Quota": {"Value": 1024.0}} class Session: def __init__(self) -> None: self.quotas = Quotas() def client(self, service: str, **_: Any): return Ec2() if service == "ec2" else self.quotas session = Session() snapshot = quota_snapshot( session, region_name="us-east-1", instance_type="m7i.xlarge", requested_instances=34, ) assert session.quotas.code == batch_module.STANDARD_QUOTA_CODE assert snapshot["quota_class"] == "Standard On-Demand" assert snapshot["used_vcpus"] == 20 assert snapshot["requested_vcpus"] == 136 assert snapshot["remaining_vcpus"] == 1004 assert snapshot["enough"] is True def test_reaper_force_terminates_matching_stragglers_without_waiting() -> None: class Ec2: def __init__(self) -> None: self.filters: list[dict[str, Any]] = [] self.terminated: list[str] = [] def describe_instances(self, *, Filters): self.filters = Filters return { "Reservations": [ { "Instances": [ {"InstanceId": "i-one"}, {"InstanceId": "i-two"}, ] } ] } def terminate_instances(self, *, InstanceIds): self.terminated = InstanceIds class Session: def __init__(self) -> None: self.ec2 = Ec2() def client(self, *_: Any, **__: Any): return self.ec2 session = Session() reaped = reap_batch_instances( session, region_name="us-east-1", batch_id="batch-test", ) assert reaped == ["i-one", "i-two"] assert session.ec2.terminated == reaped assert {"Name": "tag:BatchId", "Values": ["batch-test"]} in session.ec2.filters def test_global_controller_slots_are_shared_across_batches(tmp_path: Path) -> None: async def run() -> None: first = batch_module._GlobalControllerSlots(1, root=tmp_path) second = batch_module._GlobalControllerSlots(1, root=tmp_path) entered = asyncio.Event() async def wait_for_slot() -> None: async with second.acquire(): entered.set() async with first.acquire(): waiter = asyncio.create_task(wait_for_slot()) await asyncio.sleep(0.05) assert entered.is_set() is False await asyncio.wait_for(waiter, timeout=2) assert entered.is_set() is True asyncio.run(run()) def test_global_controller_slots_reserve_capacity_for_legacy_runners( tmp_path: Path, ) -> None: proc_root = tmp_path / "proc" legacy = proc_root / "101" managed = proc_root / "102" unrelated = proc_root / "103" for process in (legacy, managed, unrelated): process.mkdir(parents=True) (legacy / "cmdline").write_bytes(b"python\0-m\0autocad_bench.sandbox.runner\0") (legacy / "environ").write_bytes(b"PATH=/usr/bin\0") (managed / "cmdline").write_bytes(b"python\0-m\0autocad_bench.sandbox.runner\0") (managed / "environ").write_bytes( b"PATH=/usr/bin\0AUTOCAD_BENCH_CONTROLLER_SLOT=0\0" ) (unrelated / "cmdline").write_bytes(b"python\0worker.py\0") (unrelated / "environ").write_bytes(b"PATH=/usr/bin\0") slots = batch_module._GlobalControllerSlots( 2, root=tmp_path / "slots", proc_root=proc_root, ) assert slots._legacy_runner_count() == 1 async def run() -> None: async with slots.acquire() as index: assert index == 0 asyncio.run(run()) def test_global_controller_limit_is_explicit_and_validated() -> None: assert batch_module._global_controller_limit({}) is None assert ( batch_module._global_controller_limit( {"AUTOCAD_BENCH_GLOBAL_MAX_CONCURRENCY": "32"} ) == 32 ) with pytest.raises(batch_module.PreflightError, match="between 1 and 100"): batch_module._global_controller_limit( {"AUTOCAD_BENCH_GLOBAL_MAX_CONCURRENCY": "0"} ) def test_controller_failure_cancels_tasks_and_signals_child_groups(monkeypatch) -> None: class Process: pid = 1234 returncode: int | None = None async def wait(self) -> int: self.returncode = 1 return 1 async def run() -> None: task = asyncio.create_task(asyncio.sleep(60)) process = Process() signals: list[tuple[int, signal.Signals]] = [] monkeypatch.setattr( batch_module.os, "killpg", lambda pid, sent_signal: signals.append((pid, sent_signal)), ) await batch_module._stop_rollout_processes( [task], {"rollout": process}, # type: ignore[arg-type] grace_s=0.1, ) assert task.cancelled() assert signals == [(1234, signal.SIGINT)] assert process.returncode == 1 asyncio.run(run()) def test_sigterm_cancels_batch_cooperatively(monkeypatch) -> None: handlers: dict[signal.Signals, object] = {} class LoopProxy: def __init__(self, loop): self.loop = loop def add_signal_handler(self, watched_signal, callback): handlers[watched_signal] = callback def remove_signal_handler(self, watched_signal): handlers.pop(watched_signal, None) return True def __getattr__(self, name): return getattr(self.loop, name) async def blocked_main(_args): await asyncio.Event().wait() async def run() -> None: real_loop = asyncio.get_running_loop() monkeypatch.setattr(batch_module, "_main", blocked_main) monkeypatch.setattr( batch_module.asyncio, "get_running_loop", lambda: LoopProxy(real_loop), ) wrapper = asyncio.create_task( batch_module._main_with_termination_signal(object()) ) await asyncio.sleep(0) callback = handlers[signal.SIGTERM] assert callable(callback) callback() assert await wrapper == 130 assert signal.SIGTERM not in handlers asyncio.run(run())