diff --git a/litellm/llms/base_llm/rerank/transformation.py b/litellm/llms/base_llm/rerank/transformation.py index 7874201f7f0..166f876ba04 100644 --- a/litellm/llms/base_llm/rerank/transformation.py +++ b/litellm/llms/base_llm/rerank/transformation.py @@ -33,6 +33,7 @@ class BaseRerankConfig(ABC): model: str, optional_rerank_params: Dict, headers: dict, + litellm_params: Optional[dict] = None, ) -> dict: return {} diff --git a/litellm/llms/cohere/rerank/transformation.py b/litellm/llms/cohere/rerank/transformation.py index 531b94d1805..64ae8e8ffa7 100644 --- a/litellm/llms/cohere/rerank/transformation.py +++ b/litellm/llms/cohere/rerank/transformation.py @@ -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") diff --git a/litellm/llms/cohere/rerank_v2/transformation.py b/litellm/llms/cohere/rerank_v2/transformation.py index 60d22ff4be0..4c800d6455d 100644 --- a/litellm/llms/cohere/rerank_v2/transformation.py +++ b/litellm/llms/cohere/rerank_v2/transformation.py @@ -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") diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index bc2ca805e94..dc625918b98 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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 diff --git a/litellm/llms/deepinfra/rerank/transformation.py b/litellm/llms/deepinfra/rerank/transformation.py index 71e300d258c..276735f4758 100644 --- a/litellm/llms/deepinfra/rerank/transformation.py +++ b/litellm/llms/deepinfra/rerank/transformation.py @@ -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: diff --git a/litellm/llms/fireworks_ai/rerank/transformation.py b/litellm/llms/fireworks_ai/rerank/transformation.py index eb92399a058..4a7b64b9b77 100644 --- a/litellm/llms/fireworks_ai/rerank/transformation.py +++ b/litellm/llms/fireworks_ai/rerank/transformation.py @@ -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 diff --git a/litellm/llms/hosted_vllm/rerank/transformation.py b/litellm/llms/hosted_vllm/rerank/transformation.py index 8066e53afc7..60b6dc7d23d 100644 --- a/litellm/llms/hosted_vllm/rerank/transformation.py +++ b/litellm/llms/hosted_vllm/rerank/transformation.py @@ -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") diff --git a/litellm/llms/huggingface/rerank/transformation.py b/litellm/llms/huggingface/rerank/transformation.py index 3f83b8e422d..2c847b617ef 100644 --- a/litellm/llms/huggingface/rerank/transformation.py +++ b/litellm/llms/huggingface/rerank/transformation.py @@ -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") diff --git a/litellm/llms/jina_ai/rerank/transformation.py b/litellm/llms/jina_ai/rerank/transformation.py index 48d876f8ea2..ad4416925a6 100644 --- a/litellm/llms/jina_ai/rerank/transformation.py +++ b/litellm/llms/jina_ai/rerank/transformation.py @@ -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} diff --git a/litellm/llms/nvidia_nim/rerank/ranking_transformation.py b/litellm/llms/nvidia_nim/rerank/ranking_transformation.py index 757d874bf31..b9a46b8ac2b 100644 --- a/litellm/llms/nvidia_nim/rerank/ranking_transformation.py +++ b/litellm/llms/nvidia_nim/rerank/ranking_transformation.py @@ -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, ) diff --git a/litellm/llms/nvidia_nim/rerank/transformation.py b/litellm/llms/nvidia_nim/rerank/transformation.py index bd5abac60c8..fc317293acc 100644 --- a/litellm/llms/nvidia_nim/rerank/transformation.py +++ b/litellm/llms/nvidia_nim/rerank/transformation.py @@ -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. diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index b4bfde5f541..fae175612b1 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -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""" diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 533bd06d2d8..87bd4843822 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py index 1c7696d55a2..05ebd685d91 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_imagen_transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/rerank/transformation.py b/litellm/llms/vertex_ai/rerank/transformation.py index 53651839671..3b84972e946 100644 --- a/litellm/llms/vertex_ai/rerank/transformation.py +++ b/litellm/llms/vertex_ai/rerank/transformation.py @@ -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( diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 18c5ec3d839..696341598e5 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -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, ) ) diff --git a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py index 132f29987af..24396628dbd 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/vertex_embeddings/types.py b/litellm/llms/vertex_ai/vertex_embeddings/types.py index 317b9c4fb81..bf73f4d193a 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/types.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/types.py @@ -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: diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index 521dae980d5..d64450a1211 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -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} diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index 4f8e196f254..202760f68a6 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -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 diff --git a/litellm/main.py b/litellm/main.py index daa0fb063aa..0079bd750cf 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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( diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 9868634362f..e27585116ce 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -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 diff --git a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py b/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py index 6905cda0767..fe5b5a69c95 100644 --- a/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/image_generation/test_vertex_ai_image_generation_transformation.py @@ -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) diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py index d451fb24873..c2ea6f6fab9 100644 --- a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py new file mode 100644 index 00000000000..0b28fa9abc9 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_userlabels_e2e.py @@ -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 diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index 95507390df9..7cb3faf6177 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -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"}