Files
PolyGateway/src/polygateway/embedding.py
T
iomgaa 702040d1a3 feat: carry caller dimensions down the embedding chain
EmbeddingClient does not go through the chat onion: it builds its own
ChatRequest inside _emit purely to reuse the shared TelemetryEmitter, so
wiring chat() alone left every embed row without a tenant. Validate the
dimensions at the embed() entry (before batching, since anything failing
further down is degraded to a warning) and thread them through
_embed_batch -> _attempt -> _emit so every batch row carries the same
pair.
2026-08-17 09:53:37 -04:00

593 lines
22 KiB
Python

"""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, Any
from loguru import logger
from polygateway.config import EmbeddingSettings
from polygateway.errors import (
AllSourcesExhausted,
CircuitOpenError,
GovernanceBackendError,
PolyGatewayError,
RequestRejectedError,
ResultInvalidError,
SourceDeadError,
SourceNotConfiguredError,
TransientError,
)
from polygateway.middleware.breaker import BreakerGate
from polygateway.middleware.ratelimit import QuotaGate
from polygateway.middleware.retry import StallClock, _failure_reason, backoff_delay
from polygateway.middleware.telemetry import TelemetryEmitter
from polygateway.sources import SourceCooldownMemo
from polygateway.types import (
ChatRequest,
EmbeddingResponse,
LLMResponse,
strip_unsupported_extra_body,
validate_caller_dimensions,
)
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
# embed payload 硬编码 {model, input},带 extra_body 的源必须先剥离,
# 否则遥测会记录一个从未发出的采样参数(issue #4 决策 G)
self._sources = strip_unsupported_extra_body(list(sources), path="embedding")
self._selector = selector
self._quota = QuotaGate(limiter, scope=self._scope)
self._breaker = BreakerGate(breaker, scope=self._scope)
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,
tenant_id: str | None = None,
meta: Mapping[str, Any] | None = None,
) -> EmbeddingResponse:
"""一次治理 embedding 调用: 按 batch_size 切批,批间串行,全批合并返回。
`tenant_id` 与 `meta` 是调用方自定义维度,只进遥测(issue #11);它们属于
本次调用而非某一批,故每批的遥测行都带同一份维度。
"""
if not isinstance(texts, list) or any(not isinstance(t, str) for t in texts):
raise TypeError("texts 必须是 list[str](显式优于隐式,不收单条 str)")
# 必须在切批之前校验: 洋葱/链路内的一切失败都被遥测层降级成 warning
# (库铁律「遥测写失败降级不冒泡」),校验放下游等于没有校验——非法维度
# 会变成静默丢失的遥测行,而调用照常发出(issue #11 §4.2)
dimension_tenant_id, dimensions = validate_caller_dimensions(
tenant_id, meta, origin="embed(tenant_id=..., meta=...)"
)
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,
dimension_tenant_id,
dimensions,
)
)
return self._merge(outcomes)
# —— 治理循环(与 RetryMW 同构;设计 §7.1 已声明的有限重复)——
async def _embed_batch(
self,
batch: list[str],
session_id: str | None,
parent_call_id: str | None,
tenant_id: str | None,
meta: dict[str, Any],
) -> _BatchOutcome:
fails = 0
reasons: dict[str, str] = {}
# 只计非生产性等待(issue #8): 真实尝试由重试预算治理,不重复烧 stall 预算
clock = StallClock(self._now)
while True:
picked, gate_rejections = await self._pick_runnable(reasons)
if picked is None:
await self._on_no_runnable(gate_rejections, reasons, clock)
continue
async with clock.attempting():
outcome = await self._attempt(
batch, *picked, reasons, session_id, parent_call_id, tenant_id, meta
)
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], clock: StallClock
) -> 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 clock.stalled_s() > 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,
tenant_id: str | None,
meta: dict[str, Any],
) -> _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",
)
if result.usage_source == "unavailable":
# 与 RetryMW 同口径: 用量不可得时按入场预扣量结算(设计 §3.2 #9)
actual = source.effective_est_tokens()
else:
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,
tenant_id,
meta,
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,
tenant_id,
meta,
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,
tenant_id,
meta,
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:
# 保守: 失败请求可能已被网关计费(CHS 同款);与入场预扣同源取值
actual = source.effective_est_tokens()
await self._emit(
batch,
source,
call_id,
started,
session_id,
parent_call_id,
tenant_id,
meta,
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:
# M2.5 §3.1: 记成功但不计失败率窗口样本(坏结果/坏请求 ≠ 坏服务)
await self._record_quietly(self._breaker.record_success(entry, count_attempt=False))
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, SourceNotConfiguredError) 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,
tenant_id: str | None,
meta: dict[str, Any],
result: EmbeddingTransportResult | None = None,
error: object | None = None,
) -> None:
"""逐批遥测(经同一 Emitter): messages=截断 texts、向量绝不入库。"""
if self._emitter is None:
return
# 这个 ChatRequest 只为复用同一个 Emitter 而现场构造(embedding 不走 chat
# 洋葱),故调用方维度必须在这里显式填回,否则 embed 行的维度恒为空
request = ChatRequest(
messages=[{"role": "user", "content": t[:_TELEMETRY_TEXT_CAP]} for t in batch],
session_id=session_id,
parent_call_id=parent_call_id,
tenant_id=tenant_id,
meta=meta,
)
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)
# 三态合并优先级(解耦设计 §3.2 #10): 任一批不可得 → 整体不可得
sources = {o.result.usage_source for o in outcomes}
if "unavailable" in sources:
merged_source = "unavailable"
elif "estimated" in sources:
merged_source = "estimated"
else:
merged_source = "measured"
return EmbeddingResponse(
vectors=vectors,
dim=first.result.dim,
model=first.source.model,
provider=first.source.provider,
prompt_tokens=prompt_tokens,
usage_source=merged_source,
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:
"""全批成本;任一批用量不可得则整体记 NULL(解耦设计 §3.2 #11)。
逐批求和会把不可得的批当 0 计入,给出一个偏低却看似有效的金额——
与"宁可算不出成本,也不算错成本"的不变式相悖。
"""
if self._pricing is None:
return None
if any(o.result.usage_source == "unavailable" for o in outcomes):
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]