51dc8dc4f6
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
100 lines
3.2 KiB
Python
100 lines
3.2 KiB
Python
"""Dense 向量嵌入服务
|
||
|
||
提供统一的 EmbeddingService 接口,支持两种 provider:
|
||
- openai:OpenAI 兼容 API(AsyncOpenAI)
|
||
- local:本地 Ollama /api/embed 接口
|
||
"""
|
||
|
||
from typing import Protocol, runtime_checkable
|
||
|
||
import httpx
|
||
import structlog
|
||
from openai import AsyncOpenAI
|
||
|
||
from app.config import settings
|
||
|
||
logger = structlog.get_logger()
|
||
|
||
|
||
@runtime_checkable
|
||
class EmbeddingService(Protocol):
|
||
"""统一嵌入服务接口"""
|
||
|
||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||
"""批量生成文本向量
|
||
|
||
Args:
|
||
texts: 待嵌入文本列表,空列表时直接返回空列表
|
||
|
||
Returns:
|
||
与输入等长的向量列表
|
||
"""
|
||
...
|
||
|
||
|
||
def _check_dimension(vectors: list[list[float]], provider: str) -> None:
|
||
"""返回维度与配置不一致时告警(不抛错),每次调用最多提示一次"""
|
||
for vec in vectors:
|
||
if len(vec) != settings.embedding_dimension:
|
||
logger.warning(
|
||
"嵌入向量维度与配置不一致",
|
||
provider=provider,
|
||
actual=len(vec),
|
||
expected=settings.embedding_dimension,
|
||
)
|
||
return
|
||
|
||
|
||
class OpenAIEmbeddingService:
|
||
"""OpenAI 兼容 API 嵌入服务"""
|
||
|
||
def __init__(self, api_key: str, base_url: str, model: str) -> None:
|
||
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url)
|
||
self._model = model
|
||
|
||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||
if not texts:
|
||
return []
|
||
resp = await self._client.embeddings.create(model=self._model, input=texts)
|
||
vectors = [item.embedding for item in resp.data]
|
||
_check_dimension(vectors, provider="openai")
|
||
logger.debug("OpenAI 嵌入完成", model=self._model, count=len(vectors))
|
||
return vectors
|
||
|
||
|
||
class LocalEmbeddingService:
|
||
"""本地 Ollama 嵌入服务(/api/embed)"""
|
||
|
||
def __init__(self, base_url: str, model: str, timeout: float = 60.0) -> None:
|
||
self.base_url = base_url.rstrip("/")
|
||
self.model = model
|
||
self.timeout = timeout
|
||
|
||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||
if not texts:
|
||
return []
|
||
url = f"{self.base_url}/api/embed"
|
||
payload = {"model": self.model, "input": texts}
|
||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||
resp = await client.post(url, json=payload)
|
||
resp.raise_for_status()
|
||
data = resp.json()
|
||
vectors: list[list[float]] = data.get("embeddings", [])
|
||
_check_dimension(vectors, provider="local")
|
||
logger.debug("Ollama 嵌入完成", model=self.model, count=len(vectors))
|
||
return vectors
|
||
|
||
|
||
def create_embedding_service() -> EmbeddingService:
|
||
"""按 settings.embedding_provider 创建嵌入服务实例"""
|
||
if settings.embedding_provider == "local":
|
||
return LocalEmbeddingService(
|
||
base_url=settings.ollama_base_url,
|
||
model=settings.ollama_embedding_model,
|
||
)
|
||
return OpenAIEmbeddingService(
|
||
api_key=settings.openai_api_key,
|
||
base_url=settings.openai_base_url,
|
||
model=settings.embedding_model,
|
||
)
|