- 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)
149 lines
4.8 KiB
Python
149 lines
4.8 KiB
Python
"""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/超时 -> EmbedderDown(D-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
|
||
|
||
# 注入可用 embedder(monkeypatch 内部函数)
|
||
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)
|