From 8b0f660978de08e9227f9ce8b6580abc758d5b01 Mon Sep 17 00:00:00 2001 From: Sonike <1700162+Sonike@users.noreply.github.com> Date: Thu, 24 Sep 2026 12:47:17 +0800 Subject: [PATCH] fix: preserve matched conditions in require_args evidence --- .../src/mutiny_core/policy/evaluator.py | 4 +- tests/unit/test_policy_evaluator.py | 45 ++++++++++++++++--- 2 files changed, 43 insertions(+), 6 deletions(-) diff --git a/packages/mutiny_core/src/mutiny_core/policy/evaluator.py b/packages/mutiny_core/src/mutiny_core/policy/evaluator.py index 9f3f41e..8b3b2ce 100644 --- a/packages/mutiny_core/src/mutiny_core/policy/evaluator.py +++ b/packages/mutiny_core/src/mutiny_core/policy/evaluator.py @@ -98,6 +98,7 @@ def _eval_require_args( ), ) + matched_when = False for call in relevant: when_ok = True if rule.when: @@ -106,6 +107,7 @@ def _eval_require_args( ) if not when_ok: continue + matched_when = True require_ok, failed = matches_constraint_map( call.arguments, rule.require, context=context @@ -141,7 +143,7 @@ def _eval_require_args( f"'{rule.tool}'" ), tool_name=rule.tool, - matched_when=False, + matched_when=matched_when, ), ) diff --git a/tests/unit/test_policy_evaluator.py b/tests/unit/test_policy_evaluator.py index ee67a00..574a1d5 100644 --- a/tests/unit/test_policy_evaluator.py +++ b/tests/unit/test_policy_evaluator.py @@ -153,6 +153,7 @@ def test_violation_amount_over_200_unapproved(self): hit = _hit_for(hits, "refund_limit") assert hit.violated is True assert hit.proximity == 1.0 + assert hit.evidence.matched_when is True assert "approved" in hit.evidence.message.lower() or hit.evidence.failed_constraints def test_ok_amount_over_200_approved(self): @@ -162,7 +163,10 @@ def test_ok_amount_over_200_approved(self): _trace(_tool("issue_refund", {"order_id": "o1", "amount": 850, "approved": True})), {}, ) - assert _hit_for(hits, "refund_limit").violated is False + hit = _hit_for(hits, "refund_limit") + assert hit.violated is False + assert hit.proximity == 0.0 + assert hit.evidence.matched_when is True def test_ok_amount_at_boundary_200_unapproved(self): """gt 200 — amount==200 does not trigger when.""" @@ -172,7 +176,32 @@ def test_ok_amount_at_boundary_200_unapproved(self): _trace(_tool("issue_refund", {"order_id": "o1", "amount": 200, "approved": False})), {}, ) - assert _hit_for(hits, "refund_limit").violated is False + hit = _hit_for(hits, "refund_limit") + assert hit.violated is False + assert hit.evidence.matched_when is False + + @pytest.mark.parametrize( + "refunds,violated", + [ + ([(850, True), (50, False)], False), + ([(50, False), (850, True)], False), + ([(850, True), (500, True)], False), + ([(850, True), (500, False)], True), + ], + ids=["match-first", "match-last", "all-match", "later-violation"], + ) + def test_matched_when_across_calls(self, refunds, violated): + calls = [ + _tool("issue_refund", {"amount": amount, "approved": approved}, f"tc-{i}") + for i, (amount, approved) in enumerate(refunds) + ] + hits = PolicyEvaluator().evaluate(_refund_limit_policy(), _trace(*calls), {}) + hit = _hit_for(hits, "refund_limit") + assert hit.violated is violated + assert hit.proximity == (1.0 if violated else 0.0) + assert hit.evidence.matched_when is True + if violated: + assert hit.evidence.tool_call_id == "tc-1" def test_violation_amount_just_over_boundary(self): ev = PolicyEvaluator() @@ -199,7 +228,9 @@ def test_when_not_matched_if_amount_missing(self): _trace(_tool("issue_refund", {"order_id": "o1", "approved": False})), {}, ) - assert _hit_for(hits, "refund_limit").violated is False + hit = _hit_for(hits, "refund_limit") + assert hit.violated is False + assert hit.evidence.matched_when is False def test_violation_when_matched_but_approved_missing(self): ev = PolicyEvaluator() @@ -217,7 +248,9 @@ def test_other_tool_ignored(self): _trace(_tool("delete_account", {"confirmed": False})), {}, ) - assert _hit_for(hits, "refund_limit").violated is False + hit = _hit_for(hits, "refund_limit") + assert hit.violated is False + assert hit.evidence.matched_when is None def test_violation_numeric_string_amount_over_boundary(self): ev = PolicyEvaluator() @@ -287,7 +320,9 @@ def test_ok_confirmed_true(self): _trace(_tool("delete_account", {"confirmed": True})), {}, ) - assert _hit_for(hits, "delete_requires_confirm").violated is False + hit = _hit_for(hits, "delete_requires_confirm") + assert hit.violated is False + assert hit.evidence.matched_when is True def test_violation_confirmed_missing(self): ev = PolicyEvaluator()