76 lines
3.2 KiB
Python
76 lines
3.2 KiB
Python
"""Embedding 服务(T21-1 · FLOW §3「embedding_tool(Ollama bge-m3)」)。
|
||||
|
|
|
|||
|
|
技术选型(01-技术栈与版本.md §4):Ollama + bge-m3,**1024 维**,文档向量
|
|||
|
|
不出内网。调用 Ollama 批量接口 ``POST /api/embed``(一次请求携带多段文本,
|
|||
|
|
比逐条 /api/embeddings 少一个数量级的往返开销)。
|
|||
|
|
|
|||
|
|
失败口径(拍板 2026-09-07):Ollama 连不上 / 模型缺失 / 响应结构异常 /
|
|||
|
|
维度不符一律抛 EmbeddingError——**禁止静默返回零向量或截断降级**。向量
|
|||
|
|
错误会静默污染检索结果(查得慢、查得偏都难归因),明确失败比带病成功
|
|||
|
|
更可运维;调用方(build_kb 入库 / rag_service 检索)按明确错误处理。
|
|||
|
|
|
|||
|
|
测试:单测经 httpx.MockTransport 模拟 Ollama 响应,不依赖 Ollama 进程;
|
|||
|
|
真库联调归 scripts/kb/build_kb.py 与收尾验证。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import httpx
|
|||
|
|
|
|||
|
|
from app.config.settings import settings
|
|||
|
|
|
|||
|
|
|
|||
|
|
class EmbeddingError(Exception):
|
|||
|
|
"""向量化失败(连接/模型/响应/维度)——调用方明确处理,不降级假数据。"""
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _client() -> httpx.Client:
|
|||
|
|
"""HTTP 客户端工厂(测试 monkeypatch 点;超时取 settings)。"""
|
|||
|
|
return httpx.Client(timeout=settings.embed_timeout_seconds)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def embed_texts(texts: list[str]) -> list[list[float]]:
|
|||
|
|
"""批量文本 → 向量(顺序与入参一一对应)。
|
|||
|
|
|
|||
|
|
空列表直接返回 [](不发请求,防 Ollama 对空 input 报错);单条空串
|
|||
|
|
交由 Ollama 侧语义处理(bge-m3 对空串也能出向量,调用方自行决定是否
|
|||
|
|
过滤——切块层保证不为空,此处不做第二套校验)。
|
|||
|
|
"""
|
|||
|
|
if not texts:
|
|||
|
|
return []
|
|||
|
|
payload = {"model": settings.embed_model, "input": list(texts)}
|
|||
|
|
try:
|
|||
|
|
with _client() as client:
|
|||
|
|
resp = client.post(f"{settings.ollama_base_url}/api/embed", json=payload)
|
|||
|
|
except httpx.HTTPError as exc:
|
|||
|
|
raise EmbeddingError(
|
|||
|
|
f"Ollama 连接失败({settings.ollama_base_url}):{exc}"
|
|||
|
|
) from exc
|
|||
|
|
if resp.status_code != 200:
|
|||
|
|
# 模型未拉取时 Ollama 返回 404 + {"error": "model '...' not found"}
|
|||
|
|
detail = (resp.text or "")[:200]
|
|||
|
|
raise EmbeddingError(
|
|||
|
|
f"Ollama embedding 请求失败 HTTP {resp.status_code}(模型 "
|
|||
|
|
f"{settings.embed_model} 是否已 pull?):{detail}"
|
|||
|
|
)
|
|||
|
|
body = resp.json()
|
|||
|
|
embeddings = (body or {}).get("embeddings")
|
|||
|
|
if not isinstance(embeddings, list) or len(embeddings) != len(texts):
|
|||
|
|
raise EmbeddingError(
|
|||
|
|
f"Ollama 响应结构异常:embeddings 数量与输入不匹配"
|
|||
|
|
f"(输入 {len(texts)} 条)"
|
|||
|
|
)
|
|||
|
|
for vec in embeddings:
|
|||
|
|
if not isinstance(vec, list) or len(vec) != settings.embed_dim:
|
|||
|
|
raise EmbeddingError(
|
|||
|
|
f"向量维度不符:期望 {settings.embed_dim},实际 "
|
|||
|
|
f"{len(vec) if isinstance(vec, list) else '非数组'}"
|
|||
|
|
f"——检查 embed_dim 配置与 {settings.embed_model} 输出是否一致"
|
|||
|
|
)
|
|||
|
|
return embeddings
|
|||
|
|
|
|||
|
|
|
|||
|
|
def embed_text(text: str) -> list[float]:
|
|||
|
|
"""单文本便捷方法(= embed_texts([text])[0])。"""
|
|||
|
|
return embed_texts([text])[0]
|