diff --git a/sigma/correlations.py b/sigma/correlations.py index daf2c6e4..201d6529 100644 --- a/sigma/correlations.py +++ b/sigma/correlations.py @@ -522,7 +522,16 @@ def from_dict( source: SigmaRuleLocation | None = None, ) -> Self: kwargs, errors = super().from_dict_common_params(rule, collect_errors, source) - correlation_rule = rule.get("correlation", dict()) + correlation_rule: Any = rule.get("correlation", dict()) + if not isinstance(correlation_rule, dict): + errors.append( + sigma_exceptions.SigmaCorrelationRuleError( + "Sigma correlation rule 'correlation' field must be a dict", source=source + ) + ) + if not collect_errors: + raise errors[0] + correlation_rule = dict() # Correlation type correlation_type = correlation_rule.get("type") @@ -625,7 +634,9 @@ def from_dict( # Condition - can be either a dict (basic condition) or a string (extended condition) condition_value = correlation_rule.get("condition") - condition: SigmaCorrelationCondition | SigmaExtendedCorrelationCondition + condition: SigmaCorrelationCondition | SigmaExtendedCorrelationCondition = ( + SigmaCorrelationCondition(SigmaCorrelationConditionOperator.GTE, 1) + ) if condition_value is not None: if isinstance(condition_value, dict): diff --git a/tests/test_correlations.py b/tests/test_correlations.py index dd0bf225..4beab164 100644 --- a/tests/test_correlations.py +++ b/tests/test_correlations.py @@ -184,6 +184,23 @@ def test_correlation_wrong_type(): ) +@pytest.mark.parametrize("bad_correlation", [None, "not-a-dict", 123, []]) +def test_correlation_field_not_a_dict_raises(bad_correlation): + with pytest.raises(SigmaCorrelationRuleError, match="'correlation' field must be a dict"): + SigmaCorrelationRule.from_dict( + {"title": "Invalid correlation", "correlation": bad_correlation} + ) + + +@pytest.mark.parametrize("bad_correlation", [None, "not-a-dict", 123, []]) +def test_correlation_field_not_a_dict_collect_errors(bad_correlation): + rule = SigmaCorrelationRule.from_dict( + {"title": "Invalid correlation", "correlation": bad_correlation}, + collect_errors=True, + ) + assert any("'correlation' field must be a dict" in str(error) for error in rule.errors) + + def test_correlation_without_type(): with pytest.raises(SigmaCorrelationTypeError, match="Sigma correlation rule without type"): SigmaCorrelationRule.from_dict( @@ -1343,3 +1360,37 @@ def test_correlation_condition_non_dict_non_string(): - invalid - list_condition """) + + +def test_correlation_extended_condition_wrong_type_collect_errors(): + """collect_errors=True must return the error instead of raising an UnboundLocalError + when an extended (string) condition is used with a non-temporal correlation type.""" + rule = SigmaCorrelationRule.from_yaml( + """ +title: Test correlation +status: test +correlation: + type: event_count + rules: + - test_rule + timespan: 5m + condition: "count() > 5" + """, + collect_errors=True, + ) + assert {error.__class__ for error in rule.errors} == {SigmaCorrelationRuleError} + + +def test_correlation_extended_condition_wrong_type_raises(): + """Without collect_errors the error is raised instead of silently swallowed.""" + with pytest.raises(SigmaCorrelationRuleError, match="only be used with temporal"): + SigmaCorrelationRule.from_yaml(""" +title: Test correlation +status: test +correlation: + type: event_count + rules: + - test_rule + timespan: 5m + condition: "count() > 5" + """)