From 5e01dc738f661c5e944526b9d6dc5dec87024c88 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Tue, 21 Jul 2026 01:11:48 -0400 Subject: [PATCH] feat: add governed embedding client with batching --- src/polygateway/__init__.py | 11 +- src/polygateway/config.py | 45 ++++ src/polygateway/embedding.py | 492 ++++++++++++++++++++++++++++++++++ tests/e2e/test_embed_probe.py | 83 ++++++ tests/unit/test_embedding.py | 211 +++++++++++++++ 5 files changed, 840 insertions(+), 2 deletions(-) create mode 100644 src/polygateway/embedding.py create mode 100644 tests/e2e/test_embed_probe.py diff --git a/src/polygateway/__init__.py b/src/polygateway/__init__.py index 14bd9fd..d5aade0 100644 --- a/src/polygateway/__init__.py +++ b/src/polygateway/__init__.py @@ -6,7 +6,8 @@ """ 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 ( AllSourcesExhausted, CircuitOpenError, @@ -18,8 +19,9 @@ from polygateway.errors import ( SourceDeadError, TransientError, ) +from polygateway.pricing import ModelPrice, PricingTable 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" @@ -27,12 +29,17 @@ __all__ = [ "DEFAULT_PROFILES", "AllSourcesExhausted", "CircuitOpenError", + "EmbeddingClient", + "EmbeddingResponse", + "EmbeddingSettings", "GatewayClient", "GatewaySettings", "GatewayUnavailableError", "GovernanceBackendError", "LLMResponse", + "ModelPrice", "PolyGatewayError", + "PricingTable", "ProviderProfile", "RequestRejectedError", "ResultInvalidError", diff --git a/src/polygateway/config.py b/src/polygateway/config.py index 8d3f141..82db110 100644 --- a/src/polygateway/config.py +++ b/src/polygateway/config.py @@ -341,3 +341,48 @@ def _guard_stall(settings: GatewaySettings) -> None: def _load_lease_ttl(env: Mapping[str, str]) -> float: found = _first(env, "PGW_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, + ) diff --git a/src/polygateway/embedding.py b/src/polygateway/embedding.py new file mode 100644 index 0000000..ffd04c6 --- /dev/null +++ b/src/polygateway/embedding.py @@ -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"", + 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] diff --git a/tests/e2e/test_embed_probe.py b/tests/e2e/test_embed_probe.py new file mode 100644 index 0000000..de25347 --- /dev/null +++ b/tests/e2e/test_embed_probe.py @@ -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']}", + ], + ) diff --git a/tests/unit/test_embedding.py b/tests/unit/test_embedding.py index c1f1bba..a2ed5d9 100644 --- a/tests/unit/test_embedding.py +++ b/tests/unit/test_embedding.py @@ -155,3 +155,214 @@ class TestEmbedTransport: async def test_empty_texts_rejected(self): with pytest.raises(ValueError): 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"] == "" # 向量绝不入库 + 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)