Pool HTTP clients for vLLM embedding and vLLM/Jina reranking

This commit is contained in:
Yiorgis Gozadinos 2026-07-09 10:32:44 +03:00
parent 43a580afb3
commit c02bcd5dc9
No known key found for this signature in database
9 changed files with 226 additions and 51 deletions

View file

@ -6,6 +6,10 @@
- Duplicate images within a document produce a single picture chunk.
- Pictures smaller than `processing.min_picture_size` pixels on the smaller side (default 64) no longer become picture chunks; `0` disables the filter.
### Fixed
- vLLM embedding and vLLM/Jina reranking reuse one HTTP client across requests instead of opening one per request.
## [0.64.0] - 2026-07-08
### Added

View file

@ -143,6 +143,16 @@ class HaikuRAG:
async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002
"""Async context manager exit."""
await self._await_vacuum_tasks()
# Best-effort: __aexit__ may run during exception unwinding, and a
# raising close must not mask the original exception. The reranker is
# a cached_property — close it only if it was materialized.
try:
await self.embedder.aclose()
reranker = self.__dict__.get("reranker")
if reranker is not None:
await reranker.aclose()
except Exception:
logger.debug("Closing embedder/reranker failed on teardown", exc_info=True)
self.close()
return False

View file

@ -73,6 +73,10 @@ class EmbedderWrapper:
"embeddings.model.multimodal: true on a vllm, voyageai, or cohere model."
)
async def aclose(self) -> None:
"""Release resources held by the embedder. No-op by default;
embedders that own an HTTP client override this."""
def _to_data_uri(image: "bytes | PILImage.Image") -> str:
"""Render an image as a ``data:image/png;base64,...`` URI."""

View file

@ -38,6 +38,10 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
self._base_url = base_url.rstrip("/")
self._api_key = api_key
self._timeout = timeout
# One client reused across every request so the connection (and its
# name resolution) is established once and kept alive, rather than
# rebuilt per call.
self._client = httpx.AsyncClient(timeout=timeout)
def _headers(self) -> dict[str, str]:
headers = {"Content-Type": "application/json"}
@ -45,16 +49,18 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
headers["Authorization"] = f"Bearer {self._api_key}"
return headers
async def aclose(self) -> None:
await self._client.aclose()
async def _post(self, body: dict[str, Any]) -> list[list[float]]:
try:
async with httpx.AsyncClient(timeout=self._timeout) as client:
response = await client.post(
f"{self._base_url}/embeddings",
json=body,
headers=self._headers(),
)
response.raise_for_status()
payload = response.json()
response = await self._client.post(
f"{self._base_url}/embeddings",
json=body,
headers=self._headers(),
)
response.raise_for_status()
payload = response.json()
except httpx.ConnectError as e:
raise ValueError(
f"Could not connect to vLLM at {self._base_url}. "

View file

@ -18,3 +18,7 @@ class RerankerBase:
raise NotImplementedError(
"Reranker is an abstract class. Please implement the _rerank method in a subclass."
)
async def aclose(self) -> None:
"""Release resources held by the reranker. No-op by default;
rerankers that own an HTTP client override this."""

View file

@ -14,34 +14,38 @@ class JinaReranker(RerankerBase):
self._api_key = os.environ.get("JINA_API_KEY")
if not self._api_key:
raise ValueError("JINA_API_KEY environment variable required")
# One client reused across rerank calls (connection kept alive).
self._client = httpx.AsyncClient()
async def aclose(self) -> None:
await self._client.aclose()
async def _rerank(
self, query: str, chunks: list[Chunk], top_n: int = 10
) -> list[tuple[Chunk, float]]:
documents = [chunk.content for chunk in chunks]
async with httpx.AsyncClient() as client:
response = await client.post(
"https://api.jina.ai/v1/rerank",
json={
"model": self._model,
"query": query,
"documents": documents,
"top_n": top_n,
},
headers={
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
},
)
response.raise_for_status()
response = await self._client.post(
"https://api.jina.ai/v1/rerank",
json={
"model": self._model,
"query": query,
"documents": documents,
"top_n": top_n,
},
headers={
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
},
)
response.raise_for_status()
result = response.json()
result = response.json()
scored_chunks = []
for item in result.get("results", []):
index = item["index"]
score = item["relevance_score"]
scored_chunks.append((chunks[index], score))
scored_chunks = []
for item in result.get("results", []):
index = item["index"]
score = item["relevance_score"]
scored_chunks.append((chunks[index], score))
return scored_chunks
return scored_chunks

View file

@ -4,10 +4,15 @@ from haiku.rag.reranking.base import RerankerBase
from haiku.rag.store.models.chunk import Chunk
class VLLMReranker(RerankerBase): # pragma: no cover
class VLLMReranker(RerankerBase):
def __init__(self, model: str, base_url: str):
self._model = model
self._base_url = base_url
# One client reused across rerank calls (connection kept alive).
self._client = httpx.AsyncClient()
async def aclose(self) -> None:
await self._client.aclose()
async def _rerank(
self, query: str, chunks: list[Chunk], top_n: int = 10
@ -15,26 +20,25 @@ class VLLMReranker(RerankerBase): # pragma: no cover
# Prepare documents for reranking
documents = [chunk.content for chunk in chunks]
async with httpx.AsyncClient() as client:
response = await client.post(
f"{self._base_url}/v1/rerank",
json={"model": self._model, "query": query, "documents": documents},
headers={
"accept": "application/json",
"Content-Type": "application/json",
},
)
response.raise_for_status()
response = await self._client.post(
f"{self._base_url}/v1/rerank",
json={"model": self._model, "query": query, "documents": documents},
headers={
"accept": "application/json",
"Content-Type": "application/json",
},
)
response.raise_for_status()
result = response.json()
result = response.json()
# Extract scores and pair with chunks
scored_chunks = []
for item in result.get("results", []):
index = item["index"]
score = item["relevance_score"]
scored_chunks.append((chunks[index], score))
# Extract scores and pair with chunks
scored_chunks = []
for item in result.get("results", []):
index = item["index"]
score = item["relevance_score"]
scored_chunks.append((chunks[index], score))
# Sort by score (descending) and return top_n
scored_chunks.sort(key=lambda x: x[1], reverse=True)
return scored_chunks[:top_n]
# Sort by score (descending) and return top_n
scored_chunks.sort(key=lambda x: x[1], reverse=True)
return scored_chunks[:top_n]

View file

@ -154,6 +154,29 @@ async def test_client_embedder_is_store_embedder(temp_db_path):
assert client.embedder is client.store.embedder
async def test_client_aexit_closes_embedder(temp_db_path, monkeypatch):
"""__aexit__ releases the embedder's HTTP resources. A never-accessed
reranker is not materialized just to be closed; an accessed-but-None
reranker (reranking disabled) is handled."""
from haiku.rag.client import HaikuRAG
closed = []
async def _record():
closed.append(True)
async with HaikuRAG(temp_db_path, create=True) as client:
monkeypatch.setattr(client.embedder, "aclose", _record)
assert client.reranker is None # default config: reranking disabled
assert closed == [True]
async with HaikuRAG(temp_db_path) as client:
monkeypatch.setattr(client.embedder, "aclose", _record)
assert "reranker" not in client.__dict__
@pytest.mark.vcr()
async def test_embed_chunks_basic(allow_model_requests):
"""Test that embed_chunks generates embeddings for chunks."""
@ -408,6 +431,44 @@ async def test_vllm_embed_image_request_shape(monkeypatch):
assert url.startswith("data:image/png;base64,")
async def test_vllm_reuses_pooled_client(monkeypatch):
"""The embedder builds one httpx client and reuses it across requests
instead of opening a fresh connection per call; aclose releases it."""
from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder
stats = {"constructed": 0, "closed": 0}
class FakeResponse:
def raise_for_status(self):
pass
def json(self):
return {"data": [{"embedding": [0.1, 0.2]}]}
class FakeAsyncClient:
def __init__(self, *args, **kwargs):
stats["constructed"] += 1
async def post(self, url, json, headers):
return FakeResponse()
async def aclose(self):
stats["closed"] += 1
monkeypatch.setattr("httpx.AsyncClient", FakeAsyncClient)
embedder = VLLMMultimodalEmbedder(
model_name="x", vector_dim=2, base_url="http://localhost:8000/v1"
)
await embedder.embed_query("one")
await embedder.embed_query("two")
await embedder.embed_documents(["three", "four"])
assert stats["constructed"] == 1
await embedder.aclose()
assert stats["closed"] == 1
async def test_vllm_supports_images_flag():
from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder

View file

@ -274,6 +274,84 @@ class TestGetReranker:
assert getattr(result, attr) == value
class _PoolStats:
"""Fake httpx.AsyncClient factory counting constructions and closes."""
def __init__(self, response_json):
self.constructed = 0
self.closed = 0
stats = self
class FakeResponse:
def raise_for_status(self):
pass
def json(self):
return response_json
class FakeAsyncClient:
def __init__(self, *args, **kwargs):
stats.constructed += 1
async def post(self, url, json, headers):
return FakeResponse()
async def aclose(self):
stats.closed += 1
self.client_class = FakeAsyncClient
@pytest.mark.asyncio
async def test_vllm_reranker_reuses_pooled_client(monkeypatch):
"""One httpx client is built and reused across rerank calls; aclose
releases it."""
from haiku.rag.reranking.vllm import VLLMReranker
stats = _PoolStats({"results": [{"index": 0, "relevance_score": 0.9}]})
monkeypatch.setattr("httpx.AsyncClient", stats.client_class)
reranker = VLLMReranker(model="m", base_url="http://localhost:8000")
docs = [Chunk(content="a", order=0)]
await reranker.rerank("q", docs)
await reranker.rerank("q", docs)
assert stats.constructed == 1
await reranker.aclose()
assert stats.closed == 1
@pytest.mark.asyncio
async def test_jina_reranker_reuses_pooled_client(monkeypatch):
"""One httpx client is built and reused across rerank calls; aclose
releases it."""
monkeypatch.setenv("JINA_API_KEY", "test-api-key")
from haiku.rag.reranking.jina import JinaReranker
stats = _PoolStats({"results": [{"index": 0, "relevance_score": 0.9}]})
monkeypatch.setattr("httpx.AsyncClient", stats.client_class)
reranker = JinaReranker("jina-reranker-v3")
docs = [Chunk(content="a", order=0)]
await reranker.rerank("q", docs)
await reranker.rerank("q", docs)
assert stats.constructed == 1
await reranker.aclose()
assert stats.closed == 1
@pytest.mark.asyncio
async def test_reranker_base_aclose_is_noop():
"""Base aclose exists so client teardown can close any reranker."""
class Custom(RerankerBase):
async def _rerank(self, query, chunks, top_n=10):
return []
await Custom().aclose() # must not raise
def test_jina_reranker_missing_api_key(monkeypatch):
monkeypatch.delenv("JINA_API_KEY", raising=False)