Skip to content

Commit de2f04a

Browse files
feiiiiii5romanlutzCopilot
authored
FEAT: add raise-on-empty variants to true/false score aggregators (#2984)
Co-authored-by: Roman Lutz <romanlutz13@gmail.com> Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
1 parent 4beda9c commit de2f04a

4 files changed

Lines changed: 117 additions & 11 deletions

File tree

‎pyrit/score/true_false/true_false_composite_scorer.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,8 +39,9 @@ class TrueFalseCompositeScorer(TrueFalseScorer):
3939
Children are true/false scorers of any evidence kind, so a scorer over a message can be
4040
composed with one over evidence that is not a message at all.
4141
42-
Built-in AND, OR, and MAJORITY aggregators opt into order-independent evaluation
43-
identity. Duplicates remain significant; custom aggregators and execution stay ordered.
42+
Built-in AND, OR, and MAJORITY aggregators, including their raise-on-empty variants,
43+
opt into order-independent evaluation identity. Duplicates remain significant;
44+
custom aggregators and execution stay ordered.
4445
"""
4546

4647
def __init__(
@@ -87,6 +88,9 @@ def _build_identifier(self) -> ComponentIdentifier:
8788
TrueFalseScoreAggregator.AND,
8889
TrueFalseScoreAggregator.OR,
8990
TrueFalseScoreAggregator.MAJORITY,
91+
TrueFalseScoreAggregator.AND_RAISE_ON_EMPTY,
92+
TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY,
93+
TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY,
9094
)
9195
)
9296
return self._create_identifier(

‎pyrit/score/true_false/true_false_score_aggregator.py‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,7 @@ def _create_aggregator(
9898
result_func: Callable[[list[bool | None]], bool | None],
9999
true_msg: str,
100100
false_msg: str,
101+
raise_on_empty: bool = False,
101102
) -> TrueFalseAggregatorFunc:
102103
"""
103104
Create a True/False aggregator using a result function over boolean values.
@@ -109,6 +110,7 @@ def _create_aggregator(
109110
constituent scores, and a ``None`` result means no verdict was reachable.
110111
true_msg (str): Description to use when the result is True.
111112
false_msg (str): Description to use when the result is False.
113+
raise_on_empty (bool): Whether to raise ValueError when no scores are provided. Defaults to False.
112114
113115
Returns:
114116
TrueFalseAggregatorFunc: Aggregator function that reduces a sequence of true/false Scores
@@ -125,6 +127,8 @@ def aggregator(scores: Iterable[Score]) -> ScoreAggregatorResult:
125127
raise ValueError("All scores must be of type 'true_false'.")
126128

127129
if not scores_list:
130+
if raise_on_empty:
131+
raise ValueError("No scores available for aggregation")
128132
# No scores; return a neutral result
129133
return ScoreAggregatorResult(
130134
value=False,
@@ -167,6 +171,7 @@ def _create_binary_aggregator(
167171
op: BinaryBoolOp,
168172
true_msg: str,
169173
false_msg: str,
174+
raise_on_empty: bool = False,
170175
) -> TrueFalseAggregatorFunc:
171176
"""
172177
Turn a binary operator over verdicts (e.g. ``_and``) into an aggregation function.
@@ -176,6 +181,7 @@ def _create_binary_aggregator(
176181
op (BinaryBoolOp): Binary three-valued operator to apply.
177182
true_msg (str): Description to use when the result is True.
178183
false_msg (str): Description to use when the result is False.
184+
raise_on_empty (bool): Whether to raise ValueError when no scores are provided. Defaults to False.
179185
180186
Returns:
181187
TrueFalseAggregatorFunc: Aggregator function that reduces scores using the binary operator.
@@ -185,6 +191,7 @@ def _create_binary_aggregator(
185191
result_func=lambda bs, _op=op: functools.reduce(_op, bs),
186192
true_msg=true_msg,
187193
false_msg=false_msg,
194+
raise_on_empty=raise_on_empty,
188195
)
189196

190197

@@ -229,3 +236,27 @@ class TrueFalseScoreAggregator:
229236
true_msg="A strict majority of constituent scorers returned True in a MAJORITY composite scorer.",
230237
false_msg="A strict majority of constituent scorers did not return True in a MAJORITY composite scorer.",
231238
)
239+
240+
AND_RAISE_ON_EMPTY: TrueFalseAggregatorFunc = _create_binary_aggregator(
241+
"AND_RAISE_ON_EMPTY",
242+
_and,
243+
"All constituent scorers returned True in an AND composite scorer.",
244+
"At least one constituent scorer returned False in an AND composite scorer.",
245+
raise_on_empty=True,
246+
)
247+
248+
OR_RAISE_ON_EMPTY: TrueFalseAggregatorFunc = _create_binary_aggregator(
249+
"OR_RAISE_ON_EMPTY",
250+
_or,
251+
"At least one constituent scorer returned True in an OR composite scorer.",
252+
"All constituent scorers returned False in an OR composite scorer.",
253+
raise_on_empty=True,
254+
)
255+
256+
MAJORITY_RAISE_ON_EMPTY: TrueFalseAggregatorFunc = _create_aggregator(
257+
"MAJORITY_RAISE_ON_EMPTY",
258+
result_func=_majority,
259+
true_msg="A strict majority of constituent scorers returned True in a MAJORITY composite scorer.",
260+
false_msg="A strict majority of constituent scorers did not return True in a MAJORITY composite scorer.",
261+
raise_on_empty=True,
262+
)

‎tests/unit/score/test_scorer_evaluation_identifier.py‎

Lines changed: 43 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,14 @@ def test_eval_hash_matches_free_function(self):
101101
class TestCompositeEvaluationOrder:
102102
@pytest.mark.parametrize(
103103
"aggregator",
104-
[TrueFalseScoreAggregator.OR, TrueFalseScoreAggregator.AND, TrueFalseScoreAggregator.MAJORITY],
104+
[
105+
TrueFalseScoreAggregator.OR,
106+
TrueFalseScoreAggregator.AND,
107+
TrueFalseScoreAggregator.MAJORITY,
108+
TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY,
109+
TrueFalseScoreAggregator.AND_RAISE_ON_EMPTY,
110+
TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY,
111+
],
105112
)
106113
def test_permutations_preserve_eval_identity_and_content_order(self, aggregator: TrueFalseAggregatorFunc) -> None:
107114
children = [SubStringScorer(substring=value) for value in ("a", "b", "c")]
@@ -120,15 +127,25 @@ def test_permutations_preserve_eval_identity_and_content_order(self, aggregator:
120127
assert restored.hash == identifier.hash
121128
assert ScorerEvaluationIdentifier(restored).eval_hash == identifier.eval_hash
122129

123-
def test_nested_permutations_preserve_eval_identity(self) -> None:
130+
@pytest.mark.parametrize(
131+
"outer_aggregator, inner_aggregator",
132+
[
133+
(TrueFalseScoreAggregator.OR, TrueFalseScoreAggregator.AND),
134+
(TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY, TrueFalseScoreAggregator.AND_RAISE_ON_EMPTY),
135+
(TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY, TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY),
136+
],
137+
)
138+
def test_nested_permutations_preserve_eval_identity(
139+
self, *, outer_aggregator: TrueFalseAggregatorFunc, inner_aggregator: TrueFalseAggregatorFunc
140+
) -> None:
124141
a, b, c = [SubStringScorer(substring=value) for value in ("a", "b", "c")]
125142
first = TrueFalseCompositeScorer(
126-
aggregator=TrueFalseScoreAggregator.OR,
127-
scorers=[TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[a, b]), c],
143+
aggregator=outer_aggregator,
144+
scorers=[TrueFalseCompositeScorer(aggregator=inner_aggregator, scorers=[a, b]), c],
128145
)
129146
second = TrueFalseCompositeScorer(
130-
aggregator=TrueFalseScoreAggregator.OR,
131-
scorers=[c, TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.AND, scorers=[b, a])],
147+
aggregator=outer_aggregator,
148+
scorers=[c, TrueFalseCompositeScorer(aggregator=inner_aggregator, scorers=[b, a])],
132149
)
133150
assert first.get_identifier().eval_hash == second.get_identifier().eval_hash
134151
assert first.get_identifier().hash != second.get_identifier().hash
@@ -142,22 +159,39 @@ def test_configuration_aggregator_and_multiplicity_remain_distinct(self) -> None
142159
(TrueFalseScoreAggregator.MAJORITY, [a, b]),
143160
(TrueFalseScoreAggregator.MAJORITY, [a, a, b]),
144161
(TrueFalseScoreAggregator.MAJORITY, [a, b, b]),
162+
(TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY, [a, b]),
163+
(TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY, [a, c]),
164+
(TrueFalseScoreAggregator.AND_RAISE_ON_EMPTY, [a, b]),
165+
(TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY, [a, b]),
166+
(TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY, [a, a, b]),
167+
(TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY, [a, b, b]),
145168
]
146169
hashes = {
147170
TrueFalseCompositeScorer(aggregator=aggregator, scorers=scorers).get_identifier().eval_hash
148171
for aggregator, scorers in configurations
149172
}
150173
assert len(hashes) == len(configurations)
151174

152-
def test_custom_aggregator_with_builtin_name_remains_ordered(self) -> None:
175+
@pytest.mark.parametrize(
176+
"aggregator",
177+
[
178+
TrueFalseScoreAggregator.OR,
179+
TrueFalseScoreAggregator.AND,
180+
TrueFalseScoreAggregator.MAJORITY,
181+
TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY,
182+
TrueFalseScoreAggregator.AND_RAISE_ON_EMPTY,
183+
TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY,
184+
],
185+
)
186+
def test_custom_aggregator_with_builtin_name_remains_ordered(self, aggregator: TrueFalseAggregatorFunc) -> None:
153187
def first_score(scores: Iterable[Score]) -> ScoreAggregatorResult:
154188
return TrueFalseScoreAggregator.OR([next(iter(scores))])
155189

156-
first_score.__name__ = TrueFalseScoreAggregator.OR.__name__
190+
first_score.__name__ = aggregator.__name__
157191
a, b = [SubStringScorer(substring=value) for value in ("a", "b")]
158192
first = TrueFalseCompositeScorer(aggregator=first_score, scorers=[a, b]).get_identifier()
159193
second = TrueFalseCompositeScorer(aggregator=first_score, scorers=[b, a]).get_identifier()
160-
builtin = TrueFalseCompositeScorer(aggregator=TrueFalseScoreAggregator.OR, scorers=[a, b]).get_identifier()
194+
builtin = TrueFalseCompositeScorer(aggregator=aggregator, scorers=[a, b]).get_identifier()
161195

162196
assert "sub_scorers_order_independent" not in first.params
163197
assert first.eval_hash != second.eval_hash

‎tests/unit/score/test_true_false_score_aggregator.py‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -276,3 +276,40 @@ def test_generator_of_wrong_type_still_raises():
276276
)
277277
with pytest.raises(ValueError, match="must be of type 'true_false'"):
278278
TrueFalseScoreAggregator.OR(s for s in [bad])
279+
280+
281+
# Tests for raise_on_empty behavior
282+
def test_and_raise_on_empty_with_scores():
283+
"""Test that AND_RAISE_ON_EMPTY works normally when scores are present."""
284+
scores = [_mk_score(True, prr_id="1"), _mk_score(False, prr_id="1")]
285+
res = TrueFalseScoreAggregator.AND_RAISE_ON_EMPTY(scores)
286+
assert res.value is False
287+
288+
289+
def test_or_raise_on_empty_with_scores():
290+
"""Test that OR_RAISE_ON_EMPTY works normally when scores are present."""
291+
scores = [_mk_score(True, prr_id="1"), _mk_score(False, prr_id="1")]
292+
res = TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY(scores)
293+
assert res.value is True
294+
295+
296+
def test_majority_raise_on_empty_with_scores():
297+
"""Test that MAJORITY_RAISE_ON_EMPTY works normally when scores are present."""
298+
scores = [_mk_score(True, prr_id="1"), _mk_score(True, prr_id="1"), _mk_score(False, prr_id="1")]
299+
res = TrueFalseScoreAggregator.MAJORITY_RAISE_ON_EMPTY(scores)
300+
assert res.value is True
301+
302+
303+
@pytest.mark.parametrize("aggregator_name", ["AND_RAISE_ON_EMPTY", "OR_RAISE_ON_EMPTY", "MAJORITY_RAISE_ON_EMPTY"])
304+
def test_raise_on_empty_aggregators_raise_with_no_scores(aggregator_name):
305+
"""The strict-empty variants must not answer an empty input with a verdict."""
306+
aggregator = getattr(TrueFalseScoreAggregator, aggregator_name)
307+
with pytest.raises(ValueError, match="No scores available for aggregation"):
308+
aggregator([])
309+
310+
311+
def test_raise_on_empty_aggregator_accepts_generators():
312+
"""A generator with scores must not trip the empty-input guard."""
313+
values = [True, False, True]
314+
res = TrueFalseScoreAggregator.OR_RAISE_ON_EMPTY(_mk_score(v, prr_id="1") for v in values)
315+
assert res.value is True

0 commit comments

Comments
 (0)