Files
PolyGateway/tests/integration/test_governance_stack.py
T
iomgaa 9dada0be9d test: prove a rejected call's reason reaches the telemetry table
Issue #10 Task 5, the acceptance claim. Before the fix this asserted
against 'qwen_1 请求被拒: 400' and failed on the first substring - which
is exactly what the downstream batch was left with. Uses the real body
from the issue, and checks the trailing code too, since a head-only cut
would drop the one field you quote when chasing the provider.
2026-08-16 06:17:10 -04:00

318 lines
12 KiB
Python

"""全栈治理组合集成测试: 完整洋葱(遥测→缓存→结构化→重试)+ 真实内存后端。
不依赖外部服务;熔断恢复全链用注入时钟推进,取消穿透用真实 asyncio 取消。
"""
import asyncio
import dataclasses
import json
import sqlite3
import httpx
import pytest
from polygateway import (
CircuitOpenError,
GatewayClient,
RequestRejectedError,
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 TestRejectionReasonIsQueryable:
"""issue #10 的验收主张: 400 之后,网关说的话必须能在遥测表里查到。
下游一轮 1050 张影像的批处理里,1 张在读表格时收到 400 被判确定性失败,
事后"这张图到底哪里不合规"无从查起——响应体在 transport 翻译层就没了。
"""
# issue #10 原文给出的真实响应体(一字不改)
_BODY = (
'{"error":{"message":"<400> ***.***.InvalidParameter: The image format is illegal '
'and cannot be opened","type":"invalid_request_error","param":"",'
'"code":"invalid_parameter_error"}}'
)
async def test_rejected_call_leaves_the_reason_in_telemetry(self, tmp_path):
recorder = SQLiteRecorder(tmp_path / "t.db")
client = _full_client(
lambda req: httpx.Response(400, content=self._BODY.encode()), telemetry=recorder
)
with pytest.raises(RequestRejectedError):
await client.chat([{"role": "user", "content": "hi"}])
recorder.close()
rows = sqlite3.connect(tmp_path / "t.db").execute("SELECT error FROM llm_calls").fetchall()
assert rows, "400 必须留下遥测行(遥测必录)"
errors = " ".join(r[0] or "" for r in rows)
# 修复前这里只有 "qwen_1 请求被拒: 400"——诊断信息一个字都不在
assert "InvalidParameter" in errors
assert "The image format is illegal" in errors
# 尾部的 code 才是向网关方追查的凭据,头部硬切正好会丢掉它
assert "invalid_parameter_error" in errors
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 才命中