diff --git a/litellm/llms/base_llm/rerank/transformation.py b/litellm/llms/base_llm/rerank/transformation.py index 166f876ba04..6603c64142b 100644 --- a/litellm/llms/base_llm/rerank/transformation.py +++ b/litellm/llms/base_llm/rerank/transformation.py @@ -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 diff --git a/litellm/llms/cohere/rerank/guardrail_translation/handler.py b/litellm/llms/cohere/rerank/guardrail_translation/handler.py index e9a5823d2b8..0824e1cca41 100644 --- a/litellm/llms/cohere/rerank/guardrail_translation/handler.py +++ b/litellm/llms/cohere/rerank/guardrail_translation/handler.py @@ -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 diff --git a/litellm/llms/cohere/rerank/transformation.py b/litellm/llms/cohere/rerank/transformation.py index 64ae8e8ffa7..d875f420310 100644 --- a/litellm/llms/cohere/rerank/transformation.py +++ b/litellm/llms/cohere/rerank/transformation.py @@ -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 diff --git a/litellm/llms/cohere/rerank_v2/transformation.py b/litellm/llms/cohere/rerank_v2/transformation.py index 4c800d6455d..0dcb10d5664 100644 --- a/litellm/llms/cohere/rerank_v2/transformation.py +++ b/litellm/llms/cohere/rerank_v2/transformation.py @@ -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 diff --git a/litellm/llms/dashscope/rerank/transformation.py b/litellm/llms/dashscope/rerank/transformation.py index 629f3cf4af7..745e85de7e3 100644 --- a/litellm/llms/dashscope/rerank/transformation.py +++ b/litellm/llms/dashscope/rerank/transformation.py @@ -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. diff --git a/litellm/llms/deepinfra/rerank/transformation.py b/litellm/llms/deepinfra/rerank/transformation.py index e4bfbcb2513..a5c36ca2e5f 100644 --- a/litellm/llms/deepinfra/rerank/transformation.py +++ b/litellm/llms/deepinfra/rerank/transformation.py @@ -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 = {} diff --git a/litellm/llms/fireworks_ai/rerank/transformation.py b/litellm/llms/fireworks_ai/rerank/transformation.py index 4a7b64b9b77..27309780c86 100644 --- a/litellm/llms/fireworks_ai/rerank/transformation.py +++ b/litellm/llms/fireworks_ai/rerank/transformation.py @@ -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 diff --git a/litellm/llms/hosted_vllm/rerank/transformation.py b/litellm/llms/hosted_vllm/rerank/transformation.py index 60b6dc7d23d..d0c96f8b420 100644 --- a/litellm/llms/hosted_vllm/rerank/transformation.py +++ b/litellm/llms/hosted_vllm/rerank/transformation.py @@ -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) diff --git a/litellm/llms/huggingface/rerank/transformation.py b/litellm/llms/huggingface/rerank/transformation.py index 2c847b617ef..4e409f31ed2 100644 --- a/litellm/llms/huggingface/rerank/transformation.py +++ b/litellm/llms/huggingface/rerank/transformation.py @@ -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: diff --git a/litellm/llms/jina_ai/rerank/transformation.py b/litellm/llms/jina_ai/rerank/transformation.py index 56be754fc34..0d48ed5edcd 100644 --- a/litellm/llms/jina_ai/rerank/transformation.py +++ b/litellm/llms/jina_ai/rerank/transformation.py @@ -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) diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index fc317293acc..8eee188bf46 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -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. diff --git a/litellm/llms/vertex_ai/rerank/transformation.py b/litellm/llms/vertex_ai/rerank/transformation.py index 3b84972e946..d2041009efb 100644 --- a/litellm/llms/vertex_ai/rerank/transformation.py +++ b/litellm/llms/vertex_ai/rerank/transformation.py @@ -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 diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index d64450a1211..907e5b7e26b 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -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} diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index 202760f68a6..a34358a6be3 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -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 diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index e40e12e9197..3ef74d596ad 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -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}") diff --git a/litellm/rerank_api/rerank_utils.py b/litellm/rerank_api/rerank_utils.py index 38e599ef824..a8a665496fc 100644 --- a/litellm/rerank_api/rerank_utils.py +++ b/litellm/rerank_api/rerank_utils.py @@ -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, ) diff --git a/litellm/types/rerank.py b/litellm/types/rerank.py index d2c252a1e92..376d6f66603 100644 --- a/litellm/types/rerank.py +++ b/litellm/types/rerank.py @@ -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): diff --git a/tests/test_litellm/llms/cohere/rerank/test_rerank_guardrail_handler.py b/tests/test_litellm/llms/cohere/rerank/test_rerank_guardrail_handler.py index 88072cd7760..46c37e6af6c 100644 --- a/tests/test_litellm/llms/cohere/rerank/test_rerank_guardrail_handler.py +++ b/tests/test_litellm/llms/cohere/rerank/test_rerank_guardrail_handler.py @@ -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""" diff --git a/tests/test_litellm/llms/hosted_vllm/test_hosted_vllm_rerank_transformation.py b/tests/test_litellm/llms/hosted_vllm/test_hosted_vllm_rerank_transformation.py index 9e6fa608c50..6425e815db0 100644 --- a/tests/test_litellm/llms/hosted_vllm/test_hosted_vllm_rerank_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/test_hosted_vllm_rerank_transformation.py @@ -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