Merge pull request #35180 from BerriAI/litellm_lit4995_vertex_rerank_search_units

fix(rerank): bill Vertex search_units from input records and give every rerank response a unique id
This commit is contained in:
Yassin Kortam 2026-09-15 14:23:23 -07:00 committed by GitHub
commit a6127d2363
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 224 additions and 39 deletions

View file

@ -814,6 +814,7 @@ def _select_model_name_for_cost_calc(
if (
entry.get("input_cost_per_token") is not None
or entry.get("input_cost_per_second") is not None
or entry.get("input_cost_per_query") is not None
or entry.get("tiered_pricing") is not None
):
return_model = router_model_id

View file

@ -250,8 +250,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
rerank_results.append(rerank_result)
# Use model name as id if no id is provided
response_id: Final = raw_response_json.get("id") or raw_response_json.get("model") or str(uuid.uuid4())
response_id: Final = raw_response_json.get("id") or str(uuid.uuid4())
return RerankResponse(
id=response_id,

View file

@ -4,6 +4,8 @@ Translates from Cohere's `/v1/rerank` input format to Vertex AI Discovery Engine
Why separate file? Make it easy to see how transformation works
"""
import math
import uuid
from collections.abc import Mapping
from typing import Any, Final
@ -32,6 +34,8 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
Reference: https://cloud.google.com/generative-ai-app-builder/docs/ranking#rank_or_rerank_a_set_of_records_according_to_a_query
"""
MAX_RECORDS_PER_SEARCH_UNIT = 100
def __init__(self) -> None:
super().__init__()
@ -208,10 +212,11 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
RerankResponseResult(index=result["index"], relevance_score=result["relevance_score"])
)
# Create meta object
meta: Final = RerankResponseMeta(billed_units=RerankBilledUnits(search_units=len(records)))
input_record_count: Final = len(request_data.get("records", ()))
search_units: Final = math.ceil(input_record_count / self.MAX_RECORDS_PER_SEARCH_UNIT)
meta: Final = RerankResponseMeta(billed_units=RerankBilledUnits(search_units=search_units))
return RerankResponse(id=f"vertex_ai_rerank_{model}", results=rerank_results, meta=meta)
return RerankResponse(id=f"vertex_ai_rerank_{uuid.uuid4()}", results=rerank_results, meta=meta)
def get_supported_cohere_rerank_params(self, model: str) -> list:
return [

View file

@ -9,6 +9,7 @@ from typing import Any, Final
import httpx
from litellm._uuid import uuid
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.secret_managers.main import get_secret_str
@ -127,7 +128,7 @@ class VoyageRerankConfig(BaseRerankConfig):
rerank_meta: Final = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens)
return RerankResponse(
id=_json_response.get("id", f"voyage-rerank-{model}"),
id=_json_response.get("id") or str(uuid.uuid4()),
results=transformed_results,
meta=rerank_meta,
)

View file

@ -191,7 +191,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
transformed_results.append(transformed_result)
response_id: Final = raw_response_json.get("id") or raw_response_json.get("model_id") or str(uuid.uuid4())
response_id: Final = raw_response_json.get("id") or str(uuid.uuid4())
# Extract usage information
_tokens: Final = RerankTokens(

View file

@ -3,6 +3,7 @@ Tests for Fireworks AI rerank transformation functionality.
"""
import json
import uuid
from unittest.mock import MagicMock
import httpx
@ -181,8 +182,7 @@ class TestFireworksAIRerankTransform:
)
# Verify response structure
# Fireworks AI doesn't return "id", so it uses "model" as the id
assert result.id == "accounts/fireworks/models/qwen3-reranker-8b"
assert uuid.UUID(result.id).version == 4
assert len(result.results) == 2
assert result.results[0]["index"] == 0
assert result.results[0]["relevance_score"] == 0.95
@ -229,16 +229,14 @@ class TestFireworksAIRerankTransform:
logging_obj=mock_logging,
)
# Fireworks AI doesn't return "id", so it uses "model" as the id
assert result.id == "accounts/fireworks/models/qwen3-reranker-8b"
assert uuid.UUID(result.id).version == 4
assert len(result.results) == 2
assert result.results[0]["index"] == 0
assert result.results[0]["relevance_score"] == 0.95
# Document should not be present
assert "document" not in result.results[0]
def test_transform_rerank_response_missing_id(self):
"""Test response transformation when id is missing (should use model name or generate UUID)."""
def test_transform_rerank_response_missing_id_stamps_a_fresh_id_per_call(self):
response_data = {
"object": "list",
"model": "accounts/fireworks/models/qwen3-reranker-8b",
@ -248,23 +246,22 @@ class TestFireworksAIRerankTransform:
"usage": {"total_tokens": 10},
}
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.status_code = 200
mock_response.headers = {}
def transform() -> str:
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.status_code = 200
mock_response.headers = {}
return self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
).id
mock_logging = MagicMock()
model_response = RerankResponse()
first, second = transform(), transform()
result = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
)
# Should use model name when id is missing
assert result.id == "accounts/fireworks/models/qwen3-reranker-8b"
assert first != second
assert "accounts/fireworks/models/qwen3-reranker-8b" not in (first, second)
def test_transform_rerank_response_missing_results(self):
"""Test that missing results raises ValueError."""

View file

@ -104,10 +104,12 @@ class TestVertexAIRerankIntegration:
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
request_data=request_data,
)
# Verify response structure
assert result.id == f"vertex_ai_rerank_{self.model}"
assert result.id.startswith("vertex_ai_rerank_")
assert result.id != f"vertex_ai_rerank_{self.model}"
assert len(result.results) == 2
# Results should be sorted by relevance score (descending)
@ -116,8 +118,8 @@ class TestVertexAIRerankIntegration:
assert result.results[1]["index"] == 0 # Second highest score
assert result.results[1]["relevance_score"] == 0.92
# Verify metadata
assert result.meta["billed_units"]["search_units"] == 2
# Verify metadata: 4 input records bill as 1 search unit (ceil(4/100))
assert result.meta["billed_units"]["search_units"] == 1
def test_return_documents_false_flow(self):
"""Test rerank flow when return_documents=False (ID-only response)."""

View file

@ -287,10 +287,11 @@ class TestVertexAIRerankTransform:
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
request_data={"records": [{"id": "0"}, {"id": "1"}]},
)
# Verify response structure
assert result.id == f"vertex_ai_rerank_{self.model}"
assert result.id.startswith("vertex_ai_rerank_")
assert len(result.results) == 2
assert result.results[0]["index"] == 1 # Converted back to 0-based index
assert result.results[0]["relevance_score"] == 0.98
@ -298,7 +299,7 @@ class TestVertexAIRerankTransform:
assert result.results[1]["relevance_score"] == 0.64
# Verify metadata
assert result.meta["billed_units"]["search_units"] == 2
assert result.meta["billed_units"]["search_units"] == 1
def test_transform_rerank_response_with_ignore_record_details(self):
"""Test response transformation when ignoreRecordDetailsInResponse=true."""
@ -326,6 +327,96 @@ class TestVertexAIRerankTransform:
assert result.results[1]["index"] == 0
assert result.results[1]["relevance_score"] == 1.0
def _build_response(self, num_records):
response_data = {
"records": [
{"id": str(i), "score": 1.0 - i / 1000, "title": "t", "content": "c"}
for i in range(num_records)
]
}
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.text = json.dumps(response_data)
return mock_response
def test_search_units_from_input_records_not_truncated_response(self):
"""
Regression for LIT-4995 part 1: search_units must be derived from the
billable input records (ceil(input / 100)), not from the response, which
Google truncates to topN.
"""
documents = [f"doc {i}" for i in range(5)]
request_data = self.config.transform_rerank_request(
model=self.model,
optional_rerank_params={"query": "q", "documents": documents, "top_n": 2},
headers={},
)
# Google truncates the response to top_n=2 records
mock_response = self._build_response(num_records=2)
result = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
request_data=request_data,
)
assert result.meta["billed_units"]["search_units"] == 1
def test_search_units_rounds_up_per_hundred_input_records(self):
"""
Regression for LIT-4995 part 1: one query bills up to 100 input records,
so 150 input records is 2 search units regardless of the response size.
"""
documents = [f"doc {i}" for i in range(150)]
request_data = self.config.transform_rerank_request(
model=self.model,
optional_rerank_params={"query": "q", "documents": documents, "top_n": 3},
headers={},
)
mock_response = self._build_response(num_records=3)
result = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
request_data=request_data,
)
assert result.meta["billed_units"]["search_units"] == 2
def test_response_id_is_unique_per_request(self):
"""
Regression for LIT-4995 part 2: response IDs must be unique per request,
not a constant derived only from the model name.
"""
request_data = self.config.transform_rerank_request(
model=self.model,
optional_rerank_params={"query": "q", "documents": ["a", "b"]},
headers={},
)
mock_response = self._build_response(num_records=2)
first = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
request_data=request_data,
)
second = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
request_data=request_data,
)
assert first.id != second.id
assert first.id != f"vertex_ai_rerank_{self.model}"
def test_transform_rerank_response_json_error(self):
"""Test response transformation with JSON parsing error."""
mock_response = MagicMock(spec=httpx.Response)

View file

@ -3,6 +3,7 @@ Tests for Voyage AI rerank transformation functionality.
"""
import json
import uuid
from unittest.mock import MagicMock, patch
import httpx
@ -258,6 +259,33 @@ class TestVoyageRerankTransform:
assert "Failed to parse response" in str(exc_info.value)
def test_transform_rerank_response_without_id_stamps_a_fresh_id_per_call(self):
response_data = {
"object": "list",
"data": [{"relevance_score": 0.5, "index": 0}],
"model": "rerank-2.5",
"usage": {"total_tokens": 10},
}
def transform() -> str:
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.status_code = 200
mock_response.text = json.dumps(response_data)
mock_response.headers = {}
return self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
).id
first, second = transform(), transform()
assert uuid.UUID(first).version == 4
assert first != second
assert f"voyage-rerank-{self.model}" not in (first, second)
def test_get_supported_cohere_rerank_params(self):
"""Test getting supported parameters for Voyage AI rerank."""
supported_params = self.config.get_supported_cohere_rerank_params(self.model)

View file

@ -120,9 +120,7 @@ class TestIBMWatsonXRerankTransform:
logging_obj=mock_logging,
)
# Verify response structure
# IBM watsonx.ai doesn't return "id", so it uses "model" as the id
assert result.id == "watsonx/cross-encoder/ms-marco-minilm-l-12-v2"
assert uuid.UUID(result.id).version == 4
assert len(result.results) == 2
assert result.results[0]["index"] == 0
assert result.results[0]["relevance_score"] == 6.53515625
@ -172,9 +170,7 @@ class TestIBMWatsonXRerankTransform:
logging_obj=mock_logging,
)
# Verify response structure
# IBM watsonx.ai doesn't return "id", so it uses "model" as the id
assert result.id == "watsonx/cross-encoder/ms-marco-minilm-l-12-v2"
assert uuid.UUID(result.id).version == 4
assert len(result.results) == 2
assert result.results[0]["index"] == 0
@ -231,6 +227,30 @@ class TestIBMWatsonXRerankTransform:
logging_obj=mock_logging,
)
def test_transform_rerank_response_without_id_stamps_a_fresh_id_per_call(self):
response_data = {
"model_id": self.model,
"results": [{"index": 0, "score": 1.5}],
"input_token_count": 12,
}
def transform() -> str:
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.status_code = 200
mock_response.headers = {}
return self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=RerankResponse(),
logging_obj=MagicMock(),
).id
first, second = transform(), transform()
assert first != second
assert self.model not in (first, second)
def test_get_supported_cohere_rerank_params(self):
"""Test getting supported parameters for IBM watsonx.ai rerank."""
supported_params = self.config.get_supported_cohere_rerank_params(self.model)

View file

@ -1211,6 +1211,47 @@ def test_tiered_pricing_only_deployment_completion_cost_is_nonzero():
assert cost > 0
def test_per_query_priced_rerank_deployment_completion_cost_is_nonzero():
"""A rerank deployment priced only via ``input_cost_per_query`` must resolve
cost against its ``router_model_id`` entry: the shared backend alias has
custom pricing stripped, so pricing it there bills every search unit as $0.
"""
from litellm import Router
router: Final = Router(
model_list=[
{
"model_name": "semantic-ranker-default-004",
"litellm_params": {
"model": "vertex_ai/semantic-ranker-default-004",
"vertex_project": "test-project",
"vertex_location": "us-east5",
},
"model_info": {"input_cost_per_query": 0.001},
},
]
)
router_model_id: Final = router.model_list[0]["model_info"]["id"]
assert litellm.model_cost["vertex_ai/semantic-ranker-default-004"].get("input_cost_per_query") is None
response: Final = RerankResponse(
id="vertex_ai_rerank_test",
results=[{"index": 3, "relevance_score": 0.48}],
meta={"billed_units": {"search_units": 3}},
)
cost: Final = completion_cost(
completion_response=response,
model="vertex_ai/semantic-ranker-default-004",
custom_llm_provider="vertex_ai",
call_type="arerank",
custom_pricing=True,
router_model_id=router_model_id,
)
assert cost == pytest.approx(3 * 0.001)
def test_azure_realtime_cost_calculator(_local_model_cost_map):
cost = handle_realtime_stream_cost_calculation(