File indexing completed on 2026-08-12 08:24:57
0001 """Constraint integration tests across DesignConfig, SearchSpace, and Ax."""
0002
0003 import pytest
0004 from pydantic import ValidationError
0005
0006 from aid2e.optimizers.ax.config import AxOptimizerConfig
0007 from aid2e.optimizers.ax.optimizer import AxOptimizer
0008 from aid2e.optimizers.ax import optimizer as ax_optimizer_module
0009 from aid2e.optimizers.base import SearchSpace
0010 from aid2e.utilities.configurations.design_config import DesignConfig, ParameterConstraint
0011
0012
0013 AX_NODE_RUNTIME_AVAILABLE = ax_optimizer_module.AX_NODE_STRATEGY_AVAILABLE
0014
0015
0016 def _design_config_with_constraints(constraints, include_z: bool = False) -> DesignConfig:
0017 """Build a validated DesignConfig for constraint tests."""
0018 parameters = {
0019 "x": {"value": 0.5, "bounds": [0.0, 1.0]},
0020 "y": {"value": 0.5, "bounds": [0.0, 1.0]},
0021 }
0022 if include_z:
0023 parameters["z"] = {"value": 0.5, "bounds": [0.0, 1.0]}
0024
0025 return DesignConfig(
0026 design_parameters={"group1": {"parameters": parameters}},
0027 parameter_constraints=constraints,
0028 )
0029
0030
0031 class TestConstraintSyntaxValidation:
0032 """Test constraint syntax validation in DesignConfig."""
0033
0034 def test_valid_constraint_accepted(self):
0035 """Valid constraints should be accepted during DesignConfig creation."""
0036 config = _design_config_with_constraints(
0037 [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0038 )
0039 assert len(config.parameter_constraints) == 1
0040 assert config.parameter_constraints[0].name == "sum_limit"
0041
0042 def test_unknown_parameter_rejected(self):
0043 """Constraints referencing unknown parameters should be rejected."""
0044 with pytest.raises(ValidationError, match="Invalid constraint.*[Uu]nknown parameter"):
0045 DesignConfig(
0046 design_parameters={
0047 "group1": {"parameters": {"x": {"value": 0.5, "bounds": [0.0, 1.0]}}}
0048 },
0049 parameter_constraints=[
0050 {
0051 "name": "bad_constraint",
0052 "rule": "group1.x + group1.unknown <= 1.0",
0053 }
0054 ],
0055 )
0056
0057 def test_syntax_error_rejected(self):
0058 """Constraints with invalid syntax should be rejected."""
0059 with pytest.raises(ValidationError, match="Invalid constraint.*syntax"):
0060 DesignConfig(
0061 design_parameters={
0062 "group1": {"parameters": {"x": {"value": 0.5, "bounds": [0.0, 1.0]}}}
0063 },
0064 parameter_constraints=[{"name": "bad_syntax", "rule": "group1.x +* 1.0"}],
0065 )
0066
0067
0068 class TestParameterConstraintMethods:
0069 """Test ParameterConstraint helper methods."""
0070
0071 def test_extract_parameter_names(self):
0072 """Test extraction of parameter names from constraint rules."""
0073 constraint = ParameterConstraint(
0074 name="test",
0075 rule="tracker.x + magnet.y + detector.z <= 10.0",
0076 )
0077
0078 param_names = constraint.extract_parameter_names()
0079 assert param_names == {"tracker.x", "magnet.y", "detector.z"}
0080
0081 def test_validate_syntax_valid(self):
0082 """Test syntax validation for valid constraints."""
0083 constraint = ParameterConstraint(
0084 name="test",
0085 rule="group1.x + group1.y <= 1.5",
0086 )
0087
0088 valid_params = {"group1.x", "group1.y"}
0089 is_valid, error = constraint.validate_syntax(valid_params)
0090 assert is_valid
0091 assert error is None
0092
0093 def test_validate_syntax_unknown_param(self):
0094 """Test syntax validation rejects unknown parameters."""
0095 constraint = ParameterConstraint(
0096 name="test",
0097 rule="group1.x + group1.unknown <= 1.5",
0098 )
0099
0100 valid_params = {"group1.x"}
0101 is_valid, error = constraint.validate_syntax(valid_params)
0102 assert not is_valid
0103 assert "unknown parameter" in error.lower()
0104
0105 def test_evaluate_constraint(self):
0106 """Test runtime constraint evaluation."""
0107 constraint = ParameterConstraint(
0108 name="sum_limit",
0109 rule="group1.x + group1.y <= 1.5",
0110 )
0111
0112 assert constraint.evaluate({"group1.x": 0.5, "group1.y": 0.8}) is True
0113 assert constraint.evaluate({"group1.x": 1.0, "group1.y": 0.6}) is False
0114
0115
0116 class TestSearchSpaceConstraints:
0117 """Test SearchSpace constraint handling."""
0118
0119 def test_constraints_passed_from_design_config(self):
0120 """SearchSpace should receive constraints from DesignConfig."""
0121 design_config = _design_config_with_constraints(
0122 [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0123 )
0124 search_space = SearchSpace.from_design_config(design_config)
0125
0126 assert len(search_space.constraints) == 1
0127 assert search_space.constraints[0].name == "sum_limit"
0128 assert search_space.constraints[0].rule == "group1.x + group1.y <= 1.5"
0129
0130 def test_check_constraints_method(self):
0131 """Test SearchSpace.validate() for runtime validation."""
0132 design_config = _design_config_with_constraints(
0133 [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0134 )
0135 search_space = SearchSpace.from_design_config(design_config)
0136
0137 is_valid, errors = search_space.validate({"group1.x": 0.5, "group1.y": 0.8})
0138 assert is_valid
0139 assert len(errors) == 0
0140
0141 is_valid, errors = search_space.validate({"group1.x": 1.0, "group1.y": 0.8})
0142 assert not is_valid
0143 assert len(errors) > 0
0144 assert "sum_limit" in errors[0]
0145
0146
0147 class TestAxOptimizerConstraints:
0148 """Test Ax optimizer constraint integration."""
0149
0150 def test_ax_requires_node_runtime_for_constraint_integration(self):
0151 """Ax constraint integration should fail fast without node runtime support."""
0152 design_config = _design_config_with_constraints(
0153 [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0154 )
0155 search_space = SearchSpace.from_design_config(design_config)
0156 ax_config = AxOptimizerConfig(n_initial_samples=5, seed=42)
0157
0158 if AX_NODE_RUNTIME_AVAILABLE:
0159 optimizer = AxOptimizer(
0160 search_space=search_space,
0161 config=ax_config,
0162 objective_names=["f1"],
0163 )
0164 assert optimizer.search_space.constraints
0165 return
0166
0167 with pytest.raises(RuntimeError, match="node-based generation API required"):
0168 AxOptimizer(
0169 search_space=search_space,
0170 config=ax_config,
0171 objective_names=["f1"],
0172 )
0173
0174 @pytest.mark.skipif(
0175 not AX_NODE_RUNTIME_AVAILABLE,
0176 reason="Installed Ax runtime lacks required node-based generation APIs.",
0177 )
0178 def test_constraint_parsing_to_ax_format(self):
0179 """Test parsing of constraint rules to Ax ParameterConstraint format."""
0180 design_config = _design_config_with_constraints(
0181 [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0182 )
0183 search_space = SearchSpace.from_design_config(design_config)
0184
0185 optimizer = AxOptimizer(
0186 search_space=search_space,
0187 config=AxOptimizerConfig(n_initial_samples=5, seed=42),
0188 objective_names=["f1"],
0189 )
0190
0191 constraint = design_config.parameter_constraints[0]
0192 ax_constraint = optimizer._parse_constraint_to_ax(constraint)
0193
0194 assert ax_constraint is not None
0195 assert ax_constraint.constraint_dict == {"group1.x": 1.0, "group1.y": 1.0}
0196 assert ax_constraint.bound == 1.5
0197
0198 @pytest.mark.skipif(
0199 not AX_NODE_RUNTIME_AVAILABLE,
0200 reason="Installed Ax runtime lacks required node-based generation APIs.",
0201 )
0202 def test_ax_enforces_constraints(self):
0203 """Test that Ax enforces constraints during candidate generation."""
0204 design_config = _design_config_with_constraints(
0205 [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0206 )
0207 search_space = SearchSpace.from_design_config(design_config)
0208
0209 optimizer = AxOptimizer(
0210 search_space=search_space,
0211 config=AxOptimizerConfig(n_initial_samples=10, seed=42),
0212 objective_names=["f1"],
0213 )
0214
0215 candidates = optimizer.suggest_candidates(n_candidates=20)
0216 violations = []
0217 for i, candidate in enumerate(candidates):
0218 sum_val = candidate["group1.x"] + candidate["group1.y"]
0219 if sum_val > 1.5:
0220 violations.append((i, sum_val))
0221
0222 assert len(violations) == 0
0223
0224 @pytest.mark.skipif(
0225 not AX_NODE_RUNTIME_AVAILABLE,
0226 reason="Installed Ax runtime lacks required node-based generation APIs.",
0227 )
0228 def test_multiple_constraints(self):
0229 """Test Ax with multiple constraints."""
0230 design_config = _design_config_with_constraints(
0231 [
0232 {"name": "sum_xy", "rule": "group1.x + group1.y <= 1.2"},
0233 {"name": "sum_xz", "rule": "group1.x + group1.z <= 1.3"},
0234 ],
0235 include_z=True,
0236 )
0237 search_space = SearchSpace.from_design_config(design_config)
0238
0239 optimizer = AxOptimizer(
0240 search_space=search_space,
0241 config=AxOptimizerConfig(n_initial_samples=10, seed=42),
0242 objective_names=["f1"],
0243 )
0244
0245 candidates = optimizer.suggest_candidates(n_candidates=20)
0246 violations = []
0247 for i, candidate in enumerate(candidates):
0248 x = candidate["group1.x"]
0249 y = candidate["group1.y"]
0250 z = candidate["group1.z"]
0251 if x + y > 1.2:
0252 violations.append((i, "sum_xy"))
0253 if x + z > 1.3:
0254 violations.append((i, "sum_xz"))
0255
0256 assert len(violations) == 0
0257
0258 @pytest.mark.skipif(
0259 not AX_NODE_RUNTIME_AVAILABLE,
0260 reason="Installed Ax runtime lacks required node-based generation APIs.",
0261 )
0262 def test_greater_than_constraint(self):
0263 """Test Ax with >= constraint converted to negated <=."""
0264 design_config = _design_config_with_constraints(
0265 [{"name": "min_sum", "rule": "group1.x + group1.y >= 0.3"}]
0266 )
0267 search_space = SearchSpace.from_design_config(design_config)
0268
0269 optimizer = AxOptimizer(
0270 search_space=search_space,
0271 config=AxOptimizerConfig(n_initial_samples=10, seed=42),
0272 objective_names=["f1"],
0273 )
0274
0275 constraint = design_config.parameter_constraints[0]
0276 ax_constraint = optimizer._parse_constraint_to_ax(constraint)
0277
0278 assert ax_constraint is not None
0279 assert ax_constraint.constraint_dict == {"group1.x": -1.0, "group1.y": -1.0}
0280 assert ax_constraint.bound == -0.3
0281
0282 candidates = optimizer.suggest_candidates(n_candidates=20)
0283 violations = []
0284 for i, candidate in enumerate(candidates):
0285 if candidate["group1.x"] + candidate["group1.y"] < 0.3:
0286 violations.append(i)
0287
0288 assert len(violations) == 0