feat: T21-1 embedding 服务——Ollama bge-m3 批量 /api/embed(1024 维), 连接/模型/结构/维度四类失败一律 EmbeddingError 不静默降级(拍板口径), 空输入短路; settings 增 embed_dim/embed_timeout; 单测 10 例 mock HTTP 不依赖 Ollama, 386 绿
This commit is contained in:
@@ -27,6 +27,10 @@ class Settings(BaseSettings):
|
||||
|
||||
ollama_base_url: str = "http://127.0.0.1:11434"
|
||||
embed_model: str = "bge-m3"
|
||||
# 向量维度(必须与 Ollama bge-m3 输出一致;Milvus Collection 建集合时同源引用)
|
||||
embed_dim: int = 1024
|
||||
# Ollama embedding 单次请求超时(秒);本地推理首次加载模型可能较慢
|
||||
embed_timeout_seconds: float = 60.0
|
||||
|
||||
deepseek_api_key: str = ""
|
||||
deepseek_base_url: str = "https://api.deepseek.com"
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
"""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]
|
||||
@@ -0,0 +1,132 @@
|
||||
"""T21-1 embedding 服务单测:embed_texts / embed_text(全部 mock HTTP,不依赖 Ollama)。
|
||||
|
||||
覆盖四类:
|
||||
1. 成功语义:批量顺序一一对应、维度校验通过、空列表短路不发请求;
|
||||
2. 失败语义:连接拒绝 / 非 200(模型未拉取 404)/ 响应结构异常 → EmbeddingError;
|
||||
3. 维度防线:维度不符(8 维假数据)必须报错——防静默污染检索(拍板口径);
|
||||
4. 便捷方法:embed_text 等价 embed_texts([text])[0]。
|
||||
|
||||
httpx.MockTransport 注入:_client() 是 monkeypatch 点(settings 超时透传)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from app.config.settings import settings
|
||||
from app.service import embedding as emb
|
||||
|
||||
|
||||
def _mock_client(handler) -> httpx.Client:
|
||||
"""构造 mock 传输层的 httpx.Client(替换真实网络)。"""
|
||||
return httpx.Client(transport=httpx.MockTransport(handler), timeout=5)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _patch_client(monkeypatch):
|
||||
"""默认给一个「合法 1024 维响应」的 mock 客户端,各用例按需覆盖。"""
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
body = request.read()
|
||||
import json
|
||||
|
||||
payload = json.loads(body)
|
||||
vectors = [[0.1] * settings.embed_dim for _ in payload["input"]]
|
||||
return httpx.Response(200, json={"model": payload["model"], "embeddings": vectors})
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
|
||||
|
||||
class TestSuccess:
|
||||
"""成功语义。"""
|
||||
|
||||
def test_batch_order_preserved(self, monkeypatch):
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
import json
|
||||
|
||||
payload = json.loads(request.read())
|
||||
# 每条向量以首字符标记,断言顺序一一对应
|
||||
vectors = [[float(ord(t[0]))] + [0.0] * (settings.embed_dim - 1) for t in payload["input"]]
|
||||
return httpx.Response(200, json={"embeddings": vectors})
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
out = emb.embed_texts(["甲文本", "乙文本", "丙文本"])
|
||||
assert len(out) == 3
|
||||
assert out[0][0] == float(ord("甲"))
|
||||
assert out[2][0] == float(ord("丙"))
|
||||
|
||||
def test_empty_input_short_circuits(self, monkeypatch):
|
||||
# 空列表不发请求:handler 一旦被调用就 fail(返回 500 触发错误)
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
raise AssertionError("空输入不应发起 HTTP 请求")
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
assert emb.embed_texts([]) == []
|
||||
|
||||
def test_request_payload_uses_settings_model(self, monkeypatch):
|
||||
captured: dict = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
import json
|
||||
|
||||
captured.update(json.loads(request.read()))
|
||||
payload = json.loads(request.read())
|
||||
vectors = [[0.0] * settings.embed_dim for _ in payload["input"]]
|
||||
return httpx.Response(200, json={"embeddings": vectors})
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
emb.embed_texts(["测试"])
|
||||
assert captured["model"] == settings.embed_model
|
||||
assert captured["input"] == ["测试"]
|
||||
|
||||
|
||||
class TestFailure:
|
||||
"""失败语义:一律 EmbeddingError,不静默降级。"""
|
||||
|
||||
def test_connection_refused(self, monkeypatch):
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
raise httpx.ConnectError("connection refused")
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
with pytest.raises(emb.EmbeddingError, match="连接失败"):
|
||||
emb.embed_texts(["任意"])
|
||||
|
||||
def test_model_not_found_404(self, monkeypatch):
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(404, json={"error": "model 'bge-m3' not found"})
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
with pytest.raises(emb.EmbeddingError, match="HTTP 404"):
|
||||
emb.embed_texts(["任意"])
|
||||
|
||||
def test_malformed_response(self, monkeypatch):
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
# embeddings 缺失 / 数量不匹配均属结构异常
|
||||
return httpx.Response(200, json={"embeddings": [[0.0] * settings.embed_dim]})
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
with pytest.raises(emb.EmbeddingError, match="数量与输入不匹配"):
|
||||
emb.embed_texts(["一", "二"])
|
||||
|
||||
|
||||
class TestDimensionGuard:
|
||||
"""维度防线(拍板:维度不符必须报错)。"""
|
||||
|
||||
def test_wrong_dimension_rejected(self, monkeypatch):
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
import json
|
||||
|
||||
payload = json.loads(request.read())
|
||||
vectors = [[0.1] * 8 for _ in payload["input"]] # 假 8 维
|
||||
return httpx.Response(200, json={"embeddings": vectors})
|
||||
|
||||
monkeypatch.setattr(emb, "_client", lambda: _mock_client(handler))
|
||||
with pytest.raises(emb.EmbeddingError, match="维度不符"):
|
||||
emb.embed_texts(["任意"])
|
||||
|
||||
|
||||
class TestConvenience:
|
||||
def test_embed_text_single(self):
|
||||
out = emb.embed_text("单条")
|
||||
assert len(out) == settings.embed_dim
|
||||
Reference in New Issue
Block a user