Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
132 changes: 132 additions & 0 deletions server/embeddings/http_batch.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
from __future__ import annotations

import asyncio
import logging
from collections.abc import Callable, Sequence
from typing import Any

import httpx

logger = logging.getLogger(__name__)

# Shared backoff schedule for all HTTP embedding providers.
BACKOFF_DELAYS: Sequence[float] = (10, 20, 30, 40)
ATTEMPTS = 4

# Transient network failures worth another attempt. httpx.TimeoutException is a
# subclass of TransportError, but naming it keeps the intent obvious.
_RETRYABLE_EXCEPTIONS = (httpx.TransportError, httpx.TimeoutException)


def _delay(delays: Sequence[float], attempt: int) -> float:
return delays[min(attempt, len(delays) - 1)]


async def post_with_retry(
client: httpx.AsyncClient,
url: str,
body: dict,
*,
provider: str,
attempts: int = ATTEMPTS,
delays: Sequence[float] = BACKOFF_DELAYS,
) -> httpx.Response:
"""POST `body` to `url`, retrying rate limits, 5xx and transport errors.

429 honours a `Retry-After` header when the server sends one, otherwise it
falls back to the fixed backoff schedule. The final response is passed
through `raise_for_status()`, so callers only ever see a 2xx.
"""
resp: httpx.Response | None = None
for attempt in range(attempts):
try:
resp = await client.post(url, json=body)
except _RETRYABLE_EXCEPTIONS as exc:
if attempt == attempts - 1:
raise
wait = _delay(delays, attempt)
logger.warning(
"%s request failed (%s: %s) — retrying in %.0fs (attempt %d/%d)",
provider,
type(exc).__name__,
exc,
wait,
attempt + 1,
attempts,
)
await asyncio.sleep(wait)
continue

if resp.status_code == 429:
retry_after = float(resp.headers.get("Retry-After", 0))
wait = retry_after if retry_after > 0 else _delay(delays, attempt)
logger.warning(
"%s rate-limited (429) — retrying in %.0fs (attempt %d/%d)",
provider,
wait,
attempt + 1,
attempts,
)
await asyncio.sleep(wait)
continue

if resp.status_code >= 500:
wait = _delay(delays, attempt)
logger.warning(
"%s server error (%d) — retrying in %.0fs (attempt %d/%d)",
provider,
resp.status_code,
wait,
attempt + 1,
attempts,
)
await asyncio.sleep(wait)
continue

break

assert resp is not None # a transport error on the last attempt re-raises
if resp.status_code >= 400:
logger.error("%s API error %d: %s", provider, resp.status_code, resp.text[:500])
resp.raise_for_status()
return resp


async def embed_in_batches(
texts: list[str],
*,
client: httpx.AsyncClient,
url: str,
provider: str,
batch_size: int,
make_body: Callable[[list[str]], dict],
extract: Callable[[Any], list[list[float]]],
attempts: int = ATTEMPTS,
delays: Sequence[float] = BACKOFF_DELAYS,
) -> list[list[float]]:
"""Embed `texts` in chunks of `batch_size`, one retrying POST per chunk.

`make_body` builds the provider-specific request payload for a chunk and
`extract` pulls the vectors out of the decoded response.
"""
if not texts:
return []
all_vectors: list[list[float]] = []
for i in range(0, len(texts), batch_size):
batch = texts[i : i + batch_size]
resp = await post_with_retry(
client,
url,
make_body(batch),
provider=provider,
attempts=attempts,
delays=delays,
)
batch_vectors = extract(resp.json())
if len(batch_vectors) != len(batch):
raise ValueError(
f"{provider} returned {len(batch_vectors)} vectors for "
f"{len(batch)} inputs — response may be malformed"
)
all_vectors.extend(batch_vectors)
return all_vectors
46 changes: 18 additions & 28 deletions server/embeddings/jina.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,10 @@
from __future__ import annotations

import logging

import httpx

from server.config import settings
from server.embeddings.base import EmbeddingProvider

logger = logging.getLogger(__name__)
from server.embeddings.http_batch import embed_in_batches

# HuggingFace TEI uses the OpenAI-compatible /embed endpoint
_EMBED_PATH = "/embed"
Expand All @@ -26,31 +23,24 @@ def __init__(self) -> None:
def dimensions(self) -> int:
return self._dims

@staticmethod
def _extract(data) -> list[list[float]]:
# TEI returns a list of vectors directly
if isinstance(data, list):
return data
# fallback: OpenAI-style { "data": [...] }
return [item["embedding"] for item in data.get("data", [])]

async def embed_batch(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
all_vectors: list[list[float]] = []
for i in range(0, len(texts), _BATCH_SIZE):
batch = texts[i : i + _BATCH_SIZE]
resp = await self._client.post(
f"{self._base_url}{_EMBED_PATH}",
json={"inputs": batch},
)
resp.raise_for_status()
data = resp.json()
# TEI returns a list of vectors directly
if isinstance(data, list):
batch_vectors: list[list[float]] = data
else:
# fallback: OpenAI-style { "data": [...] }
batch_vectors = [item["embedding"] for item in data.get("data", [])]
if len(batch_vectors) != len(batch):
raise ValueError(
f"Embedding server returned {len(batch_vectors)} vectors for "
f"{len(batch)} inputs — response may be malformed"
)
all_vectors.extend(batch_vectors)
return all_vectors
return await embed_in_batches(
texts,
client=self._client,
url=f"{self._base_url}{_EMBED_PATH}",
provider="Embedding server",
batch_size=_BATCH_SIZE,
make_body=lambda batch: {"inputs": batch},
extract=self._extract,
)

async def embed_query(self, text: str) -> list[float]:
vectors = await self.embed_batch([text])
Expand Down
46 changes: 10 additions & 36 deletions server/embeddings/jina_api.py
Original file line number Diff line number Diff line change
@@ -1,20 +1,15 @@
from __future__ import annotations

import asyncio
import logging

import httpx

from server.config import settings
from server.embeddings.base import EmbeddingProvider

logger = logging.getLogger(__name__)
from server.embeddings.http_batch import embed_in_batches

_API_URL = "https://api.jina.ai/v1/embeddings"
# Jina's hosted API accepts up to 2048 inputs per request; 128 keeps us
# uniform with the OpenAI/Voyage providers.
_BATCH_SIZE = 128
_BACKOFF_DELAYS = [10, 20, 30, 40]

# Native output dimensions for known models. The jina-code-embeddings family
# supports Matryoshka truncation via the `dimensions` API parameter —
Expand Down Expand Up @@ -108,24 +103,6 @@ def _make_body(self, inputs: list[str], task: str) -> dict:
body["dimensions"] = self._dims_override
return body

async def _post_with_retry(self, body: dict) -> dict:
for attempt in range(4):
resp = await self._client.post(_API_URL, json=body)
if resp.status_code != 429:
break
retry_after = float(resp.headers.get("Retry-After", 0))
wait = retry_after if retry_after > 0 else _BACKOFF_DELAYS[attempt]
logger.warning(
"Jina rate-limited (429) — retrying in %.0fs (attempt %d/4)",
wait,
attempt + 1,
)
await asyncio.sleep(wait)
if resp.status_code >= 400:
logger.error("Jina API error %d: %s", resp.status_code, resp.text[:500])
resp.raise_for_status()
return resp.json()

async def _embed(self, texts: list[str], task: str) -> list[list[float]]:
if not texts:
return []
Expand All @@ -137,18 +114,15 @@ async def _embed(self, texts: list[str], task: str) -> list[list[float]]:
f"{empty_indices[:5]} of {len(sanitized)} — callers must filter "
f"empty strings before calling embed_batch/embed_query."
)
all_vectors: list[list[float]] = []
for i in range(0, len(sanitized), _BATCH_SIZE):
batch = sanitized[i : i + _BATCH_SIZE]
data = await self._post_with_retry(self._make_body(batch, task))
batch_vectors = [item["embedding"] for item in data.get("data", [])]
if len(batch_vectors) != len(batch):
raise ValueError(
f"Jina returned {len(batch_vectors)} vectors for "
f"{len(batch)} inputs — response may be malformed"
)
all_vectors.extend(batch_vectors)
return all_vectors
return await embed_in_batches(
sanitized,
client=self._client,
url=_API_URL,
provider="Jina",
batch_size=_BATCH_SIZE,
make_body=lambda batch: self._make_body(batch, task),
extract=lambda data: [item["embedding"] for item in data.get("data", [])],
)

async def close(self) -> None:
await self._client.aclose()
Expand Down
33 changes: 10 additions & 23 deletions server/embeddings/ollama.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,10 @@
from __future__ import annotations

import logging

import httpx

from server.config import settings
from server.embeddings.base import EmbeddingProvider

logger = logging.getLogger(__name__)
from server.embeddings.http_batch import embed_in_batches

_EMBED_PATH = "/api/embed"
_BATCH_SIZE = 32
Expand Down Expand Up @@ -45,25 +42,15 @@ def dimensions(self) -> int:
return self._dims

async def embed_batch(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
all_vectors: list[list[float]] = []
for i in range(0, len(texts), _BATCH_SIZE):
batch = texts[i : i + _BATCH_SIZE]
resp = await self._client.post(
f"{self._base_url}{_EMBED_PATH}",
json={"model": self._model, "input": batch},
)
resp.raise_for_status()
data = resp.json()
batch_vectors = data.get("embeddings", [])
if len(batch_vectors) != len(batch):
raise ValueError(
f"Ollama returned {len(batch_vectors)} vectors for "
f"{len(batch)} inputs — response may be malformed"
)
all_vectors.extend(batch_vectors)
return all_vectors
return await embed_in_batches(
texts,
client=self._client,
url=f"{self._base_url}{_EMBED_PATH}",
provider="Ollama",
batch_size=_BATCH_SIZE,
make_body=lambda batch: {"model": self._model, "input": batch},
extract=lambda data: data.get("embeddings", []),
)

async def embed_query(self, text: str) -> list[float]:
vectors = await self.embed_batch([text])
Expand Down
52 changes: 16 additions & 36 deletions server/embeddings/openai.py
Original file line number Diff line number Diff line change
@@ -1,20 +1,15 @@
from __future__ import annotations

import asyncio
import logging

import httpx

from server.config import settings
from server.embeddings.base import EmbeddingProvider

logger = logging.getLogger(__name__)
from server.embeddings.http_batch import embed_in_batches

_API_URL = "https://api.openai.com/v1/embeddings"
# OpenAI accepts up to 2048 inputs per request; 128 is conservative and matches
# Voyage's cap, so behavior is uniform across providers.
_BATCH_SIZE = 128
_BACKOFF_DELAYS = [10, 20, 30, 40]

_NATIVE_DIMENSIONS: dict[str, int] = {
"text-embedding-3-large": 3072,
Expand Down Expand Up @@ -56,37 +51,22 @@ def __init__(self) -> None:
def dimensions(self) -> int:
return self._dims

def _make_body(self, inputs: list[str]) -> dict:
body: dict = {"model": self._model, "input": inputs}
if self._dims_override is not None:
body["dimensions"] = self._dims_override
return body

async def embed_batch(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
all_vectors: list[list[float]] = []
for i in range(0, len(texts), _BATCH_SIZE):
batch = texts[i : i + _BATCH_SIZE]
body: dict = {"model": self._model, "input": batch}
if self._dims_override is not None:
body["dimensions"] = self._dims_override
for attempt in range(4):
resp = await self._client.post(_API_URL, json=body)
if resp.status_code != 429:
break
retry_after = float(resp.headers.get("Retry-After", 0))
wait = retry_after if retry_after > 0 else _BACKOFF_DELAYS[attempt]
logger.warning(
"OpenAI rate-limited (429) — retrying in %.0fs (attempt %d/4)",
wait,
attempt + 1,
)
await asyncio.sleep(wait)
resp.raise_for_status()
data = resp.json()
batch_vectors = [item["embedding"] for item in data.get("data", [])]
if len(batch_vectors) != len(batch):
raise ValueError(
f"OpenAI returned {len(batch_vectors)} vectors for "
f"{len(batch)} inputs — response may be malformed"
)
all_vectors.extend(batch_vectors)
return all_vectors
return await embed_in_batches(
texts,
client=self._client,
url=_API_URL,
provider="OpenAI",
batch_size=_BATCH_SIZE,
make_body=self._make_body,
extract=lambda data: [item["embedding"] for item in data.get("data", [])],
)

async def embed_query(self, text: str) -> list[float]:
vectors = await self.embed_batch([text])
Expand Down
Loading