Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
This commit is contained in:
@@ -0,0 +1,99 @@
|
||||
"""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,
|
||||
)
|
||||
Reference in New Issue
Block a user