mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
72ddbce50e
26 changed files with 425 additions and 33 deletions
|
|
@ -33,6 +33,7 @@ class BaseRerankConfig(ABC):
|
|||
model: str,
|
||||
optional_rerank_params: Dict,
|
||||
headers: dict,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> dict:
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue