feat(vertex_ai): use HTTP/2 httpx client for search_api vector store

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-10 23:17:13 +00:00
parent ef1a37795c
commit 534123a49c
7 changed files with 134 additions and 9 deletions

View file

@ -124,6 +124,9 @@ class BaseVectorStoreConfig:
def validate_create_vector_store(self) -> None:
return None
def get_httpx_client_params(self) -> Mapping[str, object]:
return MappingProxyType({})
def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]:
return []

View file

@ -566,17 +566,20 @@ class AsyncHTTPHandler:
client_alias: str | None = None, # name for client in logs
ssl_verify: VerifyTypes | None = None,
shared_session: Optional["ClientSession"] = None,
http2: bool = False,
):
self.timeout = timeout
self.event_hooks = event_hooks
self.ssl_verify = ssl_verify
self.shared_session = shared_session
self.http2 = http2
self._owns_client = True
self._client = self.create_client(
timeout=timeout,
event_hooks=event_hooks,
ssl_verify=ssl_verify,
shared_session=shared_session,
http2=http2,
)
self.client_alias = client_alias
@ -588,6 +591,7 @@ class AsyncHTTPHandler:
event_hooks=self.event_hooks,
ssl_verify=self.ssl_verify,
shared_session=self.shared_session,
http2=self.http2,
)
return self._client
@ -602,6 +606,7 @@ class AsyncHTTPHandler:
event_hooks: Mapping[str, list[Callable[..., object]]] | None,
ssl_verify: VerifyTypes | None = None,
shared_session: Optional["ClientSession"] = None,
http2: bool = False,
) -> httpx.AsyncClient:
# Get unified SSL configuration
ssl_config: Final = get_ssl_configuration(ssl_verify)
@ -614,10 +619,19 @@ class AsyncHTTPHandler:
timeout = _DEFAULT_TIMEOUT
# Create a client with a connection pool
transport: Final = AsyncHTTPHandler._create_async_transport(
ssl_context=ssl_config if isinstance(ssl_config, ssl.SSLContext) else None,
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
shared_session=shared_session,
transport: Final = (
httpx.AsyncHTTPTransport(
http2=True,
verify=ssl_config,
cert=cert,
local_address=_IPV4_LOCAL_ADDRESS if litellm.force_ipv4 else None,
)
if http2
else AsyncHTTPHandler._create_async_transport(
ssl_context=ssl_config if isinstance(ssl_config, ssl.SSLContext) else None,
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
shared_session=shared_session,
)
)
# Get default headers (User-Agent, overridable via LITELLM_USER_AGENT)
@ -625,7 +639,7 @@ class AsyncHTTPHandler:
return httpx.AsyncClient(
transport=transport,
mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=cert),
mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=cert, http2=http2),
event_hooks=event_hooks,
timeout=timeout,
verify=ssl_config,
@ -1229,11 +1243,12 @@ class AsyncHTTPHandler:
transport: LiteLLMAiohttpTransport | AsyncHTTPTransport | None,
verify: VerifyTypes,
cert: CertTypes | None,
http2: bool = False,
) -> Mapping[str, AsyncHTTPTransport | None] | None:
if not isinstance(transport, AsyncHTTPTransport):
return None
return _environment_proxy_mounts(
lambda proxy_url: AsyncHTTPTransport(proxy=proxy_url, verify=verify, cert=cert)
lambda proxy_url: AsyncHTTPTransport(proxy=proxy_url, verify=verify, cert=cert, http2=http2)
)
@ -1246,10 +1261,12 @@ class HTTPHandler:
ssl_verify: bool | str | None = None,
disable_default_headers: bool
| None = False, # arize phoenix returns different API responses when user agent header in request
http2: bool = False,
):
self.timeout = timeout
self.ssl_verify = ssl_verify
self.disable_default_headers = disable_default_headers
self.http2 = http2
self._owns_client = client is None
self._heal_lock = threading.Lock()
self._client = self.create_client() if client is None else client
@ -1269,6 +1286,7 @@ class HTTPHandler:
return httpx.Client(
transport=self._create_sync_transport(),
mounts=self._create_sync_proxy_mounts(verify=ssl_config, cert=cert),
http2=self.http2,
timeout=self.timeout if self.timeout is not None else _DEFAULT_TIMEOUT,
verify=ssl_config,
cert=cert,
@ -1549,7 +1567,7 @@ class HTTPHandler:
Some users have seen httpx ConnectionError when using ipv6 - forcing ipv4 resolves the issue for them
"""
if litellm.force_ipv4:
return HTTPTransport(local_address=_IPV4_LOCAL_ADDRESS)
return HTTPTransport(http2=self.http2, local_address=_IPV4_LOCAL_ADDRESS)
else:
return getattr(litellm, "sync_transport", None)

View file

@ -9815,7 +9815,10 @@ class BaseLLMHTTPHandler:
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
params={
"ssl_verify": litellm_params.get("ssl_verify", None),
**vector_store_provider_config.get_httpx_client_params(),
},
)
else:
async_httpx_client = client
@ -9953,7 +9956,12 @@ class BaseLLMHTTPHandler:
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)})
sync_httpx_client = _get_httpx_client(
params={
"ssl_verify": litellm_params.get("ssl_verify", None),
**vector_store_provider_config.get_httpx_client_params(),
}
)
else:
sync_httpx_client = client

View file

@ -1,4 +1,6 @@
import importlib.util
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol
import httpx
@ -51,6 +53,8 @@ VERTEX_SEARCH_DATASTORE_EXTRA_BODY_FIELDS: Final = frozenset(VertexSearchDataSto
VERTEX_SEARCH_ENGINE_EXTRA_BODY_FIELDS: Final = frozenset(VertexSearchEngineExtraBody.__annotations__)
_H2_AVAILABLE: Final = importlib.util.find_spec("h2") is not None
class VertexSearchSnippet(TypedDict, total=False):
snippet: ReadOnly[str]
@ -108,6 +112,11 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
def __init__(self):
super().__init__()
def get_httpx_client_params(self) -> Mapping[str, object]:
if _H2_AVAILABLE:
return MappingProxyType({"http2": True})
return MappingProxyType({})
@staticmethod
def get_supported_extra_body_fields(is_engine: bool = False) -> frozenset[str]:
"""

View file

@ -21,6 +21,7 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
MaskedHTTPStatusError,
_get_httpx_client,
get_async_httpx_client,
get_ssl_configuration,
)
@ -645,6 +646,52 @@ def test_get_httpx_client_applies_httpx_timeout_object_without_mocking_handler()
handler.close()
@pytest.mark.asyncio
async def test_async_http_handler_http2_transport():
monkeypatch = pytest.MonkeyPatch()
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
http2_handler = AsyncHTTPHandler(http2=True)
default_handler = AsyncHTTPHandler()
try:
assert isinstance(http2_handler.client._transport, httpx.AsyncHTTPTransport)
assert http2_handler.client._transport._pool._http2 is True
assert isinstance(default_handler.client._transport, httpx.AsyncHTTPTransport)
assert default_handler.client._transport._pool._http2 is False
finally:
await http2_handler.close()
await default_handler.close()
monkeypatch.undo()
def test_http_handler_http2_transport():
http2_handler = HTTPHandler(http2=True)
default_handler = HTTPHandler()
try:
assert http2_handler.client._transport._pool._http2 is True
assert default_handler.client._transport._pool._http2 is False
finally:
http2_handler.close()
default_handler.close()
@pytest.mark.asyncio
async def test_get_async_httpx_client_http2_cache_key():
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.types.utils import LlmProviders
monkeypatch = pytest.MonkeyPatch()
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
http2_handler = get_async_httpx_client(llm_provider=LlmProviders.VERTEX_AI, params={"http2": True})
default_handler = get_async_httpx_client(llm_provider=LlmProviders.VERTEX_AI)
try:
assert http2_handler is not default_handler
assert http2_handler.client._transport._pool._http2 is True
finally:
await http2_handler.close()
await default_handler.close()
monkeypatch.undo()
def test_sync_get_forwards_per_request_timeout():
"""HTTPHandler.get(timeout=...) must apply the timeout to that request,
overriding the client default rather than silently ignoring it."""

View file

@ -2572,6 +2572,40 @@ def test_vector_store_search_handler_direct_config_sync_skips_http():
assert pre_call_args["vector_store_id"] == "vs_direct"
@pytest.mark.asyncio
async def test_async_vector_store_search_handler_passes_provider_httpx_params():
from litellm.llms.vertex_ai.vector_stores.search_api.transformation import (
VertexSearchAPIVectorStoreConfig,
)
response = httpx.Response(200, json={"results": []})
client = AsyncMock(spec=AsyncHTTPHandler)
client.post.return_value = response
config = VertexSearchAPIVectorStoreConfig()
logging_obj = Mock(model_call_details={})
with (
patch.object(config, "validate_environment", return_value={}),
patch.object(config, "get_complete_url", return_value="https://discoveryengine.googleapis.com/v1/search"),
patch( # test-quality-ok: verifies vector-store transport configuration reaches the HTTP client factory
"litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client",
return_value=client,
) as get_client,
):
result = await BaseLLMHTTPHandler().async_vector_store_search_handler(
vector_store_id="vs",
query="q",
vector_store_search_optional_params={},
vector_store_provider_config=config,
custom_llm_provider="vertex_ai",
litellm_params=GenericLiteLLMParams(),
logging_obj=logging_obj,
)
assert result["data"] == []
assert get_client.call_args.kwargs["params"] == {"ssl_verify": None, "http2": True}
@pytest.mark.asyncio
async def test_vector_store_search_handler_direct_config_async_skips_http():
handler = BaseLLMHTTPHandler()

View file

@ -3,11 +3,17 @@ from types import SimpleNamespace
import pytest
from litellm.exceptions import BadRequestError
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.vertex_ai.vector_stores.search_api.transformation import (
VertexSearchAPIVectorStoreConfig,
)
def test_vector_store_httpx_client_params():
assert VertexSearchAPIVectorStoreConfig().get_httpx_client_params() == {"http2": True}
assert BaseVectorStoreConfig().get_httpx_client_params() == {}
def test_should_encode_vertex_search_vector_store_id_in_complete_url():
config = VertexSearchAPIVectorStoreConfig()