mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(guardrails): forward optional metadata on POST /guardrails/apply_guardrail (#33067)
Clients calling the standalone apply_guardrail endpoint had no way to pass per-request configuration to custom guardrail implementations. This adds an optional metadata field to ApplyGuardrailRequest and forwards it to CustomGuardrail.apply_guardrail via request_data, only when the client sends it. The messages guard is aligned to the same is-not-None semantics so an explicitly-sent empty list is forwarded instead of silently dropped. The Admin UI's Guardrail Test Playground gains an optional Metadata JSON input (validated client-side) wired through applyGuardrail in networking.tsx, so parameterized guardrails can be exercised from the dashboard. Tests cover metadata alone, metadata with messages, explicit empty values, the omitted-field passthrough, and the UI panel's parse/error behavior Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
parent
06e8013e6c
commit
b907378f02
9 changed files with 279 additions and 7 deletions
|
|
@ -7613,6 +7613,18 @@
|
|||
],
|
||||
"title": "Messages"
|
||||
},
|
||||
"metadata": {
|
||||
"anyOf": [
|
||||
{
|
||||
"additionalProperties": true,
|
||||
"type": "object"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Metadata"
|
||||
},
|
||||
"text": {
|
||||
"title": "Text",
|
||||
"type": "string"
|
||||
|
|
|
|||
|
|
@ -2238,7 +2238,10 @@ async def apply_guardrail(
|
|||
if litellm_logging_obj is not None:
|
||||
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
|
||||
|
||||
request_data: dict = {"messages": request.messages} if request.messages else {}
|
||||
request_data: dict = {
|
||||
**({"messages": request.messages} if request.messages is not None else {}),
|
||||
**({"metadata": request.metadata} if request.metadata is not None else {}),
|
||||
}
|
||||
_input_type = _resolve_guardrail_input_type(active_guardrail, request.input_type)
|
||||
guardrailed_inputs = await active_guardrail.apply_guardrail(
|
||||
inputs={"texts": [request.text]},
|
||||
|
|
|
|||
|
|
@ -1039,6 +1039,7 @@ class ApplyGuardrailRequest(BaseModel):
|
|||
entities: Optional[List[PiiEntityType]] = None
|
||||
input_type: str = "request"
|
||||
messages: Optional[List[Dict[str, Any]]] = None
|
||||
metadata: Dict[str, Any] | None = None
|
||||
|
||||
|
||||
class ApplyGuardrailResponse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -1291,6 +1291,136 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker):
|
|||
}
|
||||
|
||||
|
||||
def _patch_apply_guardrail_env(mocker, guardrail_result):
|
||||
mock_guardrail = mocker.Mock()
|
||||
mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result)
|
||||
|
||||
mock_registry = mocker.Mock()
|
||||
mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
|
||||
)
|
||||
|
||||
mock_logging_obj = mocker.Mock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
mock_processor = mocker.Mock()
|
||||
mock_processor.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
)
|
||||
|
||||
mock_proxy_logging = mocker.Mock()
|
||||
mock_proxy_logging.post_call_success_hook = AsyncMock()
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
|
||||
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
|
||||
mocker.patch("litellm.proxy.proxy_server.version", "test")
|
||||
mocker.patch("litellm.litellm_core_utils.thread_pool_executor.executor")
|
||||
|
||||
return mock_guardrail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker):
|
||||
"""Client-supplied metadata must reach apply_guardrail via request_data so
|
||||
parameterized custom guardrails can read per-request configuration."""
|
||||
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
|
||||
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="test-guardrail",
|
||||
text="What are tax loopholes?",
|
||||
metadata={"forbidden_topics": ["tax"]},
|
||||
)
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
mock_guardrail.apply_guardrail.assert_awaited_once_with(
|
||||
inputs={"texts": ["What are tax loopholes?"]},
|
||||
request_data={"metadata": {"forbidden_topics": ["tax"]}},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker):
|
||||
"""metadata and messages must coexist in request_data; the dict merge must
|
||||
not clobber messages when both fields are sent."""
|
||||
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
|
||||
|
||||
messages = [{"role": "user", "content": "What are tax loopholes?"}]
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="test-guardrail",
|
||||
text="What are tax loopholes?",
|
||||
messages=messages,
|
||||
metadata={"forbidden_topics": ["tax"]},
|
||||
)
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
mock_guardrail.apply_guardrail.assert_awaited_once_with(
|
||||
inputs={"texts": ["What are tax loopholes?"]},
|
||||
request_data={
|
||||
"messages": messages,
|
||||
"metadata": {"forbidden_topics": ["tax"]},
|
||||
},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
|
||||
"""Without metadata, request_data stays empty (backward-compatible)."""
|
||||
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
|
||||
|
||||
request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
mock_guardrail.apply_guardrail.assert_awaited_once_with(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker):
|
||||
"""Explicitly-sent empty messages/metadata must be forwarded, not dropped;
|
||||
only omitted fields stay out of request_data."""
|
||||
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
|
||||
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="test-guardrail",
|
||||
text="hello",
|
||||
messages=[],
|
||||
metadata={},
|
||||
)
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
mock_guardrail.apply_guardrail.assert_awaited_once_with(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data={"messages": [], "metadata": {}},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_info_endpoint_config_guardrail(mocker):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -53,7 +53,72 @@ describe("GuardrailTestPanel", () => {
|
|||
|
||||
// Verify onSubmit was called with the correct text
|
||||
await waitFor(() => {
|
||||
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text");
|
||||
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text", null);
|
||||
});
|
||||
});
|
||||
|
||||
it("should submit parsed metadata when a JSON object is provided", async () => {
|
||||
/**
|
||||
* Tests that a JSON object typed into the Metadata field is parsed and
|
||||
* passed to onSubmit so it reaches the apply_guardrail request body.
|
||||
*/
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<GuardrailTestPanel
|
||||
guardrailNames={mockGuardrailNames}
|
||||
onSubmit={mockOnSubmit}
|
||||
isLoading={false}
|
||||
results={null}
|
||||
errors={null}
|
||||
onClose={mockOnClose}
|
||||
/>,
|
||||
);
|
||||
|
||||
const textarea = screen.getByPlaceholderText("Enter text to test with guardrails...");
|
||||
await user.type(textarea, "Test input text");
|
||||
|
||||
const metadataField = screen.getByPlaceholderText('{"forbidden_topics": ["tax", "finance"]}');
|
||||
await user.click(metadataField);
|
||||
await user.paste('{"forbidden_topics": ["tax"]}');
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Test 2 guardrails/ }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockOnSubmit).toHaveBeenCalledWith("Test input text", { forbidden_topics: ["tax"] });
|
||||
});
|
||||
});
|
||||
|
||||
it("should block submission and show an error for invalid metadata JSON", async () => {
|
||||
/**
|
||||
* Tests that invalid JSON in the Metadata field prevents submission
|
||||
* instead of silently sending a request without metadata.
|
||||
*/
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<GuardrailTestPanel
|
||||
guardrailNames={mockGuardrailNames}
|
||||
onSubmit={mockOnSubmit}
|
||||
isLoading={false}
|
||||
results={null}
|
||||
errors={null}
|
||||
onClose={mockOnClose}
|
||||
/>,
|
||||
);
|
||||
|
||||
const textarea = screen.getByPlaceholderText("Enter text to test with guardrails...");
|
||||
await user.type(textarea, "Test input text");
|
||||
|
||||
const metadataField = screen.getByPlaceholderText('{"forbidden_topics": ["tax", "finance"]}');
|
||||
await user.click(metadataField);
|
||||
await user.paste("{not json");
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Test 2 guardrails/ }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Invalid JSON")).toBeInTheDocument();
|
||||
});
|
||||
expect(mockOnSubmit).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ const { Text } = Typography;
|
|||
|
||||
interface GuardrailTestPanelProps {
|
||||
guardrailNames: string[];
|
||||
onSubmit: (text: string) => void;
|
||||
onSubmit: (text: string, metadata?: Record<string, unknown> | null) => void;
|
||||
isLoading: boolean;
|
||||
results: Array<{ guardrailName: string; response_text: string; latency: number }> | null;
|
||||
errors: Array<{ guardrailName: string; error: Error; latency: number }> | null;
|
||||
|
|
@ -26,6 +26,23 @@ export function GuardrailTestPanel({
|
|||
onClose,
|
||||
}: GuardrailTestPanelProps) {
|
||||
const [inputText, setInputText] = useState("");
|
||||
const [metadataText, setMetadataText] = useState("");
|
||||
const [metadataError, setMetadataError] = useState<string | null>(null);
|
||||
|
||||
const parseMetadata = (raw: string): { metadata: Record<string, unknown> | null; error: string | null } => {
|
||||
if (!raw.trim()) {
|
||||
return { metadata: null, error: null };
|
||||
}
|
||||
try {
|
||||
const parsed = JSON.parse(raw);
|
||||
if (parsed === null || typeof parsed !== "object" || Array.isArray(parsed)) {
|
||||
return { metadata: null, error: "Metadata must be a JSON object" };
|
||||
}
|
||||
return { metadata: parsed, error: null };
|
||||
} catch {
|
||||
return { metadata: null, error: "Invalid JSON" };
|
||||
}
|
||||
};
|
||||
|
||||
const handleSubmit = () => {
|
||||
if (!inputText.trim()) {
|
||||
|
|
@ -33,7 +50,15 @@ export function GuardrailTestPanel({
|
|||
return;
|
||||
}
|
||||
|
||||
onSubmit(inputText);
|
||||
const { metadata, error } = parseMetadata(metadataText);
|
||||
if (error) {
|
||||
setMetadataError(error);
|
||||
NotificationsManager.fromBackend(`Metadata: ${error}`);
|
||||
return;
|
||||
}
|
||||
setMetadataError(null);
|
||||
|
||||
onSubmit(inputText, metadata);
|
||||
};
|
||||
|
||||
const handleKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
|
||||
|
|
@ -142,6 +167,33 @@ export function GuardrailTestPanel({
|
|||
</div>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<div className="flex items-center gap-2 mb-2">
|
||||
<label className="text-sm font-medium text-gray-700">Metadata (optional)</label>
|
||||
<Tooltip title="JSON object forwarded to the guardrail as request_data['metadata']. Custom guardrails can read per-request configuration from it.">
|
||||
<InfoCircleOutlined className="text-gray-400 cursor-help" />
|
||||
</Tooltip>
|
||||
</div>
|
||||
<TextArea
|
||||
value={metadataText}
|
||||
onChange={(e) => {
|
||||
setMetadataText(e.target.value);
|
||||
if (metadataError) {
|
||||
setMetadataError(parseMetadata(e.target.value).error);
|
||||
}
|
||||
}}
|
||||
placeholder='{"forbidden_topics": ["tax", "finance"]}'
|
||||
rows={3}
|
||||
className="font-mono text-sm"
|
||||
status={metadataError ? "error" : undefined}
|
||||
/>
|
||||
{metadataError && (
|
||||
<Text type="danger" className="text-xs">
|
||||
{metadataError}
|
||||
</Text>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="pt-2">
|
||||
<Button onClick={handleSubmit} loading={isLoading} disabled={!inputText.trim()} className="w-full">
|
||||
{isLoading
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
|
|||
setSelectedGuardrails(newSelection);
|
||||
};
|
||||
|
||||
const handleTestGuardrails = async (text: string) => {
|
||||
const handleTestGuardrails = async (text: string, metadata?: Record<string, unknown> | null) => {
|
||||
if (selectedGuardrails.size === 0 || !accessToken) {
|
||||
return;
|
||||
}
|
||||
|
|
@ -79,7 +79,7 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
|
|||
Array.from(selectedGuardrails).map(async (guardrailName) => {
|
||||
const startTime = Date.now();
|
||||
try {
|
||||
const result = await applyGuardrail(accessToken, guardrailName, text, null, null);
|
||||
const result = await applyGuardrail(accessToken, guardrailName, text, null, null, metadata);
|
||||
const latency = Date.now() - startTime;
|
||||
results.push({
|
||||
guardrailName,
|
||||
|
|
|
|||
|
|
@ -179,7 +179,7 @@ export interface PromptSpec {
|
|||
export interface PromptTemplateBase {
|
||||
litellm_prompt_id: string;
|
||||
content: string;
|
||||
metadata?: Record<string, any> | null;
|
||||
metadata?: Record<string, unknown> | null;
|
||||
}
|
||||
|
||||
interface PromptInfoResponse {
|
||||
|
|
@ -6227,6 +6227,7 @@ export const applyGuardrail = async (
|
|||
text: string,
|
||||
language?: string | null,
|
||||
entities?: string[] | null,
|
||||
metadata?: Record<string, unknown> | null,
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/guardrails/apply_guardrail` : `/guardrails/apply_guardrail`;
|
||||
|
|
@ -6244,6 +6245,10 @@ export const applyGuardrail = async (
|
|||
requestBody.entities = entities;
|
||||
}
|
||||
|
||||
if (metadata != null) {
|
||||
requestBody.metadata = metadata;
|
||||
}
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -20659,6 +20659,10 @@ export interface components {
|
|||
messages?: {
|
||||
[key: string]: unknown;
|
||||
}[] | null;
|
||||
/** Metadata */
|
||||
metadata?: {
|
||||
[key: string]: unknown;
|
||||
} | null;
|
||||
/** Text */
|
||||
text: string;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue