Merge remote-tracking branch 'origin/litellm_internal_staging' into feat-prism-provider

This commit is contained in:
rajitkhanna 2026-09-11 19:26:46 -07:00
commit 855533b9ca
6 changed files with 230 additions and 9 deletions

View file

@ -497,9 +497,14 @@ class OpenAIResponsesHandler(BaseTranslation):
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
else:
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
if len(rewritten_texts) != len(extracted.task_mappings):
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
raise UnappliableRequestRewrite(guardrail_to_apply.guardrail_name or "unknown")
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=guardrailed_inputs.get("texts") or (),
responses=rewritten_texts,
task_mappings=extracted.task_mappings,
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
@ -635,10 +640,12 @@ class OpenAIResponsesHandler(BaseTranslation):
"""
Apply guardrail responses back to input messages.
``responses`` pairs positionally with ``task_mappings``; the caller rejects
the request when the two disagree, so this never has to guess an alignment.
Override this method to customize how responses are applied.
"""
for task_idx, guardrail_response in enumerate(responses):
mapping = task_mappings[task_idx]
for guardrail_response, mapping in zip(responses, task_mappings):
msg_idx = cast(int, mapping[0])
content_idx_optional = cast(int | None, mapping[1])

View file

@ -18,6 +18,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
effective_skip_tool_message_for_guardrail,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
@ -261,6 +262,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
fail_on_error: bool | None = True,
streaming_end_of_stream_only: bool | None = None,
streaming_sampling_rate: int | None = None,
async_handler: AsyncHTTPHandler | None = None,
**kwargs,
) -> None:
"""
@ -273,9 +275,13 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
streaming_end_of_stream_only (bool | None): Scan streamed output once at end of stream instead of
every streaming_sampling_rate chunks. Defaults to False.
streaming_sampling_rate (int | None): Scan the accumulated streamed output every Nth chunk. Defaults to 5.
async_handler (AsyncHTTPHandler | None): HTTP client to call AI Guard with. Defaults to the shared
guardrail-callback client.
**kwargs: Additional arguments passed to the CustomGuardrail base class.
"""
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.fail_on_error = True if fail_on_error is None else fail_on_error
self._set_streaming_params(
CrowdStrikeAIDRGuardrailConfigModelOptionalParams(

View file

@ -58,6 +58,15 @@ class UndeliverableStreamRewrite(Exception):
self.guardrail_name: Final = guardrail_name
class UnappliableRequestRewrite(Exception):
def __init__(self, guardrail_name: str) -> None:
super().__init__(
f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, "
"so the request was rejected rather than sent unrewritten"
)
self.guardrail_name: Final = guardrail_name
def _tool_call_shape(tool_call: object) -> tuple[object, object]:
plain: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call
function: Final = plain.get("function") if isinstance(plain, Mapping) else None

View file

@ -1,3 +1,7 @@
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Final, cast
import json
from unittest.mock import patch
import httpx
@ -7,6 +11,9 @@ from pydantic import ValidationError
import litellm
from litellm.exceptions import Timeout
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import initialize_guardrail
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr.crowdstrike_aidr import (
@ -1719,3 +1726,165 @@ async def test_streaming_params_from_config_control_output_scan_cadence(
handler = _initialize_from_config(mode="post_call", **configured)
assert await _guard_calls_for_stream(handler, list("ABCDEFGHIJ")) == expected_calls
@asynccontextmanager
async def _guardrail_redacting(secret: str, replacement: str) -> AsyncIterator[CrowdStrikeAIDRHandler]:
def redacted(content: object) -> object:
if isinstance(content, str):
return content.replace(secret, replacement)
if isinstance(content, list):
return [
{**part, "text": redacted(part["text"])} if isinstance(part, dict) and "text" in part else part
for part in content
]
return content
def respond(request: httpx.Request) -> httpx.Response:
sent: Final = json.loads(request.content)["guard_input"]["messages"]
return httpx.Response(
status_code=200,
json={
"result": {
"blocked": False,
"transformed": True,
"guard_output": {
"messages": [{**message, "content": redacted(message.get("content"))} for message in sent]
},
},
},
request=request,
)
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
handler: Final = AsyncHTTPHandler()
handler.client = client
yield CrowdStrikeAIDRHandler(
mode="pre_call",
guardrail_name="crowdstrike-aidr-guard",
api_key="pts_crowdstrike_tokenid",
api_base="https://api.crowdstrike.com/aidr/aiguard",
async_handler=handler,
)
class _MessageShapedGuardrail(CustomGuardrail):
"""Returns one text per chat message and no ``structured_messages`` rewrite.
Prompt Security and friends scan messages rather than Responses text parts,
which is the shape that outnumbers the endpoint's own bookkeeping.
"""
def __init__(self, redacted: str) -> None:
super().__init__(guardrail_name="message-shaped")
self.redacted: Final = redacted
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: str,
logging_obj: object = None,
) -> GenericGuardrailAPIInputs:
messages: Final = inputs.get("structured_messages") or ()
return {"texts": [self.redacted for _ in messages]}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("case", "instructions", "responses_input"),
[
(
"instructions add a system message",
"be terse",
[{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}],
),
(
"tool items add messages that carry no text",
None,
[
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]},
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
{"type": "function_call_output", "call_id": "c1", "output": "42"},
],
),
],
)
async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
case: str,
instructions: str | None,
responses_input: list[dict[str, object]],
) -> None:
"""An unalignable rewrite must fail the request, not forward the raw prompt.
Skipping the write-back would hand the model the unredacted text, so a
guardrail could be bypassed by adding ``instructions`` or a tool call.
"""
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
if instructions is not None:
data["instructions"] = instructions
with pytest.raises(UnappliableRequestRewrite):
await OpenAIResponsesHandler().process_input_messages(
data=data,
guardrail_to_apply=_MessageShapedGuardrail("my ssn is <US_SSN>"),
)
assert "078-05-1120" in str(responses_input), case
@pytest.mark.asyncio
async def test_aligned_rewrite_is_written_back() -> None:
"""Matching counts must still redact the input in place."""
responses_input: list[dict[str, object]] = [
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}
]
await OpenAIResponsesHandler().process_input_messages(
data={"model": "gpt-4o", "input": responses_input},
guardrail_to_apply=_MessageShapedGuardrail("my ssn is <US_SSN>"),
)
assert cast(list, responses_input[0]["content"])[0]["text"] == "my ssn is <US_SSN>"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("case", "responses_input", "redacted_input"),
[
(
"instructions add a system message",
[{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}],
[{"role": "user", "content": [{"type": "input_text", "text": "my ssn is <US_SSN>"}]}],
),
(
"tool items sit between two user turns",
[
{"role": "user", "content": [{"type": "input_text", "text": "hello"}]},
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
{"type": "function_call_output", "call_id": "c1", "output": "42"},
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]},
],
[
{"role": "user", "content": [{"type": "input_text", "text": "hello"}]},
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
{"type": "function_call_output", "call_id": "c1", "output": "42"},
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is <US_SSN>"}]},
],
),
],
)
async def test_structured_rewrite_lands_on_shapes_the_flat_path_cannot_align(
case: str,
responses_input: list[dict[str, object]],
redacted_input: list[dict[str, object]],
) -> None:
data: dict[str, object] = {"model": "gpt-5.6", "instructions": "be terse", "input": responses_input}
async with _guardrail_redacting("078-05-1120", "<US_SSN>") as guardrail:
await OpenAIResponsesHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
assert data["input"] == redacted_input, case
assert data["instructions"] == "be terse", case

View file

@ -1,4 +1,4 @@
import { renderWithProviders, screen } from "../../tests/test-utils";
import { fireEvent, renderWithProviders, screen } from "../../tests/test-utils";
import { vi } from "vitest";
import { EnvCredentialLoginWarningBanner } from "./EnvCredentialLoginWarningBanner";
import type { HealthReadinessDetailsResponse } from "@/app/(dashboard)/hooks/healthReadiness/useHealthReadinessDetails";
@ -23,6 +23,22 @@ const mockRole = (userRole: string) => {
};
describe("EnvCredentialLoginWarningBanner", () => {
beforeEach(() => {
localStorage.clear();
});
it("should hide the banner when dismissed and stay hidden on remount", () => {
mockRole("Admin");
mockDetails({ status: "healthy", show_env_credential_login_warning: true });
const first = renderWithProviders(<EnvCredentialLoginWarningBanner accessToken="token" />);
fireEvent.click(screen.getByRole("button", { name: "Dismiss banner" }));
expect(first.container).toBeEmptyDOMElement();
first.unmount();
const second = renderWithProviders(<EnvCredentialLoginWarningBanner accessToken="token" />);
expect(second.container).toBeEmptyDOMElement();
});
it("should warn an admin when env-credential login is enabled", () => {
mockRole("Admin");
mockDetails({ status: "healthy", show_env_credential_login_warning: true });

View file

@ -1,26 +1,37 @@
"use client";
import React from "react";
import { TriangleAlert } from "lucide-react";
import React, { useState } from "react";
import { TriangleAlert, X } from "lucide-react";
import { Button } from "@/components/ui/button";
import { useHealthReadinessDetails } from "@/app/(dashboard)/hooks/healthReadiness/useHealthReadinessDetails";
import { useAuth } from "@/contexts/AuthContext";
import { isAdminRole } from "@/utils/roles";
const DISMISS_STORAGE_KEY = "litellm:envCredentialLoginWarningDismissed";
export const EnvCredentialLoginWarningBanner: React.FC<{ accessToken: string | null }> = ({ accessToken }) => {
const { userRole } = useAuth();
const { data: healthData } = useHealthReadinessDetails(accessToken);
const [dismissed, setDismissed] = useState(
() => typeof window !== "undefined" && localStorage.getItem(DISMISS_STORAGE_KEY) === "true",
);
if (!isAdminRole(userRole) || !healthData?.show_env_credential_login_warning) {
if (dismissed || !isAdminRole(userRole) || !healthData?.show_env_credential_login_warning) {
return null;
}
const handleDismiss = () => {
localStorage.setItem(DISMISS_STORAGE_KEY, "true");
setDismissed(true);
};
return (
<div
role="alert"
className="flex items-start gap-3 border-b border-destructive/40 bg-destructive/10 px-4 py-3 text-sm text-destructive"
>
<TriangleAlert className="mt-0.5 size-5 shrink-0" aria-hidden="true" />
<div>
<div className="min-w-0 flex-1">
<p className="font-semibold">Environment-credential login is enabled</p>
<p>
Anyone with <code className="font-mono">UI_USERNAME</code>/<code className="font-mono">UI_PASSWORD</code> (or
@ -30,6 +41,9 @@ export const EnvCredentialLoginWarningBanner: React.FC<{ accessToken: string | n
off.
</p>
</div>
<Button variant="ghost" size="icon-sm" className="shrink-0" aria-label="Dismiss banner" onClick={handleDismiss}>
<X />
</Button>
</div>
);
};