Files
projectAIpopular/tests/test_sense_embedder.py
T
tzt f783878c02 feat(sense): T-G1 Embedder(int8 量化/降级 fail-closed + /v1/embeddings 透传)
- embedder.py:embed() 调 llama-server /v1/embeddings -> 对称 per-vector
  int8 量化(scale=127/max|v|,1024 维余弦扰动 ~1e-4 << 1e-2 验收);
  超时/连接/非200/空形状 -> EmbedderDown(D-G4 fail-closed);
  quantize_int8/cosine_int8 辅助(int8 直接点积,scale 正标量不改方向)
- routes.py 拆双路由组:/sense 前缀组 + /v1 无前缀组(/v1/embeddings
  OpenAI 兼容透传,embedder 不可用 503);build_sense_routers 返回列表,
  api.py 逐个 include;build_sense_router 保留向后兼容
- 测试 +6:量化余弦误差 200 组 <1e-2/范围/零向量/int8 余弦/
  正常量化/三种 fail-closed/透传端点(503+形状+400),全量 381 passed
- 待真机项:llama-server embedder 端点冒烟(需用户配置 embedder.base_url)
2026-09-05 13:33:58 +08:00

149 lines
4.8 KiB
Python
Raw 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.
"""Embedder 测试(T-G1):int8 量化余弦误差 / 降级 fail-closed / 透传端点。"""
import asyncio
import json
import math
import random
import httpx
import pytest
pytest.importorskip("fastapi")
from fastapi import FastAPI
from fastapi.testclient import TestClient
from gateway.sense.config import EmbedderCfg, build_sense_config
from gateway.sense.embedder import cosine_int8, embed, quantize_int8
from gateway.sense.errors import EmbedderDown
def test_quantize_roundtrip_cosine_error_small():
"""200 组 1024 维随机向量:量化前后余弦误差 < 1e-2(§11 验收)。"""
random.seed(42)
worst = 0.0
for _ in range(200):
v = [random.gauss(0, 1) for _ in range(1024)]
q = quantize_int8(v)
scale = (max(abs(x) for x in v) and 127.0 / max(abs(x) for x in v)) or 1.0
dv = [x / scale for x in q]
err = 1.0 - cosine_int8(q, quantize_int8(v)) # 自反性恒 1;此行检对称
# 真误差:原向量 vs 反量化向量
dot = sum(a * b for a, b in zip(v, dv))
na = math.sqrt(sum(a * a for a in v))
nb = math.sqrt(sum(b * b for b in dv))
err = 1.0 - dot / (na * nb)
worst = max(worst, abs(err))
assert worst < 1e-2
def test_quantize_range_and_zero_vector():
q = quantize_int8([3.0, -6.0, 0.0])
assert all(-127 <= x <= 127 for x in q)
assert quantize_int8([0.0] * 4) == [0, 0, 0, 0] # 零向量不崩
def test_cosine_int8_identical_and_orthogonal():
assert cosine_int8([3, 4], [6, 8]) == pytest.approx(1.0)
assert cosine_int8([1, 0], [0, 1]) == 0.0
assert cosine_int8([], []) == 0.0
def _cfg(base_url="http://emb"):
return EmbedderCfg(base_url=base_url, model="bge-m3", timeout_s=1.0, dim=8)
def test_embed_ok_returns_int8(monkeypatch):
vec = [0.5 * i for i in range(1, 9)]
body = {"data": [{"embedding": vec}]}
def handler(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/embeddings"
payload = json.loads(request.read())
assert payload["model"] == "bge-m3" and payload["input"] == "hello"
return httpx.Response(200, content=json.dumps(body).encode())
import gateway.sense.embedder as em
orig = em._client
em._client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
try:
q = asyncio_run(embed("hello", _cfg()))
assert len(q) == 8 and all(-127 <= x <= 127 for x in q)
finally:
em._client = orig
def test_embed_down_fail_closed(monkeypatch):
"""连接拒绝/非 200/超时 -> EmbedderDownD-G4 fail-closed)。"""
from gateway.sense import embedder as em
def handler_500(request):
return httpx.Response(500)
orig = em._client
em._client = httpx.AsyncClient(transport=httpx.MockTransport(handler_500))
try:
with pytest.raises(EmbedderDown):
asyncio_run(embed("x", _cfg()))
finally:
em._client = orig
def handler_timeout(request):
raise httpx.ConnectTimeout("超时")
em._client = httpx.AsyncClient(transport=httpx.MockTransport(handler_timeout))
try:
with pytest.raises(EmbedderDown):
asyncio_run(embed("x", _cfg()))
finally:
em._client = orig
def handler_bad_shape(request):
return httpx.Response(200, content=json.dumps({"data": [{"embedding": []}]}).encode())
em._client = httpx.AsyncClient(transport=httpx.MockTransport(handler_bad_shape))
try:
with pytest.raises(EmbedderDown):
asyncio_run(embed("x", _cfg()))
finally:
em._client = orig
def test_v1_embeddings_endpoint(tmp_path):
"""enabled 独立挂载:/v1/embeddings 透传(Mock embedder 注入)。"""
tc_app = FastAPI()
cfg = build_sense_config({"sense": {"enabled": True,
"db_path": str(tmp_path / "s.sqlite3")}})
from gateway.sense.routes import build_sense_routers
app = FastAPI()
for r in build_sense_routers(cfg):
app.include_router(r)
client = TestClient(app)
# embedder 不可用 -> 503 fail-closed
r = client.post("/v1/embeddings", json={"input": "hello"})
assert r.status_code == 503
# 注入可用 embeddermonkeypatch 内部函数)
import gateway.sense.routes as sr
orig_embed = sr.embed
async def fake_embed(text, c):
return [1, 2, 3]
sr.embed = fake_embed
try:
r2 = client.post("/v1/embeddings", json={"input": "hello"})
assert r2.status_code == 200
data = r2.json()
assert data["object"] == "list"
assert data["data"][0]["embedding"] == [1, 2, 3]
# 空 input -> 400
r3 = client.post("/v1/embeddings", json={"input": ""})
assert r3.status_code == 400
finally:
sr.embed = orig_embed
def asyncio_run(coro):
return asyncio.run(coro)