Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-08-12 08:24:58

0001 """Unit tests for rule resolution and payload validation.
0002 
0003 Tests cover:
0004 - Rule template resolution with various template formats
0005 - Payload template substitution (strings, dicts, lists)
0006 - Conditional execution via payload validation
0007 - Logging and checkpointing
0008 - Error handling
0009 
0010 Project: AID2E v0.0.0
0011 """
0012 
0013 import pytest
0014 import tempfile
0015 from pathlib import Path
0016 from typing import Dict, Any
0017 
0018 from aid2e.utilities.configurations.workflow_config import JobDefinition, ArtifactSpec
0019 from aid2e.utilities.workflows.rule_resolution import (
0020     resolve_job_rule,
0021     resolve_payload_templates,
0022     validate_job_payload,
0023     RuleResolutionError,
0024     PayloadValidationError
0025 )
0026 from aid2e.utilities.workflows.execution_logger import ExecutionLogger, create_job_logger
0027 
0028 
0029 class TestPayloadTemplateResolution:
0030     """Test payload template variable resolution."""
0031     
0032     def test_simple_string_template(self):
0033         """Test resolving simple string template."""
0034         payload = {"input": "{{input_design_params}}"}
0035         context = {"input_design_params": "/path/to/design.params"}
0036         
0037         resolved = resolve_payload_templates(payload, context)
0038         assert resolved["input"] == "/path/to/design.params"
0039     
0040     def test_multiple_templates_in_payload(self):
0041         """Test multiple template variables in payload."""
0042         payload = {
0043             "input_file": "{{input_design_params}}",
0044             "output_dir": "{{output_dir}}",
0045             "job_id": "{{job_id}}"
0046         }
0047         context = {
0048             "input_design_params": "/data/design.params",
0049             "output_dir": "/tmp/out",
0050             "job_id": 0
0051         }
0052         
0053         resolved = resolve_payload_templates(payload, context)
0054         assert resolved["input_file"] == "/data/design.params"
0055         assert resolved["output_dir"] == "/tmp/out"
0056         assert resolved["job_id"] == "0"  # Templates convert to strings
0057     
0058     def test_nested_dict_resolution(self):
0059         """Test resolving templates in nested dicts."""
0060         payload = {
0061             "metadata": {
0062                 "input": "{{input_design_params}}",
0063                 "version": "1.0"
0064             }
0065         }
0066         context = {"input_design_params": "/path/to/design.params"}
0067         
0068         resolved = resolve_payload_templates(payload, context)
0069         assert resolved["metadata"]["input"] == "/path/to/design.params"
0070         assert resolved["metadata"]["version"] == "1.0"
0071     
0072     def test_list_in_payload(self):
0073         """Test resolving templates in lists."""
0074         payload = {
0075             "inputs": ["{{input_design_params}}", "/other/file.txt"]
0076         }
0077         context = {"input_design_params": "/path/to/design.params"}
0078         
0079         resolved = resolve_payload_templates(payload, context)
0080         assert resolved["inputs"][0] == "/path/to/design.params"
0081         assert resolved["inputs"][1] == "/other/file.txt"
0082     
0083     def test_undefined_template_variable(self):
0084         """Test error handling for undefined template variables."""
0085         payload = {"input": "{{undefined_variable}}"}
0086         context = {}
0087         
0088         with pytest.raises(RuleResolutionError):
0089             resolve_payload_templates(payload, context)
0090     
0091     def test_dict_access_template(self):
0092         """Test dict[key] notation in templates."""
0093         payload = {
0094             "design_params": "{{stage_outputs[preparation]}}/design_params.json"
0095         }
0096         context = {
0097             "stage_outputs": {
0098                 "preparation": "/tmp/stage1"
0099             }
0100         }
0101         
0102         resolved = resolve_payload_templates(payload, context)
0103         assert resolved["design_params"] == "/tmp/stage1/design_params.json"
0104     
0105     def test_non_string_values_pass_through(self):
0106         """Test that non-string values are not modified."""
0107         payload = {
0108             "count": 42,
0109             "enabled": True,
0110             "timeout": 3.5
0111         }
0112         context = {}
0113         
0114         resolved = resolve_payload_templates(payload, context)
0115         assert resolved["count"] == 42
0116         assert resolved["enabled"] is True
0117         assert resolved["timeout"] == 3.5
0118 
0119 
0120 class TestRuleResolution:
0121     """Test rule template resolution."""
0122     
0123     def test_simple_command_rule(self):
0124         """Test resolving simple command rule."""
0125         job = JobDefinition(
0126             name="test",
0127             command="python script.py",
0128             rule="{{command}}",
0129             payload={}
0130         )
0131         context = {}
0132         
0133         cmd = resolve_job_rule(job, context)
0134         assert cmd == "python script.py"
0135     
0136     def test_rule_with_payload_substitution(self):
0137         """Test rule with payload variable substitution."""
0138         job = JobDefinition(
0139             name="test",
0140             command="python compute.py",
0141             rule="{{command}} {{payload[input]}} {{payload[output]}}",
0142             payload={
0143                 "input": "/path/to/input.json",
0144                 "output": "/path/to/output.json"
0145             }
0146         )
0147         context = {}
0148         
0149         cmd = resolve_job_rule(job, context)
0150         assert cmd == "python compute.py /path/to/input.json /path/to/output.json"
0151     
0152     def test_rule_with_template_substitution(self):
0153         """Test rule with context template variables."""
0154         job = JobDefinition(
0155             name="test",
0156             command="python compute.py",
0157             rule="{{command}} {{payload[input_file]}} {{output_dir}} {{job_id}}",
0158             payload={
0159                 "input_file": "{{input_design_params}}"
0160             }
0161         )
0162         context = {
0163             "input_design_params": "/data/design.params",
0164             "output_dir": "/tmp/out",
0165             "job_id": 0
0166         }
0167         
0168         cmd = resolve_job_rule(job, context)
0169         assert cmd == "python compute.py /data/design.params /tmp/out 0"
0170     
0171     def test_default_rule_if_not_specified(self):
0172         """Test default rule behavior when rule is None."""
0173         job = JobDefinition(
0174             name="test",
0175             command="python script.py",
0176             rule=None,  # No rule specified
0177             payload={"ignored": "value"}
0178         )
0179         context = {}
0180         
0181         cmd = resolve_job_rule(job, context)
0182         assert cmd == "python script.py"
0183     
0184     def test_rule_with_nested_payload_access(self):
0185         """Test rule accessing simple payload values."""
0186         job = JobDefinition(
0187             name="test",
0188             command="python run.py",
0189             rule="{{command}} --input {{payload[input_file]}} --output {{payload[output_file]}}",
0190             payload={
0191                 "input_file": "/data/input.json",
0192                 "output_file": "/data/output.json"
0193             }
0194         )
0195         context = {}
0196         
0197         cmd = resolve_job_rule(job, context)
0198         assert "python run.py" in cmd
0199         assert "--input /data/input.json" in cmd
0200         assert "--output /data/output.json" in cmd
0201     
0202     def test_multiple_spaces_cleaned_up(self):
0203         """Test that multiple spaces in resolved command are cleaned up."""
0204         job = JobDefinition(
0205             name="test",
0206             command="python script.py",
0207             rule="{{command}}    {{payload[arg1]}}    {{payload[arg2]}}",
0208             payload={"arg1": "val1", "arg2": "val2"}
0209         )
0210         context = {}
0211         
0212         cmd = resolve_job_rule(job, context)
0213         assert "  " not in cmd  # No double spaces
0214         assert cmd == "python script.py val1 val2"
0215 
0216 
0217 class TestPayloadValidation:
0218     """Test payload validation."""
0219     
0220     def test_validation_with_required_keys_present(self):
0221         """Test successful validation when required keys present."""
0222         job = JobDefinition(
0223             name="test",
0224             command="python script.py",
0225             payload={"input": "/path", "output": "/path"}
0226         )
0227         
0228         is_valid, error = validate_job_payload(job, required_keys=["input", "output"])
0229         assert is_valid is True
0230         assert error is None
0231     
0232     def test_validation_fails_with_missing_required_keys(self):
0233         """Test validation fails when required keys missing."""
0234         job = JobDefinition(
0235             name="test",
0236             command="python script.py",
0237             payload={"input": "/path"}
0238         )
0239         
0240         is_valid, error = validate_job_payload(job, required_keys=["input", "output"])
0241         assert is_valid is False
0242         assert "Missing required payload keys" in error
0243         assert "output" in error
0244     
0245     def test_validation_with_template_resolution(self):
0246         """Test validation with template resolution."""
0247         job = JobDefinition(
0248             name="test",
0249             command="python script.py",
0250             payload={"input": "{{input_design_params}}"}
0251         )
0252         context = {"input_design_params": "/path/to/design.params"}
0253         
0254         is_valid, error = validate_job_payload(job, context=context)
0255         assert is_valid is True
0256         assert error is None
0257     
0258     def test_validation_fails_with_undefined_template(self):
0259         """Test validation with undefined template variables."""
0260         job = JobDefinition(
0261             name="test",
0262             command="python script.py",
0263             payload={"input": "{{undefined_var}}"}
0264         )
0265         
0266         # Test: With context, template resolution should fail if var undefined
0267         context = {"some_var": "value"}  # undefined_var not in context
0268         is_valid, error = validate_job_payload(job, context=context, logger=None)
0269         assert is_valid is False
0270         assert "Cannot resolve payload templates" in error
0271 
0272 
0273 class TestExecutionLogging:
0274     """Test execution logging and checkpointing."""
0275     
0276     @pytest.fixture
0277     def temp_output_dir(self):
0278         """Create temporary output directory for logging."""
0279         with tempfile.TemporaryDirectory() as tmpdir:
0280             yield tmpdir
0281     
0282     def test_logger_creation(self, temp_output_dir):
0283         """Test creating execution logger."""
0284         logger = ExecutionLogger(
0285             job_name="test_job",
0286             output_dir=temp_output_dir,
0287             log_level="INFO"
0288         )
0289         
0290         assert logger.job_name == "test_job"
0291         assert len(logger.checkpoints) > 0  # Initial checkpoint created
0292     
0293     def test_checkpoint_creation(self, temp_output_dir):
0294         """Test creating checkpoints."""
0295         logger = ExecutionLogger(
0296             job_name="test_job",
0297             output_dir=temp_output_dir
0298         )
0299         
0300         logger.checkpoint(
0301             stage="test_stage",
0302             status="start",
0303             message="Test checkpoint"
0304         )
0305         
0306         last_cp = logger.get_last_checkpoint()
0307         assert last_cp.stage == "test_stage"
0308         assert last_cp.status == "start"
0309         assert last_cp.message == "Test checkpoint"
0310     
0311     def test_checkpoint_file_creation(self, temp_output_dir):
0312         """Test checkpoint file is created."""
0313         logger = ExecutionLogger(
0314             job_name="test_job",
0315             output_dir=temp_output_dir,
0316             enable_checkpoint_file=True
0317         )
0318         
0319         logger.checkpoint(
0320             stage="test",
0321             status="success",
0322             message="Test"
0323         )
0324         
0325         # Check that checkpoint file was created
0326         checkpoint_file = Path(temp_output_dir) / "test_job_checkpoints.json"
0327         assert checkpoint_file.exists()
0328     
0329     def test_execution_summary(self, temp_output_dir):
0330         """Test execution summary generation."""
0331         logger = ExecutionLogger(
0332             job_name="test_job",
0333             output_dir=temp_output_dir
0334         )
0335         
0336         logger.checkpoint("stage1", "start", "Starting stage 1")
0337         logger.checkpoint("stage1", "success", "Stage 1 complete")
0338         logger.checkpoint("stage2", "error", "Stage 2 failed")
0339         
0340         summary = logger.execution_summary()
0341         assert summary["total_checkpoints"] > 0
0342         assert summary["has_errors"] is True
0343         assert "stage1" in summary["stages_executed"]
0344         assert "stage2" in summary["stages_executed"]
0345 
0346 
0347 class TestIntegrationWithLogger:
0348     """Integration tests combining rule resolution with logging."""
0349     
0350     @pytest.fixture
0351     def temp_output_dir(self):
0352         """Create temporary output directory."""
0353         with tempfile.TemporaryDirectory() as tmpdir:
0354             yield tmpdir
0355     
0356     def test_rule_resolution_with_logging(self, temp_output_dir):
0357         """Test rule resolution with execution logging."""
0358         logger = ExecutionLogger(
0359             job_name="dtlz2_eval",
0360             output_dir=temp_output_dir
0361         )
0362         
0363         job = JobDefinition(
0364             name="dtlz2_eval",
0365             command="python compute_dtlz2.py",
0366             rule="{{command}} {{payload[design_file]}} {{payload[output_dir]}} {{job_id}}",
0367             payload={
0368                 "design_file": "{{input_design_params}}",
0369                 "output_dir": "{{output_dir}}"
0370             }
0371         )
0372         
0373         context = {
0374             "input_design_params": "/data/design.params",
0375             "output_dir": "/tmp/stage_output",
0376             "job_id": 0
0377         }
0378         
0379         cmd = resolve_job_rule(job, context, logger)
0380         
0381         # Verify command was built correctly
0382         assert "python compute_dtlz2.py" in cmd
0383         assert "/data/design.params" in cmd
0384         assert "/tmp/stage_output" in cmd
0385         
0386         # Verify checkpoints were created
0387         rule_checkpoints = logger.get_checkpoint_by_stage("rule_resolution")
0388         assert len(rule_checkpoints) > 0
0389         assert any(cp.status == "success" for cp in rule_checkpoints)
0390     
0391     def test_full_job_execution_workflow(self, temp_output_dir):
0392         """Test complete job execution workflow with logging."""
0393         logger = create_job_logger(
0394             job_name="complete_test",
0395             output_dir=temp_output_dir,
0396             log_level="DEBUG"
0397         )
0398         
0399         job = JobDefinition(
0400             name="test_job",
0401             command="python script.py",
0402             rule="{{command}} {{payload[input]}} {{payload[output]}}",
0403             payload={
0404                 "input": "{{input_design_params}}",
0405                 "output": "{{output_dir}}/results.json"
0406             }
0407         )
0408         
0409         context = {
0410             "input_design_params": "/data/design.params",
0411             "output_dir": "/tmp/output"
0412         }
0413         
0414         # Validate payload
0415         is_valid, error = validate_job_payload(job, context=context, logger=logger)
0416         assert is_valid is True
0417         
0418         # Resolve rule
0419         cmd = resolve_job_rule(job, context, logger)
0420         assert "python script.py" in cmd
0421         
0422         # Check summary
0423         summary = logger.execution_summary()
0424         assert not summary["has_errors"]
0425         assert summary["status_breakdown"]["success"] > 0
0426 
0427 
0428 if __name__ == "__main__":
0429     pytest.main([__file__, "-v", "--tb=short"])