add de-duplication for image embedding and add http connection pool

This commit is contained in:
bryan davis 2026-07-08 16:17:37 -05:00
parent 0c8ff570d6
commit 9497921d9c
No known key found for this signature in database
GPG key ID: D11B4A4C0C731E5E
6 changed files with 387 additions and 25 deletions

View file

@ -143,6 +143,13 @@ class HaikuRAG:
async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002 async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002
"""Async context manager exit.""" """Async context manager exit."""
await self._await_vacuum_tasks() await self._await_vacuum_tasks()
# Release the embedder's pooled HTTP client (if any). Best-effort like
# the vacuum drain: __aexit__ runs during exception unwinding, so a
# raising close here must not mask the original exception.
try:
await self.embedder.aclose()
except Exception:
logger.debug("Embedder aclose failed on teardown", exc_info=True)
self.close() self.close()
return False return False

View file

@ -12,6 +12,7 @@ from haiku.rag.config.models import (
CircuitBreakerConfig, CircuitBreakerConfig,
ConversionOptions, ConversionOptions,
DoclingServeConfig, DoclingServeConfig,
EmbeddingHTTPConfig,
EmbeddingModelConfig, EmbeddingModelConfig,
EmbeddingsConfig, EmbeddingsConfig,
FSSourceConfig, FSSourceConfig,
@ -42,6 +43,7 @@ __all__ = [
"CircuitBreakerConfig", "CircuitBreakerConfig",
"ConversionOptions", "ConversionOptions",
"DoclingServeConfig", "DoclingServeConfig",
"EmbeddingHTTPConfig",
"EmbeddingModelConfig", "EmbeddingModelConfig",
"EmbeddingsConfig", "EmbeddingsConfig",
"FSSourceConfig", "FSSourceConfig",

View file

@ -68,9 +68,29 @@ class LanceDBConfig(BaseModel):
storage_options: dict[str, str] = Field(default_factory=dict) storage_options: dict[str, str] = Field(default_factory=dict)
class EmbeddingHTTPConfig(BaseModel):
"""HTTP transport tuning for embedding requests to OpenAI-compatible
servers (vllm, openai, ollama).
Applied to the pooled ``httpx.AsyncClient`` the embedder reuses across
every request, so a connection and its name resolution is established
once and kept warm. Providers whose SDKs manage their own transport
(voyageai, cohere, sentence-transformers) ignore these settings.
``max_connections`` bounds concurrent in-flight requests; keep it >= any
future embedding concurrency so the pool is never the limiter.
"""
timeout_s: float = 60.0
max_connections: int = 16
max_keepalive_connections: int = 16
keepalive_expiry_s: float = 300.0
class EmbeddingsConfig(BaseModel): class EmbeddingsConfig(BaseModel):
model: EmbeddingModelConfig = Field(default_factory=EmbeddingModelConfig) model: EmbeddingModelConfig = Field(default_factory=EmbeddingModelConfig)
batch_size: int = 512 batch_size: int = 512
http: EmbeddingHTTPConfig = Field(default_factory=EmbeddingHTTPConfig)
class RerankingConfig(BaseModel): class RerankingConfig(BaseModel):

View file

@ -1,7 +1,9 @@
import base64 import base64
import hashlib
import io import io
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import httpx
from pydantic_ai.embeddings import Embedder from pydantic_ai.embeddings import Embedder
from pydantic_ai.embeddings.openai import OpenAIEmbeddingModel from pydantic_ai.embeddings.openai import OpenAIEmbeddingModel
from pydantic_ai.providers.ollama import OllamaProvider from pydantic_ai.providers.ollama import OllamaProvider
@ -12,10 +14,26 @@ from haiku.rag.config import AppConfig, Config
if TYPE_CHECKING: if TYPE_CHECKING:
from PIL import Image as PILImage from PIL import Image as PILImage
from haiku.rag.config.models import EmbeddingModelConfig from haiku.rag.config.models import EmbeddingHTTPConfig, EmbeddingModelConfig
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
def build_http_client(http: "EmbeddingHTTPConfig") -> httpx.AsyncClient:
"""A pooled ``httpx.AsyncClient`` configured from ``EmbeddingHTTPConfig``.
Shared across an embedder's requests so a connection (and its name
resolution) is established once and kept warm rather than rebuilt per call.
"""
return httpx.AsyncClient(
timeout=httpx.Timeout(http.timeout_s),
limits=httpx.Limits(
max_connections=http.max_connections,
max_keepalive_connections=http.max_keepalive_connections,
keepalive_expiry=http.keepalive_expiry_s,
),
)
ImageInput = "bytes | PILImage.Image" ImageInput = "bytes | PILImage.Image"
@ -34,11 +52,17 @@ class EmbedderWrapper:
embedder: Embedder | None, embedder: Embedder | None,
vector_dim: int, vector_dim: int,
supports_images: bool | None = None, supports_images: bool | None = None,
*,
owned_http_client: "Any | None" = None,
): ):
self._embedder = embedder self._embedder = embedder
self._vector_dim = vector_dim self._vector_dim = vector_dim
if supports_images is not None: if supports_images is not None:
self.supports_images = supports_images self.supports_images = supports_images
# An httpx.AsyncClient this wrapper built and must close on teardown —
# e.g. the pooled client passed to an openai/ollama provider. None when
# the underlying SDK owns its own transport.
self._owned_http_client = owned_http_client
@property @property
def vector_dim(self) -> int: def vector_dim(self) -> int:
@ -73,6 +97,15 @@ class EmbedderWrapper:
"embeddings.model.multimodal: true on a vllm, voyageai, or cohere model." "embeddings.model.multimodal: true on a vllm, voyageai, or cohere model."
) )
async def aclose(self) -> None:
"""Release any resources held by the embedder. Closes a pooled HTTP
client this wrapper owns (openai/ollama); a no-op otherwise. Lets
callers tear down uniformly regardless of embedder type. Subclasses
that own their own client (e.g. vLLM) override this."""
if self._owned_http_client is not None:
await self._owned_http_client.aclose()
self._owned_http_client = None
def _to_data_uri(image: "bytes | PILImage.Image") -> str: def _to_data_uri(image: "bytes | PILImage.Image") -> str:
"""Render an image as a ``data:image/png;base64,...`` URI.""" """Render an image as a ``data:image/png;base64,...`` URI."""
@ -152,8 +185,27 @@ async def embed_chunks(
"embeddings.model.multimodal: true on a vllm, voyageai, or cohere " "embeddings.model.multimodal: true on a vllm, voyageai, or cohere "
"model, or omit picture chunks." "model, or omit picture chunks."
) )
# Identical image bytes embed to identical vectors, so embed each
# distinct image once and reuse the result for every chunk that shares
# it. A document that repeats one figure across many pages (header,
# watermark, logo) collapses from one request per occurrence to one
# per unique image. Keyed by a FIPS-safe content hash; order is
# preserved because we append one vector per chunk in chunk order.
embedding_cache: dict[bytes, list[float]] = {}
for chunk in picture_chunks: for chunk in picture_chunks:
picture_embeddings.append(await embedder.embed_image(chunk._picture_data)) data = chunk._picture_data
key = (
hashlib.sha256(data, usedforsecurity=False).digest()
if isinstance(data, bytes | bytearray)
else None
)
if key is not None and (cached := embedding_cache.get(key)) is not None:
picture_embeddings.append(cached)
continue
embedding = await embedder.embed_image(data)
if key is not None:
embedding_cache[key] = embedding
picture_embeddings.append(embedding)
text_iter = iter(text_embeddings) text_iter = iter(text_embeddings)
picture_iter = iter(picture_embeddings) picture_iter = iter(picture_embeddings)
@ -187,32 +239,36 @@ def get_embedder(config: AppConfig = Config) -> EmbedderWrapper:
An embedder instance configured according to the config. An embedder instance configured according to the config.
""" """
embedding_model = config.embeddings.model embedding_model = config.embeddings.model
http = config.embeddings.http
provider = embedding_model.provider provider = embedding_model.provider
model_name = embedding_model.name model_name = embedding_model.name
vector_dim = embedding_model.vector_dim vector_dim = embedding_model.vector_dim
if embedding_model.multimodal: if embedding_model.multimodal:
return _get_multimodal_embedder(embedding_model) return _get_multimodal_embedder(embedding_model, http)
if provider == "ollama": if provider == "ollama":
# Use model-level base_url if set, otherwise fall back to providers config # Use model-level base_url if set, otherwise fall back to providers config
base_url = embedding_model.base_url or config.providers.ollama.base_url base_url = embedding_model.base_url or config.providers.ollama.base_url
if not base_url.rstrip("/").endswith("/v1"): if not base_url.rstrip("/").endswith("/v1"):
base_url = base_url.rstrip("/") + "/v1" base_url = base_url.rstrip("/") + "/v1"
client = build_http_client(http)
model = OpenAIEmbeddingModel( model = OpenAIEmbeddingModel(
model_name, model_name,
provider=OllamaProvider(base_url=base_url), provider=OllamaProvider(base_url=base_url, http_client=client),
) )
return EmbedderWrapper(Embedder(model), vector_dim) return EmbedderWrapper(Embedder(model), vector_dim, owned_http_client=client)
if provider == "openai": if provider == "openai":
client = build_http_client(http)
provider_kwargs: dict[str, Any] = {"http_client": client}
if embedding_model.base_url: if embedding_model.base_url:
model = OpenAIEmbeddingModel( provider_kwargs["base_url"] = embedding_model.base_url
model_name, model = OpenAIEmbeddingModel(
provider=OpenAIProvider(base_url=embedding_model.base_url), model_name,
) provider=OpenAIProvider(**provider_kwargs),
return EmbedderWrapper(Embedder(model), vector_dim) )
return EmbedderWrapper(Embedder(f"openai:{model_name}"), vector_dim) return EmbedderWrapper(Embedder(model), vector_dim, owned_http_client=client)
if provider == "voyageai": if provider == "voyageai":
return EmbedderWrapper(Embedder(f"voyageai:{model_name}"), vector_dim) return EmbedderWrapper(Embedder(f"voyageai:{model_name}"), vector_dim)
@ -230,7 +286,7 @@ def get_embedder(config: AppConfig = Config) -> EmbedderWrapper:
base_url = _vllm_base_url(embedding_model.base_url) base_url = _vllm_base_url(embedding_model.base_url)
return VLLMMultimodalEmbedder( return VLLMMultimodalEmbedder(
model_name, vector_dim, base_url=base_url, supports_images=False model_name, vector_dim, base_url=base_url, http=http, supports_images=False
) )
raise ValueError(f"Unsupported embedding provider: {provider}") raise ValueError(f"Unsupported embedding provider: {provider}")
@ -245,6 +301,7 @@ def _vllm_base_url(base_url: str | None) -> str:
def _get_multimodal_embedder( def _get_multimodal_embedder(
embedding_model: "EmbeddingModelConfig", embedding_model: "EmbeddingModelConfig",
http: "EmbeddingHTTPConfig",
) -> EmbedderWrapper: ) -> EmbedderWrapper:
"""Build an image-capable embedder for providers that support multimodal. """Build an image-capable embedder for providers that support multimodal.
@ -260,7 +317,7 @@ def _get_multimodal_embedder(
base_url = _vllm_base_url(embedding_model.base_url) base_url = _vllm_base_url(embedding_model.base_url)
return VLLMMultimodalEmbedder( return VLLMMultimodalEmbedder(
model_name, vector_dim, base_url=base_url, supports_images=True model_name, vector_dim, base_url=base_url, http=http, supports_images=True
) )
if provider == "voyageai": if provider == "voyageai":

View file

@ -11,11 +11,13 @@ Models like ``Qwen/Qwen3-VL-Embedding-8B`` and ``jinaai/jina-embeddings-v4``
ship with chat templates that map both shapes into a shared vector space. ship with chat templates that map both shapes into a shared vector space.
""" """
import asyncio
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import httpx import httpx
from haiku.rag.embeddings import EmbedderWrapper, _to_data_uri from haiku.rag.config import EmbeddingHTTPConfig
from haiku.rag.embeddings import EmbedderWrapper, _to_data_uri, build_http_client
if TYPE_CHECKING: if TYPE_CHECKING:
from PIL import Image as PILImage from PIL import Image as PILImage
@ -28,7 +30,7 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
vector_dim: int, vector_dim: int,
base_url: str, base_url: str,
api_key: str | None = None, api_key: str | None = None,
timeout: float = 60.0, http: EmbeddingHTTPConfig | None = None,
supports_images: bool = True, supports_images: bool = True,
): ):
super().__init__( super().__init__(
@ -37,7 +39,16 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
self._model_name = model_name self._model_name = model_name
self._base_url = base_url.rstrip("/") self._base_url = base_url.rstrip("/")
self._api_key = api_key self._api_key = api_key
self._timeout = timeout # Connection-pool + timeout tuning from config (defaults preserve the
# historical 60s / 16-connection behavior when constructed directly).
self._http = http or EmbeddingHTTPConfig()
# One pooled client reused across every request (text, query, image) so
# a connection — and its name resolution — is established once and kept
# warm, rather than a fresh connect per call. Built lazily inside a
# running loop; the lock ensures concurrent first-callers create only
# one.
self._client: httpx.AsyncClient | None = None
self._client_lock = asyncio.Lock()
def _headers(self) -> dict[str, str]: def _headers(self) -> dict[str, str]:
headers = {"Content-Type": "application/json"} headers = {"Content-Type": "application/json"}
@ -45,16 +56,30 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
headers["Authorization"] = f"Bearer {self._api_key}" headers["Authorization"] = f"Bearer {self._api_key}"
return headers return headers
async def _get_client(self) -> httpx.AsyncClient:
if self._client is None:
async with self._client_lock:
if self._client is None:
self._client = build_http_client(self._http)
return self._client
async def aclose(self) -> None:
"""Close the pooled HTTP client. Idempotent and safe to call on
teardown even when no request was ever made (the client is lazy)."""
if self._client is not None:
await self._client.aclose()
self._client = None
async def _post(self, body: dict[str, Any]) -> list[list[float]]: async def _post(self, body: dict[str, Any]) -> list[list[float]]:
client = await self._get_client()
try: try:
async with httpx.AsyncClient(timeout=self._timeout) as client: response = await client.post(
response = await client.post( f"{self._base_url}/embeddings",
f"{self._base_url}/embeddings", json=body,
json=body, headers=self._headers(),
headers=self._headers(), )
) response.raise_for_status()
response.raise_for_status() payload = response.json()
payload = response.json()
except httpx.ConnectError as e: except httpx.ConnectError as e:
raise ValueError( raise ValueError(
f"Could not connect to vLLM at {self._base_url}. " f"Could not connect to vLLM at {self._base_url}. "
@ -62,7 +87,8 @@ class VLLMMultimodalEmbedder(EmbedderWrapper):
) from e ) from e
except httpx.TimeoutException as e: except httpx.TimeoutException as e:
raise ValueError( raise ValueError(
f"Request to vLLM timed out after {self._timeout}s. Error: {e}" f"Request to vLLM timed out after {self._http.timeout_s}s. "
f"Error: {e}"
) from e ) from e
except httpx.HTTPStatusError as e: except httpx.HTTPStatusError as e:
if e.response.status_code == 401: if e.response.status_code == 401:

View file

@ -228,6 +228,77 @@ async def test_embed_chunks_picture_with_text_only_embedder_raises():
await embed_chunks([chunk], get_embedder(config), config) await embed_chunks([chunk], get_embedder(config), config)
class _ImageStubEmbedder(EmbedderWrapper):
"""Multimodal stub that records each image it is asked to embed and
returns a distinct vector per unique payload."""
def __init__(self):
super().__init__(embedder=None, vector_dim=4, supports_images=True)
self.embedded: list[bytes] = []
self._vectors: dict[bytes, list[float]] = {}
async def embed_documents(self, texts):
return [[0.1] * 4 for _ in texts]
async def embed_image(self, image):
self.embedded.append(image)
# Deterministic, payload-specific vector so callers can prove that a
# reused (deduped) vector really came from the matching image.
vec = self._vectors.setdefault(image, [float(len(self._vectors))] * 4)
return list(vec)
async def test_embed_chunks_dedupes_identical_pictures():
"""Identical image bytes are embedded once and the vector is reused for
every chunk that shares them, preserving order."""
img_a = b"\x89PNG\r\n\x1a\nAAAA"
img_b = b"\x89PNG\r\n\x1a\nBBBB"
# Order: a, b, a, a, b — 5 chunks, 2 unique images.
payloads = [img_a, img_b, img_a, img_a, img_b]
chunks = []
for i, data in enumerate(payloads):
c = Chunk(id=f"pic{i}", content="x", order=i)
c._picture_data = data
chunks.append(c)
embedder = _ImageStubEmbedder()
embedded = await embed_chunks(chunks, embedder, AppConfig())
# Only the two unique images hit embed_image, in first-seen order.
assert embedder.embedded == [img_a, img_b]
# Every chunk gets a vector, and duplicates share the right one.
vecs = [c.embedding for c in embedded]
assert vecs[0] == vecs[2] == vecs[3] # all img_a
assert vecs[1] == vecs[4] # all img_b
assert vecs[0] != vecs[1]
async def test_embed_chunks_non_bytes_picture_not_deduped():
"""Picture data that isn't bytes bypasses the content-hash cache and is
embedded per-occurrence (defensive path; build_picture_chunks yields bytes
in practice)."""
payload = "not-bytes-sentinel"
chunks = []
for i in range(2):
c = Chunk(id=f"pic{i}", content="x", order=i)
c._picture_data = payload
chunks.append(c)
embedder = _ImageStubEmbedder()
embedded = await embed_chunks(chunks, embedder, AppConfig())
# No dedup for unhashable/non-bytes payloads: both occurrences embed.
assert embedder.embedded == [payload, payload]
assert embedded[0].embedding is not None
assert embedded[1].embedding is not None
async def test_embedder_aclose_without_owned_client_is_noop():
"""Base aclose is a no-op when the wrapper owns no HTTP client."""
embedder = EmbedderWrapper(embedder=None, vector_dim=4)
await embedder.aclose() # must not raise
async def test_embed_chunks_respects_configured_batch_size(monkeypatch): async def test_embed_chunks_respects_configured_batch_size(monkeypatch):
"""`embeddings.batch_size` controls how `embed_chunks` slices its input. """`embeddings.batch_size` controls how `embed_chunks` slices its input.
@ -408,6 +479,185 @@ async def test_vllm_embed_image_request_shape(monkeypatch):
assert url.startswith("data:image/png;base64,") 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."""
from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder
constructed: list[object] = []
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):
constructed.append(self)
self.closed = False
async def post(self, url, json, headers):
return FakeResponse()
async def aclose(self):
self.closed = True
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 len(constructed) == 1 # one pooled client, reused
await embedder.aclose()
assert constructed[0].closed is True
async def test_vllm_aclose_without_request_is_noop(monkeypatch):
"""aclose is safe when no request was ever made (client is lazy)."""
from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder
def fail(*args, **kwargs):
raise AssertionError("no client should be built before a request")
monkeypatch.setattr("httpx.AsyncClient", fail)
embedder = VLLMMultimodalEmbedder(
model_name="x", vector_dim=2, base_url="http://localhost:8000/v1"
)
await embedder.aclose() # must not raise or construct a client
def test_embeddings_http_config_defaults():
"""Defaults preserve the historical transport behavior."""
http = EmbeddingsConfig().http
assert http.timeout_s == 60.0
assert http.max_connections == 16
assert http.max_keepalive_connections == 16
assert http.keepalive_expiry_s == 300.0
async def test_vllm_client_built_from_http_config(monkeypatch):
"""The pooled client's timeout and pool limits come from
EmbeddingHTTPConfig, not hardcoded constants."""
from haiku.rag.config import EmbeddingHTTPConfig
from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder
captured: dict = {}
class FakeResponse:
def raise_for_status(self):
pass
def json(self):
return {"data": [{"embedding": [0.1]}]}
class FakeAsyncClient:
def __init__(self, *args, **kwargs):
captured["timeout"] = kwargs.get("timeout")
captured["limits"] = kwargs.get("limits")
async def post(self, url, json, headers):
return FakeResponse()
monkeypatch.setattr("httpx.AsyncClient", FakeAsyncClient)
http = EmbeddingHTTPConfig(
timeout_s=12.0,
max_connections=3,
max_keepalive_connections=2,
keepalive_expiry_s=45.0,
)
embedder = VLLMMultimodalEmbedder(
model_name="x", vector_dim=2, base_url="http://localhost:8000/v1", http=http
)
await embedder.embed_query("hi")
assert captured["timeout"].read == 12.0
assert captured["limits"].max_connections == 3
assert captured["limits"].max_keepalive_connections == 2
assert captured["limits"].keepalive_expiry == 45.0
async def test_openai_embedder_owns_and_closes_pooled_client(monkeypatch):
"""The openai flow builds one pooled client, hands it to the provider, and
the wrapper closes it on aclose()."""
import haiku.rag.embeddings as emb
closed = {"value": False}
class FakeClient:
async def aclose(self):
closed["value"] = True
fake_client = FakeClient()
captured: dict = {}
class FakeProvider:
def __init__(self, **kwargs):
captured.update(kwargs)
monkeypatch.setattr(emb, "build_http_client", lambda http: fake_client)
monkeypatch.setattr(emb, "OpenAIProvider", FakeProvider)
monkeypatch.setattr(
emb, "OpenAIEmbeddingModel", lambda name, provider: ("model", name, provider)
)
monkeypatch.setattr(emb, "Embedder", lambda model: ("embedder", model))
config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider="openai", name="text-embedding-3-small", vector_dim=1536
)
)
)
embedder = emb.get_embedder(config)
# The same pooled client is passed to the provider and owned by the wrapper.
assert captured["http_client"] is fake_client
await embedder.aclose()
assert closed["value"] is True
async def test_openai_embedder_forwards_base_url(monkeypatch):
"""A configured base_url is forwarded to the OpenAI provider (OpenAI-
compatible servers like vLLM/LM Studio)."""
import haiku.rag.embeddings as emb
class FakeClient:
async def aclose(self):
pass
captured: dict = {}
class FakeProvider:
def __init__(self, **kwargs):
captured.update(kwargs)
monkeypatch.setattr(emb, "build_http_client", lambda http: FakeClient())
monkeypatch.setattr(emb, "OpenAIProvider", FakeProvider)
monkeypatch.setattr(emb, "OpenAIEmbeddingModel", lambda name, provider: object())
monkeypatch.setattr(emb, "Embedder", lambda model: object())
config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider="openai",
name="text-embedding-3-small",
vector_dim=1536,
base_url="http://vllm:8000/v1",
)
)
)
emb.get_embedder(config)
assert captured["base_url"] == "http://vllm:8000/v1"
async def test_vllm_supports_images_flag(): async def test_vllm_supports_images_flag():
from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder