|
1 | | -""" |
2 | | -熵减 AI 网关 — pytest 公共夹具 |
3 | | -
|
4 | | -提供测试用的 FastAPI app、TestClient、mock Provider 等。 |
5 | | -""" |
6 | | - |
7 | | -import sys |
8 | | -import os |
9 | | -import asyncio |
10 | | -from pathlib import Path |
11 | | -from unittest.mock import AsyncMock, MagicMock |
12 | | - |
13 | | -import pytest |
14 | | - |
15 | | -# 确保 ai-gateway 根目录在 sys.path 中,使 config/errors 等可导入 |
16 | | -GATEWAY_ROOT = str(Path(__file__).resolve().parent.parent) |
17 | | -if GATEWAY_ROOT not in sys.path: |
18 | | - sys.path.insert(0, GATEWAY_ROOT) |
19 | | - |
20 | | - |
21 | | -# ──────────────────────────────────────────────────────────── |
22 | | -# Event loop(pytest-asyncio 需要) |
23 | | -# ──────────────────────────────────────────────────────────── |
24 | | - |
25 | | -@pytest.fixture(scope="session") |
26 | | -def event_loop(): |
27 | | - """创建全局事件循环,供所有 async 测试共享""" |
28 | | - loop = asyncio.new_event_loop() |
29 | | - yield loop |
30 | | - loop.close() |
31 | | - |
32 | | - |
33 | | -# ──────────────────────────────────────────────────────────── |
34 | | -# Mock Provider |
35 | | -# ──────────────────────────────────────────────────────────── |
36 | | - |
37 | | -class MockProvider: |
38 | | - """模拟 AI Provider,可自定义 generate 返回值""" |
39 | | - |
40 | | - def __init__(self, name: str = "mock", response: dict | None = None): |
41 | | - self.provider_name = name |
42 | | - self.api_key = "mock-key" |
43 | | - self._response = response or { |
44 | | - "content": "这是模拟的 AI 响应内容", |
45 | | - "tokens_used": 100, |
46 | | - "model": "mock-model", |
47 | | - "latency_ms": 50, |
48 | | - } |
49 | | - |
50 | | - async def generate(self, prompt, system_prompt="", model="", temperature=0.7, |
51 | | - max_tokens=2048, response_format=None, **kwargs): |
52 | | - return self._response.copy() |
53 | | - |
54 | | - async def health_check(self): |
55 | | - return {"status": "healthy", "latency_ms": 1.0, "error": None} |
56 | | - |
57 | | - |
58 | | -class FailingProvider: |
59 | | - """总是抛出异常的 Provider,用于测试 fallback 链""" |
60 | | - |
61 | | - def __init__(self, name: str = "failing"): |
62 | | - self.provider_name = name |
63 | | - self.api_key = "mock-key" |
64 | | - |
65 | | - async def generate(self, *args, **kwargs): |
66 | | - raise RuntimeError(f"Provider [{self.provider_name}] 模拟故障") |
67 | | - |
68 | | - async def health_check(self): |
69 | | - return {"status": "unhealthy", "latency_ms": 0, "error": "模拟故障"} |
70 | | - |
71 | | - |
72 | | -@pytest.fixture |
73 | | -def mock_provider(): |
74 | | - """返回一个可用的 mock Provider""" |
75 | | - return MockProvider() |
76 | | - |
77 | | - |
78 | | -@pytest.fixture |
79 | | -def failing_provider(): |
80 | | - """返回一个总是失败的 Provider""" |
81 | | - return FailingProvider() |
82 | | - |
83 | | - |
84 | | -# ──────────────────────────────────────────────────────────── |
85 | | -# 测试用 FastAPI 应用(绕过 JWT / RateLimit 中间件) |
86 | | -# ──────────────────────────────────────────────────────────── |
87 | | - |
88 | | -@pytest.fixture |
89 | | -def test_app(): |
90 | | - """ |
91 | | - 创建精简版 FastAPI 应用,只注册路由,不挂中间件。 |
92 | | - 在 app.state.providers 中注入 mock Provider。 |
93 | | - """ |
94 | | - from fastapi import FastAPI |
95 | | - from routers import ( |
96 | | - summarize_router, |
97 | | - generate_cards_router, |
98 | | - evaluate_router, |
99 | | - recommend_router, |
100 | | - ) |
101 | | - |
102 | | - app = FastAPI() |
103 | | - |
104 | | - # 注入 mock providers(call_with_fallback 会从 app.state.providers 读取) |
105 | | - mock = MockProvider(name="qwen") |
106 | | - app.state.providers = { |
107 | | - "qwen": mock, |
108 | | - "deepseek": MockProvider(name="deepseek"), |
109 | | - "glm": MockProvider(name="glm"), |
110 | | - "fallback": mock, |
111 | | - } |
112 | | - |
113 | | - app.include_router(summarize_router) |
114 | | - app.include_router(generate_cards_router) |
115 | | - app.include_router(evaluate_router) |
116 | | - app.include_router(recommend_router) |
117 | | - |
118 | | - return app |
119 | | - |
120 | | - |
121 | | -@pytest.fixture |
122 | | -def client(test_app): |
123 | | - """同步 TestClient,用于路由测试""" |
124 | | - from fastapi.testclient import TestClient |
125 | | - return TestClient(test_app) |
126 | | - |
127 | | - |
128 | | -# ──────────────────────────────────────────────────────────── |
129 | | -# 模拟 call_with_fallback(直接调用主 Provider,不走 fallback 链) |
130 | | -# ──────────────────────────────────────────────────────────── |
131 | | - |
132 | | -@pytest.fixture |
133 | | -def mock_call_with_fallback(monkeypatch): |
134 | | - """ |
135 | | - Patch config.call_with_fallback,直接调用第一个可用 Provider。 |
136 | | - 返回 fixture 函数:调用后获得 patch 上下文。 |
137 | | - """ |
138 | | - async def fake_call(app, feature, fn): |
139 | | - provider = list(app.state.providers.values())[0] |
140 | | - from config import MODEL_ROUTING, AI_PROVIDERS |
141 | | - routing = MODEL_ROUTING.get(feature, ("fallback", "free")) |
142 | | - model_name = AI_PROVIDERS.get(routing[0], {}).get("models", {}).get(routing[1], "mock-model") |
143 | | - result = await fn(provider, model_name) |
144 | | - return result, routing[0] |
145 | | - |
146 | | - monkeypatch.setattr("config.call_with_fallback", fake_call_with_fallback) |
147 | | - return fake_call_with_fallback |
| 1 | +""" |
| 2 | +熵减 AI 网关 — pytest 公共夹具 |
| 3 | +
|
| 4 | +提供测试用的 FastAPI app、TestClient、mock Provider 等。 |
| 5 | +""" |
| 6 | + |
| 7 | +import sys |
| 8 | +import os |
| 9 | +import asyncio |
| 10 | +from pathlib import Path |
| 11 | +from unittest.mock import AsyncMock, MagicMock |
| 12 | + |
| 13 | +import pytest |
| 14 | + |
| 15 | +# 确保 ai-gateway 根目录在 sys.path 中,使 config/errors 等可导入 |
| 16 | +GATEWAY_ROOT = str(Path(__file__).resolve().parent.parent) |
| 17 | +if GATEWAY_ROOT not in sys.path: |
| 18 | + sys.path.insert(0, GATEWAY_ROOT) |
| 19 | + |
| 20 | + |
| 21 | +# ──────────────────────────────────────────────────────────── |
| 22 | +# Event loop(pytest-asyncio 需要) |
| 23 | +# ──────────────────────────────────────────────────────────── |
| 24 | + |
| 25 | +@pytest.fixture(scope="session") |
| 26 | +def event_loop(): |
| 27 | + """创建全局事件循环,供所有 async 测试共享""" |
| 28 | + loop = asyncio.new_event_loop() |
| 29 | + yield loop |
| 30 | + loop.close() |
| 31 | + |
| 32 | + |
| 33 | +# ──────────────────────────────────────────────────────────── |
| 34 | +# Mock Provider |
| 35 | +# ──────────────────────────────────────────────────────────── |
| 36 | + |
| 37 | +class MockProvider: |
| 38 | + """模拟 AI Provider,可自定义 generate 返回值""" |
| 39 | + |
| 40 | + def __init__(self, name: str = "mock", response: dict | None = None): |
| 41 | + self.provider_name = name |
| 42 | + self.api_key = "mock-key" |
| 43 | + self._response = response or { |
| 44 | + "content": "这是模拟的 AI 响应内容", |
| 45 | + "tokens_used": 100, |
| 46 | + "model": "mock-model", |
| 47 | + "latency_ms": 50, |
| 48 | + } |
| 49 | + |
| 50 | + async def generate(self, prompt, system_prompt="", model="", temperature=0.7, |
| 51 | + max_tokens=2048, response_format=None, **kwargs): |
| 52 | + return self._response.copy() |
| 53 | + |
| 54 | + async def health_check(self): |
| 55 | + return {"status": "healthy", "latency_ms": 1.0, "error": None} |
| 56 | + |
| 57 | + |
| 58 | +class FailingProvider: |
| 59 | + """总是抛出异常的 Provider,用于测试 fallback 链""" |
| 60 | + |
| 61 | + def __init__(self, name: str = "failing"): |
| 62 | + self.provider_name = name |
| 63 | + self.api_key = "mock-key" |
| 64 | + |
| 65 | + async def generate(self, *args, **kwargs): |
| 66 | + raise RuntimeError(f"Provider [{self.provider_name}] 模拟故障") |
| 67 | + |
| 68 | + async def health_check(self): |
| 69 | + return {"status": "unhealthy", "latency_ms": 0, "error": "模拟故障"} |
| 70 | + |
| 71 | + |
| 72 | +@pytest.fixture |
| 73 | +def mock_provider(): |
| 74 | + """返回一个可用的 mock Provider""" |
| 75 | + return MockProvider() |
| 76 | + |
| 77 | + |
| 78 | +@pytest.fixture |
| 79 | +def failing_provider(): |
| 80 | + """返回一个总是失败的 Provider""" |
| 81 | + return FailingProvider() |
| 82 | + |
| 83 | + |
| 84 | +# ──────────────────────────────────────────────────────────── |
| 85 | +# 测试用 FastAPI 应用(绕过 JWT / RateLimit 中间件) |
| 86 | +# ──────────────────────────────────────────────────────────── |
| 87 | + |
| 88 | +@pytest.fixture |
| 89 | +def test_app(): |
| 90 | + """ |
| 91 | + 创建精简版 FastAPI 应用,只注册路由,不挂中间件。 |
| 92 | + 在 app.state.providers 中注入 mock Provider。 |
| 93 | + """ |
| 94 | + from fastapi import FastAPI |
| 95 | + from routers import ( |
| 96 | + summarize_router, |
| 97 | + generate_cards_router, |
| 98 | + evaluate_router, |
| 99 | + recommend_router, |
| 100 | + ) |
| 101 | + |
| 102 | + app = FastAPI() |
| 103 | + |
| 104 | + # 注入 mock providers(call_with_fallback 会从 app.state.providers 读取) |
| 105 | + mock = MockProvider(name="qwen") |
| 106 | + app.state.providers = { |
| 107 | + "qwen": mock, |
| 108 | + "deepseek": MockProvider(name="deepseek"), |
| 109 | + "glm": MockProvider(name="glm"), |
| 110 | + "fallback": mock, |
| 111 | + } |
| 112 | + |
| 113 | + app.include_router(summarize_router) |
| 114 | + app.include_router(generate_cards_router) |
| 115 | + app.include_router(evaluate_router) |
| 116 | + app.include_router(recommend_router) |
| 117 | + |
| 118 | + return app |
| 119 | + |
| 120 | + |
| 121 | +@pytest.fixture |
| 122 | +def client(test_app): |
| 123 | + """同步 TestClient,用于路由测试""" |
| 124 | + from fastapi.testclient import TestClient |
| 125 | + return TestClient(test_app) |
| 126 | + |
| 127 | + |
| 128 | +# ──────────────────────────────────────────────────────────── |
| 129 | +# 模拟 call_with_fallback(直接调用主 Provider,不走 fallback 链) |
| 130 | +# ──────────────────────────────────────────────────────────── |
| 131 | + |
| 132 | +@pytest.fixture |
| 133 | +def mock_call_with_fallback(monkeypatch): |
| 134 | + """ |
| 135 | + Patch config.call_with_fallback,直接调用第一个可用 Provider。 |
| 136 | + 返回 fixture 函数:调用后获得 patch 上下文。 |
| 137 | + """ |
| 138 | + async def fake_call(app, feature, fn): |
| 139 | + provider = list(app.state.providers.values())[0] |
| 140 | + from config import MODEL_ROUTING, AI_PROVIDERS |
| 141 | + routing = MODEL_ROUTING.get(feature, ("fallback", "free")) |
| 142 | + model_name = AI_PROVIDERS.get(routing[0], {}).get("models", {}).get(routing[1], "mock-model") |
| 143 | + result = await fn(provider, model_name) |
| 144 | + return result, routing[0] |
| 145 | + |
| 146 | + monkeypatch.setattr("config.call_with_fallback", fake_call) |
| 147 | + return fake_call |
0 commit comments