Merge pull request #25499 from BerriAI/litellm_vertex_request_metadata_labels

feat(vertex_ai): propagate metadata labels to embedding, Imagen, rerank
This commit is contained in:
Sameer Kankute 2026-05-01 08:20:55 +05:30 • committed by GitHub
commit 72ddbce50e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 425 additions and 33 deletions

View file

@ -33,6 +33,7 @@ class BaseRerankConfig(ABC):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
return {}

View file

@ -111,6 +111,7 @@ class CohereRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Cohere rerank")

View file

@ -71,6 +71,7 @@ class CohereRerankV2Config(CohereRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Cohere rerank")

View file

@ -1007,6 +1007,7 @@ class BaseLLMHTTPHandler:
api_key: Optional[str] = None,
api_base: Optional[str] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
litellm_params: Optional[Dict[str, Any]] = None,
) -> RerankResponse:
# get config from model, custom llm provider
headers = provider_config.validate_environment(
@ -1026,6 +1027,7 @@ class BaseLLMHTTPHandler:
model=model,
optional_rerank_params=optional_rerank_params,
headers=headers,
litellm_params=litellm_params,
)
## LOGGING

View file

@ -132,6 +132,7 @@ class DeepinfraRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
# Convert OptionalRerankParams to dict as expected by parent class
if optional_rerank_params is None:

View file

@ -127,6 +127,7 @@ class FireworksAIRerankConfig(FireworksAIMixin, BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
"""
Transform request to Fireworks AI rerank format

View file

@ -121,6 +121,7 @@ class HostedVLLMRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for Hosted VLLM rerank")

View file

@ -146,6 +146,7 @@ class HuggingFaceRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Union[OptionalRerankParams, dict],
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
if "query" not in optional_rerank_params:
raise ValueError("query is required for HuggingFace rerank")

View file

@ -74,7 +74,11 @@ class JinaAIRerankConfig(BaseRerankConfig):
return cleaned_base
def transform_rerank_request(
self, model: str, optional_rerank_params: Dict, headers: Dict
self,
model: str,
optional_rerank_params: Dict,
headers: Dict,
litellm_params: Optional[dict] = None,
) -> Dict:
return {"model": model, **optional_rerank_params}

View file

@ -66,6 +66,7 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
"""
Transform request, using clean model name without 'ranking/' prefix.
@ -75,4 +76,5 @@ class NvidiaNimRankingConfig(NvidiaNimRerankConfig):
model=clean_model,
optional_rerank_params=optional_rerank_params,
headers=headers,
litellm_params=litellm_params,
)

View file

@ -177,6 +177,7 @@ class NvidiaNimRerankConfig(BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
"""
Transform request to Nvidia NIM format.

View file

@ -27,6 +27,53 @@ class VertexAIError(BaseLLMException):
super().__init__(message=message, status_code=status_code, headers=headers)
def vertex_request_labels_from_litellm_params(
litellm_params: Optional[dict],
) -> Optional[Dict[str, str]]:
"""
Build Vertex/GCP billing labels from LiteLLM user metadata on ``litellm_params``:
``metadata`` (``completion(..., metadata=...)``) or ``litellm_metadata``,
using ``requester_metadata`` string key-value pairs (same convention as Gemini).
``metadata`` is tried first when both are present.
"""
if not litellm_params:
return None
for key in ("metadata", "litellm_metadata"):
if key not in litellm_params:
continue
metadata = litellm_params[key]
if metadata is None or not isinstance(metadata, dict):
continue
if "requester_metadata" not in metadata:
continue
rm = metadata["requester_metadata"]
if not isinstance(rm, dict):
continue
labels = {k: v for k, v in rm.items() if isinstance(v, str)}
if labels:
return labels
return None
def pop_vertex_request_labels(
optional_params: Optional[dict],
litellm_params: Optional[dict],
) -> Optional[Dict[str, str]]:
"""
Resolve labels from optional ``labels`` (Gemini-style) and/or
``litellm_params["metadata"]`` / ``litellm_params["litellm_metadata"]``
(``requester_metadata``). Pops ``labels`` from optional_params when present.
"""
labels: Optional[Dict[str, str]] = None
if optional_params is not None and "labels" in optional_params:
raw = optional_params.pop("labels")
if isinstance(raw, dict):
labels = {k: v for k, v in raw.items() if isinstance(v, str)}
if not labels:
labels = vertex_request_labels_from_litellm_params(litellm_params)
return labels if labels else None
class VertexAIModelRoute(str, Enum):
"""Enum for Vertex AI model routing"""

View file

@ -24,6 +24,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
response_schema_prompt,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.common_utils import pop_vertex_request_labels
from litellm.types.files import (
get_file_mime_type_for_file_type,
get_file_type_from_extension,
@ -714,16 +715,8 @@ def _transform_request_body( # noqa: PLR0915
optional_params.pop("output_config", None)
config_fields = GenerationConfig.__annotations__.keys()
# If the LiteLLM client sends Gemini-supported parameter "labels", add it
# as "labels" field to the request sent to the Gemini backend.
labels: Optional[dict[str, str]] = optional_params.pop("labels", None)
# If the LiteLLM client sends OpenAI-supported parameter "metadata", add it
# as "labels" field to the request sent to the Gemini backend.
if labels is None and "metadata" in litellm_params:
metadata = litellm_params["metadata"]
if metadata is not None and "requester_metadata" in metadata:
rm = metadata["requester_metadata"]
labels = {k: v for k, v in rm.items() if isinstance(v, str)}
# labels: optional explicit param and/or metadata.requester_metadata (OpenAI metadata)
labels = pop_vertex_request_labels(optional_params, litellm_params)
filtered_params = {
k: v

View file

@ -7,7 +7,10 @@ import litellm
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.common_utils import (
get_vertex_base_url,
pop_vertex_request_labels,
)
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
@ -203,13 +206,16 @@ class VertexAIImagenImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
"sampleCount": 1,
}
# Merge with optional params
labels = pop_vertex_request_labels(optional_params, litellm_params)
# Merge with optional params (after popping labels so they are not sent as Imagen parameters)
parameters = {**default_params, **optional_params}
request_body = {
request_body: dict = {
"instances": [{"prompt": prompt}],
"parameters": parameters,
}
if labels:
request_body["labels"] = labels
return request_body

View file

@ -11,12 +11,15 @@ import httpx
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.vertex_ai.common_utils import (
vertex_request_labels_from_litellm_params,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.secret_managers.main import get_secret_str
from litellm.types.rerank import (
RerankBilledUnits,
RerankResponse,
RerankResponseMeta,
RerankBilledUnits,
RerankResponseResult,
)
@ -109,6 +112,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
"""
Transform the request from Cohere format to Vertex AI Discovery Engine format
@ -145,6 +149,10 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase):
# When return_documents is False, we want to ignore record details (return only IDs)
request_data["ignoreRecordDetailsInResponse"] = not return_documents
user_labels = vertex_request_labels_from_litellm_params(litellm_params)
if user_labels:
request_data["userLabels"] = user_labels
return request_data
def transform_rerank_response(

View file

@ -1,4 +1,4 @@
from typing import Literal, Optional, Union
from typing import Dict, Literal, Optional, Union
import httpx
@ -44,6 +44,7 @@ class VertexEmbedding(VertexBase):
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES] = None,
gemini_api_key: Optional[str] = None,
extra_headers: Optional[dict] = None,
litellm_params: Optional[Dict] = None,
) -> EmbeddingResponse:
if aembedding is True:
return self.async_embedding( # type: ignore
@ -61,6 +62,7 @@ class VertexEmbedding(VertexBase):
vertex_credentials=vertex_credentials,
gemini_api_key=gemini_api_key,
extra_headers=extra_headers,
litellm_params=litellm_params,
)
should_use_v1beta1_features = self.is_using_v1beta1_features(
@ -92,7 +94,10 @@ class VertexEmbedding(VertexBase):
headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers)
vertex_request: VertexEmbeddingRequest = (
litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input=input, optional_params=optional_params, model=model
input=input,
optional_params=optional_params,
model=model,
litellm_params=litellm_params,
)
)
@ -156,6 +161,7 @@ class VertexEmbedding(VertexBase):
gemini_api_key: Optional[str] = None,
extra_headers: Optional[dict] = None,
encoding=None,
litellm_params: Optional[Dict] = None,
) -> EmbeddingResponse:
"""
Async embedding implementation
@ -188,7 +194,10 @@ class VertexEmbedding(VertexBase):
headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers)
vertex_request: VertexEmbeddingRequest = (
litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input=input, optional_params=optional_params, model=model
input=input,
optional_params=optional_params,
model=model,
litellm_params=litellm_params,
)
)

View file

@ -3,6 +3,7 @@ from typing import List, Literal, Optional, Union
from pydantic import BaseModel
from litellm.llms.vertex_ai.common_utils import pop_vertex_request_labels
from litellm.types.utils import EmbeddingResponse, Usage
from .types import *
@ -100,7 +101,11 @@ class VertexAITextEmbeddingConfig(BaseModel):
return optional_params
def transform_openai_request_to_vertex_embedding_request(
self, input: Union[list, str], optional_params: dict, model: str
self,
input: Union[list, str],
optional_params: dict,
model: str,
litellm_params: Optional[dict] = None,
) -> VertexEmbeddingRequest:
"""
Transforms an openai request to a vertex embedding request.
@ -108,16 +113,26 @@ class VertexAITextEmbeddingConfig(BaseModel):
# Import here to avoid circular import issues with litellm.__init__
from litellm.llms.vertex_ai.vertex_embeddings.bge import VertexBGEConfig
labels = pop_vertex_request_labels(optional_params, litellm_params)
if model.isdigit():
return self._transform_openai_request_to_fine_tuned_embedding_request(
input, optional_params, model
vertex_request = (
self._transform_openai_request_to_fine_tuned_embedding_request(
input, optional_params, model
)
)
if labels:
vertex_request["labels"] = labels
return vertex_request
if VertexBGEConfig.is_bge_model(model):
return VertexBGEConfig.transform_request(
vertex_request = VertexBGEConfig.transform_request(
input=input, optional_params=optional_params, model=model
)
if labels:
vertex_request["labels"] = labels
return vertex_request
vertex_request: VertexEmbeddingRequest = VertexEmbeddingRequest()
vertex_request = VertexEmbeddingRequest()
vertex_text_embedding_input_list: List[TextEmbeddingInput] = []
task_type: Optional[TaskType] = optional_params.get("task_type")
title = optional_params.get("title")
@ -133,6 +148,8 @@ class VertexAITextEmbeddingConfig(BaseModel):
vertex_request["instances"] = vertex_text_embedding_input_list
vertex_request["parameters"] = EmbeddingParameters(**optional_params)
if labels:
vertex_request["labels"] = labels
return vertex_request

View file

@ -3,7 +3,7 @@ Types for Vertex Embeddings Requests
"""
from enum import Enum
from typing import List, Optional, Union
from typing import Dict, List, Optional, Union
from typing_extensions import TypedDict
@ -56,6 +56,7 @@ class VertexEmbeddingRequest(TypedDict, total=False):
List[TextEmbeddingFineTunedInput],
]
parameters: Optional[Union[EmbeddingParameters, TextEmbeddingFineTunedParameters]]
labels: Optional[Dict[str, str]]
# Example usage:

View file

@ -67,7 +67,11 @@ class VoyageRerankConfig(BaseRerankConfig):
return api_base
def transform_rerank_request(
self, model: str, optional_rerank_params: Dict, headers: Dict
self,
model: str,
optional_rerank_params: Dict,
headers: Dict,
litellm_params: Optional[dict] = None,
) -> Dict:
return {"model": model, **optional_rerank_params}

View file

@ -143,6 +143,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig):
model: str,
optional_rerank_params: Dict,
headers: dict,
litellm_params: Optional[dict] = None,
) -> dict:
"""
Transform request to IBM watsonx.ai rerank format

View file

@ -5311,6 +5311,7 @@ def embedding( # noqa: PLR0915
api_key=api_key,
api_base=api_base,
client=client,
litellm_params=litellm_params_dict,
)
elif custom_llm_provider == "oobabooga":
response = oobabooga.embedding(

View file

@ -163,19 +163,21 @@ def rerank( # noqa: PLR0915
model_response = RerankResponse()
rerank_litellm_params = {
"litellm_call_id": litellm_call_id,
"proxy_server_request": proxy_server_request,
"model_info": model_info,
"preset_cache_key": None,
"stream_response": {},
**optional_params.model_dump(exclude_unset=True),
}
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(optional_rerank_params),
litellm_params={
"litellm_call_id": litellm_call_id,
"proxy_server_request": proxy_server_request,
"model_info": model_info,
"preset_cache_key": None,
"stream_response": {},
**optional_params.model_dump(exclude_unset=True),
},
litellm_params=dict(rerank_litellm_params),
custom_llm_provider=_custom_llm_provider,
)
@ -214,6 +216,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.AZURE_AI:
api_base = (
@ -235,6 +238,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.INFINITY:
# Implement Infinity rerank logic
@ -265,6 +269,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.TOGETHER_AI:
# Implement Together AI rerank logic
@ -318,6 +323,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.NVIDIA_NIM:
if dynamic_api_key is None:
@ -346,6 +352,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.BEDROCK:
api_base = (
@ -409,6 +416,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.DEEPINFRA:
@ -442,6 +450,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.FIREWORKS_AI:
api_key = (
@ -472,6 +481,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.VOYAGE:
api_key = (
@ -500,6 +510,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
elif _custom_llm_provider == litellm.LlmProviders.WATSONX:
credentials = IBMWatsonXMixin.get_watsonx_credentials(
@ -527,6 +538,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
else:
# Generic handler for all providers that use base_llm_http_handler
@ -559,6 +571,7 @@ def rerank( # noqa: PLR0915
headers=headers or litellm.headers or {},
client=client,
model_response=model_response,
litellm_params=rerank_litellm_params,
)
# Placeholder return

View file

@ -373,6 +373,20 @@ class TestVertexAIImagenImageGenerationConfig:
assert request["parameters"]["sampleCount"] == 2
assert request["parameters"]["aspectRatio"] == "16:9"
def test_transform_image_generation_request_labels_from_metadata(self):
"""Billing labels from litellm_params.metadata.requester_metadata on predict body."""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={},
litellm_params={
"metadata": {"requester_metadata": {"team": "platform", "env": "prod"}}
},
headers={},
)
assert request["labels"] == {"team": "platform", "env": "prod"}
assert "labels" not in request["parameters"]
def test_transform_image_generation_response(self):
"""Test response transformation"""
mock_response = MagicMock(spec=httpx.Response)

View file

@ -216,6 +216,22 @@ class TestVertexAIRerankTransform:
)
assert request_data_default["ignoreRecordDetailsInResponse"] == False
def test_transform_rerank_request_user_labels_from_metadata(self):
"""Discovery Engine Rank API uses userLabels (string map) for billing."""
optional_params = {
"query": "q",
"documents": ["a", "b"],
}
request_data = self.config.transform_rerank_request(
model=self.model,
optional_rerank_params=optional_params,
headers={},
litellm_params={
"metadata": {"requester_metadata": {"app": "litellm", "tier": "1"}}
},
)
assert request_data["userLabels"] == {"app": "litellm", "tier": "1"}
def test_transform_rerank_request_missing_required_params(self):
"""Test that transform_rerank_request handles missing required parameters."""
# Test missing query

View file

@ -0,0 +1,182 @@
"""
End-to-end tests for Vertex AI rerank `userLabels` propagation.
These tests go through the full `litellm.rerank()` call path with the HTTP
layer mocked, so they catch plumbing bugs (e.g. `litellm_params` losing
`metadata` between the rerank entrypoint and the Vertex transform) that
unit tests on `VertexAIRerankConfig.transform_rerank_request` miss.
"""
import asyncio
import json
import os
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import litellm
import litellm.llms.vertex_ai.rerank.transformation
def _extract_body(call_kwargs):
"""The rerank handler sends `data=json.dumps(...)`, not `json=...`."""
if "json" in call_kwargs and call_kwargs["json"] is not None:
return call_kwargs["json"]
raw = call_kwargs.get("data")
if isinstance(raw, (bytes, bytearray)):
raw = raw.decode("utf-8")
return json.loads(raw) if raw else None
def _make_mock_rank_response():
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"records": [
{"id": "0", "score": 0.9, "title": "doc 0", "content": "hello"},
{"id": "1", "score": 0.1, "title": "doc 1", "content": "world"},
]
}
mock_response.text = '{"records": []}'
return mock_response
def _make_async_mock_rank_response():
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json = MagicMock(
return_value={
"records": [
{"id": "0", "score": 0.9, "title": "doc 0", "content": "hello"},
]
}
)
mock_response.text = '{"records": []}'
return mock_response
@pytest.fixture
def clean_vertex_env():
saved = {}
for var in (
"GOOGLE_APPLICATION_CREDENTIALS",
"GOOGLE_CLOUD_PROJECT",
"VERTEXAI_PROJECT",
"VERTEXAI_CREDENTIALS",
"VERTEX_AI_CREDENTIALS",
"VERTEX_PROJECT",
"VERTEX_LOCATION",
"VERTEX_AI_PROJECT",
):
if var in os.environ:
saved[var] = os.environ.pop(var)
yield
for var, value in saved.items():
os.environ[var] = value
def _patch_vertex_auth():
return patch.object(
litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig,
"_ensure_access_token",
return_value=("test-access-token", "test-project-2049"),
)
def test_rerank_userlabels_propagates_from_metadata_sync(clean_vertex_env):
"""
`litellm.rerank(metadata={"requester_metadata": {...}})` must end up as
`userLabels` on the Discovery Engine `:rank` request body.
"""
captured = {}
def fake_post(*args, **kwargs):
captured["body"] = _extract_body(kwargs)
return _make_mock_rank_response()
with (
_patch_vertex_auth(),
patch(
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
side_effect=fake_post,
),
):
litellm.rerank(
model="vertex_ai/semantic-ranker-default@latest",
query="what is gemini?",
documents=["hello", "world"],
vertex_project="test-project-2049",
vertex_credentials='{"type": "service_account"}',
metadata={"requester_metadata": {"team": "platform", "env": "prod"}},
)
body = captured["body"]
assert body is not None, "expected POST body to be captured"
assert "userLabels" in body, (
"Vertex rerank request body is missing `userLabels` — metadata was "
"lost between litellm.rerank() and transform_rerank_request. "
f"body keys: {sorted(body.keys())}"
)
assert body["userLabels"] == {"team": "platform", "env": "prod"}
def test_rerank_userlabels_propagates_from_metadata_async(clean_vertex_env):
"""Same as the sync test, but through `litellm.arerank`."""
captured = {}
async def fake_post(*args, **kwargs):
captured["body"] = _extract_body(kwargs)
return _make_async_mock_rank_response()
with (
_patch_vertex_auth(),
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
side_effect=fake_post,
),
):
asyncio.run(
litellm.arerank(
model="vertex_ai/semantic-ranker-default@latest",
query="what is gemini?",
documents=["hello", "world"],
vertex_project="test-project-2049",
vertex_credentials='{"type": "service_account"}',
metadata={"requester_metadata": {"team": "platform"}},
)
)
body = captured["body"]
assert body is not None
assert body.get("userLabels") == {"team": "platform"}
def test_rerank_userlabels_absent_when_no_metadata(clean_vertex_env):
"""No metadata → no `userLabels` key (don't send empty maps)."""
captured = {}
def fake_post(*args, **kwargs):
captured["body"] = _extract_body(kwargs)
return _make_mock_rank_response()
with (
_patch_vertex_auth(),
patch(
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
side_effect=fake_post,
),
):
litellm.rerank(
model="vertex_ai/semantic-ranker-default@latest",
query="what is gemini?",
documents=["hello", "world"],
vertex_project="test-project-2049",
vertex_credentials='{"type": "service_account"}',
)
body = captured["body"]
assert body is not None
assert "userLabels" not in body

View file

@ -15,7 +15,9 @@ from litellm.llms.vertex_ai.common_utils import (
convert_anyof_null_to_nullable,
get_vertex_location_from_url,
get_vertex_project_id_from_url,
pop_vertex_request_labels,
set_schema_property_ordering,
vertex_request_labels_from_litellm_params,
)
@ -1444,3 +1446,65 @@ def test_add_object_type_does_not_add_type_when_anyof_present():
# Verify type was not added (anyOf handles the type)
assert "type" not in input_schema, "type should not be added when anyOf is present"
def test_vertex_request_labels_from_litellm_params_extracts_requester_metadata():
assert vertex_request_labels_from_litellm_params(None) is None
assert vertex_request_labels_from_litellm_params({}) is None
assert vertex_request_labels_from_litellm_params({"metadata": None}) is None
lp = {"metadata": {"requester_metadata": {"team": "analytics", "count": 3}}}
assert vertex_request_labels_from_litellm_params(lp) == {"team": "analytics"}
def test_vertex_request_labels_from_litellm_params_accepts_litellm_metadata():
lp = {
"litellm_metadata": {
"requester_metadata": {"team": "platform", "count": 3}
}
}
assert vertex_request_labels_from_litellm_params(lp) == {"team": "platform"}
def test_vertex_request_labels_prefers_metadata_over_litellm_metadata():
lp = {
"metadata": {"requester_metadata": {"source": "metadata"}},
"litellm_metadata": {"requester_metadata": {"source": "litellm_metadata"}},
}
assert vertex_request_labels_from_litellm_params(lp) == {"source": "metadata"}
def test_pop_vertex_request_labels_prefers_explicit_labels_then_metadata():
optional = {"labels": {"env": "prod"}}
litellm_params = {"metadata": {"requester_metadata": {"team": "x"}}}
assert pop_vertex_request_labels(optional, litellm_params) == {"env": "prod"}
assert "labels" not in optional
optional2: dict = {}
assert pop_vertex_request_labels(optional2, litellm_params) == {"team": "x"}
optional3 = {"labels": {"team": 123}}
assert pop_vertex_request_labels(optional3, litellm_params) == {"team": "x"}
def test_pop_vertex_request_labels_uses_litellm_metadata_when_metadata_absent():
optional: dict = {}
litellm_params = {
"litellm_metadata": {"requester_metadata": {"team": "from_litellm_meta"}}
}
assert pop_vertex_request_labels(optional, litellm_params) == {
"team": "from_litellm_meta"
}
def test_vertex_text_embedding_request_includes_labels_from_metadata():
import litellm
req = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input="hi",
optional_params={},
model="text-embedding-004",
litellm_params={
"metadata": {"requester_metadata": {"project_id": "cost-center-1"}}
},
)
assert req.get("labels") == {"project_id": "cost-center-1"}