from __future__ import annotations import logging import os import time from typing import Any, Dict, List, Optional, Union from huggingface_hub import InferenceClient from huggingface_hub.errors import HfHubHTTPError, InferenceTimeoutError from llama_index.core.base.embeddings.base import BaseEmbedding, Embedding from llama_index.core.bridge.pydantic import Field, PrivateAttr from llama_index.embeddings.huggingface_api.pooling import Pooling from llama_index.utils.huggingface import format_query, format_text logger = logging.getLogger(__name__) DEFAULT_MAX_CHARS = int(os.environ.get("HF_EMBED_MAX_CHARS", "3000")) DEFAULT_MIN_CHARS = int(os.environ.get("HF_EMBED_MIN_CHARS", "750")) DEFAULT_MAX_RETRIES = int(os.environ.get("HF_EMBED_RETRIES", "3")) DEFAULT_BACKOFF_SECONDS = float(os.environ.get("HF_EMBED_RETRY_BACKOFF", "1.0")) class SyncHuggingFaceInferenceEmbedding(BaseEmbedding): """Sync-only embedding adapter for HF Inference API. The upstream LlamaIndex HF wrapper uses AsyncInferenceClient internally even for sync calls, which is brittle under uvicorn/asyncio. This adapter uses only the regular InferenceClient, so indexing and retrieval can run safely inside a worker thread. """ pooling: Optional[Pooling] = Field(default=Pooling.CLS) query_instruction: Optional[str] = Field(default=None) text_instruction: Optional[str] = Field(default=None) model_name: str = Field(default="BAAI/bge-small-en-v1.5") token: Union[str, bool, None] = Field(default=None) timeout: Optional[float] = Field(default=None) headers: Optional[Dict[str, str]] = Field(default=None) cookies: Optional[Dict[str, str]] = Field(default=None) max_chars_per_request: int = Field(default=DEFAULT_MAX_CHARS, gt=0) min_chars_per_request: int = Field(default=DEFAULT_MIN_CHARS, gt=0) max_retries: int = Field(default=DEFAULT_MAX_RETRIES, ge=1) retry_backoff_seconds: float = Field(default=DEFAULT_BACKOFF_SECONDS, ge=0.0) _client: InferenceClient = PrivateAttr() def __init__(self, **kwargs: Any) -> None: super().__init__(**kwargs) self._client = InferenceClient( model=self.model_name, token=self.token, timeout=self.timeout, headers=self.headers, cookies=self.cookies, ) @classmethod def class_name(cls) -> str: return "SyncHuggingFaceInferenceEmbedding" @staticmethod def _mean_pool_vectors(vectors: List[Embedding]) -> Embedding: if not vectors: raise ValueError("Cannot average an empty list of embeddings.") if len(vectors) == 1: return vectors[0] return [sum(values) / len(values) for values in zip(*vectors)] @staticmethod def _split_text(text: str, max_chars: int) -> List[str]: stripped = text.strip() if len(stripped) <= max_chars: return [stripped] if stripped else [" "] paragraphs = [part.strip() for part in stripped.split("\n\n") if part.strip()] if not paragraphs: paragraphs = [stripped] segments: List[str] = [] current = "" for paragraph in paragraphs: pieces = [paragraph[i : i + max_chars] for i in range(0, len(paragraph), max_chars)] for piece in pieces: candidate = piece if not current else f"{current}\n\n{piece}" if len(candidate) <= max_chars: current = candidate else: if current: segments.append(current) current = piece if current: segments.append(current) return segments or [" "] def _embed_request(self, text: str) -> Embedding: embedding = self._client.feature_extraction( text, truncate=True, truncation_direction="right", ) if len(embedding.shape) == 1: return embedding.tolist() embedding = embedding.squeeze(axis=0) if len(embedding.shape) == 1: return embedding.tolist() if self.pooling is None: raise ValueError( f"Pooling is required for {self.model_name} because it returned " "a > 1-D value." ) return self.pooling(embedding).tolist() def _embed_with_retry(self, text: str) -> Embedding: last_error: Exception | None = None for attempt in range(1, self.max_retries + 1): try: return self._embed_request(text) except (HfHubHTTPError, InferenceTimeoutError) as exc: last_error = exc status_code = getattr(getattr(exc, "response", None), "status_code", None) retryable = isinstance(exc, InferenceTimeoutError) or status_code in (429, 500, 502, 503, 504) if not retryable or attempt == self.max_retries: break delay = self.retry_backoff_seconds * (2 ** (attempt - 1)) logger.warning( "HF embedding request failed for %s (status=%s, attempt %d/%d). Retrying in %.1fs.", self.model_name, status_code, attempt, self.max_retries, delay, ) time.sleep(delay) if last_error is not None: raise last_error raise RuntimeError("Embedding request failed without an exception.") def _embed_single(self, text: str, *, max_chars: Optional[int] = None) -> Embedding: stripped = text.strip() or " " current_max_chars = max_chars or self.max_chars_per_request segments = self._split_text(stripped, current_max_chars) if len(segments) > 1: embeddings = [self._embed_single(segment, max_chars=current_max_chars) for segment in segments] return self._mean_pool_vectors(embeddings) try: return self._embed_with_retry(segments[0]) except (HfHubHTTPError, InferenceTimeoutError) as exc: if current_max_chars <= self.min_chars_per_request or len(stripped) <= self.min_chars_per_request: raise exc smaller_max_chars = max(self.min_chars_per_request, current_max_chars // 2) if smaller_max_chars >= current_max_chars: raise exc logger.warning( "HF embedding request for %s failed after retries. Falling back to smaller text segments (%d -> %d chars).", self.model_name, current_max_chars, smaller_max_chars, ) return self._embed_single(stripped, max_chars=smaller_max_chars) def _get_query_embedding(self, query: str) -> Embedding: return self._embed_single( format_query(query, self.model_name, self.query_instruction) ) async def _aget_query_embedding(self, query: str) -> Embedding: return self._get_query_embedding(query) def _get_text_embedding(self, text: str) -> Embedding: return self._embed_single( format_text(text, self.model_name, self.text_instruction) ) async def _aget_text_embedding(self, text: str) -> Embedding: return self._get_text_embedding(text)