Files
PolyGateway/tests/integration/test_governance_stack.py
T

279 lines
10 KiB
Python

"""全栈治理组合集成测试: 完整洋葱(遥测→缓存→结构化→重试)+ 真实内存后端。
不依赖外部服务;熔断恢复全链用注入时钟推进,取消穿透用真实 asyncio 取消。
"""
import asyncio
import dataclasses
import json
import sqlite3
import httpx
import pytest
from polygateway import CircuitOpenError, GatewayClient, TransientError
from polygateway.backends.memory.breaker import InMemoryGate
from polygateway.backends.memory.cache import InMemoryCache
from polygateway.backends.memory.limiter import InMemoryLimiter
from polygateway.sources import RoundRobinSelector
from polygateway.structured.json_repair import JsonRepairStrategy
from polygateway.telemetry.sqlite import SQLiteRecorder
from polygateway.transports.openai_compat import OpenAICompatTransport
from polygateway.types import (
BackpressurePolicy,
BreakerConfig,
GlobalLimits,
RetryPolicy,
SourceConfig,
)
from tests.contracts.conftest import FakeClock
_BREAKER = BreakerConfig(fail_threshold=2, cooldown_s=60.0, probe_ttl_s=120.0)
def _source(name="qwen_1"):
return SourceConfig(
name=name,
provider="qwen",
base_url="https://gw.example/v1",
api_key="sk",
model="qwen-max",
timeout_s=5.0,
)
def _sse(content="ok"):
chunk = json.dumps({"choices": [{"delta": {"content": content}}]})
usage = json.dumps({"choices": [], "usage": {"prompt_tokens": 3, "completion_tokens": 4}})
body = f"data: {chunk}\n\ndata: {usage}\n\ndata: [DONE]\n\n"
return httpx.Response(200, content=body.encode(), headers={"content-type": "text/event-stream"})
async def _noop_sleep(seconds):
return None
def _full_client(handler, *, clock=None, telemetry=None, cache=None):
clock = clock or FakeClock()
src = _source()
return GatewayClient(
scope="llm",
sources=[src],
selector=RoundRobinSelector(),
limiter=InMemoryLimiter(
scope="llm",
sources={src.name: src},
global_limits=GlobalLimits(0, 0, 0),
now=clock,
),
breaker=InMemoryGate(config=_BREAKER, now=clock),
transport=OpenAICompatTransport(
client_factory=lambda s: httpx.AsyncClient(transport=httpx.MockTransport(handler))
),
retry=RetryPolicy(2, 2.0, 30.0),
backpressure=BackpressurePolicy(300.0, 0.01),
telemetry=telemetry,
cache=cache,
cache_namespace="itest" if cache else None,
cache_ttl_s=3600 if cache else None,
structured_strategy=JsonRepairStrategy(),
now=clock,
sleep=_noop_sleep,
)
class TestBreakerRecoveryFullChain:
async def test_open_cooldown_probe_close_cycle(self):
"""开路 → 冷却 → 半开探针 → 恢复闭路,经完整 client 洋葱走通。"""
clock = FakeClock()
state = {"fail": True}
def handler(request):
if state["fail"]:
return httpx.Response(503, content=b"{}")
return _sse("recovered")
client = _full_client(handler, clock=clock)
# 2 次尝试全 503 → retry_exhausted;熔断计 2 次失败达阈值开路
with pytest.raises(Exception) as ei:
await client.chat([{"role": "user", "content": "hi"}])
assert "retry_exhausted" in str(ei.value)
# 开路期间: 直接 CircuitOpenError,不打网关
with pytest.raises(CircuitOpenError) as open_err:
await client.chat([{"role": "user", "content": "hi"}])
assert open_err.value.retry_after_s > 0
# 冷却到期 + 网关恢复 → 探针成功闭路
clock.advance(_BREAKER.cooldown_s + 1)
state["fail"] = False
resp = await client.chat([{"role": "user", "content": "hi"}])
assert resp.content == "recovered"
# 闭路后正常服务
resp2 = await client.chat([{"role": "user", "content": "hi2"}])
assert resp2.content == "recovered"
class TestCancellationThroughStack:
async def test_cancel_mid_request_releases_and_records(self, tmp_path):
recorder = SQLiteRecorder(tmp_path / "t.db")
entered = asyncio.Event()
async def hanging_handler(request):
entered.set()
await asyncio.sleep(30)
limiter_probe = {}
client = _full_client(hanging_handler, telemetry=recorder)
limiter_probe["limiter"] = client._handler # noqa: SLF001 — 仅为断言持引用
task = asyncio.ensure_future(client.chat([{"role": "user", "content": "hi"}]))
await entered.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
recorder.close()
rows = sqlite3.connect(tmp_path / "t.db").execute("SELECT error FROM llm_calls").fetchall()
# 尽力而为遥测: 取消路径留痕(尝试级与最外层各一行,均 error=cancelled)
assert rows and all(r[0] == "cancelled" for r in rows)
class TestTelemetryAcrossPaths:
async def test_success_cache_hit_and_failure_rows(self, tmp_path):
recorder = SQLiteRecorder(tmp_path / "t.db")
client = _full_client(lambda req: _sse(), telemetry=recorder, cache=InMemoryCache())
await client.chat([{"role": "user", "content": "hi"}]) # 成功(尝试行)
await client.chat([{"role": "user", "content": "hi"}]) # 缓存命中行
recorder.close()
conn = sqlite3.connect(tmp_path / "t.db")
(hits,) = conn.execute("SELECT COUNT(*) FROM llm_calls WHERE cache_hit=1").fetchone()
(total,) = conn.execute("SELECT COUNT(*) FROM llm_calls").fetchone()
assert hits == 1 and total == 2
async def test_transient_attempts_each_recorded(self, tmp_path):
recorder = SQLiteRecorder(tmp_path / "t.db")
calls = {"n": 0}
def flaky(request):
calls["n"] += 1
if calls["n"] == 1:
return httpx.Response(503, content=b"{}")
return _sse()
client = _full_client(flaky, telemetry=recorder)
resp = await client.chat([{"role": "user", "content": "hi"}])
assert resp.content == "ok"
recorder.close()
rows = (
sqlite3.connect(tmp_path / "t.db")
.execute("SELECT error IS NULL, call_id FROM llm_calls ORDER BY created_at")
.fetchall()
)
assert len(rows) == 2 # 失败尝试 + 成功尝试各一行
assert {ok for ok, _ in rows} == {0, 1}
assert len({cid for _, cid in rows}) == 2 # call_id 逐次独立
class TestStructuredThroughStack:
async def test_feedback_reask_passes_through_governance(self):
"""重问经过内层治理: 第二次真实请求同样被限流/熔断记账。"""
from pydantic import BaseModel
class Out(BaseModel):
answer: int
contents = ['{"answer": "bad"}', '{"answer": 7}']
def handler(request):
return _sse(contents.pop(0))
client = _full_client(handler)
resp = await client.chat([{"role": "user", "content": "hi"}], structured=Out)
assert resp.structured_data.answer == 7
class TestTransientErrorExport:
async def test_business_side_catches_top_level_errors(self):
"""迁移承诺: 业务侧 (TimeoutError, OSError) 元组换成库异常后可捕获。"""
def always_503(request):
return httpx.Response(503, content=b"{}")
client = _full_client(always_503)
with pytest.raises(Exception) as ei:
await client.chat([{"role": "user", "content": "hi"}])
import polygateway
assert isinstance(ei.value, polygateway.AllSourcesExhausted)
assert isinstance(ei.value.__cause__, TransientError)
class TestSamplingThroughStack:
"""issue #4: 采样参数经完整洋葱到达请求体,且缓存/遥测口径一致。"""
async def test_reaches_wire_and_lands_in_telemetry(self, tmp_path):
seen = []
def handler(request):
seen.append(json.loads(request.content))
return _sse()
db = tmp_path / "t.db"
recorder = SQLiteRecorder(db)
client = _full_client(handler, telemetry=recorder)
await client.chat([{"role": "user", "content": "hi"}], overlay={"seed": 42})
recorder.close()
assert seen[0]["seed"] == 42 # 穿过全栈到达线上
rows = sqlite3.connect(db).execute("SELECT sampling FROM llm_calls").fetchall()
assert json.loads(rows[0][0]) == {"seed": 42}
async def test_config_level_merges_and_records(self, tmp_path):
"""源级 extra_body 只有 emit_attempt 记得到(唯一有生效源的入口)。"""
seen = []
def handler(request):
seen.append(json.loads(request.content))
return _sse()
src = dataclasses.replace(_source(), extra_body={"temperature": 0})
db = tmp_path / "t.db"
recorder = SQLiteRecorder(db)
client = GatewayClient(
scope="llm",
sources=[src],
selector=RoundRobinSelector(),
limiter=InMemoryLimiter(
scope="llm", sources={src.name: src}, global_limits=GlobalLimits(0, 0, 0)
),
breaker=InMemoryGate(config=_BREAKER),
transport=OpenAICompatTransport(
client_factory=lambda s: httpx.AsyncClient(transport=httpx.MockTransport(handler))
),
retry=RetryPolicy(2, 2.0, 30.0),
backpressure=BackpressurePolicy(300.0, 0.01),
telemetry=recorder,
structured_strategy=JsonRepairStrategy(),
sleep=_noop_sleep,
)
await client.chat([{"role": "user", "content": "hi"}], overlay={"seed": 1})
recorder.close()
assert seen[0]["temperature"] == 0 and seen[0]["seed"] == 1
rows = sqlite3.connect(db).execute("SELECT sampling FROM llm_calls").fetchall()
assert json.loads(rows[0][0]) == {"seed": 1, "temperature": 0}
async def test_differing_seed_bypasses_cache_end_to_end(self):
"""issue 场景全栈回归: 逐 rollout 变 seed 必须真的回源。"""
calls = []
def handler(request):
calls.append(json.loads(request.content)["seed"])
return _sse()
client = _full_client(handler, cache=InMemoryCache())
await client.chat([{"role": "user", "content": "hi"}], overlay={"seed": 1})
await client.chat([{"role": "user", "content": "hi"}], overlay={"seed": 2})
second_same = await client.chat([{"role": "user", "content": "hi"}], overlay={"seed": 1})
assert calls == [1, 2] # 两个不同 seed 各自回源
assert second_same.cache_hit is True # 同 seed 才命中