feat: pass through optional instruction field in the rerank API (vLLM/Qwen3-Reranker) (#30757)

* Add optional `instruction` passthrough to the rerank API

vLLM's /v1/rerank and /v1/score accept an optional top-level `instruction`
field (folded into the model's chat_template_kwargs and consumed by the
chat template — e.g. Qwen3-Reranker). LiteLLM's managed rerank route silently
dropped it: RerankRequest / OptionalRerankParams had no such field, so the
outgoing body was rebuilt without it.

Thread an opt-in `instruction: Optional[str]` through rerank()/arerank(),
get_optional_rerank_params, and the hosted_vllm transformation into the
request body, only when non-None. When callers omit it, model_dump(exclude_none)
drops the field and the outgoing request is byte-for-byte unchanged — fully
backward-compatible. (DeepInfra already forwards `instruction` via
non_default_params; this formalizes the field in the shared types.)

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>

* Address review: thread `instruction` as a typed param + cover rerank_utils

Per PR review (greptile P2 + codecov):

- Make `instruction` a typed, named argument on the rerank provider interface
  instead of recovering it from the opaque `non_default_params` blob. Adds
  `instruction: Optional[str] = None` to `BaseRerankConfig.map_cohere_rerank_params`
  and every provider override, and forwards it explicitly from
  `get_optional_rerank_params`. hosted_vllm now reads the named param directly.
  It is still also surfaced in `non_default_params` so providers that read it
  there (e.g. DeepInfra) keep working now that `rerank()` consumes `instruction`
  as a named param rather than leaving it in **kwargs.
- Add get_optional_rerank_params unit tests (present + absent) to cover the
  previously-uncovered threading line flagged by codecov.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>

* 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.

* test: narrow Optional results before len() to satisfy basedpyright budget

The lint gate (basedpyright delta-vs-base budget) flagged one new
reportArgumentType: len(result.results) where results is
List[RerankResponseResult] | None. Assert results is not None first to
narrow the type before len()/indexing.

* fix: read rerank `instruction` from kwargs to satisfy basedpyright budget

The basedpyright delta-vs-base gate flagged one new reportArgumentType: the
Router forwards rerank calls via an untyped `**kwargs` unpack
(`litellm.arerank(**{**data, **kwargs})`), and declaring `instruction` as a
typed named param on the public `rerank`/`arerank` entrypoints made pyright
check that key against `str | None`, adding an error at router.py with no real
safety gain. Read `instruction` from kwargs in `rerank` instead.

It remains fully typed where it matters - threaded as a typed argument through
`get_optional_rerank_params` and each provider's `map_cohere_rerank_params`
(the original Greptile P2 ask). Whole-repo reportArgumentType is back to the
base count (net 0); rerank hosted_vllm + cohere guardrail suites pass; ruff clean.

---------

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Jim Smith 2026-06-24 07:15:38 -04:00 • committed by GitHub
parent 1eb6bdde9c
commit e1187c0462
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 242 additions and 38 deletions

View file

@ -85,6 +85,7 @@ class BaseRerankConfig(ABC):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
pass

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

@ -57,6 +57,7 @@ class CohereRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
"""
Map Cohere rerank params

View file

@ -49,6 +49,7 @@ class CohereRerankV2Config(CohereRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
"""
Map Cohere rerank params

View file

@ -116,6 +116,7 @@ class DashScopeRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
# qwen3-rerank accepts query/documents/top_n/return_documents. The
# rest (rank_fields, max_*_per_doc) are silently dropped.

View file

@ -104,6 +104,7 @@ class DeepinfraRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
# Start with the basic parameters
optional_rerank_params = {}

View file

@ -67,6 +67,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict[str, Any]:
"""
Map Cohere rerank params to Fireworks AI rerank params

View file

@ -61,6 +61,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
"top_n",
"rank_fields",
"return_documents",
"instruction",
]
def map_cohere_rerank_params(
@ -76,6 +77,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
"""
Map parameters for Hosted VLLM rerank
@ -83,16 +85,22 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
if max_chunks_per_doc is not None:
raise ValueError("Hosted VLLM does not support max_chunks_per_doc")
return dict(
OptionalRerankParams(
query=query,
documents=documents,
top_n=top_n,
rank_fields=rank_fields,
return_documents=return_documents,
)
mapped_params = OptionalRerankParams(
query=query,
documents=documents,
top_n=top_n,
rank_fields=rank_fields,
return_documents=return_documents,
)
# `instruction` is a vLLM-supported passthrough (folded into the model's
# chat_template_kwargs). Only forward it when explicitly set so omitting
# it leaves the request unchanged.
if instruction is not None:
mapped_params["instruction"] = instruction
return dict(mapped_params)
def validate_environment(
self,
headers: dict,
@ -135,6 +143,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
top_n=optional_rerank_params.get("top_n", None),
rank_fields=optional_rerank_params.get("rank_fields", None),
return_documents=optional_rerank_params.get("return_documents", None),
instruction=optional_rerank_params.get("instruction", None),
)
return rerank_request.model_dump(exclude_none=True)

View file

@ -100,6 +100,7 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
optional_rerank_params = {}
if non_default_params is not None:

View file

@ -45,6 +45,7 @@ class JinaAIRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
optional_params = {}
supported_params = self.get_supported_cohere_rerank_params(model)

View file

@ -117,6 +117,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
"""
Map Cohere/OpenAI rerank params to Nvidia NIM format.

View file

@ -242,6 +242,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
"""
Map Cohere rerank params to Vertex AI format

View file

@ -39,6 +39,7 @@ class VoyageRerankConfig(BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
# Voyage AI uses 'top_k' instead of 'top_n'
optional_params: Dict[str, Any] = {"query": query, "documents": documents}

View file

@ -104,6 +104,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
) -> Dict:
"""
Map Cohere rerank params to IBM watsonx.ai rerank params

View file

@ -103,6 +103,11 @@ def rerank(
"""
Reranks a list of documents based on their relevance to the query
"""
# `instruction` is read from kwargs rather than declared as a named param.
# The router forwards rerank calls via an untyped `**kwargs` unpack, and a
# typed named param there would trip the basedpyright budget gate without
# adding real safety; it stays typed downstream via get_optional_rerank_params.
instruction: Optional[str] = kwargs.get("instruction", None)
headers: Optional[dict] = kwargs.get("headers") # type: ignore
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
@ -155,6 +160,7 @@ def rerank(
return_documents=return_documents,
max_chunks_per_doc=max_chunks_per_doc,
max_tokens_per_doc=max_tokens_per_doc,
instruction=instruction,
non_default_params=kwargs,
)
verbose_logger.info(f"optional_rerank_params: {optional_rerank_params}")

View file

@ -15,6 +15,7 @@ def get_optional_rerank_params(
return_documents: Optional[bool] = True,
max_chunks_per_doc: Optional[int] = None,
max_tokens_per_doc: Optional[int] = None,
instruction: Optional[str] = None,
non_default_params: Optional[dict] = None,
) -> Dict:
all_non_default_params = non_default_params or {}
@ -30,6 +31,11 @@ def get_optional_rerank_params(
all_non_default_params["max_chunks_per_doc"] = max_chunks_per_doc
if max_tokens_per_doc is not None:
all_non_default_params["max_tokens_per_doc"] = max_tokens_per_doc
if instruction is not None:
# Also surfaced in non_default_params so providers that read it from
# there (e.g. DeepInfra) keep working now that `rerank()` consumes
# `instruction` as a named param instead of leaving it in **kwargs.
all_non_default_params["instruction"] = instruction
return rerank_provider_config.map_cohere_rerank_params(
model=model,
drop_params=drop_params,
@ -41,5 +47,6 @@ def get_optional_rerank_params(
return_documents=return_documents,
max_chunks_per_doc=max_chunks_per_doc,
max_tokens_per_doc=max_tokens_per_doc,
instruction=instruction,
non_default_params=all_non_default_params,
)

View file

@ -19,6 +19,10 @@ class RerankRequest(BaseModel):
return_documents: Optional[bool] = None
max_chunks_per_doc: Optional[int] = None
max_tokens_per_doc: Optional[int] = None
# Optional task/query instruction passed through to providers that support it
# (e.g. hosted vLLM / Qwen3-Reranker, DeepInfra). Omitted from the outgoing
# request when None, so this is fully backward-compatible.
instruction: Optional[str] = None
class OptionalRerankParams(TypedDict, total=False):
@ -29,6 +33,7 @@ class OptionalRerankParams(TypedDict, total=False):
return_documents: Optional[bool]
max_chunks_per_doc: Optional[int]
max_tokens_per_doc: Optional[int]
instruction: Optional[str]
class RerankBilledUnits(TypedDict, total=False):

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

View file

@ -4,6 +4,7 @@ import sys
import pytest
from litellm.llms.hosted_vllm.rerank.transformation import HostedVLLMRerankConfig
from litellm.rerank_api.rerank_utils import get_optional_rerank_params
from litellm.types.rerank import (
OptionalRerankParams,
RerankBilledUnits,
@ -37,6 +38,54 @@ class TestHostedVLLMRerankTransform:
assert params["rank_fields"] == ["field1"]
assert params["return_documents"] is True
def test_map_cohere_rerank_params_omits_instruction_when_absent(self):
# Backward-compat: when no instruction is supplied, it must not appear
# in the mapped params (and therefore not in the outgoing request body).
params = self.config.map_cohere_rerank_params(
non_default_params=None,
model=self.model,
drop_params=False,
query="test query",
documents=["doc1", "doc2"],
)
assert "instruction" not in params
def test_map_cohere_rerank_params_passes_instruction_when_set(self):
params = self.config.map_cohere_rerank_params(
non_default_params=None,
model=self.model,
drop_params=False,
query="test query",
documents=["doc1", "doc2"],
instruction="Rank by relevance to genomics",
)
assert params["instruction"] == "Rank by relevance to genomics"
def test_transform_request_includes_instruction_when_set(self):
body = self.config.transform_rerank_request(
model=self.model,
optional_rerank_params={
"query": "test query",
"documents": ["doc1", "doc2"],
"instruction": "Rank by relevance to genomics",
},
headers={},
)
assert body["instruction"] == "Rank by relevance to genomics"
def test_transform_request_omits_instruction_when_absent(self):
# exclude_none must drop the field entirely so the body matches the
# pre-existing (instruction-less) shape exactly.
body = self.config.transform_rerank_request(
model=self.model,
optional_rerank_params={
"query": "test query",
"documents": ["doc1", "doc2"],
},
headers={},
)
assert "instruction" not in body
def test_map_cohere_rerank_params_raises_on_max_chunks_per_doc(self):
with pytest.raises(
ValueError, match="Hosted VLLM does not support max_chunks_per_doc"
@ -74,6 +123,7 @@ class TestHostedVLLMRerankTransform:
}
result = self.config._transform_response(response_dict)
assert result.id == "abc123"
assert result.results is not None
assert len(result.results) == 2
assert result.results[0]["index"] == 0
assert result.results[0]["relevance_score"] == 0.9
@ -94,3 +144,32 @@ class TestHostedVLLMRerankTransform:
}
with pytest.raises(ValueError, match="Missing required fields in the result="):
self.config._transform_response(response_dict)
class TestGetOptionalRerankParamsInstruction:
"""`instruction` is threaded through get_optional_rerank_params only when set."""
def setup_method(self):
self.config = HostedVLLMRerankConfig()
self.model = "hosted-vllm-model"
def test_instruction_threaded_when_set(self):
params = get_optional_rerank_params(
rerank_provider_config=self.config,
model=self.model,
drop_params=False,
query="test query",
documents=["doc1", "doc2"],
instruction="Rank by relevance to genomics",
)
assert params["instruction"] == "Rank by relevance to genomics"
def test_instruction_absent_when_not_set(self):
params = get_optional_rerank_params(
rerank_provider_config=self.config,
model=self.model,
drop_params=False,
query="test query",
documents=["doc1", "doc2"],
)
assert "instruction" not in params