Skip to content

Commit 56f5202

Browse files
committed
feat: add RetryPolicy for automatic RETRY failure mode handling
1 parent e133e53 commit 56f5202

4 files changed

Lines changed: 238 additions & 1 deletion

File tree

gateframe/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from gateframe.core.contract import ValidationContract, ValidationResult
33
from gateframe.core.escalation import EscalationRoute, EscalationRouter
44
from gateframe.core.failure import FailureMode, FailureResult
5+
from gateframe.core.retry import RetryPolicy, RetryResult
56
from gateframe.rules.boundary import AllowedValues, BoundaryRule
67
from gateframe.rules.confidence import ConfidenceRule
78
from gateframe.rules.semantic import LlmJudge, SemanticRule
@@ -18,6 +19,8 @@
1819
"LlmJudge",
1920
"SemanticRule",
2021
"StructuralRule",
22+
"RetryPolicy",
23+
"RetryResult",
2124
"ValidationContract",
2225
"ValidationResult",
2326
"WorkflowContext",

gateframe/core/retry.py

Lines changed: 106 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,106 @@
1+
from __future__ import annotations
2+
3+
from collections.abc import Callable
4+
from dataclasses import dataclass, field
5+
from typing import Any
6+
7+
from gateframe.core.contract import ValidationContract, ValidationResult
8+
from gateframe.core.failure import FailureMode
9+
10+
11+
@dataclass
12+
class RetryResult:
13+
"""Outcome of a RetryPolicy.run() call.
14+
15+
Attributes:
16+
final: The last ValidationResult produced (pass or exhausted retries).
17+
attempts: Total number of validate calls made (1 = no retry needed).
18+
succeeded: True if any attempt passed validation.
19+
"""
20+
21+
final: ValidationResult
22+
attempts: int
23+
succeeded: bool
24+
history: list[ValidationResult] = field(default_factory=list)
25+
26+
27+
class RetryPolicy:
28+
"""Executes a ValidationContract with automatic retry on RETRY failure mode.
29+
30+
The caller supplies a ``prompt_fn`` callable that returns a fresh output
31+
on each invocation (e.g. re-calls the LLM). On each attempt, if the
32+
result contains a RETRY failure, ``prompt_fn`` is called again with the
33+
previous failures injected as ``retry_failures`` context so the caller
34+
can craft a follow-up prompt.
35+
36+
Args:
37+
max_retries: Maximum number of *additional* attempts after the first.
38+
Total attempts = max_retries + 1.
39+
40+
Example::
41+
42+
policy = RetryPolicy(max_retries=2)
43+
result = policy.run(
44+
contract,
45+
prompt_fn=lambda **ctx: call_llm(ctx.get("retry_failures")),
46+
)
47+
if result.succeeded:
48+
use(result.final)
49+
"""
50+
51+
def __init__(self, max_retries: int = 3) -> None:
52+
if max_retries < 0:
53+
raise ValueError("max_retries must be >= 0")
54+
self.max_retries = max_retries
55+
56+
def run(
57+
self,
58+
contract: ValidationContract,
59+
prompt_fn: Callable[..., Any],
60+
**context: Any,
61+
) -> RetryResult:
62+
"""Run *contract* against output from *prompt_fn*, retrying on RETRY failures.
63+
64+
Args:
65+
contract: The ValidationContract to evaluate.
66+
prompt_fn: Called each attempt to produce the output to validate.
67+
Receives ``**context`` plus ``retry_failures`` (list of
68+
FailureResult) on retry attempts so callers can adjust prompts.
69+
**context: Extra keyword arguments forwarded to both ``prompt_fn``
70+
and ``contract.validate()``.
71+
72+
Returns:
73+
A :class:`RetryResult` with the final result and attempt metadata.
74+
"""
75+
history: list[ValidationResult] = []
76+
retry_failures: list = []
77+
78+
for attempt in range(self.max_retries + 1):
79+
call_context = dict(context)
80+
if retry_failures:
81+
call_context["retry_failures"] = retry_failures
82+
83+
output = prompt_fn(**call_context)
84+
result = contract.validate(output, **call_context)
85+
history.append(result)
86+
87+
if result.passed:
88+
return RetryResult(
89+
final=result,
90+
attempts=attempt + 1,
91+
succeeded=True,
92+
history=history,
93+
)
94+
95+
has_retry = any(f.failure_mode is FailureMode.RETRY for f in result.failures)
96+
if not has_retry or attempt == self.max_retries:
97+
break
98+
99+
retry_failures = [f for f in result.failures if f.failure_mode is FailureMode.RETRY]
100+
101+
return RetryResult(
102+
final=history[-1],
103+
attempts=len(history),
104+
succeeded=False,
105+
history=history,
106+
)

tests/core/test_context.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -194,7 +194,6 @@ def test_reset_restores_initial_confidence_and_clears_history(self) -> None:
194194
assert ctx.history == []
195195
assert ctx.threshold_breached is False
196196

197-
198197
def test_concurrent_updates_do_not_corrupt_confidence(self) -> None:
199198
ctx = WorkflowContext("wf_concurrent", soft_fail_penalty=0.01)
200199
result = _make_result(passed=False, failures=[_soft_failure()])

tests/core/test_retry.py

Lines changed: 129 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,129 @@
1+
from typing import Any
2+
3+
import pytest
4+
5+
from gateframe.core.contract import ValidationContract
6+
from gateframe.core.failure import FailureMode, FailureResult
7+
from gateframe.core.retry import RetryPolicy
8+
from gateframe.core.rule import Rule
9+
10+
11+
class _PassingRule(Rule):
12+
def validate(self, output: Any, **context: Any) -> FailureResult | None:
13+
return None
14+
15+
16+
class _FailingRule(Rule):
17+
def __init__(self, name: str = "failing", mode: FailureMode = FailureMode.RETRY) -> None:
18+
super().__init__(name)
19+
self._mode = mode
20+
21+
def validate(self, output: Any, **context: Any) -> FailureResult | None:
22+
return FailureResult(rule_name=self.name, failure_mode=self._mode, message="rejected")
23+
24+
25+
class _PassOnAttempt(Rule):
26+
"""Passes on the Nth call (1-indexed)."""
27+
28+
def __init__(self, pass_on: int) -> None:
29+
super().__init__("pass_on_attempt")
30+
self._pass_on = pass_on
31+
self._calls = 0
32+
33+
def validate(self, output: Any, **context: Any) -> FailureResult | None:
34+
self._calls += 1
35+
if self._calls >= self._pass_on:
36+
return None
37+
return FailureResult(rule_name=self.name, failure_mode=FailureMode.RETRY, message="not yet")
38+
39+
40+
class TestRetryPolicy:
41+
def test_passes_on_first_attempt(self) -> None:
42+
contract = ValidationContract("c", [_PassingRule("r")])
43+
policy = RetryPolicy(max_retries=2)
44+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
45+
assert result.succeeded is True
46+
assert result.attempts == 1
47+
48+
def test_retries_until_pass(self) -> None:
49+
rule = _PassOnAttempt(pass_on=3)
50+
contract = ValidationContract("c", [rule])
51+
policy = RetryPolicy(max_retries=3)
52+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
53+
assert result.succeeded is True
54+
assert result.attempts == 3
55+
56+
def test_exhausts_retries_and_fails(self) -> None:
57+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.RETRY)])
58+
policy = RetryPolicy(max_retries=2)
59+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
60+
assert result.succeeded is False
61+
assert result.attempts == 3 # 1 initial + 2 retries
62+
63+
def test_does_not_retry_on_hard_fail(self) -> None:
64+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.HARD_FAIL)])
65+
policy = RetryPolicy(max_retries=3)
66+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
67+
assert result.succeeded is False
68+
assert result.attempts == 1
69+
70+
def test_does_not_retry_on_soft_fail(self) -> None:
71+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.SOFT_FAIL)])
72+
policy = RetryPolicy(max_retries=3)
73+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
74+
assert result.succeeded is False
75+
assert result.attempts == 1
76+
77+
def test_retry_failures_passed_to_prompt_fn(self) -> None:
78+
received: list = []
79+
80+
def prompt_fn(**ctx: Any) -> dict:
81+
received.append(ctx.get("retry_failures"))
82+
return {}
83+
84+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.RETRY)])
85+
policy = RetryPolicy(max_retries=1)
86+
policy.run(contract, prompt_fn=prompt_fn)
87+
88+
assert received[0] is None # first call has no retry context
89+
assert received[1] is not None # second call has failures injected
90+
assert received[1][0].failure_mode is FailureMode.RETRY
91+
92+
def test_context_forwarded_to_contract(self) -> None:
93+
class _ContextRule(Rule):
94+
def validate(self, output: Any, **context: Any) -> FailureResult | None:
95+
if context.get("role") != "admin":
96+
return FailureResult(
97+
rule_name=self.name,
98+
failure_mode=FailureMode.HARD_FAIL,
99+
message="forbidden",
100+
)
101+
return None
102+
103+
contract = ValidationContract("c", [_ContextRule("role_check")])
104+
policy = RetryPolicy()
105+
result = policy.run(contract, prompt_fn=lambda **ctx: {}, role="admin")
106+
assert result.succeeded is True
107+
108+
def test_history_contains_all_attempts(self) -> None:
109+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.RETRY)])
110+
policy = RetryPolicy(max_retries=2)
111+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
112+
assert len(result.history) == 3
113+
114+
def test_max_retries_zero_means_single_attempt(self) -> None:
115+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.RETRY)])
116+
policy = RetryPolicy(max_retries=0)
117+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
118+
assert result.attempts == 1
119+
assert result.succeeded is False
120+
121+
def test_negative_max_retries_raises(self) -> None:
122+
with pytest.raises(ValueError, match="max_retries"):
123+
RetryPolicy(max_retries=-1)
124+
125+
def test_retry_result_final_is_last_result(self) -> None:
126+
contract = ValidationContract("c", [_FailingRule(mode=FailureMode.RETRY)])
127+
policy = RetryPolicy(max_retries=1)
128+
result = policy.run(contract, prompt_fn=lambda **ctx: {})
129+
assert result.final is result.history[-1]

0 commit comments

Comments
 (0)