Files
hair/gateway/pool.py
T
UbuntuandClaude bb9f55e93c feat(gateway): 全局串行化(并发=1) + 记录入参/出参/耗时日志
并发模型从「每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>
2026-07-23 22:56:48 +08:00

345 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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 下线: %sHTTP %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}