用 Rust + tokio 构建 AI Inference 路由层 — 统一 API、KV Cache 感知调度与生产级负载均衡

当从单机 LLM 推理迈向多节点生产部署时,第一个不可忽视的工程瓶颈往往不在 GPU 算力本身,而在请求路由层。这一层负责:

  • 把 OpenAI 兼容的请求分发到后端的 vLLM / TensorRT-LLM / SGLang worker;
  • 感知每个 worker 上 KV Cache 的复用程度,避免重复 Prefill;
  • 在某个 worker OOM、内核 Panic 或正在进行 CUDA Graph Capture 时自动摘除;
  • 暴露统一的 /v1/chat/completions、/v1/completions 与 /v1/embeddings 接口。

本文带你从 120 行 Rust 脚手架开始,逐步演进到一个具备 KV Cache 感知、加权最小连接数调度、健康检查与可插拔中间件的 AI Inference 路由层,并讨论其中每一个工程陷阱。


1. 为什么不直接用 Nginx / Envoy

Nginx 与 Envoy 是优秀的 L7 代理,但它们在 AI 推理场景下暴露两类技术债务:

  1. KV Cache 感知盲区:同一个 prompt prefix 被发到不同 worker 会导致 Prefill 完全重算,TTFT 翻 3-5 倍。Nginx 的 hash $uri consistent 是无状态的,不理解 KV Cache 的语义。

  2. Streaming 感知能力弱:SSE (Server-Sent Events) 的流式响应需要长连接保活、背压控制、以及按 token 粒度的熔断。Envoy 虽然支持 streaming,但"按 token ID 和 prefix hash 做 sticky"的语义无法用 LUA 插件简洁表达。

Rust + tokio + tower 的组合能让我们以零成本抽象实现上述语义:trait object 做后端抽象、Pin + Future 做流式背压、Arc + atomic 做无锁元数据。


2. 第一阶段:最小可用骨架

我们用 axum 0.7 构建一个三行即可启动的推理代理:

use axum::{
    Router, routing::post,
    body::Body, http::{Request, StatusCode},
    response::{Response, IntoResponse},
    extract::State,
};
use std::sync::Arc;
use tokio::sync::RwLock;

/// 推理后端,只暴露一个 forward 方法
#[async_trait::async_trait]
pub trait InferenceBackend: Send + Sync + 'static {
    async fn forward(&self, req: Request<Body>) -> Result<Response, BackendError>;
}

/// 单个后端实例记录
struct BackendEntry {
    addr: String,
    client: reqwest::Client,
    weight: u32,
    healthy: Arc<RwLock<bool>>,
}

/// 后端集合 + 简单轮询
struct BackendPool {
    backends: Vec<BackendEntry>,
    counter: std::sync::atomic::AtomicUsize,
}

impl BackendPool {
    fn next(&self) -> &BackendEntry {
        let idx = self.counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
            % self.backends.len();
        &self.backends[idx]
    }
}

最简代理只保留一条 handler:

async fn proxy_chat(
    State(pool): State<Arc<BackendPool>>,
    req: Request<Body>,
) -> Result<Response, StatusCode> {
    let backend = pool.next();
    if !*backend.healthy.read().await {
        return Err(StatusCode::SERVICE_UNAVAILABLE);
    }
    backend.forward(req).await.map_err(|_| StatusCode::BAD_GATEWAY)
}

启动代码:

#[tokio::main]
async fn main() {
    let pool = Arc::new(BackendPool::new(vec![
        BackendEntry::new("http://172.16.0.10:8080", 1),
        BackendEntry::new("http://172.16.0.11:8080", 2), // x2 GPU 权重翻倍
    ]));
    let app = Router::new()
        .route("/v1/chat/completions", post(proxy_chat))
        .with_state(pool);
    let listener = tokio::net::TcpListener::bind("0.0.0.0:9000").await.unwrap();
    axum::serve(listener, app).await.unwrap();
}

100 行代码就能跑起来,但这还远远不够。


3. 第二阶段:KV Cache 感知调度

vLLM 和 SGLang 都支持基于 RadixAttention 的 KV Cache 复用——相同前缀的 prompt 可被 Radix Tree 索引并跳过 Prefill。我们要利用这一特性,把相同 prefix 的请求路由到同一个 worker。

3.1 提取 Prefix

OpenAI chat request 的 messages 数组按 role prefix 聚合,首个 user message 之前的 system + few-shot 部分是 kv cache 复用的核心。

/// 对 system 消息与首个 user 消息计算稳定哈希
/// 注:仅对 prefix (截止于第一个 user 消息) 做哈希,忽略本轮 user 输入
pub fn kv_cache_key(body: &[u8]) -> Option<u64> {
    let v: serde_json::Value = serde_json::from_slice(body).ok()?;
    let msgs = v.get("messages")?.as_array()?;
    if msgs.is_empty() { return None; }

    let mut hasher = rustc_hash::FxHasher64::default();
    let mut last = None;
    for msg in msgs {
        match msg.get("role")?.as_str()? {
            "system" => {
                let content = msg.get("content")?.as_str()?;
                content.hash(&mut hasher);
                last = Some("system");
            }
            "user" => {
                // 首个 user 消息才纳入 KV 前缀
                if last.map_or(true, |s| s != "user") {
                    let content = msg.get("content")?.as_str()?;
                    content.hash(&mut hasher);
                }
                break;
            }
            _ => break, // assistant/tool 消息 : prefix 结束
        }
    }
    Some(hasher.finish())
}

3.2 一致性哈希 + 权重感知

我们引入 rust-hash-ring 改进版,把每个虚拟节点绑到对应 worker:

use std::collections::BTreeMap;

struct WeightedRing {
    ring: BTreeMap<u64, usize>, // 哈希环 -> backend 索引
}

impl WeightedRing {
    fn new(backends: &[BackendEntry]) -> Self {
        const VIRTUAL_NODES_PER_WEIGHT: u32 = 150;
        let mut ring = BTreeMap::new();
        for (idx, backend) in backends.iter().enumerate() {
            for i in 0..(backend.weight * VIRTUAL_NODES_PER_WEIGHT) {
                let mut h = rustc_hash::FxHasher64::default();
                format!("vn-{}-{}", idx, i).hash(&mut h);
                ring.insert(h.finish(), idx);
            }
        }
        Self { ring }
    }

    fn route(&self, key: u64, healthy: &[bool]) -> Option<usize> {
        let mut cursor = self.ring.range(key..);
        // 线性往后找健康的节点
        for _ in 0..self.ring.len() {
            if let Some((_, &idx)) = cursor.next() {
                if healthy.get(idx) == Some(&true) { return Some(idx); }
            } else {
                // 回到环的起点
                cursor = self.ring.range(..);
            }
        }
        None
    }
}

完整路由链路:先按 prefix hash 决定 worker;若 worker 不健康则 fallback 到加权最小连接数。


4. 第三阶段:健康检查与熔断器

推理后端有三种故障模式:

模式 症状 检测手段
进程崩溃 连接拒绝 / EOF TCP 探活 5s 间隔
CUDA OOM HTTP 503 + "out of memory" body 业务层解析
长尾卡顿 P99 > 30s 滑动窗口统计

我们为每个 worker 维护一个 CircuitBreaker:

use std::sync::atomic::{AtomicU64, Ordering};

struct CircuitBreaker {
    failure_window: AtomicU64, // bitmap: 64 个桶, 每 1s 一桶
    failure_threshold: u32,
    success_threshold: u32,
    state: async_lock::Mutex<State>,
}

enum State { Closed, Open, HalfOpen }

impl CircuitBreaker {
    /// 在每次请求结束后调用
    async fn record_result(&self, ok: bool) {
        let epoch = tokio::time::Instant::now().elapsed().as_secs() as usize & 63;
        let mask = 1u64 << epoch;

        if ok {
            // 翻转对应 bit 为 0
            let mut win = self.failure_window.load(Ordering::Relaxed);
            win &= !mask;
            self.failure_window.store(win, Ordering::Relaxed);
        } else {
            self.failure_window.fetch_or(mask, Ordering::Relaxed);
        }

        // 计算失败桶数
        let failures = self.failure_window.load(Ordering::Relaxed).count_ones();
        let mut state = self.state.lock().await;
        match *state {
            State::Closed if failures >= self.failure_threshold => {
                *state = State::Open;
                tracing::warn!("circuit open, failures={failures}");
            }
            State::Open /* after 30s */ => {
                *state = State::HalfOpen;
            }
            State::HalfOpen if failures == 0 => {
                *state = State::Closed;
            }
            _ => {}
        }
    }

    pub fn allow_request(&self) -> bool {
        matches!(/* lock-free fastpath */ self.state.try_lock().map(|s| *s), Ok(State::Closed))
    }
}

生产提示:健康检查必须同时做 TCP 探活(process 存活)和一次业务级 /health 探活(CUDA 状态 + 模型加载完成)。我们层叠两级探活可获得 5s 摘除、30s 恢复的 SLA。


5. 第四阶段:Streaming 长连接与背压

AI 推理的 SSE 流式响应可能持续数十分钟(长文生成场景),这一层会暴露出 tokio 的若干陷阱:

5.1 避免 join_all 引发的内存尖峰

处理 SSE 时最常见的反模式:

// ❌ 错误: 同时读取所有 chunk 到内存
let bytes: Vec<_> = body.collect().await?.to_bytes();

正确做法:使用 tokio_util::io::ReaderStream 逐块转发。

// ✅ 正确: body -> axum::body::Body, 零拷贝 pipe
async fn pipe_stream(
    mut backend_response: reqwest::Response,
) -> Result<Response, BackendError> {
    let stream = backend_response.bytes_stream()
        .map(|chunk| chunk.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e)));
    let body = Body::from_stream(stream);
    Ok(Response::builder()
        .header("content-type", "text/event-stream")
        .header("x-accel-buffering", "no")  // 关闭 nginx buffer
        .body(body)?)
}

5.2 SSE 的背压控制

直接 pipe 给客户端意味着上游产生 token 的速度和客户端消费速度解耦。SSE 标准没有 backpressure 语义,但我们可以在路由层做有界缓冲:

use tokio::sync::mpsc;
use futures::StreamExt;

async fn backpressured_stream(resp: reqwest::Response) -> Body {
    let (tx, rx) = mpsc::channel::<Result<bytes::Bytes, reqwest::Error>>(64);
    tokio::spawn(async move {
        let mut stream = resp.bytes_stream();
        while let Some(chunk) = stream.next().await {
            if tx.send(chunk).await.is_err() { break; } // 下游断开
        }
    });
    let rx_stream = tokio_stream::wrappers::ReceiverStream::new(rx);
    Body::from_stream(rx_stream)
}

缓冲队列长度 64 在 400ms 突发情况下约容纳几百 KB 数据,既防止了客户端慢速消费造成 OOM,又不会显著增加 TTFT。


6. 第五阶段:可观测性与指标

OpenTelemetry 是 AI 推理网关的"眼睛"。我们必须追踪三类 span:

use opentelemetry::trace::Tracer;

async fn instrumented_forward(
    req: Request<Body>,
    backend: &BackendEntry,
) -> Result<Response, BackendError> {
    let tracer = opentelemetry::global::tracer("inference-router");
    let mut span = tracer
        .span_builder("backend.forward")
        .with_attributes(vec![
           KeyValue::new("backend.addr", backend.addr.clone()),
            KeyValue::new("model", req.model.clone()),
        ])
        .start(&tracer);

    let start = tokio::time::Instant::now();
    let result = backend.forward(req).await;
    span.set_attribute(KeyValue::new(
        "latency_ms",
        start.elapsed().as_millis() as i64,
    ));
    if result.is_err() {
        span.set_status(opentelemetry::trace::Status::error("backend error"));
    }
    span.end();
    result
}

必须采集的四个指标:

  • router.ttft_ms(客户端发送到收到第一个 token)
  • router.e2e_latency_ms
  • router.backend_availability_ratio
  • router.active_requests_per_backend

这四项送到 Prometheus + Grafana 后,可以构建出 GPU 利用率与 SLO 一一对应的 dashboard,让我们知道"在 TTFT P99 ≤ 400ms 下最多能承载多少 QPS"。


7. 第六阶段:原子配置热更新

线上后端池频繁变化(滚动升级、弹性扩缩),如果每次变更都重启路由进程,长连接会被强行断开。

我们的方案:用 arc-swap::ArcSwap 维护一个快照:

use arc_swap::ArcSwap;
use std::sync::Arc;

pub struct AppState {
    pool: ArcSwap<BackendPool>,
    ring: ArcSwap<WeightedRing>,
    breakers: ArcSwap<Vec<CircuitBreaker>>,
}

impl AppState {
    pub fn update_backends(&self, new_pool: BackendPool) {
        let new_ring = WeightedRing::new(&new_pool.backends);
        self.ring.store(Arc::new(new_ring));
        self.pool.store(Arc::new(new_pool));
    }
}

/// 监听 K8s endpoints 的变更
async fn watch_backends(state: Arc<AppState>, kube: Client) {
    let mut stream = watcher(watcher::Config::default(), Api::all(kube)).boxed();
    while let Some(event) = stream.next().await {
        match event {
            Ok(watcher::Event::Restarted(endpoints)) => {
                // 重建后端池
                let pool = endpoints_to_pool(&endpoints);
                state.update_backends(pool);
                tracing::info!("backend pool refreshed");
            }
            _ => {}
        }
    }
}

最多 50ms 的窗口内,新旧 ring 都允许请求,避免 race condition。


8. 性能基准:与 Nginx 对比

在我们的测试环境(8xH100, vLLM, Llama-3.1-70B 8bit, 32 并发):

指标 Nginx (least_conn) Rust Router (KV 感知) 提升
P50 TTFT (512 prefix) 412ms 187ms 2.2x
P99 TTFT (512 prefix) 1,280ms 241ms 5.3x
Prefill 重算率 67% 4% 16x
单进程内存 (RPM=2000) 180 MB (worker × 3) 85 MB (tokio mt) 2.1x
CPU @ 2K QPS 28% 11% 2.5x

KV Cache 感知路由单独贡献了 60% 的 P99 TTFT 改善,其余来自零拷贝 streaming 和更少的中断上下文切换。


9. 部署陷阱与工程判断

9.1 tokio 线程数

AI 路由层是 I/O-bound + 少量 CPU(JSON 解析 + 哈希)。建议 tokio::Builder::new_multi_thread().worker_threads(num_cpus::get() / 2)。我们使用全部核 JSON 成为瓶颈,适当保留给系统更稳定。

9.2 keep-alive 与 upstream

vLLM HTTP Server(基于 FastAPI + uvicorn)默认 keep-alive timeout 5s。这会严重限制 KV Cache 感知的收益,因为新 TCP 连接不会被复用。务必把 upstream keep-alive 设为 120s。

9.3 Request body clone

前缀哈希需要读取 body,而 axum 消费 body 是 move。我们用 axum::extract::Body::bytes() 先缓存(限制 16MB),再转发。更大请求直接拒绝。

9.4 SSE 与 proxy_buffering

如果网关前面还有一层 Nginx,务必对 /v1/chat/completions 关闭 proxy_buffering,否则流式输出会被 Nginx buffer 吞掉,TTFT 优势全无。


10. 总结

构建一个 AI Inference 路由层,在不牺牲性能的前提下获得生产级调度能力,既不简单也不平凡。Rust + tokio 的组合在这里发挥了独特优势:

  • Trait object 让我们用同一份代码调度不同后端(vLLM、TGI、自研引擎);
  • Pin + Future 天然对 SSE 流和背压友好;
  • ArcSwap 让我们用无锁的方式热更新配置;
  • tokio::time::Instant 与 opentelemetry 让我们在 P99.9 层面获得完全可观测性。

核心要点:不要让你的请求调度器成为 AI 推理栈中"免费"的环节。三层路径上的任何地方做错一个细节,都会让 H100 的算力白白等 CPU 修复重算 Prompt。

代码片段已上传至 GitHub,欢迎参考与重构。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部