feat: add retry configuration for guardrails in pipeline flow editor

- Add num_retries field to PipelineStep (0-10, default 0) with exponential backoff
- Update PipelineExecutor to retry failed guardrail steps before applying on_fail action
- Add retries_attempted to PipelineStepResult for observability
- Add RETRIES section to flow builder StepCard with InputNumber control and dynamic helper text
- Show retry count in read-only PipelineInfoDisplay and test result panel
- Add comprehensive tests for retry behavior (transient failures, exhaustion, pass-through)

Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
Cursor Agent 2026-03-11 01:08:10 +00:00
parent d9e6758655
commit 07e6c634ff
6 changed files with 366 additions and 4 deletions

View file

@ -5,6 +5,7 @@ Runs guardrails sequentially per pipeline step definitions, handling
pass/fail actions (allow, block, next, modify_response) and data forwarding.
"""
import asyncio
import time
from typing import Any, List, Optional
@ -72,6 +73,27 @@ class PipelineExecutor:
call_type=call_type,
)
retries_attempted = 0
if outcome != "pass" and step.num_retries > 0:
for retry in range(step.num_retries):
retries_attempted = retry + 1
verbose_proxy_logger.debug(
f"Pipeline '{policy_name}' step {i}: retrying guardrail "
f"'{step.guardrail}' (attempt {retries_attempted}/{step.num_retries})"
)
await asyncio.sleep(0.1 * retries_attempted)
outcome, modified_data, error_detail = (
await PipelineExecutor._run_step(
step=step,
mode=mode,
data=working_data,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
)
)
if outcome == "pass":
break
duration = time.perf_counter() - start_time
action = step.on_pass if outcome == "pass" else step.on_fail
@ -83,12 +105,13 @@ class PipelineExecutor:
modified_data=modified_data,
error_detail=error_detail,
duration_seconds=round(duration, 4),
retries_attempted=retries_attempted,
)
step_results.append(step_result)
verbose_proxy_logger.debug(
f"Pipeline '{policy_name}' step {i}: guardrail={step.guardrail}, "
f"outcome={outcome}, action={action}"
f"outcome={outcome}, action={action}, retries={retries_attempted}"
)
# Forward modified data to next step if pass_data is True

View file

@ -38,6 +38,12 @@ class PipelineStep(BaseModel):
default=None,
description="Custom message for modify_response action.",
)
num_retries: int = Field(
default=0,
ge=0,
le=10,
description="Number of times to retry the guardrail on failure before applying the on_fail action.",
)
model_config = ConfigDict(extra="forbid")
@ -86,6 +92,10 @@ class PipelineStepResult(BaseModel):
modified_data: Optional[Dict[str, Any]] = None
error_detail: Optional[str] = None
duration_seconds: Optional[float] = None
retries_attempted: int = Field(
default=0,
description="Number of retries that were attempted before the final outcome.",
)
class PipelineExecutionResult(BaseModel):

View file

@ -46,6 +46,28 @@ class AlwaysFailGuardrail(CustomGuardrail):
raise HTTPException(status_code=400, detail="Content policy violation")
class FailThenPassGuardrail(CustomGuardrail):
"""Mock guardrail that fails N times then passes."""
def __init__(self, guardrail_name: str, fail_count: int = 2):
super().__init__(
guardrail_name=guardrail_name,
event_hook="pre_call",
default_on=True,
)
self.calls = 0
self.fail_count = fail_count
def should_run_guardrail(self, data, event_type) -> bool:
return True
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
self.calls += 1
if self.calls <= self.fail_count:
raise HTTPException(status_code=400, detail="Transient failure")
return None
class AlwaysPassGuardrail(CustomGuardrail):
"""Mock guardrail that always passes."""
@ -482,3 +504,220 @@ async def test_step_results_include_duration():
assert result.step_results[0].duration_seconds >= 0
finally:
litellm.callbacks = original_callbacks
# ─────────────────────────────────────────────────────────────────────────────
# Retry Tests
# ─────────────────────────────────────────────────────────────────────────────
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_retry_succeeds_after_transient_failure():
"""
Guardrail fails twice then passes. With num_retries=3, it should succeed.
"""
guard = FailThenPassGuardrail(guardrail_name="flaky-guard", fail_count=2)
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(
guardrail="flaky-guard",
on_fail="block",
on_pass="allow",
num_retries=3,
)
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="retry-test",
)
assert guard.calls == 3 # 1 initial + 2 retries
assert result.terminal_action == "allow"
assert result.step_results[0].outcome == "pass"
assert result.step_results[0].retries_attempted == 2
finally:
litellm.callbacks = original_callbacks
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_retry_exhausted_still_fails():
"""
Guardrail fails 4 times. With num_retries=2, it should exhaust retries and fail.
"""
guard = FailThenPassGuardrail(guardrail_name="stubborn-guard", fail_count=4)
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(
guardrail="stubborn-guard",
on_fail="block",
on_pass="allow",
num_retries=2,
)
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="retry-test",
)
assert guard.calls == 3 # 1 initial + 2 retries
assert result.terminal_action == "block"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].retries_attempted == 2
finally:
litellm.callbacks = original_callbacks
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_no_retries_when_zero():
"""
With num_retries=0 (default), guardrail failure triggers on_fail immediately.
"""
guard = AlwaysFailGuardrail(guardrail_name="no-retry")
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(
guardrail="no-retry",
on_fail="block",
on_pass="allow",
num_retries=0,
)
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="retry-test",
)
assert guard.calls == 1
assert result.terminal_action == "block"
assert result.step_results[0].retries_attempted == 0
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_no_retries_on_pass():
"""
When a guardrail passes on the first try, no retries should be attempted.
"""
guard = AlwaysPassGuardrail(guardrail_name="pass-guard")
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(
guardrail="pass-guard",
on_fail="block",
on_pass="allow",
num_retries=3,
)
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard]
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="retry-test",
)
assert guard.calls == 1
assert result.terminal_action == "allow"
assert result.step_results[0].retries_attempted == 0
finally:
litellm.callbacks = original_callbacks
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
@pytest.mark.asyncio
async def test_retry_with_on_fail_next_still_retries():
"""
Retries happen even when on_fail is 'next'. After retries exhausted, the
on_fail action is taken.
"""
guard = AlwaysFailGuardrail(guardrail_name="retry-then-next")
fallback = AlwaysPassGuardrail(guardrail_name="fallback")
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(
guardrail="retry-then-next",
on_fail="next",
on_pass="allow",
num_retries=1,
),
PipelineStep(
guardrail="fallback",
on_fail="block",
on_pass="allow",
),
],
)
original_callbacks = litellm.callbacks.copy()
litellm.callbacks = [guard, fallback]
try:
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data={"messages": [{"role": "user", "content": "test"}]},
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="retry-test",
)
assert guard.calls == 2 # 1 initial + 1 retry
assert fallback.calls == 1
assert result.terminal_action == "allow"
assert result.step_results[0].outcome == "fail"
assert result.step_results[0].action_taken == "next"
assert result.step_results[0].retries_attempted == 1
assert result.step_results[1].outcome == "pass"
finally:
litellm.callbacks = original_callbacks

View file

@ -23,6 +23,7 @@ def test_pipeline_step_defaults():
assert step.on_pass == "allow"
assert step.pass_data is False
assert step.modify_response_message is None
assert step.num_retries == 0
def test_pipeline_step_valid_actions():
@ -150,3 +151,42 @@ def test_pipeline_extra_fields_rejected():
steps=[PipelineStep(guardrail="g")],
unknown="value",
)
def test_pipeline_step_num_retries():
step = PipelineStep(guardrail="g", num_retries=3)
assert step.num_retries == 3
def test_pipeline_step_num_retries_max():
step = PipelineStep(guardrail="g", num_retries=10)
assert step.num_retries == 10
def test_pipeline_step_num_retries_negative_rejected():
with pytest.raises(ValidationError):
PipelineStep(guardrail="g", num_retries=-1)
def test_pipeline_step_num_retries_exceeds_max_rejected():
with pytest.raises(ValidationError):
PipelineStep(guardrail="g", num_retries=11)
def test_pipeline_step_result_retries_attempted():
result = PipelineStepResult(
guardrail_name="g1",
outcome="pass",
action_taken="allow",
retries_attempted=2,
)
assert result.retries_attempted == 2
def test_pipeline_step_result_retries_attempted_default():
result = PipelineStepResult(
guardrail_name="g1",
outcome="pass",
action_taken="allow",
)
assert result.retries_attempted == 0

View file

@ -1,5 +1,5 @@
import React, { useState } from "react";
import { Select, Typography, message, Spin } from "antd";
import { Select, Typography, message, Spin, InputNumber } from "antd";
import { Button, TextInput } from "@tremor/react";
import { ArrowLeftIcon, PlusIcon } from "@heroicons/react/outline";
import { DotsVerticalIcon } from "@heroicons/react/solid";
@ -46,6 +46,7 @@ function createDefaultStep(): PipelineStep {
on_fail: "block",
pass_data: false,
modify_response_message: null,
num_retries: 0,
};
}
@ -92,6 +93,7 @@ function derivePipelineFromPolicy(policy: Policy | null | undefined): GuardrailP
on_fail: "block" as const,
pass_data: false,
modify_response_message: null,
num_retries: 0,
})),
};
}
@ -154,6 +156,13 @@ const FailIcon: React.FC = () => (
</svg>
);
const RetryIcon: React.FC = () => (
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="#f59e0b" strokeWidth="2.5" strokeLinecap="round" strokeLinejoin="round" style={{ flexShrink: 0 }}>
<polyline points="23 4 23 10 17 10" />
<path d="M20.49 15a9 9 0 1 1-2.12-9.36L23 10" />
</svg>
);
// ─────────────────────────────────────────────────────────────────────────────
// Connector
// ─────────────────────────────────────────────────────────────────────────────
@ -348,6 +357,35 @@ const StepCard: React.FC<StepCardProps> = ({
</div>
)}
</div>
{/* RETRIES section */}
<div style={{ borderTop: "1px solid #f0f0f0", padding: "14px 20px" }}>
<div className="flex items-center gap-2" style={{ marginBottom: 8 }}>
<RetryIcon />
<span style={{ fontSize: 13, fontWeight: 600, color: "#374151" }}>RETRIES</span>
</div>
<div className="flex items-center gap-3">
<div style={{ flex: 1 }}>
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
Retry on failure
</label>
<InputNumber
min={0}
max={10}
value={step.num_retries ?? 0}
onChange={(value) => onChange({ num_retries: value ?? 0 })}
style={{ width: "100%" }}
/>
</div>
<div style={{ flex: 2 }}>
<span style={{ fontSize: 12, color: "#9ca3af", lineHeight: 1.4, display: "block", marginTop: 18 }}>
{(step.num_retries ?? 0) === 0
? "No retries — on_fail action runs immediately."
: `Retry up to ${step.num_retries} time${(step.num_retries ?? 0) > 1 ? "s" : ""} before applying the on_fail action.`}
</span>
</div>
</div>
</div>
</div>
);
};
@ -565,14 +603,19 @@ export const PipelineInfoDisplay: React.FC<PipelineInfoDisplayProps> = ({ pipeli
{/* Divider */}
<div style={{ borderTop: "1px solid #f3f4f6", marginBottom: 10 }} />
{/* Pass / Fail */}
<div className="flex items-center gap-6" style={{ fontSize: 13, color: "#374151" }}>
{/* Pass / Fail / Retries */}
<div className="flex items-center gap-6" style={{ fontSize: 13, color: "#374151", flexWrap: "wrap" }}>
<span className="flex items-center gap-1.5">
<PassIcon /> Pass &#8594; {ACTION_LABELS[step.on_pass] || step.on_pass}
</span>
<span className="flex items-center gap-1.5">
<FailIcon /> Fail &#8594; {ACTION_LABELS[step.on_fail] || step.on_fail}
</span>
{(step.num_retries ?? 0) > 0 && (
<span className="flex items-center gap-1.5">
<RetryIcon /> {step.num_retries} {step.num_retries === 1 ? "retry" : "retries"}
</span>
)}
</div>
</div>
</React.Fragment>
@ -844,6 +887,11 @@ const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
({(step.duration_seconds * 1000).toFixed(0)}ms)
</span>
)}
{(step.retries_attempted ?? 0) > 0 && (
<span style={{ marginLeft: 8, color: "#f59e0b" }}>
({step.retries_attempted} {step.retries_attempted === 1 ? "retry" : "retries"})
</span>
)}
</div>
{step.error_detail && (
<div style={{ fontSize: 12, color: "#dc2626", marginTop: 4 }}>

View file

@ -26,6 +26,7 @@ export interface PipelineStep {
on_pass: "allow" | "block" | "next" | "modify_response";
pass_data?: boolean;
modify_response_message?: string | null;
num_retries?: number;
}
export interface GuardrailPipeline {
@ -99,6 +100,7 @@ export interface PipelineStepResult {
modified_data: Record<string, any> | null;
error_detail: string | null;
duration_seconds: number | null;
retries_attempted?: number;
}
export interface PipelineTestResult {