Download tests/test_batch.py from windows2t2/cad-bench: direct link, hf CLI and curl.
- Browser
- Download file 35.8 kB
-
https://huggingface.co/windows2t2/cad-bench/resolve/main/tests/test_batch.py
- Command line
-
hf download hf://windows2t2/cad-bench/tests/test_batch.py
-
curl -L -o test_batch.py https://huggingface.co/windows2t2/cad-bench/resolve/main/tests/test_batch.py
35.8 kB
| 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()) | |