mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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>
This commit is contained in:
parent
c546b58c09
commit
77016a623a
5 changed files with 76 additions and 8 deletions
|
|
@ -61,6 +61,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
|
|||
"top_n",
|
||||
"rank_fields",
|
||||
"return_documents",
|
||||
"instruction",
|
||||
]
|
||||
|
||||
def map_cohere_rerank_params(
|
||||
|
|
@ -83,16 +84,23 @@ 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). It arrives via non_default_params; only forward
|
||||
# it when explicitly set so omitting it leaves the request unchanged.
|
||||
instruction = (non_default_params or {}).get("instruction")
|
||||
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)
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ async def arerank(
|
|||
rank_fields: Optional[List[str]] = None,
|
||||
return_documents: Optional[bool] = None,
|
||||
max_chunks_per_doc: Optional[int] = None,
|
||||
instruction: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]:
|
||||
"""
|
||||
|
|
@ -58,6 +59,7 @@ async def arerank(
|
|||
rank_fields,
|
||||
return_documents,
|
||||
max_chunks_per_doc,
|
||||
instruction=instruction,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -98,6 +100,7 @@ def rerank(
|
|||
return_documents: Optional[bool] = True,
|
||||
max_chunks_per_doc: Optional[int] = None,
|
||||
max_tokens_per_doc: Optional[int] = None,
|
||||
instruction: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[RerankResponse, Coroutine[Any, Any, RerankResponse]]:
|
||||
"""
|
||||
|
|
@ -155,6 +158,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}")
|
||||
|
|
|
|||
|
|
@ -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,8 @@ 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:
|
||||
all_non_default_params["instruction"] = instruction
|
||||
return rerank_provider_config.map_cohere_rerank_params(
|
||||
model=model,
|
||||
drop_params=drop_params,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -37,6 +37,53 @@ 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={"instruction": "Rank by relevance to genomics"},
|
||||
model=self.model,
|
||||
drop_params=False,
|
||||
query="test query",
|
||||
documents=["doc1", "doc2"],
|
||||
)
|
||||
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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue