""" Embedding 服务 通过 OpenAI 兼容的 Embeddings API 将文本转换为向量 支持所有兼容 OpenAI 接口的提供商(OpenAI、Qwen、Ollama 等) """ import logging import math from typing import List, Optional from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession logger = logging.getLogger(__name__) # 单次批量请求的默认最大文本数 DEFAULT_BATCH_SIZE = 50 # 不同提供商的批次大小限制 PROVIDER_BATCH_SIZE = { 'qwen': 10, # 阿里云 DashScope 限制 10 'dashscope': 10, 'siliconflow': 10, # 硅基流动限制较小 'ollama': 1, # Ollama 通常逐条处理 } class EmbeddingService: """ Embedding 服务 通过模型 ID 获取对应的提供商,调用 OpenAI 兼容的 Embeddings API """ def __init__(self, db: AsyncSession): self._db = db self._client_cache = {} async def _get_client_and_model(self, model_id: str): """ 根据模型 ID 获取异步客户端和模型名称 Returns: (async_client, model_name, max_tokens, provider_type) """ from ai_platform.models import LLMModel, LLMProvider result = await self._db.execute( select(LLMModel).where( LLMModel.id == model_id, LLMModel.is_active == True, LLMModel.is_deleted == False ) ) model = result.scalar_one_or_none() if not model: raise ValueError(f'Embedding 模型不存在或已禁用: {model_id}') if model.model_type != 'embedding': raise ValueError(f'模型 {model.display_name} 不是 Embedding 类型') provider_result = await self._db.execute( select(LLMProvider).where( LLMProvider.id == model.provider_id, LLMProvider.is_active == True, LLMProvider.is_deleted == False ) ) provider = provider_result.scalar_one_or_none() if not provider: raise ValueError('Embedding 模型对应的提供商不存在或已禁用') cache_key = str(provider.id) if cache_key not in self._client_cache: import httpx from openai import AsyncOpenAI # 根据提供商类型确定 base_url if provider.provider_type == 'ollama': base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1' else: from ai_platform.providers.registry import ProviderRegistry provider_class = ProviderRegistry.get(provider.provider_type) default_base = getattr(provider_class, 'DEFAULT_API_BASE', 'https://api.openai.com/v1') if provider_class else 'https://api.openai.com/v1' base_url = provider.api_base or default_base self._client_cache[cache_key] = AsyncOpenAI( api_key=provider.api_key or 'ollama', base_url=base_url, timeout=httpx.Timeout(120.0, connect=30.0), max_retries=5, ) return self._client_cache[cache_key], model.model_name, model.context_window or 8191, provider.provider_type async def embed_text( self, model_id: str, text: str, dimensions: Optional[int] = None, ) -> List[float]: """ 将单个文本转换为向量 Args: model_id: Embedding 模型 ID text: 文本内容 dimensions: 向量维度(可选,部分模型支持) Returns: 向量列表 List[float] """ results = await self.embed_texts(model_id, [text], dimensions) return results[0] async def embed_texts( self, model_id: str, texts: List[str], dimensions: Optional[int] = None, ) -> List[List[float]]: """ 批量将文本转换为向量 Args: model_id: Embedding 模型 ID texts: 文本列表 dimensions: 向量维度(可选) Returns: 向量列表 List[List[float]] """ if not texts: return [] client, model_name, max_tokens, provider_type = await self._get_client_and_model(model_id) # 根据提供商类型确定批次大小 batch_size = PROVIDER_BATCH_SIZE.get(provider_type, DEFAULT_BATCH_SIZE) # 预处理:截断超长文本 processed_texts = [] for text in texts: if not text or not text.strip(): processed_texts.append(" ") else: # 粗略估算 token 数(中文约 1 字 = 1.5 token,英文约 4 字符 = 1 token) # 保守截断到 max_tokens * 2 个字符 max_chars = max_tokens * 2 if len(text) > max_chars: processed_texts.append(text[:max_chars]) else: processed_texts.append(text) # 分批处理 all_embeddings = [None] * len(processed_texts) total_batches = math.ceil(len(processed_texts) / batch_size) for batch_idx in range(total_batches): start = batch_idx * batch_size end = min(start + batch_size, len(processed_texts)) batch_texts = processed_texts[start:end] try: kwargs = { 'model': model_name, 'input': batch_texts, } # 仅对明确支持 dimensions 参数的模型传递该参数 if dimensions: model_lower = model_name.lower() # OpenAI text-embedding-3 系列原生支持任意 dimensions if 'text-embedding-3' in model_lower: kwargs['dimensions'] = dimensions # DashScope text-embedding-v3 只接受 [64,128,256,512,768,1024] elif 'text-embedding-v3' in model_lower and dimensions in (64, 128, 256, 512, 768, 1024): kwargs['dimensions'] = dimensions response = await client.embeddings.create(**kwargs) for item in response.data: all_embeddings[start + item.index] = item.embedding except Exception as e: logger.exception(f'Embedding 批次 {batch_idx + 1}/{total_batches} 失败: {e}') raise ValueError(f'Embedding 调用失败: {str(e)}') # 检查是否所有向量都已生成 for i, emb in enumerate(all_embeddings): if emb is None: raise ValueError(f'第 {i} 个文本的向量未生成') return all_embeddings async def get_embedding_dimensions(self, model_id: str) -> int: """ 获取模型的向量维度(通过嵌入一个测试文本来检测) Args: model_id: Embedding 模型 ID Returns: 向量维度 """ test_embedding = await self.embed_text(model_id, "test") return len(test_embedding)