diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 6bb7885810a..8f459d05d03 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -26,6 +26,52 @@ else: LiteLLMLoggingObj = Any +VERTEX_SEARCH_TARGET_SELECTING_FIELDS = frozenset( + { + "dataStoreSpecs", + "branch", + "servingConfig", + "entity", + } +) + +VERTEX_SEARCH_SUPPORTED_EXTRA_BODY_FIELDS = frozenset( + { + "query", + "pageSize", + "pageToken", + "offset", + "oneBoxPageSize", + "numResultsPerDataStore", + "pageCategories", + "imageQuery", + "filter", + "canonicalFilter", + "orderBy", + "userInfo", + "languageCode", + "facetSpecs", + "boostSpec", + "params", + "queryExpansionSpec", + "spellCorrectionSpec", + "userPseudoId", + "contentSearchSpec", + "rankingExpression", + "rankingExpressionBackend", + "safeSearch", + "userLabels", + "naturalLanguageQueryUnderstandingSpec", + "searchAsYouTypeSpec", + "displaySpec", + "crowdingSpecs", + "relevanceThreshold", + "relevanceScoreSpec", + "customRankingParams", + } +) + + class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): """ Configuration for Vertex AI Search API Vector Store @@ -36,6 +82,42 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def __init__(self): super().__init__() + @staticmethod + def get_supported_extra_body_fields() -> frozenset: + """Native SearchRequest fields callers may forward via ``extra_body``.""" + return VERTEX_SEARCH_SUPPORTED_EXTRA_BODY_FIELDS + + @classmethod + def _filter_extra_body(cls, extra_body: Dict[str, Any]) -> Dict[str, Any]: + """ + Validate ``extra_body`` against the supported-field allowlist. + + Raises ``ValueError`` if the caller includes a target-selecting field + (e.g. ``dataStoreSpecs``) or any field not on the allowlist, so the + request fails loudly instead of silently searching the wrong store. + """ + supported = cls.get_supported_extra_body_fields() + filtered = { + key: value for key, value in extra_body.items() if value is not None + } + + target_selecting = set(filtered) & VERTEX_SEARCH_TARGET_SELECTING_FIELDS + if target_selecting: + raise ValueError( + "Vertex AI Search extra_body may not set target-selecting fields " + f"{sorted(target_selecting)}: the data store is scoped by " + "vector_store_id / vertex_engine_id and cannot be overridden per request." + ) + + unsupported = set(filtered) - supported + if unsupported: + raise ValueError( + f"Unsupported Vertex AI Search extra_body fields {sorted(unsupported)}. " + f"Supported fields: {sorted(supported)}." + ) + + return filtered + def get_auth_credentials( self, litellm_params: dict ) -> BaseVectorStoreAuthCredentials: @@ -136,9 +218,14 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): Transform a search request for the Vertex AI Search (Discovery Engine) API. Per-request params pass through to the engine: max_num_results maps to - pageSize, and extra_body is merged in raw and takes precedence, so callers - can send native Discovery Engine fields such as dataStoreSpecs, filter, + pageSize, and extra_body fields on the supported allowlist + (`get_supported_extra_body_fields`) are merged in with precedence, so + callers can send native Discovery Engine tuning fields such as filter, boostSpec, or contentSearchSpec. + + Target-selecting fields (e.g. dataStoreSpecs, branch) are rejected: the + data store is scoped by the URL path (vector_store_id / vertex_engine_id) + and must not be overridable per request. """ if isinstance(query, list): query = " ".join(query) @@ -150,7 +237,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): if max_num_results is not None: request_body["pageSize"] = max_num_results if isinstance(extra_body, dict): - request_body.update(extra_body) + request_body.update(self._filter_extra_body(extra_body)) litellm_logging_obj.model_call_details["query"] = request_body.get( "query", query diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py index 3e7a1c33d58..1a22e49778a 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py @@ -166,18 +166,47 @@ def test_search_request_maps_max_num_results_to_pagesize(): assert body["pageSize"] == 25 -def test_search_request_passes_datastorespecs_through_extra_body(): +def test_search_request_rejects_datastorespecs_in_extra_body(): specs = [ { "dataStore": "projects/p/locations/global/collections/default_collection/dataStores/ds-beta" } ] - _, body = _search_request(extra_body={"dataStoreSpecs": specs}) + with pytest.raises(ValueError, match="target-selecting"): + _search_request(extra_body={"dataStoreSpecs": specs}) - assert body["dataStoreSpecs"] == specs + +@pytest.mark.parametrize( + "field", ["dataStoreSpecs", "branch", "servingConfig", "entity"] +) +def test_search_request_rejects_target_selecting_fields(field): + with pytest.raises(ValueError, match="target-selecting"): + _search_request(extra_body={field: "x"}) + + +def test_search_request_rejects_unsupported_extra_body_field(): + with pytest.raises(ValueError, match="Unsupported Vertex AI Search extra_body"): + _search_request(extra_body={"notARealField": True}) + + +def test_search_request_forwards_supported_extra_body_fields(): + _, body = _search_request( + extra_body={ + "filter": 'category: ANY("docs")', + "boostSpec": {"conditionBoostSpecs": []}, + } + ) + + assert body["filter"] == 'category: ANY("docs")' + assert body["boostSpec"] == {"conditionBoostSpecs": []} assert body["query"] == "hello" - assert body["pageSize"] == 10 + + +def test_search_request_ignores_none_valued_extra_body_fields(): + _, body = _search_request(extra_body={"filter": None}) + + assert "filter" not in body def test_search_request_extra_body_takes_precedence_over_defaults():