feat: add governed embedding client with batching
This commit is contained in:
@@ -6,7 +6,8 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from polygateway.client import GatewayClient, gather_bounded
|
from polygateway.client import GatewayClient, gather_bounded
|
||||||
from polygateway.config import GatewaySettings
|
from polygateway.config import EmbeddingSettings, GatewaySettings
|
||||||
|
from polygateway.embedding import EmbeddingClient
|
||||||
from polygateway.errors import (
|
from polygateway.errors import (
|
||||||
AllSourcesExhausted,
|
AllSourcesExhausted,
|
||||||
CircuitOpenError,
|
CircuitOpenError,
|
||||||
@@ -18,8 +19,9 @@ from polygateway.errors import (
|
|||||||
SourceDeadError,
|
SourceDeadError,
|
||||||
TransientError,
|
TransientError,
|
||||||
)
|
)
|
||||||
|
from polygateway.pricing import ModelPrice, PricingTable
|
||||||
from polygateway.providers import DEFAULT_PROFILES, ProviderProfile, register_provider
|
from polygateway.providers import DEFAULT_PROFILES, ProviderProfile, register_provider
|
||||||
from polygateway.types import LLMResponse, SourceConfig
|
from polygateway.types import EmbeddingResponse, LLMResponse, SourceConfig
|
||||||
|
|
||||||
__version__ = "0.1.0"
|
__version__ = "0.1.0"
|
||||||
|
|
||||||
@@ -27,12 +29,17 @@ __all__ = [
|
|||||||
"DEFAULT_PROFILES",
|
"DEFAULT_PROFILES",
|
||||||
"AllSourcesExhausted",
|
"AllSourcesExhausted",
|
||||||
"CircuitOpenError",
|
"CircuitOpenError",
|
||||||
|
"EmbeddingClient",
|
||||||
|
"EmbeddingResponse",
|
||||||
|
"EmbeddingSettings",
|
||||||
"GatewayClient",
|
"GatewayClient",
|
||||||
"GatewaySettings",
|
"GatewaySettings",
|
||||||
"GatewayUnavailableError",
|
"GatewayUnavailableError",
|
||||||
"GovernanceBackendError",
|
"GovernanceBackendError",
|
||||||
"LLMResponse",
|
"LLMResponse",
|
||||||
|
"ModelPrice",
|
||||||
"PolyGatewayError",
|
"PolyGatewayError",
|
||||||
|
"PricingTable",
|
||||||
"ProviderProfile",
|
"ProviderProfile",
|
||||||
"RequestRejectedError",
|
"RequestRejectedError",
|
||||||
"ResultInvalidError",
|
"ResultInvalidError",
|
||||||
|
|||||||
@@ -341,3 +341,48 @@ def _guard_stall(settings: GatewaySettings) -> None:
|
|||||||
def _load_lease_ttl(env: Mapping[str, str]) -> float:
|
def _load_lease_ttl(env: Mapping[str, str]) -> float:
|
||||||
found = _first(env, "PGW_LEASE_TTL_S")
|
found = _first(env, "PGW_LEASE_TTL_S")
|
||||||
return float(_cast(found[1], "float", found[0])) if found else _DEFAULT_LEASE_TTL_S
|
return float(_cast(found[1], "float", found[0])) if found else _DEFAULT_LEASE_TTL_S
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class EmbeddingSettings:
|
||||||
|
"""Embedding scope 装配配置(M2 §7): 复用 GatewaySettings + embedding 专用键。
|
||||||
|
|
||||||
|
专用键不进 GatewaySettings(LLM scope 不受影响): `{SCOPE}__BATCH_SIZE`
|
||||||
|
必填(分批是行为关键,不设默认)、`{SCOPE}__NORMALIZE`/`{SCOPE}__EXPECTED_DIM`
|
||||||
|
可选。cache/structured 键对 embedding 无意义,装配时忽略。
|
||||||
|
"""
|
||||||
|
|
||||||
|
gateway: GatewaySettings
|
||||||
|
batch_size: int
|
||||||
|
normalize: bool = False
|
||||||
|
expected_dim: int | None = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_env(
|
||||||
|
cls,
|
||||||
|
scope: str = "EMBED",
|
||||||
|
env: Mapping[str, str] | None = None,
|
||||||
|
*,
|
||||||
|
env_file: str = ".env",
|
||||||
|
) -> EmbeddingSettings:
|
||||||
|
if env is None:
|
||||||
|
env = {
|
||||||
|
k: v for k, v in {**dotenv_values(env_file), **os.environ}.items() if v is not None
|
||||||
|
}
|
||||||
|
scope_u = scope.upper()
|
||||||
|
gateway = GatewaySettings.from_env(scope_u, env=env)
|
||||||
|
key, raw = _require(env, f"{scope_u}__BATCH_SIZE")
|
||||||
|
batch_size = int(_cast(raw, "int", key))
|
||||||
|
if batch_size < 1:
|
||||||
|
raise ValueError(f"{scope_u}__BATCH_SIZE 必须 ≥ 1")
|
||||||
|
norm = _first(env, f"{scope_u}__NORMALIZE")
|
||||||
|
dim = _first(env, f"{scope_u}__EXPECTED_DIM")
|
||||||
|
expected_dim = int(_cast(dim[1], "int", dim[0])) if dim else None
|
||||||
|
if expected_dim is not None and expected_dim < 1:
|
||||||
|
raise ValueError(f"{scope_u}__EXPECTED_DIM 必须 ≥ 1")
|
||||||
|
return cls(
|
||||||
|
gateway=gateway,
|
||||||
|
batch_size=batch_size,
|
||||||
|
normalize=bool(_cast(norm[1], "bool", norm[0])) if norm else False,
|
||||||
|
expected_dim=expected_dim,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,492 @@
|
|||||||
|
"""EmbeddingClient: 治理化 embedding 调用(M2 设计 §7,方案 G2)。
|
||||||
|
|
||||||
|
独立精简治理循环,**复用**库的算法件: `RateLimiter`/`ProviderGate` 端口与
|
||||||
|
两种后端、错误四分类、`backoff_delay` 退避公式、`SourceCooldownMemo`、
|
||||||
|
`TelemetryEmitter`(遥测单一 helper 铁律)。选源/等待循环与 RetryMW 同构
|
||||||
|
——这是设计 §7.1 已声明的有限重复(chat 循环含流式/结构化/缓存分支,
|
||||||
|
强行合一才是复制);行为口径(stall 双条件、记账降级、取消穿透)与 chat
|
||||||
|
完全一致。
|
||||||
|
|
||||||
|
对参考实现的已声明裁决(设计 §7.3): async httpx;返回 list[list[float]]
|
||||||
|
(核心不依赖 numpy);normalize 开关(VT 语义,防除零 max(norm,1e-12));
|
||||||
|
必填 batch_size 批间串行;GovDoc 自研退避与 on_usage 回调、VT 同步接口
|
||||||
|
均**有意放弃**。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import math
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from polygateway.config import EmbeddingSettings
|
||||||
|
from polygateway.errors import (
|
||||||
|
AllSourcesExhausted,
|
||||||
|
CircuitOpenError,
|
||||||
|
GovernanceBackendError,
|
||||||
|
PolyGatewayError,
|
||||||
|
RequestRejectedError,
|
||||||
|
ResultInvalidError,
|
||||||
|
SourceDeadError,
|
||||||
|
TransientError,
|
||||||
|
)
|
||||||
|
from polygateway.middleware.breaker import BreakerGate
|
||||||
|
from polygateway.middleware.ratelimit import QuotaGate
|
||||||
|
from polygateway.middleware.retry import _failure_reason, backoff_delay
|
||||||
|
from polygateway.middleware.telemetry import TelemetryEmitter
|
||||||
|
from polygateway.sources import SourceCooldownMemo
|
||||||
|
from polygateway.types import ChatRequest, EmbeddingResponse, LLMResponse
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Awaitable, Callable, Mapping
|
||||||
|
|
||||||
|
from polygateway.ports import (
|
||||||
|
EmbeddingTransport,
|
||||||
|
GateDecision,
|
||||||
|
Permit,
|
||||||
|
ProviderGate,
|
||||||
|
RateLimiter,
|
||||||
|
SourceSelector,
|
||||||
|
TelemetryRecorder,
|
||||||
|
)
|
||||||
|
from polygateway.pricing import PricingTable
|
||||||
|
from polygateway.types import (
|
||||||
|
BackpressurePolicy,
|
||||||
|
EmbeddingTransportResult,
|
||||||
|
RetryPolicy,
|
||||||
|
SourceConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
_TELEMETRY_TEXT_CAP = 200 # 遥测行每条 text 截断长度(原文不整段入库,VT R12)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _FailedBatch:
|
||||||
|
exc: PolyGatewayError
|
||||||
|
immediate: bool
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _BatchOutcome:
|
||||||
|
result: EmbeddingTransportResult
|
||||||
|
source: SourceConfig
|
||||||
|
call_id: str
|
||||||
|
latency_ms: int
|
||||||
|
|
||||||
|
|
||||||
|
class EmbeddingClient:
|
||||||
|
"""治理化 embedding 入口;与 GatewayClient 共享后端实例即共享全局闸。"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
scope: str,
|
||||||
|
sources: list[SourceConfig],
|
||||||
|
selector: SourceSelector,
|
||||||
|
limiter: RateLimiter,
|
||||||
|
breaker: ProviderGate,
|
||||||
|
transport: EmbeddingTransport,
|
||||||
|
retry: RetryPolicy,
|
||||||
|
backpressure: BackpressurePolicy,
|
||||||
|
quota_full: str = "wait",
|
||||||
|
telemetry: TelemetryRecorder | None = None,
|
||||||
|
pricing: PricingTable | None = None,
|
||||||
|
batch_size: int,
|
||||||
|
normalize: bool = False,
|
||||||
|
expected_dim: int | None = None,
|
||||||
|
now: Callable[[], float] = time.monotonic,
|
||||||
|
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
|
||||||
|
rng: Callable[[], float] = random.random,
|
||||||
|
) -> None:
|
||||||
|
if batch_size < 1:
|
||||||
|
raise ValueError("batch_size 必须 ≥ 1")
|
||||||
|
if quota_full not in ("wait", "fail_fast"):
|
||||||
|
raise ValueError(f"quota_full 必须是 wait|fail_fast: {quota_full!r}")
|
||||||
|
if expected_dim is not None and expected_dim < 1:
|
||||||
|
raise ValueError("expected_dim 必须 ≥ 1")
|
||||||
|
self._scope = scope
|
||||||
|
self._sources = list(sources)
|
||||||
|
self._selector = selector
|
||||||
|
self._quota = QuotaGate(limiter)
|
||||||
|
self._breaker = BreakerGate(breaker)
|
||||||
|
self._transport = transport
|
||||||
|
self._retry = retry
|
||||||
|
self._bp = backpressure
|
||||||
|
self._quota_full = quota_full
|
||||||
|
self._emitter = TelemetryEmitter(telemetry, pricing=pricing) if telemetry else None
|
||||||
|
self._telemetry = telemetry
|
||||||
|
self._pricing = pricing
|
||||||
|
self._batch_size = batch_size
|
||||||
|
self._normalize = normalize
|
||||||
|
self._expected_dim = expected_dim
|
||||||
|
self._memo = SourceCooldownMemo(now=now)
|
||||||
|
self._now = now
|
||||||
|
self._sleep = sleep
|
||||||
|
self._rng = rng
|
||||||
|
self._closed = False
|
||||||
|
|
||||||
|
async def embed(
|
||||||
|
self,
|
||||||
|
texts: list[str],
|
||||||
|
*,
|
||||||
|
session_id: str | None = None,
|
||||||
|
parent_call_id: str | None = None,
|
||||||
|
) -> EmbeddingResponse:
|
||||||
|
"""一次治理 embedding 调用: 按 batch_size 切批,批间串行,全批合并返回。"""
|
||||||
|
if not isinstance(texts, list) or any(not isinstance(t, str) for t in texts):
|
||||||
|
raise TypeError("texts 必须是 list[str](显式优于隐式,不收单条 str)")
|
||||||
|
if not texts:
|
||||||
|
return EmbeddingResponse(
|
||||||
|
vectors=[],
|
||||||
|
dim=0,
|
||||||
|
model="",
|
||||||
|
provider="",
|
||||||
|
prompt_tokens=0,
|
||||||
|
usage_source="measured",
|
||||||
|
latency_ms=0,
|
||||||
|
call_id=str(uuid.uuid4()),
|
||||||
|
source_name="",
|
||||||
|
)
|
||||||
|
if not self._sources:
|
||||||
|
raise AllSourcesExhausted(scope=self._scope, reason="no_sources", retry_after_s=0.0)
|
||||||
|
outcomes = []
|
||||||
|
for start in range(0, len(texts), self._batch_size):
|
||||||
|
outcomes.append(
|
||||||
|
await self._embed_batch(texts[start : start + self._batch_size], session_id, parent_call_id)
|
||||||
|
)
|
||||||
|
return self._merge(outcomes)
|
||||||
|
|
||||||
|
# —— 治理循环(与 RetryMW 同构;设计 §7.1 已声明的有限重复)——
|
||||||
|
|
||||||
|
async def _embed_batch(
|
||||||
|
self, batch: list[str], session_id: str | None, parent_call_id: str | None
|
||||||
|
) -> _BatchOutcome:
|
||||||
|
fails = 0
|
||||||
|
reasons: dict[str, str] = {}
|
||||||
|
entered_at = self._now()
|
||||||
|
while True:
|
||||||
|
picked, gate_rejections = await self._pick_runnable(reasons)
|
||||||
|
if picked is None:
|
||||||
|
await self._on_no_runnable(gate_rejections, reasons, entered_at)
|
||||||
|
continue
|
||||||
|
outcome = await self._attempt(batch, *picked, reasons, session_id, parent_call_id)
|
||||||
|
if isinstance(outcome, _BatchOutcome):
|
||||||
|
return outcome
|
||||||
|
fails += 1
|
||||||
|
if fails >= self._retry.max_attempts:
|
||||||
|
raise AllSourcesExhausted(
|
||||||
|
scope=self._scope,
|
||||||
|
reason="retry_exhausted",
|
||||||
|
retry_after_s=self._retry.backoff_base_s,
|
||||||
|
per_source_reasons=reasons,
|
||||||
|
) from outcome.exc
|
||||||
|
if not outcome.immediate:
|
||||||
|
await self._sleep(backoff_delay(self._retry, fails, outcome.exc, self._rng))
|
||||||
|
|
||||||
|
async def _pick_runnable(
|
||||||
|
self, reasons: dict[str, str]
|
||||||
|
) -> tuple[tuple[SourceConfig, Permit, GateDecision] | None, int]:
|
||||||
|
stats = {s.name: await self._quota.stats(s) for s in self._sources}
|
||||||
|
gate_rejections = 0
|
||||||
|
for cand in self._selector.order(self._sources, stats):
|
||||||
|
if self._memo.active(cand.name):
|
||||||
|
gate_rejections += 1
|
||||||
|
reasons[cand.name] = "cooldown"
|
||||||
|
continue
|
||||||
|
permit = await self._quota.try_acquire(cand)
|
||||||
|
if permit is None:
|
||||||
|
reasons.setdefault(cand.name, "rate_limited")
|
||||||
|
continue
|
||||||
|
entry = None
|
||||||
|
try:
|
||||||
|
entry = await self._breaker.try_enter(cand, uuid.uuid4().hex)
|
||||||
|
finally:
|
||||||
|
if entry is None:
|
||||||
|
await self._settle_and_release(permit, 0)
|
||||||
|
if entry.allowed:
|
||||||
|
return (cand, permit, entry), gate_rejections
|
||||||
|
gate_rejections += 1
|
||||||
|
reasons[cand.name] = "circuit_open"
|
||||||
|
self._memo.set_until(cand.name, self._now() + entry.retry_after_s)
|
||||||
|
await self._settle_and_release(permit, 0)
|
||||||
|
return None, gate_rejections
|
||||||
|
|
||||||
|
async def _on_no_runnable(
|
||||||
|
self, gate_rejections: int, reasons: dict[str, str], entered_at: float
|
||||||
|
) -> None:
|
||||||
|
if gate_rejections == len(self._sources):
|
||||||
|
names = tuple(s.name for s in self._sources)
|
||||||
|
raise CircuitOpenError(
|
||||||
|
scope=self._scope,
|
||||||
|
retry_after_s=await self._breaker.retry_after_s(names),
|
||||||
|
per_source_reasons=reasons,
|
||||||
|
)
|
||||||
|
if self._quota_full == "fail_fast":
|
||||||
|
raise AllSourcesExhausted(
|
||||||
|
scope=self._scope,
|
||||||
|
reason="quota_exhausted",
|
||||||
|
retry_after_s=self._bp.poll_interval_s,
|
||||||
|
per_source_reasons=reasons,
|
||||||
|
)
|
||||||
|
stall = self._bp.stall_window_s
|
||||||
|
if self._now() - entered_at > stall and await self._quota.progress_age_s() > stall:
|
||||||
|
names = tuple(s.name for s in self._sources)
|
||||||
|
raise AllSourcesExhausted(
|
||||||
|
scope=self._scope,
|
||||||
|
reason="stalled",
|
||||||
|
retry_after_s=await self._breaker.retry_after_s(names),
|
||||||
|
per_source_reasons=reasons,
|
||||||
|
)
|
||||||
|
await self._sleep(self._bp.poll_interval_s * (0.5 + 0.5 * self._rng()))
|
||||||
|
|
||||||
|
async def _attempt(
|
||||||
|
self,
|
||||||
|
batch: list[str],
|
||||||
|
source: SourceConfig,
|
||||||
|
permit: Permit,
|
||||||
|
entry: GateDecision,
|
||||||
|
reasons: dict[str, str],
|
||||||
|
session_id: str | None,
|
||||||
|
parent_call_id: str | None,
|
||||||
|
) -> _BatchOutcome | _FailedBatch:
|
||||||
|
call_id = str(uuid.uuid4())
|
||||||
|
started = self._now()
|
||||||
|
actual = 0
|
||||||
|
try:
|
||||||
|
result = await self._transport.embed(texts=batch, source=source, call_id=call_id)
|
||||||
|
if self._expected_dim is not None and result.dim != self._expected_dim:
|
||||||
|
raise ResultInvalidError(
|
||||||
|
f"{source.name} 维度 {result.dim} 不符期望 {self._expected_dim}",
|
||||||
|
source_name=source.name,
|
||||||
|
operation="embedding",
|
||||||
|
)
|
||||||
|
actual = result.prompt_tokens
|
||||||
|
await self._record_quietly(self._breaker.record_success(entry))
|
||||||
|
await self._record_quietly(self._quota.mark_progress())
|
||||||
|
latency_ms = int((self._now() - started) * 1000)
|
||||||
|
await self._emit(batch, source, call_id, started, session_id, parent_call_id, result)
|
||||||
|
return _BatchOutcome(result, source, call_id, latency_ms)
|
||||||
|
except (RequestRejectedError, ResultInvalidError) as exc:
|
||||||
|
await self._gate_on_terminal(exc, entry)
|
||||||
|
await self._emit(batch, source, call_id, started, session_id, parent_call_id, error=exc)
|
||||||
|
raise
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
if entry.is_probe:
|
||||||
|
await self._record_quietly(self._breaker.release_probe(entry))
|
||||||
|
await self._emit(
|
||||||
|
batch, source, call_id, started, session_id, parent_call_id, error="cancelled"
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
except (SourceDeadError, TransientError) as exc:
|
||||||
|
dead = isinstance(exc, SourceDeadError)
|
||||||
|
reason = _failure_reason(exc)
|
||||||
|
reasons[source.name] = reason
|
||||||
|
await self._record_quietly(self._breaker.record_failure(entry, reason, dead))
|
||||||
|
if not dead:
|
||||||
|
actual = source.est_tokens # 保守: 失败请求可能已被网关计费(CHS 同款)
|
||||||
|
await self._emit(batch, source, call_id, started, session_id, parent_call_id, error=exc)
|
||||||
|
return _FailedBatch(exc, immediate=dead)
|
||||||
|
finally:
|
||||||
|
await self._settle_and_release(permit, actual)
|
||||||
|
|
||||||
|
# —— 辅助 ——
|
||||||
|
|
||||||
|
async def _gate_on_terminal(self, exc: PolyGatewayError, entry: GateDecision) -> None:
|
||||||
|
"""终态异常的门控写回(与 RetryMW 同口径): 坏结果/网关健康拒绝 ≠ 坏服务
|
||||||
|
→ 记成功;网关没响应的拒绝若持探针则归还。"""
|
||||||
|
if isinstance(exc, ResultInvalidError) or exc.status_code is not None:
|
||||||
|
await self._record_quietly(self._breaker.record_success(entry))
|
||||||
|
elif entry.is_probe:
|
||||||
|
await self._record_quietly(self._breaker.release_probe(entry))
|
||||||
|
|
||||||
|
async def _record_quietly(self, write_back: Awaitable[object]) -> None:
|
||||||
|
"""记账侧写回降级(与 RetryMW._record_quietly 同口径,设计 §10)。"""
|
||||||
|
try:
|
||||||
|
await write_back
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except GovernanceBackendError as exc:
|
||||||
|
logger.warning("embedding 治理记账写回降级(不冒泡): {}", exc)
|
||||||
|
|
||||||
|
async def _settle_and_release(self, permit: Permit, actual: int) -> None:
|
||||||
|
try:
|
||||||
|
try:
|
||||||
|
await permit.settle(actual)
|
||||||
|
finally:
|
||||||
|
await permit.release()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("embedding permit 结算/释放失败(不掩盖主异常): {}", exc)
|
||||||
|
|
||||||
|
async def _emit(
|
||||||
|
self,
|
||||||
|
batch: list[str],
|
||||||
|
source: SourceConfig,
|
||||||
|
call_id: str,
|
||||||
|
started: float,
|
||||||
|
session_id: str | None,
|
||||||
|
parent_call_id: str | None,
|
||||||
|
result: EmbeddingTransportResult | None = None,
|
||||||
|
error: object | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""逐批遥测(经同一 Emitter): messages=截断 texts、向量绝不入库。"""
|
||||||
|
if self._emitter is None:
|
||||||
|
return
|
||||||
|
request = ChatRequest(
|
||||||
|
messages=[{"role": "user", "content": t[:_TELEMETRY_TEXT_CAP]} for t in batch],
|
||||||
|
session_id=session_id,
|
||||||
|
parent_call_id=parent_call_id,
|
||||||
|
)
|
||||||
|
response = None
|
||||||
|
if result is not None:
|
||||||
|
response = LLMResponse(
|
||||||
|
content=f"<vectors n={len(result.vectors)} dim={result.dim}>",
|
||||||
|
thinking="",
|
||||||
|
model=source.model,
|
||||||
|
provider=source.provider,
|
||||||
|
prompt_tokens=result.prompt_tokens,
|
||||||
|
completion_tokens=0,
|
||||||
|
latency_ms=int((self._now() - started) * 1000),
|
||||||
|
ttft_ms=None,
|
||||||
|
max_inter_token_ms=None,
|
||||||
|
cache_hit=False,
|
||||||
|
call_id=call_id,
|
||||||
|
source_name=source.name,
|
||||||
|
usage_source=result.usage_source,
|
||||||
|
)
|
||||||
|
await self._emitter.emit_attempt(
|
||||||
|
request=request,
|
||||||
|
source=source,
|
||||||
|
call_id=call_id,
|
||||||
|
latency_ms=int((self._now() - started) * 1000),
|
||||||
|
response=response,
|
||||||
|
error=None if error is None else str(error),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _merge(self, outcomes: list[_BatchOutcome]) -> EmbeddingResponse:
|
||||||
|
"""全批合并(设计 §7.3): vectors 拼接、tokens/latency 求和、保守 usage_source。"""
|
||||||
|
vectors = [v for o in outcomes for v in o.result.vectors]
|
||||||
|
if self._normalize:
|
||||||
|
vectors = [_l2_normalize(v) for v in vectors]
|
||||||
|
first = outcomes[0]
|
||||||
|
prompt_tokens = sum(o.result.prompt_tokens for o in outcomes)
|
||||||
|
estimated = any(o.result.usage_source == "estimated" for o in outcomes)
|
||||||
|
return EmbeddingResponse(
|
||||||
|
vectors=vectors,
|
||||||
|
dim=first.result.dim,
|
||||||
|
model=first.source.model,
|
||||||
|
provider=first.source.provider,
|
||||||
|
prompt_tokens=prompt_tokens,
|
||||||
|
usage_source="estimated" if estimated else "measured",
|
||||||
|
latency_ms=sum(o.latency_ms for o in outcomes),
|
||||||
|
call_id=first.call_id,
|
||||||
|
source_name=first.source.name,
|
||||||
|
cost=self._total_cost(outcomes),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _total_cost(self, outcomes: list[_BatchOutcome]) -> float | None:
|
||||||
|
if self._pricing is None:
|
||||||
|
return None
|
||||||
|
costs = [
|
||||||
|
self._pricing.cost(o.source.model, o.result.prompt_tokens, 0) for o in outcomes
|
||||||
|
]
|
||||||
|
known = [c for c in costs if c is not None]
|
||||||
|
return sum(known) if known else None
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
"""幂等释放 transport 连接池与遥测连接(与 GatewayClient 对称)。"""
|
||||||
|
if self._closed:
|
||||||
|
return
|
||||||
|
self._closed = True
|
||||||
|
transport_aclose = getattr(self._transport, "aclose", None)
|
||||||
|
if transport_aclose is not None:
|
||||||
|
await transport_aclose()
|
||||||
|
telemetry_aclose = getattr(self._telemetry, "aclose", None)
|
||||||
|
if telemetry_aclose is not None:
|
||||||
|
await telemetry_aclose()
|
||||||
|
else:
|
||||||
|
telemetry_close = getattr(self._telemetry, "close", None)
|
||||||
|
if telemetry_close is not None:
|
||||||
|
telemetry_close()
|
||||||
|
|
||||||
|
async def __aenter__(self) -> EmbeddingClient:
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __aexit__(self, *exc_info: object) -> None:
|
||||||
|
await self.aclose()
|
||||||
|
|
||||||
|
# —— 工厂(与 GatewayClient 对称)——
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_settings(
|
||||||
|
cls,
|
||||||
|
settings: EmbeddingSettings,
|
||||||
|
*,
|
||||||
|
limiter: RateLimiter | None = None,
|
||||||
|
breaker: ProviderGate | None = None,
|
||||||
|
telemetry: TelemetryRecorder | None = None,
|
||||||
|
registry: Mapping[str, object] | None = None,
|
||||||
|
) -> EmbeddingClient:
|
||||||
|
"""按配置装配;显式传入的后端实例即共享(与 chat scope 共享全局闸)。"""
|
||||||
|
from polygateway.client import (
|
||||||
|
_build_breaker,
|
||||||
|
_build_limiter,
|
||||||
|
_build_selector,
|
||||||
|
_build_telemetry,
|
||||||
|
)
|
||||||
|
from polygateway.pricing import PricingTable
|
||||||
|
from polygateway.transports.openai_compat import OpenAICompatTransport
|
||||||
|
|
||||||
|
gw = settings.gateway
|
||||||
|
sources = list(gw.sources)
|
||||||
|
return cls(
|
||||||
|
scope=gw.scope,
|
||||||
|
sources=sources,
|
||||||
|
selector=_build_selector(gw.selector),
|
||||||
|
limiter=limiter or _build_limiter(gw, sources),
|
||||||
|
breaker=breaker or _build_breaker(gw),
|
||||||
|
transport=OpenAICompatTransport(registry=registry),
|
||||||
|
retry=gw.retry,
|
||||||
|
backpressure=gw.backpressure,
|
||||||
|
quota_full=gw.quota_full,
|
||||||
|
telemetry=telemetry if telemetry is not None else _build_telemetry(gw),
|
||||||
|
pricing=PricingTable.from_file(gw.pricing_path)
|
||||||
|
if gw.pricing_path is not None
|
||||||
|
else None,
|
||||||
|
batch_size=settings.batch_size,
|
||||||
|
normalize=settings.normalize,
|
||||||
|
expected_dim=settings.expected_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_env(
|
||||||
|
cls,
|
||||||
|
scope: str = "EMBED",
|
||||||
|
*,
|
||||||
|
limiter: RateLimiter | None = None,
|
||||||
|
breaker: ProviderGate | None = None,
|
||||||
|
telemetry: TelemetryRecorder | None = None,
|
||||||
|
registry: Mapping[str, object] | None = None,
|
||||||
|
env: Mapping[str, str] | None = None,
|
||||||
|
) -> EmbeddingClient:
|
||||||
|
"""从 .env/环境变量装配一个 embedding scope 的 client。"""
|
||||||
|
return cls.from_settings(
|
||||||
|
EmbeddingSettings.from_env(scope, env=env),
|
||||||
|
limiter=limiter,
|
||||||
|
breaker=breaker,
|
||||||
|
telemetry=telemetry,
|
||||||
|
registry=registry,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _l2_normalize(vector: list[float]) -> list[float]:
|
||||||
|
"""L2 归一化;`max(norm, 1e-12)` 防除零(VT embedding.py:167-170 语义)。"""
|
||||||
|
norm = max(math.sqrt(sum(x * x for x in vector)), 1e-12)
|
||||||
|
return [x / norm for x in vector]
|
||||||
@@ -0,0 +1,83 @@
|
|||||||
|
"""真实网关 /embeddings 端点探测(M2 设计 §11.6;人类默认口径: 实现时探测)。
|
||||||
|
|
||||||
|
对 .env 的 LLM 源网关发一次真实 embeddings 请求: 支持则记录向量证据,
|
||||||
|
不支持(404/翻译为领域错误)则 skip 并把响应记录进 tests/outputs/
|
||||||
|
(降级证据)。无 EMBED scope 配置时复用 LLM 源的 base_url/api_key。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
import os
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from dotenv import dotenv_values
|
||||||
|
|
||||||
|
from polygateway.errors import PolyGatewayError
|
||||||
|
from polygateway.transports.openai_compat import OpenAICompatTransport
|
||||||
|
from polygateway.types import SourceConfig
|
||||||
|
|
||||||
|
_ENV = {k: v for k, v in {**dotenv_values(".env"), **os.environ}.items() if v is not None}
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.skipif(
|
||||||
|
"LLM__MINIMAX__1__BASE_URL" not in _ENV, reason="缺真实网关配置(.env)"
|
||||||
|
)
|
||||||
|
|
||||||
|
_OUT = Path("tests/outputs/embedding")
|
||||||
|
|
||||||
|
|
||||||
|
def _record(name: str, lines: list[str]) -> Path:
|
||||||
|
_OUT.mkdir(parents=True, exist_ok=True)
|
||||||
|
path = _OUT / f"{name}_{datetime.now():%Y%m%d_%H%M%S}.md"
|
||||||
|
path.write_text("\n".join(lines) + "\n", encoding="utf-8")
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
|
async def test_probe_real_gateway_embeddings():
|
||||||
|
source = SourceConfig(
|
||||||
|
name="probe_1",
|
||||||
|
provider="minimax",
|
||||||
|
base_url=_ENV["LLM__MINIMAX__1__BASE_URL"],
|
||||||
|
api_key=_ENV["LLM__MINIMAX__1__API_KEY"],
|
||||||
|
model=_ENV.get("PGW_EMBED_PROBE_MODEL", "text-embedding-v1"),
|
||||||
|
timeout_s=30.0,
|
||||||
|
est_tokens=8,
|
||||||
|
)
|
||||||
|
transport = OpenAICompatTransport()
|
||||||
|
try:
|
||||||
|
result = await transport.embed(
|
||||||
|
texts=["polygateway embedding probe"], source=source, call_id="probe"
|
||||||
|
)
|
||||||
|
except PolyGatewayError as exc:
|
||||||
|
path = _record(
|
||||||
|
"probe_unsupported",
|
||||||
|
[
|
||||||
|
"# Embedding 端点探测: 网关不支持",
|
||||||
|
f"- base_url: {source.base_url}",
|
||||||
|
f"- model: {source.model}",
|
||||||
|
f"- 错误分类: {type(exc).__name__}",
|
||||||
|
f"- status_code: {exc.status_code}",
|
||||||
|
f"- 详情: {exc}",
|
||||||
|
"",
|
||||||
|
"结论: e2e 按设计 §11.6 降级,embedding 行为由 unit 全覆盖。",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
await transport.aclose()
|
||||||
|
pytest.skip(f"网关不支持 embeddings({type(exc).__name__}),证据: {path}")
|
||||||
|
else:
|
||||||
|
await transport.aclose()
|
||||||
|
assert result.dim > 0 and len(result.vectors) == 1
|
||||||
|
_record(
|
||||||
|
"probe_supported",
|
||||||
|
[
|
||||||
|
"# Embedding 端点探测: 网关支持",
|
||||||
|
f"- base_url: {source.base_url}",
|
||||||
|
f"- model: {source.model}",
|
||||||
|
f"- dim: {result.dim}",
|
||||||
|
f"- usage: {result.prompt_tokens}({result.usage_source})",
|
||||||
|
f"- 向量前 5 维: {result.vectors[0][:5]}",
|
||||||
|
f"- raw: {dataclasses.asdict(result)['raw']}",
|
||||||
|
],
|
||||||
|
)
|
||||||
@@ -155,3 +155,214 @@ class TestEmbedTransport:
|
|||||||
async def test_empty_texts_rejected(self):
|
async def test_empty_texts_rejected(self):
|
||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
await _transport_with(lambda r: None).embed(texts=[], source=_src(), call_id="c")
|
await _transport_with(lambda r: None).embed(texts=[], source=_src(), call_id="c")
|
||||||
|
|
||||||
|
|
||||||
|
# ═══════════ T9: EmbeddingClient 治理循环 ═══════════
|
||||||
|
|
||||||
|
import asyncio # noqa: E402
|
||||||
|
|
||||||
|
from polygateway.backends.memory.breaker import InMemoryGate # noqa: E402
|
||||||
|
from polygateway.backends.memory.limiter import InMemoryLimiter # noqa: E402
|
||||||
|
from polygateway.config import EmbeddingSettings # noqa: E402
|
||||||
|
from polygateway.embedding import EmbeddingClient # noqa: E402
|
||||||
|
from polygateway.sources import RoundRobinSelector # noqa: E402
|
||||||
|
from polygateway.types import ( # noqa: E402
|
||||||
|
BackpressurePolicy,
|
||||||
|
BreakerConfig,
|
||||||
|
GlobalLimits,
|
||||||
|
RetryPolicy,
|
||||||
|
)
|
||||||
|
|
||||||
|
_BREAKER = BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0)
|
||||||
|
_NO_GLOBAL = GlobalLimits(max_concurrency=0, rpm=0, tpm=0)
|
||||||
|
|
||||||
|
|
||||||
|
def _vec_for(texts):
|
||||||
|
"""确定性向量: 每条 text 一个 [len(text)] 一维向量,便于断言保序。"""
|
||||||
|
return EmbeddingTransportResult(
|
||||||
|
vectors=[[float(len(t))] for t in texts],
|
||||||
|
dim=1,
|
||||||
|
prompt_tokens=len(texts),
|
||||||
|
usage_source="measured",
|
||||||
|
raw={},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ScriptedEmbedTransport:
|
||||||
|
"""按脚本响应: 条目为 Exception / "ok"(按输入生成) / EmbeddingTransportResult / "hang"。"""
|
||||||
|
|
||||||
|
def __init__(self, script):
|
||||||
|
self.script = list(script)
|
||||||
|
self.calls = []
|
||||||
|
|
||||||
|
async def embed(self, *, texts, source, call_id):
|
||||||
|
self.calls.append((source.name, list(texts), call_id))
|
||||||
|
action = self.script.pop(0)
|
||||||
|
if isinstance(action, Exception):
|
||||||
|
raise action
|
||||||
|
if action == "hang":
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
if action == "ok":
|
||||||
|
return _vec_for(texts)
|
||||||
|
return action
|
||||||
|
|
||||||
|
|
||||||
|
class _MemoryRecorder:
|
||||||
|
def __init__(self):
|
||||||
|
self.rows = []
|
||||||
|
|
||||||
|
async def record_llm_call(self, **fields):
|
||||||
|
self.rows.append(fields)
|
||||||
|
|
||||||
|
|
||||||
|
def _embed_client(sources, script, *, batch_size=2, telemetry=None, **overrides):
|
||||||
|
limiter = InMemoryLimiter(
|
||||||
|
scope="embed",
|
||||||
|
sources={s.name: s for s in sources},
|
||||||
|
global_limits=_NO_GLOBAL,
|
||||||
|
lease_ttl_s=100.0,
|
||||||
|
)
|
||||||
|
kwargs = {
|
||||||
|
"scope": "embed",
|
||||||
|
"sources": sources,
|
||||||
|
"selector": RoundRobinSelector(),
|
||||||
|
"limiter": limiter,
|
||||||
|
"breaker": InMemoryGate(config=_BREAKER),
|
||||||
|
"transport": ScriptedEmbedTransport(script),
|
||||||
|
"retry": RetryPolicy(max_attempts=3, backoff_base_s=0.001, backoff_max_s=0.01),
|
||||||
|
"backpressure": BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.001),
|
||||||
|
"batch_size": batch_size,
|
||||||
|
"telemetry": telemetry,
|
||||||
|
}
|
||||||
|
kwargs.update(overrides)
|
||||||
|
client = EmbeddingClient(**kwargs)
|
||||||
|
return client, limiter
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbedBatching:
|
||||||
|
async def test_batches_sequential_and_order_preserved(self):
|
||||||
|
texts = ["a", "bb", "ccc", "dddd", "eeeee"]
|
||||||
|
client, _ = _embed_client([_src()], ["ok", "ok", "ok"], batch_size=2)
|
||||||
|
resp = await client.embed(texts)
|
||||||
|
transport = client._transport
|
||||||
|
assert [len(batch) for _, batch, _ in transport.calls] == [2, 2, 1]
|
||||||
|
assert resp.vectors == [[1.0], [2.0], [3.0], [4.0], [5.0]] # 全批拼接保序
|
||||||
|
assert resp.prompt_tokens == 5 and resp.dim == 1
|
||||||
|
|
||||||
|
async def test_empty_input_short_circuits(self):
|
||||||
|
client, _ = _embed_client([_src()], [])
|
||||||
|
resp = await client.embed([])
|
||||||
|
assert resp.vectors == [] and resp.prompt_tokens == 0
|
||||||
|
assert client._transport.calls == []
|
||||||
|
|
||||||
|
async def test_usage_source_aggregates_conservatively(self):
|
||||||
|
estimated = EmbeddingTransportResult(
|
||||||
|
vectors=[[1.0], [1.0]], dim=1, prompt_tokens=9, usage_source="estimated", raw={}
|
||||||
|
)
|
||||||
|
client, _ = _embed_client([_src()], ["ok", estimated], batch_size=2)
|
||||||
|
resp = await client.embed(["a", "b", "c", "d"])
|
||||||
|
assert resp.usage_source == "estimated" # 任一批 estimated 则整体 estimated
|
||||||
|
assert resp.prompt_tokens == 2 + 9
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbedPostProcess:
|
||||||
|
async def test_normalize_l2(self):
|
||||||
|
raw = EmbeddingTransportResult(
|
||||||
|
vectors=[[3.0, 4.0]], dim=2, prompt_tokens=1, usage_source="measured", raw={}
|
||||||
|
)
|
||||||
|
client, _ = _embed_client([_src()], [raw], normalize=True)
|
||||||
|
resp = await client.embed(["x"])
|
||||||
|
assert resp.vectors[0] == pytest.approx([0.6, 0.8])
|
||||||
|
|
||||||
|
async def test_zero_vector_normalize_no_nan(self):
|
||||||
|
raw = EmbeddingTransportResult(
|
||||||
|
vectors=[[0.0, 0.0]], dim=2, prompt_tokens=1, usage_source="measured", raw={}
|
||||||
|
)
|
||||||
|
client, _ = _embed_client([_src()], [raw], normalize=True)
|
||||||
|
resp = await client.embed(["x"])
|
||||||
|
assert resp.vectors[0] == [0.0, 0.0] # max(norm, 1e-12) 防除零(VT 语义)
|
||||||
|
|
||||||
|
async def test_expected_dim_violation_is_result_invalid(self):
|
||||||
|
client, _ = _embed_client([_src()], ["ok"], expected_dim=768)
|
||||||
|
with pytest.raises(ResultInvalidError):
|
||||||
|
await client.embed(["x"])
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbedGovernance:
|
||||||
|
async def test_transient_retries_then_succeeds(self):
|
||||||
|
client, _ = _embed_client(
|
||||||
|
[_src()], [TransientError("boom", status_code=500), "ok"], batch_size=8
|
||||||
|
)
|
||||||
|
resp = await client.embed(["a", "b"])
|
||||||
|
assert resp.vectors == [[1.0], [1.0]]
|
||||||
|
assert len(client._transport.calls) == 2
|
||||||
|
|
||||||
|
async def test_source_dead_switches_source(self):
|
||||||
|
s1, s2 = _src(name="e1"), _src(name="e2")
|
||||||
|
client, _ = _embed_client(
|
||||||
|
[s1, s2], [SourceDeadError("401", status_code=401), "ok"], batch_size=8
|
||||||
|
)
|
||||||
|
await client.embed(["a"])
|
||||||
|
assert [name for name, _, _ in client._transport.calls] == ["e1", "e2"]
|
||||||
|
|
||||||
|
async def test_cancel_releases_permit(self):
|
||||||
|
client, limiter = _embed_client([_src(max_concurrency=1)], ["hang"])
|
||||||
|
task = asyncio.create_task(client.embed(["a"]))
|
||||||
|
while not (await limiter.source_stats("e1")).inflight:
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
task.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
assert (await limiter.source_stats("e1")).inflight == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbedTelemetry:
|
||||||
|
async def test_per_batch_rows_with_digest(self):
|
||||||
|
rec = _MemoryRecorder()
|
||||||
|
client, _ = _embed_client([_src()], ["ok", "ok"], batch_size=1, telemetry=rec)
|
||||||
|
await client.embed(["hello", "x" * 5000], session_id="sess", parent_call_id="pc")
|
||||||
|
assert len(rec.rows) == 2 # 每批一行
|
||||||
|
row = rec.rows[0]
|
||||||
|
assert row["session_id"] == "sess" and row["parent_call_id"] == "pc"
|
||||||
|
assert row["completion_tokens"] == 0
|
||||||
|
assert row["response"] == "<vectors n=1 dim=1>" # 向量绝不入库
|
||||||
|
assert len(rec.rows[1]["messages"]) < 1000 # 长文本截断后入库
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmbeddingSettings:
|
||||||
|
_ENV = {
|
||||||
|
"EMBED__QWEN__1__BASE_URL": "https://gw.example/v1",
|
||||||
|
"EMBED__QWEN__1__API_KEY": "sk-a",
|
||||||
|
"EMBED__QWEN__1__MODEL": "text-embedding-v3",
|
||||||
|
"EMBED__QWEN__1__TIMEOUT_S": "60",
|
||||||
|
"EMBED__RETRY__MAX_ATTEMPTS": "3",
|
||||||
|
"EMBED__RETRY__BACKOFF_BASE_S": "1.0",
|
||||||
|
"EMBED__RETRY__BACKOFF_MAX_S": "10.0",
|
||||||
|
"EMBED__BREAKER__FAIL_THRESHOLD": "5",
|
||||||
|
"EMBED__BREAKER__COOLDOWN_S": "60",
|
||||||
|
"PGW_CACHE_BACKEND": "none",
|
||||||
|
"PGW_TELEMETRY_BACKEND": "none",
|
||||||
|
"EMBED__BATCH_SIZE": "64",
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_loads_scope_and_batch(self):
|
||||||
|
s = EmbeddingSettings.from_env("EMBED", env=self._ENV)
|
||||||
|
assert s.gateway.sources[0].model == "text-embedding-v3"
|
||||||
|
assert s.batch_size == 64 and s.normalize is False and s.expected_dim is None
|
||||||
|
|
||||||
|
def test_batch_size_required_and_positive(self):
|
||||||
|
env = {k: v for k, v in self._ENV.items() if k != "EMBED__BATCH_SIZE"}
|
||||||
|
with pytest.raises(ValueError, match="BATCH_SIZE"):
|
||||||
|
EmbeddingSettings.from_env("EMBED", env=env)
|
||||||
|
with pytest.raises(ValueError, match="BATCH_SIZE"):
|
||||||
|
EmbeddingSettings.from_env("EMBED", env={**self._ENV, "EMBED__BATCH_SIZE": "0"})
|
||||||
|
|
||||||
|
def test_optional_normalize_and_dim(self):
|
||||||
|
env = {**self._ENV, "EMBED__NORMALIZE": "true", "EMBED__EXPECTED_DIM": "768"}
|
||||||
|
s = EmbeddingSettings.from_env("EMBED", env=env)
|
||||||
|
assert s.normalize is True and s.expected_dim == 768
|
||||||
|
|
||||||
|
def test_from_settings_assembles_client(self):
|
||||||
|
s = EmbeddingSettings.from_env("EMBED", env=self._ENV)
|
||||||
|
client = EmbeddingClient.from_settings(s)
|
||||||
|
assert isinstance(client, EmbeddingClient)
|
||||||
|
|||||||
Reference in New Issue
Block a user