mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: scan rerank instruction through request guardrails
The rerank guardrail translation (CohereRerankHandler.process_input_messages) only scanned `query`, so the newly added `instruction` field reached the backend model unscanned. Since instruction-aware rerankers (hosted vLLM / Qwen3-Reranker) fold `instruction` into the prompt, an authenticated caller could place content there to bypass configured rerank request guardrails. Generalize the handler to scan every user-controlled text field (`query` and `instruction`) in one apply_guardrail call and write each sanitized value back by index. Query-only requests are unchanged (single-element list at index 0); non-string fields are left untouched. Adds tests covering instruction scanning, PII masking write-back, and the non-string case. Addresses the Veria AI security review on PR #30757.
This commit is contained in:
parent
8f2108d222
commit
8f90e9ea44
2 changed files with 116 additions and 30 deletions
|
|
@ -26,11 +26,18 @@ class CohereRerankHandler(BaseTranslation):
|
|||
|
||||
The handler specifically processes:
|
||||
- The 'query' parameter (string)
|
||||
- The 'instruction' parameter (string), when present
|
||||
|
||||
Note: Documents are not processed by guardrails as they are the corpus
|
||||
being searched, not user input.
|
||||
"""
|
||||
|
||||
# User-controlled free-text fields that reach the model and must be
|
||||
# scanned. 'instruction' is folded into the prompt by instruction-aware
|
||||
# rerankers (e.g. hosted vLLM / Qwen3-Reranker), so it is as sensitive as
|
||||
# 'query'; omitting it would let a caller smuggle content past guardrails.
|
||||
_SCANNED_FIELDS = ("query", "instruction")
|
||||
|
||||
async def process_input_messages(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -38,42 +45,55 @@ class CohereRerankHandler(BaseTranslation):
|
|||
litellm_logging_obj: Optional[Any] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
Process input query by applying guardrails.
|
||||
Process input text fields ('query' and 'instruction') by applying
|
||||
guardrails and writing the sanitized values back.
|
||||
|
||||
Args:
|
||||
data: Request data dictionary containing 'query'
|
||||
data: Request data dictionary containing 'query' and optionally
|
||||
'instruction'
|
||||
guardrail_to_apply: The guardrail instance to apply
|
||||
|
||||
Returns:
|
||||
Modified data with guardrails applied to query only
|
||||
Modified data with guardrails applied to query/instruction only
|
||||
"""
|
||||
# Process query only
|
||||
query = data.get("query")
|
||||
if query is not None and isinstance(query, str):
|
||||
inputs = GenericGuardrailAPIInputs(texts=[query])
|
||||
# Include model information if available
|
||||
model = data.get("model")
|
||||
if model:
|
||||
inputs["model"] = model
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
# Collect every scannable text field in a stable order so the
|
||||
# guardrailed results can be written back to the right key by index.
|
||||
fields_to_scan = [
|
||||
(key, data[key])
|
||||
for key in self._SCANNED_FIELDS
|
||||
if isinstance(data.get(key), str)
|
||||
]
|
||||
if not fields_to_scan:
|
||||
verbose_proxy_logger.debug(
|
||||
"Rerank: No query/instruction to process or not strings"
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["query"] = guardrailed_texts[0] if guardrailed_texts else query
|
||||
return data
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Rerank: Applied guardrail to query. "
|
||||
"Original length: %d, New length: %d",
|
||||
len(query),
|
||||
len(data["query"]),
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Rerank: No query to process or query is not a string"
|
||||
)
|
||||
inputs = GenericGuardrailAPIInputs(texts=[value for _, value in fields_to_scan])
|
||||
# Include model information if available
|
||||
model = data.get("model")
|
||||
if model:
|
||||
inputs["model"] = model
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
for idx, (key, original) in enumerate(fields_to_scan):
|
||||
# Defensive: only write back when the guardrail returned a value for
|
||||
# this index; otherwise keep the original (never forward unscanned).
|
||||
if idx < len(guardrailed_texts):
|
||||
data[key] = guardrailed_texts[idx]
|
||||
verbose_proxy_logger.debug(
|
||||
"Rerank: Applied guardrail to %s. "
|
||||
"Original length: %d, New length: %d",
|
||||
key,
|
||||
len(original),
|
||||
len(data[key]),
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -2,10 +2,8 @@
|
|||
Unit tests for Cohere Rerank Guardrail Translation Handler
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -94,6 +92,74 @@ class TestInputProcessing:
|
|||
"id": "doc2",
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_query_and_instruction(self):
|
||||
"""Both query and instruction are guardrailed; documents untouched"""
|
||||
handler = CohereRerankHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="test")
|
||||
|
||||
data = {
|
||||
"model": "qwen3-reranker",
|
||||
"query": "What is machine learning?",
|
||||
"instruction": "Rank by relevance to ML research",
|
||||
"documents": ["Doc 1", "Doc 2"],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
# Both user-controlled text fields are scanned and written back
|
||||
assert result["query"] == "What is machine learning? [GUARDRAILED]"
|
||||
assert result["instruction"] == "Rank by relevance to ML research [GUARDRAILED]"
|
||||
# Documents unchanged
|
||||
assert result["documents"] == ["Doc 1", "Doc 2"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_instruction_masked_with_pii(self):
|
||||
"""A masking guardrail rewrites instruction, not just query"""
|
||||
|
||||
class PIIMaskingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self, inputs: dict, request_data: dict, input_type: str, **kwargs
|
||||
) -> dict:
|
||||
texts = inputs.get("texts", [])
|
||||
return {"texts": [t.replace("John Doe", "[NAME_REDACTED]") for t in texts]}
|
||||
|
||||
handler = CohereRerankHandler()
|
||||
guardrail = PIIMaskingGuardrail(guardrail_name="mask_pii")
|
||||
|
||||
data = {
|
||||
"model": "qwen3-reranker",
|
||||
"query": "find records",
|
||||
"instruction": "prioritize anything authored by John Doe",
|
||||
"documents": ["Doc 1"],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
# The sensitive value in instruction is sanitized before forwarding
|
||||
assert "John Doe" not in result["instruction"]
|
||||
assert "[NAME_REDACTED]" in result["instruction"]
|
||||
assert result["documents"] == ["Doc 1"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_string_instruction_not_scanned(self):
|
||||
"""A non-string instruction is left as-is (only strings are scanned)"""
|
||||
handler = CohereRerankHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="test")
|
||||
|
||||
data = {
|
||||
"model": "qwen3-reranker",
|
||||
"query": "hello",
|
||||
"instruction": 12345, # invalid type; backend will reject it
|
||||
"documents": ["Doc 1"],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
# Query still guardrailed; non-string instruction untouched
|
||||
assert result["query"] == "hello [GUARDRAILED]"
|
||||
assert result["instruction"] == 12345
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_no_query(self):
|
||||
"""Test processing when query is missing"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue