Skip to content

Commit 1aeeac5

Browse files
committed
fix(gateway): resolve high-risk issues #1-#4
- #1: JWT algorithm now injected via SUPABASE_JWT_ALGORITHM (default HS256, aligned with Supabase default signing); placeholder detection prevents 'configured' false impression causing whole-site 401 - #2: remove double rollback on streaming rate-limit overrun - the Lua script already atomically DECRs, manual rollback let users drain their own quotas into negative counts - #3: error_pattern chain/router defensive validation - missing fields from LLM JSON output filtered instead of 500 (ValidationError) - #4: GLM generate_vision returns actual clamped max_tokens; vision chain logs truncation warning when output approaches the limit
1 parent b82668d commit 1aeeac5

9 files changed

Lines changed: 152 additions & 22 deletions

File tree

‎server/ai-gateway/chains/error_pattern_chain.py‎

Lines changed: 18 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -57,9 +57,25 @@ async def run(
5757
}
5858
"""
5959
# 构建输入文本
60+
# GW-2#3: 防御性取值——原实现直接键索引 e['correctAnswer']/e['userAnswer'],
61+
# 调用方传入缺字段 dict 时 KeyError 导致 500;改为 get() + str() 归一化,
62+
# 并过滤空值项(避免 prompt 中出现无意义的空错误记录)
63+
def _safe_text(value: Any) -> str:
64+
return str(value) if value is not None else ""
65+
66+
normalized_errors: list[tuple[str, str]] = []
67+
for e in golden_errors[:20]: # 最多分析前 20 条
68+
if not isinstance(e, dict):
69+
continue
70+
correct = _safe_text(e.get("correctAnswer"))
71+
user = _safe_text(e.get("userAnswer"))
72+
if not correct and not user:
73+
continue
74+
normalized_errors.append((correct, user))
75+
6076
errors_text = "\n".join([
61-
f"【错误 #{i+1}】\n正确答案:{e['correctAnswer']}\n用户回答:{e['userAnswer']}\n---"
62-
for i, e in enumerate(golden_errors[:20]) # 最多分析前 20 条
77+
f"【错误 #{i+1}】\n正确答案:{correct}\n用户回答:{user}\n---"
78+
for i, (correct, user) in enumerate(normalized_errors)
6379
])
6480

6581
template = self._load_prompt_template()

‎server/ai-gateway/chains/vision_extract_chain.py‎

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -278,8 +278,23 @@ async def run(
278278
_feature="vision_extract",
279279
)
280280

281+
# GW-2#4: 显式截断检测——GLM 等 provider 可能把 max_tokens clamp 到
282+
# 低于请求值(full 模式请求 4096 被 clamp 到 1024),输出在接近上限时
283+
# 大概率被截断。原实现静默返回残缺 JSON(结构化字段丢失)且照常缓存,
284+
# 用户无感知;此处记录 truncated 告警供排查与降级重试决策
285+
content = result["content"]
286+
actual_max_tokens = result.get("max_tokens")
287+
if actual_max_tokens and len(content) > actual_max_tokens * 3:
288+
# 1 token ≈ 3 字符的保守估算(中文 1 token ≈ 1-2 字符,英文 ≈ 4)
289+
logger.warning(
290+
"视觉提取疑似截断: mode=%s, model=%s, content=%d 字符, "
291+
"max_tokens=%d(provider clamp 后),结构化字段可能缺失",
292+
effective_mode, result.get("model", self.model),
293+
len(content), actual_max_tokens,
294+
)
295+
281296
# 解析结构化结果
282-
structured = self._parse_response(result["content"])
297+
structured = self._parse_response(content)
283298

284299
return {
285300
"content": result["content"],

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

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,8 @@
44
@ai-context: 汇聚运行环境、CORS、JWT 与中间件连接串等应用级配置。
55
Supabase JWKS 端点从 SUPABASE_URL 推导(标准路径含 /auth/v1 前缀),
66
亦可经 SUPABASE_JWKS_URL 显式覆盖。所有值均经环境变量注入,无硬编码密钥。
7+
@ai-context: GW-2#1——jwt_algorithm 由 SUPABASE_JWT_ALGORITHM 注入,默认 HS256
8+
(Supabase 默认对称密钥签发机制);ES256 仅在开启自定义 JWT 时选用。
79
"""
810

911
import os
@@ -25,7 +27,9 @@
2527
if origin.strip()
2628
],
2729
"jwt_secret": os.getenv("SUPABASE_JWT_SECRET", ""),
28-
"jwt_algorithm": "ES256",
30+
# GW-2#1: 算法由环境变量注入(默认 HS256 与 Supabase 默认签发机制对齐);
31+
# 原硬编码 ES256 导致默认 HS256 项目(无 JWKS 端点)全站 401
32+
"jwt_algorithm": os.getenv("SUPABASE_JWT_ALGORITHM", "HS256"),
2933
"supabase_url": _supabase_url,
3034
"supabase_jwks_url": os.getenv(
3135
"SUPABASE_JWKS_URL",

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

Lines changed: 37 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -33,20 +33,40 @@
3333
# 不需要认证的路径白名单
3434
PUBLIC_PATHS = {"/health", "/health/quick", "/health/live"}
3535

36+
# GW-2#1: 占位符识别——SUPABASE_URL 模板中的示例项目域名与常见示例密钥
37+
# 被视为"未配置",避免"已配置"假象导致认证链路断裂后全站 401 且难排查
38+
_PLACEHOLDER_PROJECT_ID = "your-project-id"
39+
_PLACEHOLDER_SECRETS = {
40+
"change-this-to-a-random-string-in-production",
41+
"your-jwt-secret",
42+
"placeholder",
43+
"supabase_jwt_secret",
44+
}
45+
46+
47+
def _is_placeholder_text(value: str) -> bool:
48+
"""判断配置值是否为占位符/示例值"""
49+
return value.strip().lower() in _PLACEHOLDER_SECRETS
50+
51+
3652
def _jwt_verification_configured() -> bool:
3753
"""
3854
检查当前算法是否具备验证所需的密钥材料。
3955
40-
- HS256/RS256:需要 SUPABASE_JWT_SECRET(对称密钥或 PEM 公钥)
41-
- ES256:需要 SUPABASE_JWKS_URL 或 SUPABASE_URL(用于获取 JWKS 公钥)
56+
- HS256/RS256:需要 SUPABASE_JWT_SECRET(对称密钥或 PEM 公钥),
57+
占位符/示例值视为未配置;
58+
- ES256:需要 SUPABASE_JWKS_URL 或 SUPABASE_URL(用于获取 JWKS 公钥),
59+
URL 含 your-project-id 占位符域名时视为未配置。
4260
"""
4361
alg = APP_CONFIG.get("jwt_algorithm", "HS256")
4462
if alg == "ES256":
45-
return bool(
63+
jwks_url = (
4664
APP_CONFIG.get("supabase_jwks_url", "")
4765
or APP_CONFIG.get("supabase_url", "")
4866
)
49-
return bool(APP_CONFIG.get("jwt_secret", ""))
67+
return bool(jwks_url) and _PLACEHOLDER_PROJECT_ID not in jwks_url
68+
secret = APP_CONFIG.get("jwt_secret", "")
69+
return bool(secret) and not _is_placeholder_text(secret)
5070

5171

5272
# 启动时检查密钥配置
@@ -71,8 +91,12 @@ def _jwt_verification_configured() -> bool:
7191
stacklevel=2,
7292
)
7393

74-
# 启动日志:输当前 JWT 验证算法
75-
logger.info("JWT 验证算法: %s", APP_CONFIG["jwt_algorithm"])
94+
# 启动日志:输当前 JWT 验证算法与配置状态(GW-2#1: 避免"已配置"假象)
95+
logger.info(
96+
"JWT 验证算法: %s (验证材料已配置=%s)",
97+
APP_CONFIG["jwt_algorithm"],
98+
_jwt_verification_configured(),
99+
)
76100

77101

78102
def _normalize_pem_key(raw: str) -> str:
@@ -136,6 +160,13 @@ def _resolve_jwks_url() -> str:
136160
"""解析 JWKS 端点 URL:优先 SUPABASE_JWKS_URL,其次从 SUPABASE_URL 推导"""
137161
jwks_url = APP_CONFIG.get("supabase_jwks_url", "")
138162
if jwks_url:
163+
if _PLACEHOLDER_PROJECT_ID in jwks_url:
164+
# GW-2#1: 占位符域名推导出的 JWKS 端点必然不可达,显式告警
165+
#(此分支通常不会被走到——_jwt_verification_configured 已将其判为未配置)
166+
logger.warning(
167+
"ES256 验签 JWKS URL 含占位符域名 %s,请配置真实的 SUPABASE_URL/SUPABASE_JWKS_URL",
168+
jwks_url,
169+
)
139170
return jwks_url
140171
supabase_url = APP_CONFIG.get("supabase_url", "")
141172
if supabase_url:

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

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -229,7 +229,13 @@ async def _lua_check_limit(key: str, limit: int, ttl: int) -> int:
229229

230230

231231
async def rollback_rate_limit(user_id: str, feature: str) -> None:
232-
"""回退频率限制计数器(原子 INCR 超出限额或请求失败时调用)"""
232+
"""回退频率限制计数器。
233+
234+
⚠️ 仅允许在 check_rate_limit 返回 True(本次 INCR 已生效)之后、
235+
请求未成功消费配额(非 2xx)时调用——此时功能级与全局计数均已 +1,
236+
回滚两层是准确的。超限路径(check 返回 False)Lua 脚本已原子回滚,
237+
禁止调用本函数,否则双重 DECR 会把配额刷穿(GW-2#2)。
238+
"""
233239
cache = get_cache()
234240
if not cache._client:
235241
return

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

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -249,6 +249,10 @@ async def generate_vision(
249249
"tokens_used": tokens_used,
250250
"model": model,
251251
"latency_ms": latency_ms,
252+
# GW-2#4: clamp 后实际生效的 max_tokens,供 chain 侧截断检测使用
253+
#(与 generate_vision_multi 保持一致;原实现缺此字段,
254+
# full 模式请求 4096 被 clamp 到 1024 后截断不可感知)
255+
"max_tokens": max_tokens,
252256
}
253257

254258
except Exception as e:

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

Lines changed: 54 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -66,6 +66,39 @@ class ErrorPatternResponse(BaseModel):
6666
# ============================================================
6767

6868

69+
def _validate_pattern_item(item: Any) -> dict | None:
70+
"""校验 LLM 输出的错误模式项:缺失必填字段/类型非法时返回 None(过滤)。
71+
72+
GW-2#3: LLM JSON 输出截断或格式漂移时,PatternItem(**p) 严格构造会抛
73+
ValidationError 导致 500——与 quiz_gen_chain._validate_question 相同的
74+
逐项校验+过滤模式,非法项丢弃而非崩溃。
75+
"""
76+
if not isinstance(item, dict):
77+
return None
78+
required = {"type", "keywords", "explanation", "suggestion"}
79+
if not all(k in item and item[k] not in (None, "") for k in required):
80+
return None
81+
keywords = item.get("keywords")
82+
if not isinstance(keywords, list) or not all(isinstance(k, str) for k in keywords):
83+
return None
84+
return item
85+
86+
87+
def _validate_top_offender(item: Any) -> dict | None:
88+
"""校验 LLM 输出的高频错误卡片项:缺 flashcardId 时返回 None(过滤)。"""
89+
if not isinstance(item, dict):
90+
return None
91+
flashcard_id = item.get("flashcardId")
92+
if not isinstance(flashcard_id, str) or not flashcard_id:
93+
return None
94+
try:
95+
count = int(item.get("count", 0))
96+
except (TypeError, ValueError):
97+
# GW-2#3: count 非数字时容错为 0,而非让 pydantic 抛 500
98+
count = 0
99+
return {"flashcardId": flashcard_id, "count": count}
100+
101+
69102
@router.post("/error-pattern", response_model=ErrorPatternResponse, summary="分析黄金错误模式")
70103
async def error_pattern(request: Request, body: ErrorPatternRequest) -> ErrorPatternResponse:
71104
"""
@@ -126,12 +159,27 @@ async def _run_chain(provider, model_name):
126159
await cache.set_ai_cache(cache_key, chain_result, expire=3600)
127160

128161
# 构建响应对象
129-
patterns = [
130-
PatternItem(**p) for p in patterns_data
131-
]
132-
top_offenders = [
133-
TopOffender(**o) for o in top_offenders_data
134-
]
162+
# GW-2#3: LLM 输出逐项校验+过滤,缺字段项不导致 500(截断/格式漂移降级)
163+
patterns = []
164+
dropped_patterns = 0
165+
for p in patterns_data:
166+
validated = _validate_pattern_item(p)
167+
if validated is not None:
168+
patterns.append(PatternItem(**validated))
169+
else:
170+
dropped_patterns += 1
171+
172+
top_offenders = []
173+
for o in top_offenders_data:
174+
validated = _validate_top_offender(o)
175+
if validated is not None:
176+
top_offenders.append(TopOffender(**validated))
177+
178+
if dropped_patterns > 0:
179+
logger.warning(
180+
"错误模式分析: %d 条模式项因缺字段被过滤(LLM 输出质量问题)",
181+
dropped_patterns,
182+
)
135183

136184
return ErrorPatternResponse(
137185
patterns=patterns,

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

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@
2424
from pydantic import BaseModel, Field
2525

2626
from config import call_with_fallback_stream
27-
from middleware.rate_limit import check_rate_limit, rollback_rate_limit
27+
from middleware.rate_limit import check_rate_limit
2828

2929
logger = logging.getLogger(__name__)
3030
router = APIRouter(prefix="/api/v1/ai", tags=["流式输出"])
@@ -278,7 +278,9 @@ async def stream_ai(request: Request, feature: str, body: StreamRequest):
278278
# 此处复用与中间件一致的 Redis 限流逻辑(含功能级 + 全局每日总量)
279279
is_allowed, detail = await check_rate_limit(user_id, config_key)
280280
if not is_allowed:
281-
await rollback_rate_limit(user_id, config_key)
281+
# GW-2#2: 超限时禁止再次 rollback——Lua 脚本已在超限分支内原子 DECR,
282+
# 路由层再 rollback 会双重回滚(feature 减到负数、global 被无条件 DECR),
283+
# 恶意用户可反复刷超限请求把自己的配额刷低甚至清零
282284
raise HTTPException(status_code=429, detail=detail)
283285

284286
try:

‎server/ai-gateway/tests/test_config.py‎

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
- _resolve_model_name() 在 fallback 链中正确选择模型名
1010
"""
1111

12+
import os
1213
import sys
1314
from pathlib import Path
1415

@@ -254,9 +255,12 @@ def test_evaluate_returns_glm_model(self):
254255
class TestConfigAssertions:
255256
"""配置值精确断言(JWT 算法、API Key 校验、超时配置、Fallback 链尾)"""
256257

257-
def test_jwt_algorithm_is_es256(self):
258-
"""APP_CONFIG jwt_algorithm 必须为 ES256(与 Supabase JWKS 一致)"""
259-
assert APP_CONFIG["jwt_algorithm"] == "ES256"
258+
def test_jwt_algorithm_defaults_to_hs256(self):
259+
"""APP_CONFIG jwt_algorithm 默认 HS256(GW-2#1: 与 Supabase 默认对称密钥
260+
签发机制对齐;原硬编码 ES256 导致无 JWKS 的默认项目全站 401)。
261+
显式配置 SUPABASE_JWT_ALGORITHM 时使用配置值。"""
262+
expected = os.getenv("SUPABASE_JWT_ALGORITHM", "HS256")
263+
assert APP_CONFIG["jwt_algorithm"] == expected
260264

261265
def test_is_valid_api_key_rejects_placeholder(self):
262266
"""占位符 API Key 应被识别为无效"""

0 commit comments

Comments
 (0)