This commit is contained in:
Soumyajit Ghosh 2026-09-28 19:27:07 -04:00 • committed by GitHub
commit 2414cc245f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 177 additions and 20 deletions

View file

@ -22,6 +22,7 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.rerank import (
RerankBilledUnits,
RerankResponse,
RerankResponseDocument,
RerankResponseMeta,
RerankResponseResult,
)
@ -176,6 +177,17 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
except Exception as e:
raise ValueError(f"Failed to parse response: {e}")
# Determine whether to return documents (defaults to True)
return_documents = True
if "return_documents" in optional_params and optional_params["return_documents"] is not None:
return_documents = bool(optional_params["return_documents"])
elif "return_documents" in request_data and request_data["return_documents"] is not None:
return_documents = bool(request_data["return_documents"])
elif "ignoreRecordDetailsInResponse" in request_data:
return_documents = not bool(request_data["ignoreRecordDetailsInResponse"])
elif "return_documents" in litellm_params and litellm_params["return_documents"] is not None:
return_documents = bool(litellm_params["return_documents"])
# Extract records from response
records: Final = raw_response_json.get("records", [])
@ -183,23 +195,16 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
results: Final = []
for record in records:
# Handle both cases: with full details and with only IDs
if "score" in record:
# Full response with score and details
results.append(
{
"index": int(record["id"]),
"relevance_score": record.get("score", 0.0),
}
)
else:
# Response with only IDs (when ignoreRecordDetailsInResponse=true)
# We can't provide a relevance score, so we'll use a default
results.append(
{
"index": int(record["id"]),
"relevance_score": 1.0, # Default score when details are ignored
}
)
score_val = record.get("score", 0.0) if "score" in record else 1.0
doc_text = record.get("content")
result_item = {
"index": int(record["id"]),
"relevance_score": score_val,
}
if return_documents and doc_text is not None:
result_item["document"] = RerankResponseDocument(text=doc_text)
results.append(result_item)
# Sort by relevance score (descending)
results.sort(key=lambda x: x["relevance_score"], reverse=True)
@ -208,9 +213,10 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
# Convert results to proper RerankResponseResult objects
rerank_results: Final = []
for result in results:
rerank_results.append(
RerankResponseResult(index=result["index"], relevance_score=result["relevance_score"])
)
rerank_result = RerankResponseResult(index=result["index"], relevance_score=result["relevance_score"])
if "document" in result:
rerank_result["document"] = result["document"]
rerank_results.append(rerank_result)
input_record_count: Final = len(request_data.get("records", ()))
search_units: Final = math.ceil(input_record_count / self.MAX_RECORDS_PER_SEARCH_UNIT)

View file

@ -115,8 +115,16 @@ class TestVertexAIRerankIntegration:
# Results should be sorted by relevance score (descending)
assert result.results[0]["index"] == 3 # Highest score
assert result.results[0]["relevance_score"] == 0.95
assert (
result.results[0]["document"]["text"]
== "Google's Gemini AI model represents a significant advancement in artificial intelligence technology."
)
assert result.results[1]["index"] == 0 # Second highest score
assert result.results[1]["relevance_score"] == 0.92
assert (
result.results[1]["document"]["text"]
== "Gemini is a cutting edge large language model created by Google."
)
# Verify metadata: 4 input records bill as 1 search unit (ceil(4/100))
assert result.meta["billed_units"]["search_units"] == 1
@ -159,6 +167,7 @@ class TestVertexAIRerankIntegration:
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
request_data=request_data,
)
# Verify response structure with default scores
@ -168,6 +177,7 @@ class TestVertexAIRerankIntegration:
result_item["relevance_score"] == 1.0
) # Default score when details are ignored
assert "index" in result_item
assert "document" not in result_item
def test_document_title_generation(self):
"""Test that document titles are generated correctly from content."""

View file

@ -295,12 +295,153 @@ class TestVertexAIRerankTransform:
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
assert (
result.results[0]["document"]["text"]
== "The sky appears blue due to a phenomenon called Rayleigh scattering."
)
assert result.results[1]["index"] == 0
assert result.results[1]["relevance_score"] == 0.64
assert (
result.results[1]["document"]["text"]
== "A canvas stretched across the day, Where sunlight learns to dance and play."
)
# Verify metadata
assert result.meta["billed_units"]["search_units"] == 1
def test_transform_rerank_response_return_documents_true_populates_document_text(self):
"""Test that return_documents=True populates document with {'text': record['content']}."""
response_data = {
"records": [
{
"id": "1",
"score": 0.95,
"title": "Doc 1",
"content": "Content of document 1",
},
{
"id": "0",
"score": 0.80,
"title": "Doc 0",
"content": "Content of document 0",
},
]
}
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.text = json.dumps(response_data)
mock_logging = MagicMock()
model_response = RerankResponse()
# Test with optional_params={"return_documents": True}
result = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
optional_params={"return_documents": True},
)
assert len(result.results) == 2
assert result.results[0]["index"] == 1
assert result.results[0]["relevance_score"] == 0.95
assert result.results[0]["document"] == {"text": "Content of document 1"}
assert result.results[0]["document"]["text"] == "Content of document 1"
assert result.results[1]["index"] == 0
assert result.results[1]["relevance_score"] == 0.80
assert result.results[1]["document"] == {"text": "Content of document 0"}
assert result.results[1]["document"]["text"] == "Content of document 0"
def test_transform_rerank_response_return_documents_false_omits_document_text(self):
"""Test that return_documents=False does not populate document field."""
response_data = {
"records": [
{
"id": "1",
"score": 0.95,
"title": "Doc 1",
"content": "Content of document 1",
},
{
"id": "0",
"score": 0.80,
"title": "Doc 0",
"content": "Content of document 0",
},
]
}
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = response_data
mock_response.text = json.dumps(response_data)
mock_logging = MagicMock()
model_response = RerankResponse()
# Test with optional_params={"return_documents": False}
result = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
optional_params={"return_documents": False},
)
assert len(result.results) == 2
assert result.results[0]["index"] == 1
assert result.results[0]["relevance_score"] == 0.95
assert "document" not in result.results[0]
assert result.results[1]["index"] == 0
assert result.results[1]["relevance_score"] == 0.80
assert "document" not in result.results[1]
# Test with request_data={"ignoreRecordDetailsInResponse": True}
result_request_data = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
request_data={"ignoreRecordDetailsInResponse": True},
)
assert "document" not in result_request_data.results[0]
assert "document" not in result_request_data.results[1]
# Test with litellm_params={"return_documents": True}
result_litellm_params = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
litellm_params={"return_documents": True},
)
assert result_litellm_params.results[0]["document"]["text"] == "Content of document 1"
# Test with request_data={"return_documents": True}
result_req_data_true = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
model_response=model_response,
logging_obj=mock_logging,
request_data={"return_documents": True},
)
assert result_req_data_true.results[0]["document"]["text"] == "Content of document 1"
# Test with records missing content
no_content_response_data = {"records": [{"id": "0", "score": 0.9}]}
mock_no_content = MagicMock(spec=httpx.Response)
mock_no_content.json.return_value = no_content_response_data
mock_no_content.text = json.dumps(no_content_response_data)
result_no_content = self.config.transform_rerank_response(
model=self.model,
raw_response=mock_no_content,
model_response=model_response,
logging_obj=mock_logging,
optional_params={"return_documents": True},
)
assert "document" not in result_no_content.results[0]
def test_transform_rerank_response_with_ignore_record_details(self):
"""Test response transformation when ignoreRecordDetailsInResponse=true."""
# Mock response with only IDs (when ignoreRecordDetailsInResponse=true)