feat: add health-aware P2C selector as default
Score is success-rate EWMA over (1 + inflight) with a 0.05 exploration floor so quarantined sources can prove recovery; EWMA climb doubles as slow-start. New optional OutcomeAwareSelector port feeds attempt outcomes.
This commit is contained in:
@@ -24,7 +24,12 @@ from polygateway.middleware.structured import StructuredMW
|
|||||||
from polygateway.middleware.telemetry import TelemetryEmitter, TelemetryMW
|
from polygateway.middleware.telemetry import TelemetryEmitter, TelemetryMW
|
||||||
from polygateway.pricing import PricingTable
|
from polygateway.pricing import PricingTable
|
||||||
from polygateway.providers import get_provider
|
from polygateway.providers import get_provider
|
||||||
from polygateway.sources import LeastInflightSelector, RoundRobinSelector, SourceCooldownMemo
|
from polygateway.sources import (
|
||||||
|
HealthAwareSelector,
|
||||||
|
LeastInflightSelector,
|
||||||
|
RoundRobinSelector,
|
||||||
|
SourceCooldownMemo,
|
||||||
|
)
|
||||||
from polygateway.transports.openai_compat import OpenAICompatTransport
|
from polygateway.transports.openai_compat import OpenAICompatTransport
|
||||||
from polygateway.types import ChatRequest, LLMResponse
|
from polygateway.types import ChatRequest, LLMResponse
|
||||||
|
|
||||||
@@ -193,6 +198,7 @@ class GatewayClient:
|
|||||||
cache: CacheBackend | None = None,
|
cache: CacheBackend | None = None,
|
||||||
telemetry: TelemetryRecorder | None = None,
|
telemetry: TelemetryRecorder | None = None,
|
||||||
registry: Mapping[str, ProviderProfile] | None = None,
|
registry: Mapping[str, ProviderProfile] | None = None,
|
||||||
|
rng: Any = random.random,
|
||||||
) -> GatewayClient:
|
) -> GatewayClient:
|
||||||
"""按配置装配;显式传入的后端实例即共享(None 项按配置自建私有实例)。"""
|
"""按配置装配;显式传入的后端实例即共享(None 项按配置自建私有实例)。"""
|
||||||
sources = list(settings.sources)
|
sources = list(settings.sources)
|
||||||
@@ -201,7 +207,7 @@ class GatewayClient:
|
|||||||
return cls(
|
return cls(
|
||||||
scope=settings.scope,
|
scope=settings.scope,
|
||||||
sources=sources,
|
sources=sources,
|
||||||
selector=_build_selector(settings.selector),
|
selector=_build_selector(settings.selector, rng=rng),
|
||||||
limiter=limiter or _build_limiter(settings, sources),
|
limiter=limiter or _build_limiter(settings, sources),
|
||||||
breaker=breaker or _build_breaker(settings),
|
breaker=breaker or _build_breaker(settings),
|
||||||
transport=OpenAICompatTransport(registry=registry),
|
transport=OpenAICompatTransport(registry=registry),
|
||||||
@@ -272,8 +278,12 @@ def _build_breaker(settings: GatewaySettings) -> ProviderGate:
|
|||||||
return InMemoryGate(config=settings.breaker)
|
return InMemoryGate(config=settings.breaker)
|
||||||
|
|
||||||
|
|
||||||
def _build_selector(name: str) -> SourceSelector:
|
def _build_selector(name: str, *, rng: Any = random.random) -> SourceSelector:
|
||||||
return RoundRobinSelector() if name == "round_robin" else LeastInflightSelector()
|
if name == "round_robin":
|
||||||
|
return RoundRobinSelector()
|
||||||
|
if name == "least_inflight":
|
||||||
|
return LeastInflightSelector()
|
||||||
|
return HealthAwareSelector(rng=rng)
|
||||||
|
|
||||||
|
|
||||||
def _build_cache(settings: GatewaySettings) -> CacheBackend | None:
|
def _build_cache(settings: GatewaySettings) -> CacheBackend | None:
|
||||||
|
|||||||
@@ -181,6 +181,18 @@ class SourceSelector(Protocol):
|
|||||||
) -> list[SourceConfig]: ...
|
) -> list[SourceConfig]: ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class OutcomeAwareSelector(Protocol):
|
||||||
|
"""可选选源器扩展(M2.5 设计 §3.2): 消费尝试结果以维护健康视图。
|
||||||
|
|
||||||
|
RetryMW 构造时 isinstance 判定一次;非本 Protocol 的选源器不受影响。
|
||||||
|
喂数口径: 真实成功 ok=True;Transient/SourceDead/429 ok=False;
|
||||||
|
ResultInvalid 与"网关健康拒坏请求"不喂(坏结果 ≠ 坏服务)。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def record_outcome(self, source_name: str, ok: bool) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
@runtime_checkable
|
||||||
class StructuredOutputStrategy(Protocol):
|
class StructuredOutputStrategy(Protocol):
|
||||||
"""结构化输出策略(D7/D14): 请求侧叠加 + 响应侧解析。"""
|
"""结构化输出策略(D7/D14): 请求侧叠加 + 响应侧解析。"""
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ RetryMW 的进程本地状态——熔断开路的源在本地记冷却截止,
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import random
|
||||||
import time
|
import time
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
@@ -15,6 +16,9 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
from polygateway.types import SourceConfig, SourceStats
|
from polygateway.types import SourceConfig, SourceStats
|
||||||
|
|
||||||
|
_EWMA_ALPHA = 0.2 # 健康 EWMA 步长: 约 10 次成功从谷底爬回 0.9(天然 slow-start)
|
||||||
|
_SCORE_FLOOR = 0.05 # 探索地板: 塌陷源保有微量被选概率,恢复靠真实成功自证
|
||||||
|
|
||||||
|
|
||||||
class RoundRobinSelector:
|
class RoundRobinSelector:
|
||||||
"""轮转起点后移(CHS selector.py:20 同款);单 client 内游标推进。"""
|
"""轮转起点后移(CHS selector.py:20 同款);单 client 内游标推进。"""
|
||||||
@@ -41,6 +45,43 @@ class LeastInflightSelector:
|
|||||||
return sorted(sources, key=lambda s: stats[s.name].inflight if s.name in stats else 0)
|
return sorted(sources, key=lambda s: stats[s.name].inflight if s.name in stats else 0)
|
||||||
|
|
||||||
|
|
||||||
|
class HealthAwareSelector:
|
||||||
|
"""健康感知选源(M2.5 设计 §3.2;蓝本 Envoy least-request + gRPC WRR)。
|
||||||
|
|
||||||
|
score = max(ewma_success, 地板) / (1 + inflight)。头名经 P2C(随机取
|
||||||
|
两源比分,高者先)引入探索;其余按分数降序。健康态为进程本地(业界
|
||||||
|
共识: Envoy/Finagle/gRPC 全本地),属 client 实例,不违反纯 asyncio 中立。
|
||||||
|
已知取舍(设计 §3.2): 仅头名随机化,多 worker 溢出会集中到同一次优源。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, *, rng: Callable[[], float] = random.random) -> None:
|
||||||
|
self._rng = rng
|
||||||
|
self._ewma: dict[str, float] = {}
|
||||||
|
|
||||||
|
def record_outcome(self, source_name: str, ok: bool) -> None:
|
||||||
|
"""尝试结果喂数(OutcomeAwareSelector 端口);初始 1.0 乐观起步。"""
|
||||||
|
prev = self._ewma.get(source_name, 1.0)
|
||||||
|
self._ewma[source_name] = prev + _EWMA_ALPHA * ((1.0 if ok else 0.0) - prev)
|
||||||
|
|
||||||
|
def _score(self, name: str, stats: dict[str, SourceStats]) -> float:
|
||||||
|
ewma = max(self._ewma.get(name, 1.0), _SCORE_FLOOR)
|
||||||
|
inflight = stats[name].inflight if name in stats else 0
|
||||||
|
return ewma / (1.0 + inflight)
|
||||||
|
|
||||||
|
def order(
|
||||||
|
self, sources: list[SourceConfig], stats: dict[str, SourceStats]
|
||||||
|
) -> list[SourceConfig]:
|
||||||
|
if len(sources) < 2:
|
||||||
|
return list(sources)
|
||||||
|
ranked = sorted(sources, key=lambda s: self._score(s.name, stats), reverse=True)
|
||||||
|
# P2C: 随机取两源比分,胜者提为头名(平分取采样序首位)
|
||||||
|
i = int(self._rng() * len(sources)) % len(sources)
|
||||||
|
j = int(self._rng() * len(sources)) % len(sources)
|
||||||
|
a, b = sources[i], sources[j]
|
||||||
|
head = a if self._score(a.name, stats) >= self._score(b.name, stats) else b
|
||||||
|
return [head] + [s for s in ranked if s.name != head.name]
|
||||||
|
|
||||||
|
|
||||||
class SourceCooldownMemo:
|
class SourceCooldownMemo:
|
||||||
"""进程本地的源冷却备忘(CHS governance.py:107 同款)。
|
"""进程本地的源冷却备忘(CHS governance.py:107 同款)。
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
"""HealthAwareSelector 单测(M2.5 设计 §3.2): EWMA×在途 P2C,地板探索。"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from polygateway.ports import OutcomeAwareSelector
|
||||||
|
from polygateway.sources import HealthAwareSelector, RoundRobinSelector
|
||||||
|
from polygateway.types import SourceConfig, SourceStats
|
||||||
|
|
||||||
|
|
||||||
|
def _src(name: str) -> SourceConfig:
|
||||||
|
return SourceConfig(
|
||||||
|
name=name,
|
||||||
|
provider="p",
|
||||||
|
base_url="https://gw.example/v1",
|
||||||
|
api_key="sk-x",
|
||||||
|
model="m",
|
||||||
|
timeout_s=60.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _stats(**inflight: int) -> dict[str, SourceStats]:
|
||||||
|
return {name: SourceStats(inflight=n, rpm_used=0, tpm_used=0) for name, n in inflight.items()}
|
||||||
|
|
||||||
|
|
||||||
|
class TestHealthAwareSelector:
|
||||||
|
def test_implements_outcome_protocol(self):
|
||||||
|
assert isinstance(HealthAwareSelector(), OutcomeAwareSelector)
|
||||||
|
assert not isinstance(RoundRobinSelector(), OutcomeAwareSelector)
|
||||||
|
|
||||||
|
def test_unhealthy_source_demoted(self):
|
||||||
|
# s2 连续失败 → EWMA 塌陷 → 排序永远在健康源之后(rng 定值消除 P2C 随机性)
|
||||||
|
sel = HealthAwareSelector(rng=lambda: 0.0)
|
||||||
|
sources = [_src("s1"), _src("s2")]
|
||||||
|
for _ in range(10):
|
||||||
|
sel.record_outcome("s2", ok=False)
|
||||||
|
order = sel.order(sources, _stats(s1=0, s2=0))
|
||||||
|
assert [s.name for s in order] == ["s1", "s2"]
|
||||||
|
|
||||||
|
def test_floor_keeps_exploration_possible(self):
|
||||||
|
# 地板 0.05: 塌陷源分数不归零——EWMA 十次失败后仍 > 0,P2C 采样到时可胜平局
|
||||||
|
sel = HealthAwareSelector(rng=lambda: 0.0)
|
||||||
|
for _ in range(50):
|
||||||
|
sel.record_outcome("s2", ok=False)
|
||||||
|
assert sel._score("s2", _stats(s2=0)) == pytest.approx(0.05)
|
||||||
|
|
||||||
|
def test_ewma_climbs_back_on_recovery(self):
|
||||||
|
# 恢复源连续成功,EWMA α=0.2 自然爬升(即天然 slow-start)
|
||||||
|
sel = HealthAwareSelector(rng=lambda: 0.0)
|
||||||
|
for _ in range(50):
|
||||||
|
sel.record_outcome("s2", ok=False)
|
||||||
|
low = sel._score("s2", _stats(s2=0))
|
||||||
|
for _ in range(10):
|
||||||
|
sel.record_outcome("s2", ok=True)
|
||||||
|
high = sel._score("s2", _stats(s2=0))
|
||||||
|
assert high > 0.85 > low
|
||||||
|
|
||||||
|
def test_inflight_suppresses_score(self):
|
||||||
|
sel = HealthAwareSelector(rng=lambda: 0.0)
|
||||||
|
assert sel._score("s1", _stats(s1=0)) == pytest.approx(1.0)
|
||||||
|
assert sel._score("s1", _stats(s1=3)) == pytest.approx(0.25)
|
||||||
|
|
||||||
|
def test_p2c_head_randomized_rest_by_score(self):
|
||||||
|
# rng 驱动 P2C 取样: 三源同分时头名由 rng 决定;其余按分数降序
|
||||||
|
sources = [_src("s1"), _src("s2"), _src("s3")]
|
||||||
|
sel = HealthAwareSelector(rng=iter([0.9, 0.0]).__next__) # 采样 s3 与 s1 比分
|
||||||
|
for _ in range(5):
|
||||||
|
sel.record_outcome("s3", ok=False) # s3 塌陷
|
||||||
|
order = sel.order(sources, _stats(s1=0, s2=0, s3=0))
|
||||||
|
assert order[0].name == "s1" # 两候选中 s1 胜出
|
||||||
|
assert order[-1].name == "s3" # 塌陷源垫底
|
||||||
|
|
||||||
|
def test_missing_stats_defaults_to_zero_inflight(self):
|
||||||
|
sel = HealthAwareSelector(rng=lambda: 0.0)
|
||||||
|
assert [s.name for s in sel.order([_src("s1")], {})] == ["s1"]
|
||||||
Reference in New Issue
Block a user