Skip to content

Commit f0b92a9

Browse files
committed
fix(gateway): resolve medium-risk issues #5-#10
- #5: Gemini generate_stream iteration moved into thread pool via asyncio.to_thread, unblocking the event loop during streaming - #6: providers return real input/output token split; fallback cost tracking prefers it and falls back to 60/40 estimate instead of the misleading 50/50 split (output priced 2-3x higher) - #7: balance queries use Key pool primary key (plural env vars), DeepSeek balance no longer silently skipped - #8: JWKS fetch distinguishes network failures (stale cache reuse with 60s retry backoff) from HTTP/key errors (fail-closed), preventing whole-site 401 on transient Supabase hiccups - #9: import_concept registered in TIMEOUT_CONFIG/RATE_LIMITS plus startup validation warning for unregistered features - #10: error_pattern cache key includes user_id and hashes full content, preventing cross-user result reuse
1 parent 1aeeac5 commit f0b92a9

10 files changed

Lines changed: 117 additions & 9 deletions

File tree

‎server/ai-gateway/config/fallback.py‎

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -214,12 +214,22 @@ async def _do_call():
214214
try:
215215
tokens_used = result.get("tokens_used", 0)
216216
model_name = result.get("model", model_name)
217+
# GW-2#6: 优先使用 provider 返回的真实 input/output 拆分
218+
#(OpenAI 兼容 usage 均有 prompt/completion tokens);
219+
# 缺失时按经验比例估算(output 占比 60%)——原对半拆分
220+
# 系统性低估费用(output 单价普遍 2-3 倍于 input),
221+
# 导致超预算用户未被拦截
222+
input_tokens = result.get("input_tokens", 0)
223+
output_tokens = result.get("output_tokens", 0)
224+
if not input_tokens and not output_tokens and tokens_used:
225+
output_tokens = int(tokens_used * 0.6)
226+
input_tokens = tokens_used - output_tokens
217227
await get_cost_tracker().record(
218228
user_id=user_id,
219229
feature=feature,
220230
model=model_name,
221-
input_tokens=tokens_used // 2 if tokens_used else 0,
222-
output_tokens=tokens_used // 2 if tokens_used else 0,
231+
input_tokens=input_tokens,
232+
output_tokens=output_tokens,
223233
)
224234
except Exception as cost_err:
225235
logger.debug("CostTracker 记录失败(可忽略): %s", cost_err)

‎server/ai-gateway/config/limits.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,8 @@
5050
"conflict_detect": 30,
5151
# E1: 概念预检(JSON Mode,1-2 个探测问题)
5252
"concept_precheck": 30,
53+
# 知识入籍概念化(JSON Mode,切块文本→概念候选,与同批 JSON Mode 链一致)
54+
"import_concept": 30,
5355
}
5456

5557
# ============================================================
@@ -100,4 +102,6 @@
100102
"conflict_detect": 5,
101103
# E1: 概念预检(生成成本高,严格限频)
102104
"concept_precheck": 5,
105+
# 知识入籍概念化(批量导入场景,适度限频)
106+
"import_concept": 10,
103107
}

‎server/ai-gateway/middleware/auth.py‎

Lines changed: 33 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -143,6 +143,14 @@ def _normalize_pem_key(raw: str) -> str:
143143
_jwks_cache: Optional[dict] = None
144144
_jwks_cache_time: float = 0.0
145145
_JWKS_CACHE_TTL: int = 3600 # 缓存有效期(秒),默认 1 小时
146+
# GW-2#8: 网络类失败时允许复用过期缓存的宽限窗口(秒)——密钥轮换的越权
147+
# 风险与网络抖动导致的可用性风险权衡:JWKS 1 小时才刷新一次,期间 Supabase
148+
# 任何一次网络故障若 fail-closed 会让全部请求 401(认证风暴)
149+
_JWKS_STALE_GRACE_SECONDS: int = 300
150+
# GW-2#8: 网络失败后的重试退避(秒),避免每次请求都打 JWKS 端点
151+
_JWKS_FAIL_RETRY_SECONDS: int = 60
152+
# 网络失败后下次刷新尝试的最早时间戳(0 表示无退避)
153+
_jwks_retry_after: float = 0.0
146154
# GW-M3: 防缓存击穿锁 + 共享连接池客户端(避免每次请求新建 TCP/TLS 连接)
147155
_jwks_lock: asyncio.Lock = asyncio.Lock()
148156
_jwks_client: Optional[httpx.AsyncClient] = None
@@ -181,11 +189,21 @@ async def _fetch_jwks(jwks_url: str) -> dict:
181189
GW-M3: 网络失败时 fail-closed(拒绝验证)而非返回过期缓存——
182190
密钥轮换后过期缓存会让旧 token 在验证窗口内继续有效(越权窗口)。
183191
asyncio.Lock 防止冷启动多请求并发刷新(缓存击穿)。
192+
GW-2#8: 区分两类失败——HTTP 状态错误(密钥轮换/端点失效)保持
193+
fail-closed;网络类失败(超时/DNS/TLS/连接)短暂复用过期缓存
194+
(最多 5 分钟宽限)并退避 60s 重试,避免 Supabase 短暂抖动全站 401。
184195
"""
185-
global _jwks_cache, _jwks_cache_time
196+
global _jwks_cache, _jwks_cache_time, _jwks_retry_after
186197
now = time.time()
187198
if _jwks_cache is not None and (now - _jwks_cache_time) < _JWKS_CACHE_TTL:
188199
return _jwks_cache
200+
# 缓存过期但在宽限窗口内且未到重试时间:降级复用(网络抖动自愈窗口)
201+
if (
202+
_jwks_cache is not None
203+
and (now - _jwks_cache_time) < _JWKS_CACHE_TTL + _JWKS_STALE_GRACE_SECONDS
204+
and now < _jwks_retry_after
205+
):
206+
return _jwks_cache
189207
async with _jwks_lock:
190208
# 双重检查:等待锁期间其他协程可能已刷新缓存
191209
now = time.time()
@@ -196,10 +214,23 @@ async def _fetch_jwks(jwks_url: str) -> dict:
196214
resp.raise_for_status()
197215
_jwks_cache = resp.json()
198216
_jwks_cache_time = time.time()
217+
_jwks_retry_after = 0.0
199218
logger.info("JWKS 获取成功,缓存已更新: %s", jwks_url)
200219
return _jwks_cache
220+
except httpx.HTTPStatusError as e:
221+
# HTTP 状态错误(401/404/5xx):密钥轮换或端点失效,fail-closed
222+
logger.error("JWKS 获取失败: %s, status=%s", jwks_url, e.response.status_code)
223+
raise
201224
except Exception as e:
202-
logger.error("JWKS 获取失败: %s, error=%s", jwks_url, str(e))
225+
# 网络类失败(超时/DNS/TLS/连接中断):短暂复用过期缓存降级
226+
logger.error("JWKS 获取失败(网络): %s, error=%s", jwks_url, str(e))
227+
if _jwks_cache is not None:
228+
_jwks_retry_after = time.time() + _JWKS_FAIL_RETRY_SECONDS
229+
logger.warning(
230+
"JWKS 网络失败,降级复用过期缓存(%ds 内重试)",
231+
_JWKS_FAIL_RETRY_SECONDS,
232+
)
233+
return _jwks_cache
203234
raise
204235

205236

‎server/ai-gateway/middleware/rate_limit.py‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from starlette.middleware.base import BaseHTTPMiddleware
1919

2020
from cache.redis_cache import get_cache
21-
from config import RATE_LIMITS
21+
from config import RATE_LIMITS, TIMEOUT_CONFIG
2222

2323
logger = logging.getLogger(__name__)
2424

@@ -108,6 +108,19 @@
108108
# 仅受各自功能级上限约束
109109
GLOBAL_EXEMPT_FEATURES: frozenset[str] = frozenset({"transcribe", "chat"})
110110

111+
# GW-2#9: 启动校验——PATH_TO_FEATURE 登记的 feature 必须同时在
112+
# TIMEOUT_CONFIG 与 RATE_LIMITS 登记,否则静默兜底(300s 超时/默认 10 次限流)
113+
# 会掩盖配置缺失(import_concept 曾因此落入 300s 超时预算)
114+
_MISSING_CONFIG = sorted(
115+
f for f in set(PATH_TO_FEATURE.values())
116+
if f not in RATE_LIMITS or f not in TIMEOUT_CONFIG
117+
)
118+
if _MISSING_CONFIG:
119+
logger.warning(
120+
"以下 feature 缺少 TIMEOUT_CONFIG/RATE_LIMITS 登记(将使用兜底值): %s",
121+
_MISSING_CONFIG,
122+
)
123+
111124

112125
class RateLimitMiddleware(BaseHTTPMiddleware):
113126
"""频率限制中间件 — 基于 Redis 滑动窗口"""

‎server/ai-gateway/providers/deepseek_provider.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -97,8 +97,14 @@ async def generate(
9797
# 提取结果
9898
content = response.choices[0].message.content or ""
9999
tokens_used = 0
100+
input_tokens = 0
101+
output_tokens = 0
100102
if response.usage:
101103
tokens_used = response.usage.total_tokens or 0
104+
# GW-2#6: 提取真实 input/output 拆分供成本记账(OpenAI 兼容
105+
# usage 字段),fallback 链不再对半估算
106+
input_tokens = getattr(response.usage, "prompt_tokens", 0) or 0
107+
output_tokens = getattr(response.usage, "completion_tokens", 0) or 0
102108

103109
# DeepSeek 缓存命中信息(如果有的话)
104110
cache_hit_tokens = 0
@@ -132,6 +138,8 @@ async def generate(
132138
return {
133139
"content": content,
134140
"tokens_used": tokens_used,
141+
"input_tokens": input_tokens,
142+
"output_tokens": output_tokens,
135143
"effective_tokens": effective_tokens,
136144
"cache_hit": cache_hit,
137145
"cache_hit_tokens": cache_hit_tokens,

‎server/ai-gateway/providers/gemini_provider.py‎

Lines changed: 20 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -91,11 +91,17 @@ async def generate(
9191
latency_ms = int((time.monotonic() - start_time) * 1000)
9292
content = response.text or ""
9393
tokens_used = 0
94+
input_tokens = 0
95+
output_tokens = 0
9496
if response.usage_metadata:
9597
tokens_used = response.usage_metadata.total_token_count or 0
98+
# GW-2#6: Gemini usage_metadata 的拆分字段供成本记账
99+
input_tokens = getattr(response.usage_metadata, "prompt_token_count", 0) or 0
100+
output_tokens = getattr(response.usage_metadata, "candidates_token_count", 0) or 0
96101
logger.info("GeminiProvider.generate 成功: model=%s, tokens=%d, latency=%dms",
97102
model, tokens_used, latency_ms)
98103
return {"content": content, "tokens_used": tokens_used,
104+
"input_tokens": input_tokens, "output_tokens": output_tokens,
99105
"model": model, "latency_ms": latency_ms}
100106
except Exception as e:
101107
logger.error("GeminiProvider.generate 失败: %s", str(e))
@@ -294,7 +300,20 @@ async def generate_stream(
294300
contents=prompt,
295301
config=config,
296302
)
297-
for chunk in response:
303+
304+
# GW-2#5: 同步 SDK 生成器的迭代(含底层阻塞网络 IO)直接在事件循环
305+
# 线程执行会卡死整个网关(流式期间全站延迟升高)。改为逐块
306+
# asyncio.to_thread 搬进线程池,chunk 间让出控制权
307+
def _next_chunk(iterator):
308+
try:
309+
return next(iterator)
310+
except StopIteration:
311+
return None
312+
313+
while True:
314+
chunk = await asyncio.to_thread(_next_chunk, response)
315+
if chunk is None:
316+
break
298317
if chunk.text:
299318
yield chunk.text
300319
# 让出事件循环控制权

‎server/ai-gateway/providers/glm_provider.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -102,8 +102,14 @@ async def generate(
102102
# 提取结果
103103
content = response.choices[0].message.content or ""
104104
tokens_used = 0
105+
input_tokens = 0
106+
output_tokens = 0
105107
if response.usage:
106108
tokens_used = response.usage.total_tokens or 0
109+
# GW-2#6: 提取真实 input/output 拆分供成本记账(OpenAI 兼容
110+
# usage 字段),fallback 链不再对半估算
111+
input_tokens = getattr(response.usage, "prompt_tokens", 0) or 0
112+
output_tokens = getattr(response.usage, "completion_tokens", 0) or 0
107113

108114
logger.info(
109115
"GLMProvider 调用成功: model=%s, tokens=%d, latency=%dms",
@@ -113,6 +119,8 @@ async def generate(
113119
return {
114120
"content": content,
115121
"tokens_used": tokens_used,
122+
"input_tokens": input_tokens,
123+
"output_tokens": output_tokens,
116124
"model": model,
117125
"latency_ms": latency_ms,
118126
}

‎server/ai-gateway/providers/qwen_provider.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,8 +87,14 @@ async def generate(
8787
# 提取结果
8888
content = response.choices[0].message.content or ""
8989
tokens_used = 0
90+
input_tokens = 0
91+
output_tokens = 0
9092
if response.usage:
9193
tokens_used = response.usage.total_tokens or 0
94+
# GW-2#6: 提取真实 input/output 拆分供成本记账(OpenAI 兼容
95+
# usage 字段),fallback 链不再对半估算
96+
input_tokens = getattr(response.usage, "prompt_tokens", 0) or 0
97+
output_tokens = getattr(response.usage, "completion_tokens", 0) or 0
9298

9399
logger.info(
94100
"QwenProvider 调用成功: model=%s, tokens=%d, latency=%dms",
@@ -98,6 +104,8 @@ async def generate(
98104
return {
99105
"content": content,
100106
"tokens_used": tokens_used,
107+
"input_tokens": input_tokens,
108+
"output_tokens": output_tokens,
101109
"model": model,
102110
"latency_ms": latency_ms,
103111
}

‎server/ai-gateway/routers/balance.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525

2626
from cache.redis_cache import get_cache
2727
from config import AI_PROVIDERS, is_valid_api_key, ALIYUN_ACCESS_KEY_ID, ALIYUN_ACCESS_KEY_SECRET
28+
from config.key_pool import get_primary_key
2829

2930
logger = logging.getLogger(__name__)
3031
router = APIRouter(prefix="/api/v1/ai", tags=["余额查询"])
@@ -301,7 +302,10 @@ async def _query_all_balances(cache, start_time: float) -> dict:
301302

302303
for provider_key, (display_name, query_fn) in _PROVIDER_QUERIES.items():
303304
cfg = AI_PROVIDERS.get(provider_key, {})
304-
api_key = cfg.get("api_key", "")
305+
# GW-2#7: 优先从 Key 池取主 Key(复数环境变量 DEEPSEEK_API_KEYS 的
306+
# 首 Key)——原实现只读 AI_PROVIDERS 的单数变量 DEEPSEEK_API_KEY,
307+
# 部署仅配置复数变量时余额查询被静默跳过(余额面板缺 DeepSeek 条目)
308+
api_key = get_primary_key(provider_key) or cfg.get("api_key", "")
305309

306310
# 跳过未配置的 Provider
307311
if not is_valid_api_key(api_key):

‎server/ai-gateway/routers/error_pattern.py‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -111,11 +111,14 @@ async def error_pattern(request: Request, body: ErrorPatternRequest) -> ErrorPat
111111
logger.info("错误模式分析请求: user=%s, errors_count=%d", user_id, len(body.goldenErrors))
112112

113113
# 生成 AI 响应缓存键(基于输入错误记录)
114+
# GW-2#10: 缓存键加入 user_id 隔离——错误模式分析结果基于用户私有
115+
# 学习数据,原键不含用户维度,内容相同的不同用户会串用彼此的分析结论
116+
#(隐私泄露 + 正确性错误);同时不再截断 50 字符(sha256 本身压缩长度)
114117
errors_str = "|".join([
115-
f"{e.flashcardId}:{e.correctAnswer[:50]}:{e.userAnswer[:50]}"
118+
f"{e.flashcardId}:{e.correctAnswer}:{e.userAnswer}"
116119
for e in body.goldenErrors
117120
])
118-
cache_key = hashlib.sha256(errors_str.encode()).hexdigest()
121+
cache_key = hashlib.sha256(f"{user_id}:{errors_str}".encode()).hexdigest()
119122

120123
# 检查 Redis AI 响应缓存
121124
cache = get_cache()

0 commit comments

Comments
 (0)