mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
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:
commit
a6127d2363
11 changed files with 224 additions and 39 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 [
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue