Skip to content

Commit 8e78b6f

Browse files
committed
fix(gateway): AI网关高危修复——并发信号量/预算记账/流式限流/SSE资源释放/事件循环阻塞/请求体限制/降级链容错
- GW-H1: 启用 Phase3 并发信号量(ai_semaphore/ai_heavy_semaphore 统一入口保护,视频/多模态走 heavy 上限) - GW-H2: 预算记账 user_id 从 request.state 取真实用户,BudgetMiddleware 日限额恢复生效 - GW-H3: 流式通配路径 /{feature}/stream 复用中间件 Redis 限流逻辑,堵住免限流漏洞 - GW-H4: SSE 所有退出路径(超时/断连/异常)关闭上游生成器 aclose,消除连接泄漏与幽灵计费;streaming.py 补 is_disconnected 检测 - GW-H5: gemini generate_vision_multi 同步阻塞改专用线程池执行,事件循环不再冻结 - GW-H6: Provider SDK 专用线程池隔离慢调用 + google-genai HttpOptions 120s 客户端超时 - GW-H7: 多模态/ASR/视觉端点纳入输入校验(64MB/32MB/32MB 前缀上限),base64/列表字段加 max_length - GW-H8: 视频上传线程内流式计数复制(chunked 绕过 content-length 时由计数兜底),文件句柄 with 管理 - GW-H9: ffprobe/ffmpeg 改 create_subprocess_exec 异步执行,超时强杀防僵尸进程 - GW-H10: FallbackProvider.generate/generate_stream 补 **kwargs,云端全挂时降级内容正常返回 测试: python -m pytest tests/ -q → 193 passed
1 parent 197a7f2 commit 8e78b6f

11 files changed

Lines changed: 322 additions & 137 deletions

File tree

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

Lines changed: 38 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
使用 ffmpeg 均匀抽取关键帧后转为多图分析。
1414
"""
1515

16+
import asyncio
1617
import base64
1718
import logging
1819
import tempfile
@@ -30,8 +31,32 @@
3031
# 降级抽帧配置:均匀抽取的帧数上限
3132
_FALLBACK_FRAME_COUNT = 10
3233

34+
# GW-H9: 子进程超时(秒)——ffprobe 时长探测与 ffmpeg 单帧抽取
35+
_PROBE_TIMEOUT = 30.0
36+
_FFMPEG_FRAME_TIMEOUT = 15.0
3337

34-
def _extract_frames_ffmpeg(video_path: str, frame_count: int) -> list[str]:
38+
39+
async def _run_subprocess(cmd: list[str], timeout: float) -> tuple[int, bytes, bytes]:
40+
"""异步执行子进程并等待完成(带超时,超时后强杀防僵尸进程)。
41+
42+
GW-H9: 原 subprocess.run 同步阻塞事件循环(单次最长 30s),
43+
并发视频分析会相互拖垮;改为 create_subprocess_exec 非阻塞执行。
44+
"""
45+
proc = await asyncio.create_subprocess_exec(
46+
*cmd,
47+
stdout=asyncio.subprocess.PIPE,
48+
stderr=asyncio.subprocess.PIPE,
49+
)
50+
try:
51+
stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout=timeout)
52+
except asyncio.TimeoutError:
53+
proc.kill()
54+
await proc.wait()
55+
raise TimeoutError(f"子进程超时({timeout:.0f}s): {cmd[0]}")
56+
return proc.returncode or 0, stdout, stderr
57+
58+
59+
async def _extract_frames_ffmpeg(video_path: str, frame_count: int) -> list[str]:
3560
"""
3661
使用 ffmpeg 从视频中均匀抽取关键帧(返回 base64 列表)
3762
@@ -44,8 +69,6 @@ def _extract_frames_ffmpeg(video_path: str, frame_count: int) -> list[str]:
4469
Raises:
4570
RuntimeError: ffmpeg 不可用或抽帧失败
4671
"""
47-
import subprocess
48-
4972
try:
5073
# 获取视频时长
5174
probe_cmd = [
@@ -54,10 +77,10 @@ def _extract_frames_ffmpeg(video_path: str, frame_count: int) -> list[str]:
5477
"-of", "default=noprint_wrappers=1:nokey=1",
5578
video_path,
5679
]
57-
result = subprocess.run(probe_cmd, capture_output=True, text=True, timeout=30)
58-
if result.returncode != 0:
59-
raise RuntimeError(f"ffprobe 获取时长失败: {result.stderr}")
60-
duration = float(result.stdout.strip())
80+
code, stdout, stderr = await _run_subprocess(probe_cmd, _PROBE_TIMEOUT)
81+
if code != 0:
82+
raise RuntimeError(f"ffprobe 获取时长失败: {stderr.decode(errors='replace')}")
83+
duration = float(stdout.decode(errors="replace").strip())
6184

6285
# 均匀间隔抽帧
6386
interval = max(duration / (frame_count + 1), 1.0)
@@ -75,7 +98,11 @@ def _extract_frames_ffmpeg(video_path: str, frame_count: int) -> list[str]:
7598
str(out_path),
7699
"-y",
77100
]
78-
subprocess.run(cmd, capture_output=True, timeout=15, check=True)
101+
code, _, stderr = await _run_subprocess(cmd, _FFMPEG_FRAME_TIMEOUT)
102+
if code != 0:
103+
raise RuntimeError(
104+
f"ffmpeg 抽帧失败: {stderr.decode(errors='replace')}"
105+
)
79106
if out_path.exists():
80107
frames.append(base64.b64encode(out_path.read_bytes()).decode())
81108

@@ -87,6 +114,8 @@ def _extract_frames_ffmpeg(video_path: str, frame_count: int) -> list[str]:
87114

88115
except FileNotFoundError as e:
89116
raise RuntimeError("ffmpeg 未安装,无法进行视频抽帧降级") from e
117+
except asyncio.TimeoutError as e:
118+
raise RuntimeError(f"视频抽帧超时: {e}") from e
90119
except Exception as e:
91120
raise RuntimeError(f"视频抽帧失败: {e}") from e
92121

@@ -171,7 +200,7 @@ async def run(
171200
)
172201

173202
try:
174-
frames = _extract_frames_ffmpeg(str(video_input), _FALLBACK_FRAME_COUNT)
203+
frames = await _extract_frames_ffmpeg(str(video_input), _FALLBACK_FRAME_COUNT)
175204
except RuntimeError as e:
176205
raise RuntimeError(f"视频分析失败:Provider 不支持视频且抽帧不可用: {e}") from e
177206

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

Lines changed: 53 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,31 @@
2222

2323
_FIRST_TOKEN_RETRY_TIMEOUT = 5.0 # 首 token 探测重试的超时预算(秒)
2424

25+
# 高耗时功能(走独立并发上限 ai_heavy_semaphore)
26+
_HEAVY_CONCURRENCY_FEATURES: frozenset[str] = frozenset(
27+
{"video_analyze", "multimodal_analyze"}
28+
)
29+
30+
31+
def _get_concurrency_semaphore(app, feature: str):
32+
"""获取并发信号量;未初始化(如测试环境)时返回 None 表示不限流。
33+
34+
@ai-context: Phase3 并发控制——主信号量限制全部 AI 调用,
35+
heavy 信号量进一步限制高耗时功能(视频/多模态分析)。
36+
"""
37+
if feature in _HEAVY_CONCURRENCY_FEATURES:
38+
return getattr(app.state, "ai_heavy_semaphore", None)
39+
return getattr(app.state, "ai_semaphore", None)
40+
41+
42+
async def _run_under_semaphore(app, feature: str, coro_factory: Callable[[], Awaitable[Any]]) -> Any:
43+
"""在信号量保护下执行 AI 调用(信号量缺失时直通)。"""
44+
sem = _get_concurrency_semaphore(app, feature)
45+
if sem is None:
46+
return await coro_factory()
47+
async with sem:
48+
return await coro_factory()
49+
2550
# ============================================================
2651
# Provider Fallback 链:主 Provider 失败时依次尝试的备选
2752
# ============================================================
@@ -109,6 +134,7 @@ async def call_with_fallback(
109134
app,
110135
feature: str,
111136
fn: Callable[..., Awaitable[dict[str, Any]]],
137+
user_id: str = "system",
112138
) -> tuple[dict[str, Any], str]:
113139
"""
114140
使用 Provider fallback 链执行 AI 调用(非流式路径)。
@@ -122,6 +148,7 @@ async def call_with_fallback(
122148
app: FastAPI 应用实例(通过 app.state.providers 获取 Provider)
123149
feature: 功能标识,对应 PROVIDER_FALLBACK_CHAIN 的 key
124150
fn: 异步可调用对象,签名为 async fn(provider, model_name) -> dict
151+
user_id: 记账归属用户(预算控制按此维度计数)
125152
126153
Returns:
127154
tuple: (result_dict, provider_key)
@@ -168,18 +195,22 @@ async def _run_fallback_chain():
168195
)
169196

170197
try:
171-
_FEATURE_CONTEXT.set(feature)
172-
result = await fn(provider, model_name)
198+
async def _do_call():
199+
_FEATURE_CONTEXT.set(feature)
200+
return await fn(provider, model_name)
201+
202+
# Phase3: 并发信号量保护(视频/多模态走 heavy 上限)
203+
result = await _run_under_semaphore(app, feature, _do_call)
173204
# 调用成功:通知熔断器重置失败计数,恢复 Provider 健康状态
174205
cb = get_circuit(provider_key)
175206
if cb:
176207
await cb.on_success()
177-
# 记录成本
208+
# 记录成本(按真实用户记账,预算中间件按 user_id 维度查询)
178209
try:
179210
tokens_used = result.get("tokens_used", 0)
180211
model_name = result.get("model", model_name)
181212
await get_cost_tracker().record(
182-
user_id="system",
213+
user_id=user_id,
183214
feature=feature,
184215
model=model_name,
185216
input_tokens=tokens_used // 2 if tokens_used else 0,
@@ -235,7 +266,10 @@ async def call_with_fallback_for_request(
235266
Raises:
236267
RuntimeError: 所有 Provider 均不可用
237268
"""
238-
result, provider_key = await call_with_fallback(app, feature, fn)
269+
# GW-H2: 预算记账必须归属真实用户,预算中间件(BudgetMiddleware)
270+
# 按 request.state.user_id 查询日用量,写死 system 会导致限额永不触发
271+
user_id = getattr(request.state, "user_id", "system")
272+
result, provider_key = await call_with_fallback(app, feature, fn, user_id=user_id)
239273
return result, provider_key, False
240274

241275

@@ -313,11 +347,20 @@ async def _run_stream_fallback_chain():
313347
probe_timeout = min(_FIRST_TOKEN_PROBE_TIMEOUT, remaining)
314348
first_chunk = await asyncio.wait_for(agen.__anext__(), timeout=probe_timeout)
315349

316-
# 首 token 成功:包装生成器,先吐首 token 再吐剩余
317-
async def _wrapped_gen(a=agen, first=first_chunk):
318-
yield first
319-
async for c in a:
320-
yield c
350+
# 首 token 成功:包装生成器,先吐首 token 再吐剩余;
351+
# 并发信号量在生成器整个生命周期内持有,防止流式长连接堆积
352+
sem = _get_concurrency_semaphore(app, feature)
353+
if sem is None:
354+
async def _wrapped_gen(a=agen, first=first_chunk):
355+
yield first
356+
async for c in a:
357+
yield c
358+
else:
359+
async def _wrapped_gen(a=agen, first=first_chunk, s=sem):
360+
async with s:
361+
yield first
362+
async for c in a:
363+
yield c
321364

322365
# 通知熔断器成功
323366
cb = get_circuit(provider_key)

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

Lines changed: 21 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -16,11 +16,16 @@
1616

1717
logger = logging.getLogger(__name__)
1818

19-
# 需要输入验证的路径前缀
20-
VALIDATED_PATHS = ("/api/v1/ai/",)
19+
# 需要输入验证的路径前缀 → 各前缀允许的最大请求体字节数
20+
# GW-H7: 多模态/ASR/视觉端点原不在校验范围,base64 大载荷可致内存 DoS;
21+
# 各前缀按业务实际载荷设定独立上限(文本端点 1MB,多模态 64MB,音频/视觉 32MB)
22+
VALIDATED_PATHS: dict[str, int] = {
23+
"/api/v1/ai/": 1 * 1024 * 1024, # 1MB
24+
"/api/v1/multimodal/": 64 * 1024 * 1024, # 64MB(多图联合分析,100 帧 × ≤10MB 编码)
25+
"/api/v1/asr/": 32 * 1024 * 1024, # 32MB(音频转写)
26+
"/api/v1/vision/": 32 * 1024 * 1024, # 32MB(视觉提取)
27+
}
2128

22-
# 限制常量
23-
MAX_CONTENT_LENGTH = 1 * 1024 * 1024 # 1MB
2429
MAX_TEXT_FIELD_LENGTH = 50000 # 字符
2530

2631

@@ -29,27 +34,33 @@ class InputValidationMiddleware(BaseHTTPMiddleware):
2934

3035
async def dispatch(self, request: Request, call_next):
3136
# 仅对 AI 功能 API 进行输入验证
32-
if not request.url.path.startswith(VALIDATED_PATHS):
37+
max_content_length = None
38+
for prefix, limit in VALIDATED_PATHS.items():
39+
if request.url.path.startswith(prefix):
40+
max_content_length = limit
41+
break
42+
if max_content_length is None:
3343
return await call_next(request)
3444

3545
# 仅检查有请求体的方法
3646
if request.method in ("POST", "PUT", "PATCH"):
37-
# 检查 Content-Length 头
47+
# 检查 Content-Length 头(chunked 编码无此头时跳过,由 JSON 解析兜底)
3848
content_length = request.headers.get("content-length")
3949
if content_length:
4050
try:
41-
if int(content_length) > MAX_CONTENT_LENGTH:
51+
if int(content_length) > max_content_length:
4252
return JSONResponse(
43-
status_code=422,
53+
status_code=413,
4454
content={
45-
"detail": f"request body exceeds {MAX_CONTENT_LENGTH} bytes (max 1MB)",
55+
"detail": f"request body exceeds {max_content_length} bytes",
4656
},
4757
)
4858
except (ValueError, TypeError):
4959
# 畸形 Content-Length 头,忽略该检查(后续 JSON 解析会进一步校验)
5060
pass
5161

52-
# 解析 JSON 请求体并检查文本字段
62+
# 解析 JSON 请求体并检查文本字段(body 全量读入内存前先受
63+
# Content-Length 限制;chunked 大载荷由路由层 Pydantic max_length 兜底)
5364
content_type = request.headers.get("content-type", "")
5465
if "application/json" in content_type:
5566
try:

0 commit comments

Comments
 (0)