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:
Jim Smith 2026-06-24 06:08:33 -04:00
parent 8f2108d222
commit 8f90e9ea44
2 changed files with 116 additions and 30 deletions

View file

@ -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

View file

@ -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"""