from __future__ import annotations import os import logging import time from dataclasses import dataclass from typing import List, Sequence import numpy as np from openai import OpenAI from langchain_core.embeddings import Embeddings logger = logging.getLogger(__name__) @dataclass class EmbeddingConfig: api_base_url: str = "https://api.siliconflow.com/v1" model: str = "Qwen/Qwen3-Embedding-8B" dimension: int = 4096 batch_size: int = 16 _embed_config: EmbeddingConfig | None = None def get_embedding_config() -> EmbeddingConfig: global _embed_config if _embed_config is None: _embed_config = EmbeddingConfig() return _embed_config class QwenEmbeddings(Embeddings): def __init__(self, config: EmbeddingConfig | None = None): self.config = config or get_embedding_config() api_key = os.getenv("SILICONFLOW_API_KEY", "").strip() if not api_key: raise ValueError("Missing SILICONFLOW_API_KEY environment variable") self._client = OpenAI( api_key=api_key, base_url=self.config.api_base_url, ) logger.info(f"Initialized QwenEmbeddings: {self.config.model}") def embed_query(self, text: str) -> List[float]: return self._embed_texts([text])[0] def embed_documents(self, texts: List[str]) -> List[List[float]]: return self._embed_texts(texts) def _embed_texts(self, texts: Sequence[str]) -> List[List[float]]: if not texts: return [] all_embeddings: List[List[float]] = [] batch_size = self.config.batch_size max_retries = 3 # Process in batches for i in range(0, len(texts), batch_size): batch = list(texts[i:i + batch_size]) # Retry logic for rate limits for attempt in range(max_retries): try: response = self._client.embeddings.create( model=self.config.model, input=batch, ) for item in response.data: all_embeddings.append(item.embedding) break except Exception as e: # Rate limit -> wait and retry if "rate" in str(e).lower() and attempt < max_retries - 1: wait_time = 2 ** attempt logger.warning(f"Rate limited, waiting {wait_time}s...") time.sleep(wait_time) else: raise return all_embeddings def embed_texts_np(self, texts: Sequence[str]) -> np.ndarray: return np.asarray(self._embed_texts(list(texts)), dtype=np.float32) # Backward compatibility aliases SiliconFlowConfig = EmbeddingConfig get_config = get_embedding_config