158 lines
5.1 KiB
Python
158 lines
5.1 KiB
Python
"""Embedding 类型/端口/transport 测试(M2 设计 §7;T8)。
|
|
|
|
蓝本审计: GovDoc retrieval/embedding.py(分批/index 排序/维度校验)与
|
|
VT adapters/embedding.py(归一化);库裁决见设计 §7.3 表。
|
|
"""
|
|
|
|
import dataclasses
|
|
import json
|
|
|
|
import httpx
|
|
import pytest
|
|
|
|
from polygateway.errors import (
|
|
RequestRejectedError,
|
|
ResultInvalidError,
|
|
SourceDeadError,
|
|
TransientError,
|
|
)
|
|
from polygateway.ports import EmbeddingTransport
|
|
from polygateway.transports.openai_compat import OpenAICompatTransport
|
|
from polygateway.types import EmbeddingResponse, EmbeddingTransportResult, SourceConfig
|
|
|
|
|
|
def _src(**overrides):
|
|
base = {
|
|
"name": "e1",
|
|
"provider": "openai",
|
|
"base_url": "https://gw.example/v1",
|
|
"api_key": "sk",
|
|
"model": "embed-1",
|
|
"timeout_s": 10.0,
|
|
"est_tokens": 7,
|
|
}
|
|
base.update(overrides)
|
|
return SourceConfig(**base)
|
|
|
|
|
|
class TestTypes:
|
|
def test_embedding_response_frozen_with_defaults(self):
|
|
resp = EmbeddingResponse(
|
|
vectors=[[0.1, 0.2]],
|
|
dim=2,
|
|
model="m",
|
|
provider="p",
|
|
prompt_tokens=3,
|
|
usage_source="measured",
|
|
latency_ms=10,
|
|
call_id="c",
|
|
source_name="e1",
|
|
)
|
|
assert resp.cost is None
|
|
with pytest.raises(dataclasses.FrozenInstanceError):
|
|
resp.dim = 3
|
|
|
|
def test_transport_result_frozen(self):
|
|
r = EmbeddingTransportResult(
|
|
vectors=[[1.0]], dim=1, prompt_tokens=1, usage_source="measured", raw={}
|
|
)
|
|
with pytest.raises(dataclasses.FrozenInstanceError):
|
|
r.dim = 2
|
|
|
|
|
|
class _DummyEmbedTransport:
|
|
async def embed(self, *, texts, source, call_id):
|
|
raise NotImplementedError
|
|
|
|
|
|
def test_embedding_transport_protocol_runtime_checkable():
|
|
assert isinstance(_DummyEmbedTransport(), EmbeddingTransport)
|
|
assert isinstance(OpenAICompatTransport(), EmbeddingTransport)
|
|
|
|
|
|
def _transport_with(handler):
|
|
return OpenAICompatTransport(
|
|
client_factory=lambda source: httpx.AsyncClient(transport=httpx.MockTransport(handler))
|
|
)
|
|
|
|
|
|
def _ok_body(vectors, *, usage=None, shuffle=False):
|
|
data = [{"index": i, "embedding": v} for i, v in enumerate(vectors)]
|
|
if shuffle:
|
|
data = list(reversed(data))
|
|
body = {"data": data}
|
|
if usage is not None:
|
|
body["usage"] = usage
|
|
return body
|
|
|
|
|
|
class TestEmbedTransport:
|
|
async def test_sorts_by_index_and_measures_usage(self):
|
|
def handler(request):
|
|
assert request.url.path.endswith("/embeddings")
|
|
payload = json.loads(request.content)
|
|
assert payload == {"model": "embed-1", "input": ["a", "b"]}
|
|
return httpx.Response(
|
|
200,
|
|
json=_ok_body(
|
|
[[1.0, 0.0], [0.0, 1.0]], usage={"prompt_tokens": 5}, shuffle=True
|
|
),
|
|
)
|
|
|
|
result = await _transport_with(handler).embed(
|
|
texts=["a", "b"], source=_src(), call_id="c"
|
|
)
|
|
assert result.vectors == [[1.0, 0.0], [0.0, 1.0]] # 乱序响应按 index 重排
|
|
assert result.dim == 2
|
|
assert result.prompt_tokens == 5 and result.usage_source == "measured"
|
|
|
|
async def test_missing_usage_falls_back_estimated(self):
|
|
def handler(request):
|
|
return httpx.Response(200, json=_ok_body([[1.0]]))
|
|
|
|
result = await _transport_with(handler).embed(texts=["a"], source=_src(), call_id="c")
|
|
assert result.prompt_tokens == 7 and result.usage_source == "estimated" # est_tokens
|
|
|
|
@pytest.mark.parametrize(
|
|
("status", "exc_type"),
|
|
[(401, SourceDeadError), (400, RequestRejectedError), (500, TransientError)],
|
|
)
|
|
async def test_http_errors_translate(self, status, exc_type):
|
|
def handler(request):
|
|
return httpx.Response(status, text="boom")
|
|
|
|
with pytest.raises(exc_type):
|
|
await _transport_with(handler).embed(texts=["a"], source=_src(), call_id="c")
|
|
|
|
async def test_network_error_is_transient(self):
|
|
def handler(request):
|
|
raise httpx.ConnectError("refused")
|
|
|
|
with pytest.raises(TransientError):
|
|
await _transport_with(handler).embed(texts=["a"], source=_src(), call_id="c")
|
|
|
|
@pytest.mark.parametrize(
|
|
"body",
|
|
[
|
|
{"data": []}, # 空 data
|
|
{"data": [{"index": 0, "embedding": [1.0]}]}, # 数量与输入不符(输入 2 条)
|
|
{
|
|
"data": [
|
|
{"index": 0, "embedding": [1.0, 2.0]},
|
|
{"index": 1, "embedding": [1.0]}, # 维度不一致
|
|
]
|
|
},
|
|
{"nope": True}, # 缺 data
|
|
],
|
|
)
|
|
async def test_malformed_payload_is_result_invalid(self, body):
|
|
def handler(request):
|
|
return httpx.Response(200, json=body)
|
|
|
|
with pytest.raises(ResultInvalidError):
|
|
await _transport_with(handler).embed(texts=["a", "b"], source=_src(), call_id="c")
|
|
|
|
async def test_empty_texts_rejected(self):
|
|
with pytest.raises(ValueError):
|
|
await _transport_with(lambda r: None).embed(texts=[], source=_src(), call_id="c")
|