mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(vertex_ai): derive rerank search_units from input records and use unique response id
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
56d51bc32e
commit
ec1d1efc4b
2 changed files with 102 additions and 5 deletions
|
|
@ -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 typing import Any, Dict, List, Union
|
||||
|
||||
import httpx
|
||||
|
|
@ -31,6 +33,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__()
|
||||
|
||||
|
|
@ -206,10 +210,12 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
|
|||
RerankResponseResult(index=result["index"], relevance_score=result["relevance_score"])
|
||||
)
|
||||
|
||||
# Create meta object
|
||||
meta = RerankResponseMeta(billed_units=RerankBilledUnits(search_units=len(records)))
|
||||
input_record_count = len(request_data.get("records", []))
|
||||
search_units = math.ceil(input_record_count / self.MAX_RECORDS_PER_SEARCH_UNIT)
|
||||
|
||||
return RerankResponse(id=f"vertex_ai_rerank_{model}", results=rerank_results, meta=meta)
|
||||
meta = RerankResponseMeta(billed_units=RerankBilledUnits(search_units=search_units))
|
||||
|
||||
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 [
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue