AI Inference Gateway 生产级架构:请求路由、多模型治理与智能限流系统设计

AI Inference Gateway 生产级架构:请求路由、多模型治理与智能限流系统设计

在大模型从实验走向生产的过程中,一个被严重低估的组件正在成为架构瓶颈:AI Inference Gateway。它不是传统 API Gateway 的简单扩展,而是需要重新理解 LLM 推理请求特征的全新基础设施层。

1. 为什么传统 API Gateway 不够用

Nginx、Kong、Envoy 等传统 API Gateway 围绕 HTTP 请求-响应模型设计,假设后端是无状态、可均匀分发的。但 LLM 推理请求有几个关键差异:

维度传统 API 请求LLM 推理请求
延迟毫秒级(p99 < 50ms)秒级到分钟级(p99 可达 30s+)
成本结构CPU 时间为主GPU 秒 × Token 数,成本可量化
后端异构性同质服务实例不同模型、不同 GPU 规格共存
请求可预测性固定开销输入/输出 Token 数差异 10-100x
重试语义幂等请求可安全重试长推理中断后重试成本极高

这些差异意味着:你需要在网关层理解 Token 这一核心计量单位,并基于 Token Budget 而非简单的 QPS 做限流、路由和降级。

2. 架构设计:六层推理网关模型

一个生产级 AI Inference Gateway 通常包含以下六层:

┌─────────────────────────────────────────────────────────────┐
│  Layer 6: 可观测性 (Metrics / Tracing / Cost Attribution)  │
├─────────────────────────────────────────────────────────────┤
│  Layer 5: 降级与Fallback (Fallback Chain / Circuit Breaker) │
├─────────────────────────────────────────────────────────────┤
│  Layer 4: 智能限流 (Token Rate Limiting / Budget Control)   │
├─────────────────────────────────────────────────────────────┤
│  Layer 3: 多模型路由 (Model Routing / A-B Testing)         │
├─────────────────────────────────────────────────────────────┤
│  Layer 2: 负载均衡 (GPU-Aware Load Balancing)              │
├─────────────────────────────────────────────────────────────┤
│  Layer 1: 协议适配 (OpenAI API / Anthropic / Custom)       │
└─────────────────────────────────────────────────────────────┘

下面我们逐层深入,重点讲解 Layer 3-5 的工程实现。

3. 多模型路由:从规则到智能调度

3.1 路由维度设计

LLM 网关的路由远不只是 "把 /v1/chat/completions 转发到某个后端"。实际场景中的路由维度包括:

  • 模型维度:用户请求 gpt-4o 是否可以用 gpt-4o-mini 替代?
  • 租户维度:企业 A 的预算池已用完,是否路由到降级模型?
  • 延迟维度:当前后端队列深度超过阈值,是否切换到备用集群?
  • 成本维度:同等质量下,哪个后端成本更低?(自部署 vs 第三方 API)
  • 地域维度:数据主权要求某些请求必须在特定区域处理

3.2 实现:基于特征的多维路由表

from dataclasses import dataclass, field
from typing import Optional
from enum import Enum

class RoutingStrategy(Enum):
    DIRECT = "direct"           # 精确匹配
    CHEAPER = "cheaper"         # 成本优先
    FASTER = "faster"           # 延迟优先
    FALLBACK = "fallback"       # 降级链路

@dataclass
class RouteRule:
    """单条路由规则"""
    model_name: str
    target_backends: list
    max_input_tokens: int = 32768
    max_output_tokens: int = 8192
    allow_fallback: bool = True
    cost_budget_usd: float = 0.01
    timeout_seconds: float = 60.0
    tenant_whitelist: list = field(default_factory=list)

class ModelRouter:
    def __init__(self):
        self.rules: dict[str, RouteRule] = {}
        self.fallback_chains: dict[str, list] = {}
        # 分级降级链路
        self.fallback_chains = {
            "gpt-4o": ["gpt-4o-2024-08-06", "gpt-4o-mini", "qwen2.5-72b"],
            "claude-3.5-sonnet": ["claude-3-haiku", "gpt-4o-mini"],
        }

    def resolve(self, requested_model: str, tenant: str = "default",
                estimated_tokens: int = 0) -> RouteRule:
        """根据请求特征决定路由目标"""
        rule = self.rules.get(requested_model)

        # 1. 精确匹配
        if rule and not rule.tenant_whitelist:
            if estimated_tokens <= rule.max_input_tokens:
                return rule

        # 2. 检查是否有降级链路
        chain = self.fallback_chains.get(requested_model, [])
        for fallback_model in chain:
            fb_rule = self.rules.get(fallback_model)
            if fb_rule and estimated_tokens <= fb_rule.max_input_tokens:
                return fb_rule

        # 3. 最终兜底
        return self.rules.get("default")

    def select_backend(self, rule: RouteRule, strategy: RoutingStrategy) -> dict:
        """从候选后端中选择一个"""
        if strategy == RoutingStrategy.CHEAPER:
            return min(rule.target_backends, key=lambda b: b["cost_per_1m_tokens"])
        elif strategy == RoutingStrategy.FASTER:
            return min(rule.target_backends, key=lambda b: b["p99_latency_ms"])
        else:
            return rule.target_backends[0]

4. Token 级限流:超越 QPS 的计量模型

4.1 为什么 QPS 限流不够

假设你限制一个租户 100 QPS。如果每个请求只生成 10 个 Token,后端几乎无压力;如果每个请求要生成 8192 个 Token(长文总结),GPU 会被长时间占用。因此 Token 级的计量是必需的。

核心限流维度:

  • TPM (Tokens Per Minute):每分钟可处理的 Token 总量
  • RPM (Requests Per Minute):每分钟请求数
  • TPD (Tokens Per Day):每日 Token 预算(防止费用爆炸)
  • Concurrency:并发推理请求数(受 GPU 显存约束)

4.2 滑动窗口 Token Bucket 实现

import time
import threading
from collections import deque

class TokenRateLimiter:
    """
    多维度 Token 限流器
    支持 TPM / RPM / TPD 三种计量
    滑动窗口实现,避免固定窗口的突刺问题
    """

    def __init__(self, tpm_limit: int, rpm_limit: int, tpd_limit: int):
        self.tpm_limit = tpm_limit
        self.rpm_limit = rpm_limit
        self.tpd_limit = tpd_limit
        self.request_window = deque()
        self.daily_token_total = 0
        self.day_start = time.time()
        self.lock = threading.Lock()
        self._last_estimated = 0

    def _cleanup_window(self):
        """清除 60 秒外的记录"""
        cutoff = time.time() - 60
        while self.request_window and self.request_window[0][0] < cutoff:
            self.request_window.popleft()

    def _reset_daily(self):
        """跨天重置日配额"""
        if time.time() - self.day_start > 86400:
            self.daily_token_total = 0
            self.day_start = time.time()

    def allow_request(self, estimated_input_tokens: int = 0,
                      estimated_output_tokens: int = 0):
        """
        判断请求是否允许通过。
        返回: (是否允许, 拒绝原因)
        """
        with self.lock:
            self._cleanup_window()
            self._reset_daily()

            estimated_total = estimated_input_tokens + estimated_output_tokens
            self._last_estimated = estimated_total

            # 日配额检查
            if self.daily_token_total + estimated_total > self.tpd_limit:
                return False, f"Daily token budget exceeded"

            # TPM 检查(滑动窗口内累计)
            tpm_total = sum(t[1] + t[2] for t in self.request_window)
            if tpm_total + estimated_total > self.tpm_limit:
                return False, f"TPM limit reached"

            # RPM 检查
            if len(self.request_window) >= self.rpm_limit:
                return False, f"RPM limit reached"

            # 通过:记录本次请求
            now = time.time()
            self.request_window.append(
                (now, estimated_input_tokens, estimated_output_tokens)
            )
            self.daily_token_total += estimated_total
            return True, ""

    def report_actual_tokens(self, input_tokens: int, output_tokens: int):
        """请求完成后用实际 token 数修正计量"""
        diff = (input_tokens + output_tokens) - self._last_estimated
        with self.lock:
            self.daily_token_total += diff


class GatewayRateLimitManager:
    """Gateway 级别的租户限流器管理"""

    def __init__(self):
        self.tenant_limiters: dict[str, TokenRateLimiter] = {}
        self.tier_config = {
            "free":       {"tpm": 100_000,  "rpm": 60,   "tpd": 10_000_000},
            "pro":        {"tpm": 1_000_000,"rpm": 600,  "tpd": 500_000_000},
            "enterprise": {"tpm": 10_000_000,"rpm": 6000, "tpd": 5_000_000_000},
        }

    def get_limiter(self, tenant_id: str, tier: str) -> TokenRateLimiter:
        if tenant_id not in self.tenant_limiters:
            cfg = self.tier_config.get(tier, self.tier_config["free"])
            self.tenant_limiters[tenant_id] = TokenRateLimiter(
                cfg["tpm"], cfg["rpm"], cfg["tpd"]
            )
        return self.tenant_limiters[tenant_id]
实践提示: 在流式响应 (SSE) 场景下,output_tokens 需要预估。常见做法是根据历史请求的 input/output 比例预估(通常 1:1 到 1:4),或设置保守上界。请求结束后用实际 token 数做 diff 修正。

5. GPU 感知的负载均衡

5.1 传统 LB vs GPU 感知 LB

Nginx 的 Round-Robin 或 Least-Connections 对 LLM 推理无效:

  • 不同请求的 GPU 占用差异巨大(1K tokens vs 100K tokens)
  • KV Cache 的显存占用会随序列长度动态变化
  • 批处理状态下新请求的排队时间不可预测

5.2 后端健康状态模型

@dataclass
class BackendHealth:
    """推理后端的实时健康状态"""
    url: str
    model: str
    gpu_type: str                    # e.g., "A100-80G", "H100"

    # 实时指标
    queue_depth: int = 0
    active_requests: int = 0
    kv_cache_usage_pct: float = 0.0
    tokens_per_second: float = 0.0

    # 成本
    cost_per_1m_input: float = 0.0
    cost_per_1m_output: float = 0.0

    # 可用性
    last_health_check: float = 0.0
    consecutive_failures: int = 0
    is_healthy: bool = True

    @property
    def estimated_wait_ms(self) -> float:
        """预估排队时间(ms)"""
        if self.tokens_per_second > 0:
            return (self.queue_depth * 512 / self.tokens_per_second) * 1000
        return 99999.0

    @property
    def available_slots(self) -> int:
        """估算当前可接受的并发数"""
        if self.kv_cache_usage_pct > 0.95:
            return 0
        if self.kv_cache_usage_pct > 0.85:
            return max(0, 4 - self.active_requests)
        return max(0, 32 - self.active_requests)


class GPULoadBalancer:
    """GPU 感知的负载均衡器"""

    def __init__(self, backends: list):
        self.backends = {b.model: b for b in backends}

    def select(self, model: str, strategy: str = "least_queued"):
        """选择最优后端"""
        candidates = [b for b in self.backends.values()
                      if b.model == model and b.is_healthy]
        if not candidates:
            return None

        if strategy == "least_queued":
            return min(candidates, key=lambda b: b.queue_depth)
        elif strategy == "least_latency":
            return min(candidates, key=lambda b: b.estimated_wait_ms)
        elif strategy == "most_available":
            return max(candidates, key=lambda b: b.available_slots)
        elif strategy == "cost_optimized":
            return min(candidates, key=lambda b: b.cost_per_1m_output)
        else:
            return candidates[0]

6. 降级与容灾:六层熔断模型

LLM 服务的容灾比传统微服务更复杂:我们希望部分可用(用降级模型也要返回结果),而非直接失败。

6.1 Fallback Chain 设计

class AllBackendsExhausted(Exception):
    pass

class FallbackChain:
    """
    降级链路管理器
    优先级: 精确模型 → 同系列新版 → 同系列小模型 → 跨系列替代 → 本地兜底 → 报错
    """

    FALLBACK_GRAPH = {
        "gpt-4o": [
            {"model": "gpt-4o-2024-08-06", "reason": "version_fallback"},
            {"model": "gpt-4o-mini",    "reason": "capability_downgrade"},
            {"model": "qwen2.5-72b",     "reason": "cross_provider_fallback"},
            {"model": "llama-3.1-70b-local", "reason": "onprem_fallback"},
        ],
        "claude-3.5-sonnet": [
            {"model": "claude-3-haiku",  "reason": "capability_downgrade"},
            {"model": "gpt-4o",          "reason": "cross_provider_fallback"},
            {"model": "qwen2.5-72b",     "reason": "cross_provider_fallback"},
        ],
    }

    async def execute_with_fallback(self, request: dict) -> dict:
        """带降级链的执行"""
        requested_model = request["model"]
        chain = self.FALLBACK_GRAPH.get(requested_model, [])
        attempted = [requested_model] + [s["model"] for s in chain]

        for model in attempted:
            try:
                timeout = 60 if model == requested_model else 30
                result = await self._call_backend(model, request, timeout)
                if model != requested_model:
                    result["_fallback_info"] = {
                        "original_model": requested_model,
                        "actual_model": model,
                        "reason": self._get_fallback_reason(
                            requested_model, model
                        ),
                    }
                return result
            except Exception as e:
                print(f"Fallback: {model} failed ({e}), trying next")
                continue

        raise AllBackendsExhausted(
            f"All fallbacks failed. Attempted: {attempted}"
        )

6.2 基于滑动窗口的熔断器

class CircuitBreaker:
    """
    针对后端的熔断器
    状态: CLOSED(正常) → OPEN(熔断) → HALF_OPEN(探测)
    """

    def __init__(self, failure_threshold: int = 5,
                 recovery_timeout: float = 30.0,
                 success_threshold: int = 3):
        self.failure_threshold = failure_threshold
        self.recovery_timeout = recovery_timeout
        self.success_threshold = success_threshold

        self.state = "CLOSED"
        self.failure_count = 0
        self.success_count = 0
        self.last_failure_time = 0.0

    def record_success(self):
        if self.state == "HALF_OPEN":
            self.success_count += 1
            if self.success_count >= self.success_threshold:
                self.state = "CLOSED"
                self.failure_count = 0
        else:
            self.failure_count = max(0, self.failure_count - 1)

    def record_failure(self):
        self.failure_count += 1
        self.last_failure_time = time.time()
        if self.state == "HALF_OPEN":
            self.state = "OPEN"
            self.success_count = 0
        elif self.failure_count >= self.failure_threshold:
            self.state = "OPEN"

    def allow_request(self) -> bool:
        if self.state == "CLOSED":
            return True
        if self.state == "OPEN":
            if time.time() - self.last_failure_time > self.recovery_timeout:
                self.state = "HALF_OPEN"
                self.success_count = 0
                return True
            return False
        return True

7. 可观测性:推理网关的黄金指标

LLM 网关的可观测性不能只看 HTTP Status Code。核心指标:

  • TTFT (Time To First Token):首 Token 延迟,决定用户感知响应速度
  • TPOT (Time Per Output Token):每 Token 生成间隔,决定流式体验
  • Token Throughput:每秒处理 Token 数
  • GPU KV Cache Utilization:KV Cache 显存利用率
  • Fallback Rate:触发降级的请求比例
  • Cost per 1K Tokens:每千 Token 成本

这些指标需要在网关层打点,并与具体租户、模型、后端关联。OpenTelemetry + Prometheus 是常见栈:每个请求在 Gateway 侧生成一个 Span,携带 model_name、tier、tenant_id、backend_endpoint、fallback_used 等标签。

8. 实战经验总结

基于生产环境中的实际经验,总结几点关键设计决策:

  1. Token 预估是核心痛点:流式请求无法提前知道输出 Token 数。方案:用 tiktoken 快速计算输入长度,输出按历史 P50 预估,事后用实际值修正计费。
  2. 限流管道化:将限流、鉴权、路由、审计做成中间件管道,每层独立可测。每一层都能快速 fail-fast,避免无效的 GPU 调用。
  3. 后端健康检查要区分层次:Liveness 检查(进程存活)→ Readiness 检查(模型加载完成)→ Deep Health Check(GPU 可用显存大于阈值)→ Load Check(队列深度小于阈值)。
  4. 流式场景下的超时设计:不要用整体超时(60s),而是用 Chunk 间隔超时(如果 10s 没有新 Chunk 到达即判死)。这能更精准地检测后端假死。
  5. 费用防护是必须的:设置租户级别的 hard cap(如每月 $5000),超出后进入只读模式或降级到免费模型。没有预算防护的网关等于给 GPU 火灾买了保险但不装烟雾报警器。
  6. 架构演进路径:初创期用 LiteLLM(开源,支持 100+ 后端)→ 规模商用时自研(需要定制化限流/路由)→ 最终形态是控制面(Go/Rust)+ 数据面(Rust/eBPF)分离。

9. 结语

AI Inference Gateway 不是 API Gateway 的升级版,而是理解计算范式变迁后的重新设计。它要求我们:

  • 从 "请求-响应" 思维转向 "Token-计算" 思维
  • 从 "均匀分发" 转向 "异构感知调度"
  • 从 "最大化可用性" 转向 "成本约束下的最优体验"

在大模型成本仍然是 AI 落地最大瓶颈的当下,一个精心设计的推理网关往往能带来 30-60% 的成本优化空间——这不是锦上添花,而是 AI 基础设施的核心竞争力。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部