mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
d9e6758655
commit
07e6c634ff
6 changed files with 366 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 → {ACTION_LABELS[step.on_pass] || step.on_pass}
|
||||
</span>
|
||||
<span className="flex items-center gap-1.5">
|
||||
<FailIcon /> Fail → {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 }}>
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue