test: add integration suite for governance stack and redis cache
This commit is contained in:
@@ -0,0 +1,206 @@
|
||||
"""全栈治理组合集成测试: 完整洋葱(遥测→缓存→结构化→重试)+ 真实内存后端。
|
||||
|
||||
不依赖外部服务;熔断恢复全链用注入时钟推进,取消穿透用真实 asyncio 取消。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
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,
|
||||
),
|
||||
gate=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)
|
||||
@@ -0,0 +1,174 @@
|
||||
"""真实 Redis 缓存集成测试(实验室远程实例;CLAUDE.md: Redis 相关不 mock)。
|
||||
|
||||
前置: `.env`/环境变量提供 `REDIS_URL`;缺失时整组 skip(T13 门在验收前必须真跑)。
|
||||
隔离: 所有 key 落在一次性命名空间 `pgw:test:{uuid}`,teardown 精确清理,
|
||||
绝不 SCAN/FLUSH 全库,不触碰在用数据。
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from dotenv import dotenv_values
|
||||
|
||||
from polygateway.backends.redis_cache import RedisCache
|
||||
from polygateway.middleware.cache import CacheMW, build_cache_key
|
||||
from polygateway.types import ChatRequest, LLMResponse
|
||||
|
||||
_REDIS_URL = os.environ.get("REDIS_URL") or dotenv_values(".env").get("REDIS_URL")
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _REDIS_URL, reason="需实验室远程 Redis: 在 .env 设置 REDIS_URL(M1 验收前必须真跑)"
|
||||
)
|
||||
|
||||
_MSGS = [{"role": "user", "content": "integration-hi"}]
|
||||
|
||||
|
||||
def _resp(content="cached"):
|
||||
return LLMResponse(content, "", "m", "p", 1, 2, 30, None, None, False, "orig")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def redis_cache():
|
||||
import redis.asyncio as aioredis
|
||||
|
||||
client = aioredis.from_url(
|
||||
_REDIS_URL, decode_responses=True, socket_connect_timeout=3.0, socket_timeout=3.0
|
||||
)
|
||||
cache = RedisCache(client)
|
||||
used_keys: list[str] = []
|
||||
original_set = cache.set
|
||||
|
||||
async def tracking_set(key, value, ttl_s):
|
||||
used_keys.append(key)
|
||||
await original_set(key, value, ttl_s)
|
||||
|
||||
cache.set = tracking_set # 记录写入的 key,teardown 精确删除
|
||||
yield cache
|
||||
if used_keys:
|
||||
await client.delete(*used_keys)
|
||||
await client.aclose()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def namespace():
|
||||
return f"pgw:test:{uuid.uuid4().hex}"
|
||||
|
||||
|
||||
class TestRealRedisRoundtrip:
|
||||
async def test_set_get_ttl(self, redis_cache, namespace):
|
||||
key = build_cache_key("m", _MSGS, namespace, None)
|
||||
assert await redis_cache.get(key) is None
|
||||
await redis_cache.set(key, json.dumps({"v": 1}), 60)
|
||||
assert json.loads(await redis_cache.get(key)) == {"v": 1}
|
||||
|
||||
async def test_cache_mw_miss_then_hit(self, redis_cache, namespace):
|
||||
mw = CacheMW(
|
||||
backend=redis_cache,
|
||||
model_fingerprint="m",
|
||||
default_namespace=namespace,
|
||||
ttl_s=60,
|
||||
)
|
||||
calls = {"n": 0}
|
||||
|
||||
async def terminal(request):
|
||||
calls["n"] += 1
|
||||
return _resp()
|
||||
|
||||
first = await mw(ChatRequest(messages=_MSGS), terminal)
|
||||
second = await mw(ChatRequest(messages=_MSGS), terminal)
|
||||
assert first.cache_hit is False and second.cache_hit is True
|
||||
assert calls["n"] == 1
|
||||
|
||||
async def test_namespace_isolation_on_shared_instance(self, redis_cache, namespace):
|
||||
"""多项目共用同一 Redis 时,namespace 不同绝不互相命中(防毒化铁律)。"""
|
||||
other = f"pgw:test:{uuid.uuid4().hex}"
|
||||
k1 = build_cache_key("m", _MSGS, namespace, None)
|
||||
k2 = build_cache_key("m", _MSGS, other, None)
|
||||
assert k1 != k2
|
||||
await redis_cache.set(k1, json.dumps({"ns": 1}), 60)
|
||||
assert await redis_cache.get(k2) is None
|
||||
|
||||
|
||||
class TestDegradationAgainstDeadRedis:
|
||||
async def test_unreachable_redis_degrades_to_miss(self):
|
||||
"""断连方向: 缓存后端挂 → 静默降级,调用照常(铁律,与限流相反)。"""
|
||||
import redis.asyncio as aioredis
|
||||
|
||||
dead = RedisCache(
|
||||
aioredis.from_url(
|
||||
"redis://127.0.0.1:1/0",
|
||||
socket_connect_timeout=0.2,
|
||||
socket_timeout=0.2,
|
||||
decode_responses=True,
|
||||
)
|
||||
)
|
||||
mw = CacheMW(
|
||||
backend=dead, model_fingerprint="m", default_namespace="pgw:test:dead", ttl_s=60
|
||||
)
|
||||
|
||||
async def terminal(request):
|
||||
return _resp("alive")
|
||||
|
||||
resp = await mw(ChatRequest(messages=_MSGS), terminal)
|
||||
assert resp.content == "alive" and resp.cache_hit is False
|
||||
|
||||
|
||||
def _sse_response():
|
||||
chunk = json.dumps({"choices": [{"delta": {"content": "hello"}}]})
|
||||
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"})
|
||||
|
||||
|
||||
class TestClientWithRealRedis:
|
||||
async def test_from_settings_style_assembly_with_redis_cache(self, redis_cache, namespace):
|
||||
"""GatewayClient 组装 + 真实 Redis 缓存的端到端命中。"""
|
||||
from polygateway import GatewayClient
|
||||
from polygateway.backends.memory.breaker import InMemoryGate
|
||||
from polygateway.backends.memory.limiter import InMemoryLimiter
|
||||
from polygateway.sources import RoundRobinSelector
|
||||
from polygateway.structured.json_repair import JsonRepairStrategy
|
||||
from polygateway.transports.openai_compat import OpenAICompatTransport
|
||||
from polygateway.types import (
|
||||
BackpressurePolicy,
|
||||
BreakerConfig,
|
||||
GlobalLimits,
|
||||
RetryPolicy,
|
||||
SourceConfig,
|
||||
)
|
||||
|
||||
src = SourceConfig(
|
||||
name="qwen_1",
|
||||
provider="qwen",
|
||||
base_url="https://gw.example/v1",
|
||||
api_key="sk",
|
||||
model="qwen-max",
|
||||
timeout_s=5.0,
|
||||
)
|
||||
client = GatewayClient(
|
||||
scope="llm",
|
||||
sources=[src],
|
||||
selector=RoundRobinSelector(),
|
||||
limiter=InMemoryLimiter(
|
||||
scope="llm", sources={src.name: src}, global_limits=GlobalLimits(0, 0, 0)
|
||||
),
|
||||
gate=InMemoryGate(config=BreakerConfig(5, 60.0, 120.0)),
|
||||
transport=OpenAICompatTransport(
|
||||
client_factory=lambda s: httpx.AsyncClient(
|
||||
transport=httpx.MockTransport(lambda req: _sse_response())
|
||||
)
|
||||
),
|
||||
retry=RetryPolicy(3, 2.0, 30.0),
|
||||
backpressure=BackpressurePolicy(300.0, 0.01),
|
||||
cache=redis_cache,
|
||||
cache_namespace=namespace,
|
||||
cache_ttl_s=60,
|
||||
structured_strategy=JsonRepairStrategy(),
|
||||
)
|
||||
first = await client.chat(_MSGS)
|
||||
second = await client.chat(_MSGS)
|
||||
assert first.cache_hit is False and second.cache_hit is True
|
||||
assert second.content == "hello"
|
||||
Reference in New Issue
Block a user