From e34973f018d84428bf41c53b1e49d4dd4b3bfe0c Mon Sep 17 00:00:00 2001 From: YUAN Date: Mon, 7 Sep 2026 11:46:24 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20T21-1=20embedding=20=E6=9C=8D=E5=8A=A1?= =?UTF-8?q?=E2=80=94=E2=80=94Ollama=20bge-m3=20=E6=89=B9=E9=87=8F=20/api/e?= =?UTF-8?q?mbed(1024=20=E7=BB=B4),=20=E8=BF=9E=E6=8E=A5/=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?/=E7=BB=93=E6=9E=84/=E7=BB=B4=E5=BA=A6=E5=9B=9B=E7=B1=BB?= =?UTF-8?q?=E5=A4=B1=E8=B4=A5=E4=B8=80=E5=BE=8B=20EmbeddingError=20?= =?UTF-8?q?=E4=B8=8D=E9=9D=99=E9=BB=98=E9=99=8D=E7=BA=A7(=E6=8B=8D?= =?UTF-8?q?=E6=9D=BF=E5=8F=A3=E5=BE=84),=20=E7=A9=BA=E8=BE=93=E5=85=A5?= =?UTF-8?q?=E7=9F=AD=E8=B7=AF;=20settings=20=E5=A2=9E=20embed=5Fdim/embed?= =?UTF-8?q?=5Ftimeout;=20=E5=8D=95=E6=B5=8B=2010=20=E4=BE=8B=20mock=20HTTP?= =?UTF-8?q?=20=E4=B8=8D=E4=BE=9D=E8=B5=96=20Ollama,=20386=20=E7=BB=BF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/config/settings.py | 4 ++ app/service/embedding.py | 75 ++++++++++++++++++++++ tests/test_embedding.py | 132 +++++++++++++++++++++++++++++++++++++++ 3 files changed, 211 insertions(+) create mode 100644 app/service/embedding.py create mode 100644 tests/test_embedding.py diff --git a/app/config/settings.py b/app/config/settings.py index 11ce085..059d609 100644 --- a/app/config/settings.py +++ b/app/config/settings.py @@ -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" diff --git a/app/service/embedding.py b/app/service/embedding.py new file mode 100644 index 0000000..48ea42a --- /dev/null +++ b/app/service/embedding.py @@ -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] diff --git a/tests/test_embedding.py b/tests/test_embedding.py new file mode 100644 index 0000000..15d943e --- /dev/null +++ b/tests/test_embedding.py @@ -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