并发模型从「每worker并发1 + 多worker并行」改为全局串行:同一时间 只处理1个请求,其余排队;多 worker 仅作热备(主 worker 坏了才用备机)。 接口4(face/features) 走豆包、不占 GPU,不纳入串行。 - pool: asyncio.Semaphore(max_global_concurrency=1) + acquire/release_global_slot (依赖单进程 uvicorn 部署,已在注释中标注) - forward: proxy_request 最外层 acquire 全局槽、try/finally 全路径释放; 入参 multipart 解析挂 request.state;原重试/故障转移逻辑抽到 _dispatch_with_retries - reqlog(新): 标量入参保留;图片(file/base64)存盘转URL,绝不内嵌base64; 出参递归摘要截断(landmarks/长串/大数组) - logging_middleware: RequestLogEntry 加 request_params/response_data, jsonl 全量记录;_load_from_logfile 同步映射防重启丢字段;get_stats recent 暴露 - /gateway-health 暴露 global_max/global_busy/global_waiting - config: dispatch 新增 max_global_concurrency / max_queue_wait_seconds - tests: test_reqlog + test_gateway_serialization(9 用例) 顺带提交此前未提交的网关统计(daily stats 接口)与耗时看板(api_timing_dashboard.html)。 Co-Authored-By: Claude <noreply@anthropic.com>
345 lines
13 KiB
Python
345 lines
13 KiB
Python
"""Worker 健康池 + 空闲派发。
|
||
|
||
- 后台 asyncio 任务周期性探测每个 worker 的 /health
|
||
- 连续失败 N 次 → 下线;连续成功 N 次 → 上线
|
||
- 每个 worker 一个 busy 标志(per_worker_concurrency=1)
|
||
- acquire_worker() 从在线池挑空闲 worker;全忙排队;池空→NoWorkerAvailable
|
||
- release_worker() 释放 worker
|
||
"""
|
||
|
||
import asyncio
|
||
import logging
|
||
import time
|
||
from dataclasses import dataclass, field
|
||
from typing import Dict, List, Optional
|
||
|
||
import httpx
|
||
|
||
logger = logging.getLogger("gateway.pool")
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 异常
|
||
# ---------------------------------------------------------------------------
|
||
|
||
class NoWorkerAvailable(Exception):
|
||
"""无可用 worker(池空或全忙排队超时)。"""
|
||
pass
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 数据结构
|
||
# ---------------------------------------------------------------------------
|
||
|
||
@dataclass
|
||
class WorkerState:
|
||
"""单个 worker 的运行时状态。"""
|
||
url: str
|
||
online: bool = False # 初始 offline,等健康检查通过后上线
|
||
busy: bool = False
|
||
consecutive_failures: int = 0
|
||
consecutive_successes: int = 0
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 全局状态
|
||
# ---------------------------------------------------------------------------
|
||
|
||
_workers: Dict[str, WorkerState] = {}
|
||
_pool_condition: Optional[asyncio.Condition] = None
|
||
_health_task: Optional[asyncio.Task] = None
|
||
_shutdown_event: Optional[asyncio.Event] = None
|
||
|
||
# 全局串行槽:同一时间最多 max_global_concurrency 个请求进入 worker 派发,其余排队。
|
||
# ⚠️ 依赖单进程 uvicorn 部署(hair-gateway.service 无 --workers)。
|
||
# 若改多 worker / gunicorn,进程内信号量会失效,需换成跨进程锁(文件锁/Redis)。
|
||
_global_sem: Optional[asyncio.Semaphore] = None
|
||
_global_max: int = 1
|
||
_global_busy: int = 0
|
||
_global_waiting: int = 0
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 健康检查后台任务
|
||
# ---------------------------------------------------------------------------
|
||
|
||
async def _check_worker_health(
|
||
client: httpx.AsyncClient,
|
||
w: WorkerState,
|
||
cfg: dict,
|
||
) -> None:
|
||
"""探测单个 worker 的 /health,更新上下线状态。"""
|
||
hc_cfg = cfg["health_check"]
|
||
url = f"{w.url}{hc_cfg['path']}"
|
||
token = cfg["shared_password"]
|
||
|
||
try:
|
||
resp = await client.get(
|
||
url,
|
||
headers={"X-Internal-Token": token},
|
||
timeout=hc_cfg["timeout_seconds"],
|
||
)
|
||
if resp.status_code == 200:
|
||
w.consecutive_failures = 0
|
||
w.consecutive_successes += 1
|
||
if w.consecutive_successes >= hc_cfg["healthy_threshold"]:
|
||
if not w.online:
|
||
w.online = True
|
||
logger.info("Worker 上线: %s(连续成功 %d 次)", w.url, w.consecutive_successes)
|
||
else:
|
||
_mark_failure(w, f"HTTP {resp.status_code}")
|
||
except Exception as exc:
|
||
_mark_failure(w, str(exc))
|
||
|
||
|
||
def _mark_failure(w: WorkerState, reason: str) -> None:
|
||
"""记录一次失败,达到阈值后下线。"""
|
||
w.consecutive_successes = 0
|
||
w.consecutive_failures += 1
|
||
threshold = w.consecutive_failures # used in log
|
||
hc_threshold = 2 # default, will be overridden
|
||
if w.consecutive_failures >= hc_threshold:
|
||
# 实际 threshold 从配置读取,这里先做基本判断
|
||
pass
|
||
logger.debug("Worker %s 健康检查失败 (%d/%d): %s", w.url, w.consecutive_failures, 99, reason)
|
||
|
||
|
||
async def _health_check_loop(cfg: dict) -> None:
|
||
"""后台循环:周期性探测所有 worker 健康状态。"""
|
||
hc_cfg = cfg["health_check"]
|
||
interval = hc_cfg["interval_seconds"]
|
||
unhealthy_threshold = hc_cfg["unhealthy_threshold"]
|
||
healthy_threshold = hc_cfg["healthy_threshold"]
|
||
token = cfg["shared_password"]
|
||
|
||
logger.info(
|
||
"健康检查循环启动 | 间隔=%ds | 下线阈值=%d | 上线阈值=%d | workers=%d",
|
||
interval, unhealthy_threshold, healthy_threshold, len(_workers),
|
||
)
|
||
|
||
async with httpx.AsyncClient() as client:
|
||
while not _shutdown_event.is_set():
|
||
for w in _workers.values():
|
||
url = f"{w.url}{hc_cfg['path']}"
|
||
try:
|
||
resp = await client.get(
|
||
url,
|
||
headers={"X-Internal-Token": token},
|
||
timeout=hc_cfg["timeout_seconds"],
|
||
)
|
||
if resp.status_code == 200:
|
||
w.consecutive_failures = 0
|
||
w.consecutive_successes += 1
|
||
if w.consecutive_successes >= healthy_threshold and not w.online:
|
||
w.online = True
|
||
logger.info("✅ Worker 上线: %s", w.url)
|
||
else:
|
||
w.consecutive_successes = 0
|
||
w.consecutive_failures += 1
|
||
if w.consecutive_failures >= unhealthy_threshold and w.online:
|
||
w.online = False
|
||
logger.warning("⚠ Worker 下线: %s(HTTP %d,连续失败 %d 次)",
|
||
w.url, resp.status_code, w.consecutive_failures)
|
||
except Exception as exc:
|
||
w.consecutive_successes = 0
|
||
w.consecutive_failures += 1
|
||
if w.consecutive_failures >= unhealthy_threshold and w.online:
|
||
w.online = False
|
||
logger.warning("⚠ Worker 下线: %s(%s,连续失败 %d 次)",
|
||
w.url, exc, w.consecutive_failures)
|
||
|
||
# 等待下一次探测(支持快速关闭)
|
||
try:
|
||
await asyncio.wait_for(_shutdown_event.wait(), timeout=interval)
|
||
break # shutdown signaled
|
||
except asyncio.TimeoutError:
|
||
pass # 正常的 interval 到期
|
||
|
||
logger.info("健康检查循环已停止")
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 初始化 / 关闭
|
||
# ---------------------------------------------------------------------------
|
||
|
||
async def init_pool(cfg: dict) -> None:
|
||
"""初始化 worker 池并启动健康检查后台任务。"""
|
||
global _workers, _pool_condition, _health_task, _shutdown_event
|
||
|
||
_workers = {url: WorkerState(url=url) for url in cfg["workers"]}
|
||
_pool_condition = asyncio.Condition()
|
||
_shutdown_event = asyncio.Event()
|
||
|
||
# 全局串行槽(独立于 worker busy 标志)
|
||
init_global_slot(cfg)
|
||
|
||
# 立即做一轮健康检查以快速上线
|
||
hc_cfg = cfg["health_check"]
|
||
token = cfg["shared_password"]
|
||
async with httpx.AsyncClient() as client:
|
||
for w in _workers.values():
|
||
try:
|
||
resp = await client.get(
|
||
f"{w.url}{hc_cfg['path']}",
|
||
headers={"X-Internal-Token": token},
|
||
timeout=hc_cfg["timeout_seconds"],
|
||
)
|
||
if resp.status_code == 200:
|
||
w.consecutive_successes = 1
|
||
if hc_cfg["healthy_threshold"] <= 1:
|
||
w.online = True
|
||
else:
|
||
w.consecutive_failures = 1
|
||
except Exception:
|
||
w.consecutive_failures = 1
|
||
|
||
# 对于 healthy_threshold == 1 的情况,已在上面的检查中上线
|
||
# 标记初始在线状态
|
||
online_count = sum(1 for w in _workers.values() if w.online)
|
||
logger.info("Worker 池初始化完成 | 总数=%d | 在线=%d", len(_workers), online_count)
|
||
|
||
# 启动后台健康检查
|
||
_health_task = asyncio.create_task(_health_check_loop(cfg))
|
||
|
||
|
||
async def shutdown_pool() -> None:
|
||
"""关闭健康检查任务,释放资源。"""
|
||
global _shutdown_event, _health_task
|
||
if _shutdown_event:
|
||
_shutdown_event.set()
|
||
if _health_task:
|
||
_health_task.cancel()
|
||
try:
|
||
await _health_task
|
||
except asyncio.CancelledError:
|
||
pass
|
||
_health_task = None
|
||
logger.info("Worker 池已关闭")
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 全局串行槽
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def init_global_slot(cfg: dict) -> None:
|
||
"""初始化全局串行信号量。
|
||
|
||
max_global_concurrency=1 → 同一时间只有一个请求进入 worker 派发,其余在信号量上排队。
|
||
依赖单进程 uvicorn 部署;多进程需换跨进程锁。
|
||
"""
|
||
global _global_sem, _global_max, _global_busy, _global_waiting
|
||
n = int(cfg.get("dispatch", {}).get("max_global_concurrency", 1))
|
||
if n < 1:
|
||
n = 1
|
||
_global_sem = asyncio.Semaphore(n)
|
||
_global_max = n
|
||
_global_busy = 0
|
||
_global_waiting = 0
|
||
logger.info("全局并发槽初始化: max=%d(单进程生效)", n)
|
||
|
||
|
||
async def acquire_global_slot(cfg: dict) -> bool:
|
||
"""获取一个全局处理槽。True=获得;False=排队超时。
|
||
|
||
超时由 dispatch.max_queue_wait_seconds 控制(需 < nginx proxy_read_timeout 600s)。
|
||
"""
|
||
global _global_busy, _global_waiting
|
||
if _global_sem is None:
|
||
return True # 未初始化,不限流
|
||
timeout = float(cfg.get("dispatch", {}).get("max_queue_wait_seconds", 590))
|
||
_global_waiting += 1
|
||
try:
|
||
await asyncio.wait_for(_global_sem.acquire(), timeout=timeout)
|
||
_global_busy += 1
|
||
return True
|
||
except asyncio.TimeoutError:
|
||
logger.warning("全局排队超时(%.1fs),拒绝请求 | waiting=%d", timeout, _global_waiting)
|
||
return False
|
||
finally:
|
||
_global_waiting -= 1
|
||
|
||
|
||
def release_global_slot() -> None:
|
||
"""释放全局处理槽。"""
|
||
global _global_busy
|
||
if _global_sem is not None:
|
||
_global_busy -= 1
|
||
_global_sem.release()
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 派发
|
||
# ---------------------------------------------------------------------------
|
||
|
||
async def mark_worker_unhealthy(w: WorkerState) -> None:
|
||
"""被动标记:转发失败时立即将该 worker 下线。"""
|
||
async with _pool_condition:
|
||
if w.online:
|
||
w.online = False
|
||
w.consecutive_failures += 1
|
||
logger.warning("⚠ Worker 被动下线: %s(转发失败)", w.url)
|
||
|
||
|
||
async def acquire_worker(cfg: dict) -> WorkerState:
|
||
"""从在线池中获取一个空闲 worker。
|
||
|
||
全忙则排队等待,最多 queue_wait_seconds;池空立即抛异常。
|
||
"""
|
||
queue_wait = cfg["dispatch"]["queue_wait_seconds"]
|
||
deadline = time.monotonic() + queue_wait
|
||
|
||
async with _pool_condition:
|
||
while True:
|
||
# 在线 + 空闲
|
||
idle = [w for w in _workers.values() if w.online and not w.busy]
|
||
|
||
if not idle:
|
||
online = [w for w in _workers.values() if w.online]
|
||
if not online:
|
||
raise NoWorkerAvailable("后端服务暂不可用,请稍后重试")
|
||
|
||
# 全忙,等待
|
||
remaining = deadline - time.monotonic()
|
||
if remaining <= 0:
|
||
raise NoWorkerAvailable("后端服务繁忙,请稍后重试")
|
||
|
||
logger.info("所有 worker 全忙(%d 在线),等待 %.1fs...", len(online), remaining)
|
||
try:
|
||
await asyncio.wait_for(_pool_condition.wait(), timeout=remaining)
|
||
except asyncio.TimeoutError:
|
||
raise NoWorkerAvailable("后端服务繁忙,请稍后重试")
|
||
continue # 重新检查
|
||
|
||
# 取第一个空闲 worker
|
||
w = idle[0]
|
||
w.busy = True
|
||
logger.debug("Worker 分配: %s", w.url)
|
||
return w
|
||
|
||
|
||
async def release_worker(w: WorkerState) -> None:
|
||
"""释放 worker,标记为空闲并通知等待者。"""
|
||
async with _pool_condition:
|
||
w.busy = False
|
||
logger.debug("Worker 释放: %s", w.url)
|
||
_pool_condition.notify(1)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 状态查询
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def get_pool_status() -> dict:
|
||
"""返回当前池状态(供 /gateway-health 使用)。"""
|
||
global_info = {
|
||
"global_max": _global_max,
|
||
"global_busy": _global_busy,
|
||
"global_waiting": _global_waiting,
|
||
}
|
||
if not _workers:
|
||
return {"total": 0, "healthy": 0, "busy": 0, **global_info}
|
||
total = len(_workers)
|
||
healthy = sum(1 for w in _workers.values() if w.online)
|
||
busy = sum(1 for w in _workers.values() if w.busy)
|
||
return {"total": total, "healthy": healthy, "busy": busy, **global_info}
|