Files
kplam 51dc8dc4f6 Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
2026-07-29 21:24:40 +08:00

100 lines
3.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Dense 向量嵌入服务
提供统一的 EmbeddingService 接口,支持两种 provider
- openaiOpenAI 兼容 APIAsyncOpenAI
- 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,
)