Files
PolyGateway/tests/integration/test_redis_cross_connection.py
T

249 lines
8.5 KiB
Python

"""跨连接共享治理状态验证(M2 设计 §11.3;双连接池 = 多 worker 等价,人类认可)。
移植 CHS tests/integration/test_redis_limiter.py 的 2 个跨连接用例
(:127 全局并发、:166 进度可见)并把 :105 的全局 RPM(CHS 原版单连接)
升级为跨连接;再加熔断共享、双 client 联合 RPM 不超配、取消释放与
Redis 掉线 fail-closed 方向。真多进程验证在 soak harness --workers。
"""
from __future__ import annotations
import asyncio
from uuid import uuid4
import pytest
from polygateway.backends.redis.breaker import RedisGate
from polygateway.backends.redis.limiter import RedisLimiter
from polygateway.client import GatewayClient
from polygateway.errors import AllSourcesExhausted, GovernanceBackendError
from polygateway.sources import RoundRobinSelector
from polygateway.types import (
BackpressurePolicy,
BreakerConfig,
GlobalLimits,
RetryPolicy,
TransportResult,
)
from tests.contracts.conftest import (
await_window_headroom,
make_source,
redis_url_from_env,
)
_CFG = BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0)
_BP = BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.02)
_RETRY = RetryPolicy(max_attempts=2, backoff_base_s=0.01, backoff_max_s=0.05)
@pytest.fixture
async def clients():
"""两个独立连接池(模拟两进程)。"""
url = redis_url_from_env()
if url is None:
pytest.skip("REDIS_URL 未配置")
import redis.asyncio as aioredis
a, b = aioredis.from_url(url), aioredis.from_url(url)
try:
yield a, b
finally:
await a.aclose()
await b.aclose()
def _limiter(client, scope, sources, global_limits, **kwargs) -> RedisLimiter:
return RedisLimiter(
scope=scope,
sources={s.name: s for s in sources},
global_limits=global_limits,
redis=client,
**kwargs,
)
class ScriptedTransport:
"""返回固定成功结果(或挂起)的 transport;记录调用数(参照 test_retry FakeTransport)。"""
def __init__(self, hang: bool = False):
self.hang = hang
self.calls: list[str] = []
async def complete(self, *, messages, source, stream, overlay, call_id):
self.calls.append(source.name)
if self.hang:
await asyncio.Event().wait()
return TransportResult(
content="ok",
thinking="",
prompt_tokens=1,
completion_tokens=1,
usage_source="measured",
ttft_ms=None,
max_inter_token_ms=None,
raw={},
)
def _client(scope, sources, limiter, gate, transport, *, quota_full="wait") -> GatewayClient:
return GatewayClient(
scope=scope,
sources=sources,
selector=RoundRobinSelector(),
limiter=limiter,
breaker=gate,
transport=transport,
retry=_RETRY,
backpressure=_BP,
quota_full=quota_full,
)
# —— CHS 移植 3 例 ——
async def test_cross_connection_global_concurrency(clients):
"""CHS :127: 两连接池共享全局并发闸,第三个 acquire 跨连接被拒。"""
a_cli, b_cli = clients
scope = f"t{uuid4().hex[:8]}"
sources = [make_source(max_concurrency=100)]
limits = GlobalLimits(max_concurrency=2, rpm=0, tpm=0)
a = _limiter(a_cli, scope, sources, limits)
b = _limiter(b_cli, scope, sources, limits)
p1 = await a.try_acquire("s1", 1)
p2 = await b.try_acquire("s1", 1)
assert p1 is not None and p2 is not None
assert await a.try_acquire("s1", 1) is None # 全局 2 满,跨连接生效
await p1.release()
await p2.release()
async def test_cross_connection_global_rpm(clients):
"""CHS :105 升级为跨连接: 两池各消费全局 RPM 名额,第三个在任一连接都被拒。"""
a_cli, b_cli = clients
await await_window_headroom(a_cli)
scope = f"t{uuid4().hex[:8]}"
sources = [make_source("s1", rpm=100), make_source("s2", rpm=100)]
limits = GlobalLimits(max_concurrency=0, rpm=2, tpm=0)
a = _limiter(a_cli, scope, sources, limits)
b = _limiter(b_cli, scope, sources, limits)
p1 = await a.try_acquire("s1", 0)
p2 = await b.try_acquire("s2", 0)
assert p1 is not None and p2 is not None
await p1.release()
await p2.release()
assert await a.try_acquire("s1", 0) is None # RPM 不随 release 归还
assert await b.try_acquire("s2", 0) is None
async def test_progress_visible_across_connections(clients):
"""CHS :166: A mark 后 B(另一连接)立即读到 fresh age——背压活性是跨进程的。"""
a_cli, b_cli = clients
scope = f"t{uuid4().hex[:8]}"
sources = [make_source()]
a = _limiter(a_cli, scope, sources, GlobalLimits(0, 0, 0))
b = _limiter(b_cli, scope, sources, GlobalLimits(0, 0, 0))
assert await b.progress_age_s() == float("inf")
await a.mark_progress()
assert await b.progress_age_s() < 5.0
# —— 联合验证 ——
async def test_breaker_state_shared_across_connections(clients):
"""A 连接把源打开路,B 连接的 try_enter 立即被拒(熔断状态共享)。"""
a_cli, b_cli = clients
scope = f"t{uuid4().hex[:8]}"
gate_a = RedisGate(config=_CFG, redis=a_cli, scope=scope)
gate_b = RedisGate(config=_CFG, redis=b_cli, scope=scope)
for _ in range(_CFG.fail_threshold):
entry = await gate_a.try_enter("s1", "wa")
await gate_a.record_failure(entry, "network_error", False)
decision = await gate_b.try_enter("s1", "wb")
assert not decision.allowed and decision.retry_after_s > 0
async def test_two_clients_global_rpm_not_exceeded(clients):
"""双 GatewayClient(双连接池、同 scope)并发打满: 真实通过数 ≤ 全局 RPM。"""
a_cli, b_cli = clients
await await_window_headroom(a_cli)
scope = f"t{uuid4().hex[:8]}"
sources = [make_source(rpm=100, est_tokens=1)]
limits = GlobalLimits(max_concurrency=0, rpm=4, tpm=0)
transport = ScriptedTransport()
ca = _client(
scope,
sources,
_limiter(a_cli, scope, sources, limits),
RedisGate(config=_CFG, redis=a_cli, scope=scope),
transport,
quota_full="fail_fast",
)
cb = _client(
scope,
sources,
_limiter(b_cli, scope, sources, limits),
RedisGate(config=_CFG, redis=b_cli, scope=scope),
transport,
quota_full="fail_fast",
)
msgs = [{"role": "user", "content": "hi"}]
results = await asyncio.gather(
*(c.chat(msgs) for c in (ca, cb) for _ in range(4)), return_exceptions=True
)
ok = [r for r in results if not isinstance(r, BaseException)]
rejected = [r for r in results if isinstance(r, AllSourcesExhausted)]
assert len(transport.calls) == 4 # 8 并发中恰 4 个穿过全局 RPM 闸
assert len(ok) == 4 and len(rejected) == 4
assert all(r.reason == "quota_exhausted" for r in rejected)
async def test_cancel_in_flight_releases_lease(clients):
"""in-flight 取消 → permit 在 finally 释放,租约不泄漏(取消穿透铁律)。"""
a_cli, _ = clients
scope = f"t{uuid4().hex[:8]}"
sources = [make_source(max_concurrency=1)]
limiter = _limiter(a_cli, scope, sources, GlobalLimits(0, 0, 0))
client = _client(
scope,
sources,
limiter,
RedisGate(config=_CFG, redis=a_cli, scope=scope),
ScriptedTransport(hang=True),
)
task = asyncio.create_task(client.chat([{"role": "user", "content": "hi"}]))
while not (await limiter.source_stats("s1")).inflight:
await asyncio.sleep(0.02)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert (await limiter.source_stats("s1")).inflight == 0
# —— 掉线方向(fail-closed 集成证据)——
async def test_redis_down_admission_fails_closed():
"""Redis 不可达 → 准入侧抛 GovernanceBackendError,绝不放行(库铁律)。"""
import redis.asyncio as aioredis
dead = aioredis.from_url(
"redis://127.0.0.1:1/0", socket_connect_timeout=0.3, socket_timeout=0.3
)
try:
limiter = RedisLimiter(
scope="t-dead",
sources={"s1": make_source()},
global_limits=GlobalLimits(0, 0, 0),
redis=dead,
lease_ttl_s=30.0,
)
gate = RedisGate(config=_CFG, redis=dead, scope="t-dead")
with pytest.raises(GovernanceBackendError):
await limiter.try_acquire("s1", 0)
with pytest.raises(GovernanceBackendError):
await gate.try_enter("s1", "w")
finally:
await dead.aclose()