为什么用 Rust 做 AI 网关?
AI 推理服务有独特的性能要求:
| 需求 | 挑战 | Rust 优势 |
|---|---|---|
| 流式响应 (SSE) | 长连接、大量并发 | 异步零成本,万级连接无压力 |
| Token 计数 | 每个请求都要算 | 比 Python 快 10-20x |
| 成本追踪 | 高频写入 | 低延迟原子操作 |
| 多模型路由 | 复杂路由逻辑 | 零运行时开销的类型状态机 |
LLM 推理网关架构
客户端
│ POST /v1/chat/completions (OpenAI 兼容接口)
▼
┌─────────────────────────────────────────┐
│ Rust AI Gateway │
│ │
│ ┌──────────┐ ┌──────────┐ ┌────────┐ │
│ │ 认证限流 │→│ 路由选择 │→│ 成本追踪│ │
│ └──────────┘ └────┬─────┘ └────────┘ │
│ │ │
│ ┌───────────┼───────────┐ │
│ ▼ ▼ ▼ │
│ [OpenAI] [Anthropic] [本地 vLLM] │
└─────────────────────────────────────────┘流式响应代理(SSE)
use axum::response::Sse;
use futures::StreamExt;
use tokio_stream::wrappers::ReceiverStream;
pub async fn chat_completions(
State(state): State<Arc<GatewayState>>,
Json(req): Json<ChatRequest>,
) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {
let (tx, rx) = tokio::sync::mpsc::channel(32);
let model = state.router.select_model(&req).await;
tokio::spawn(async move {
let mut stream = model.stream_completion(req).await.unwrap();
while let Some(chunk) = stream.next().await {
let delta = chunk.choices[0].delta.content.as_deref().unwrap_or("");
// 追踪 token 使用
state.metrics.record_tokens(delta);
let data = serde_json::to_string(&chunk).unwrap();
let _ = tx.send(Event::default().data(data)).await;
}
let _ = tx.send(Event::default().data("[DONE]")).await;
});
Sse::new(ReceiverStream::new(rx).map(Ok))
.keep_alive(KeepAlive::default())
}SSE (Server-Sent Events) 是 LLM 流式输出的标准协议。Tokio 的 MPSC channel 解耦了"从上游接收"和"向客户端发送"两个任务,背压机制(bounded channel)防止内存爆炸。
智能路由策略
#[derive(Clone)]
pub struct ModelRouter {
models: Vec<ModelConfig>,
strategy: RoutingStrategy,
}
#[derive(Clone)]
pub enum RoutingStrategy {
CostOptimized, // 优先选最便宜的能处理当前请求的模型
LatencyOptimized, // 优先选响应最快的
LoadBalanced, // 轮询
Canary { new_model: String, percentage: u8 }, // 金丝雀发布
}
impl ModelRouter {
pub async fn select_model(&self, req: &ChatRequest) -> &ModelConfig {
match &self.strategy {
RoutingStrategy::CostOptimized => {
// 根据输入 token 数估算成本,选最便宜能满足上下文的模型
let token_estimate = count_tokens(&req.messages);
self.models.iter()
.filter(|m| m.context_window >= token_estimate)
.min_by_key(|m| m.cost_per_1k_tokens)
.unwrap()
}
RoutingStrategy::Canary { new_model, percentage } => {
// 金丝雀:percentage% 流量到新模型
if rand::random::<u8>() < *percentage * 255 / 100 {
self.models.iter().find(|m| &m.name == new_model).unwrap()
} else {
&self.models[0]
}
}
// ...
}
}
}Token 计数与成本追踪
use tiktoken_rs::cl100k_base;
use std::sync::atomic::{AtomicU64, Ordering};
pub struct CostTracker {
total_input_tokens: AtomicU64,
total_output_tokens: AtomicU64,
total_cost_cents: AtomicU64, // 以分为单位,避免浮点数精度问题
}
impl CostTracker {
pub fn record_usage(&self, model: &str, input: u64, output: u64) {
let (input_price, output_price) = model_pricing(model);
self.total_input_tokens.fetch_add(input, Ordering::Relaxed);
self.total_output_tokens.fetch_add(output, Ordering::Relaxed);
let cost = (input * input_price + output * output_price) / 1000;
self.total_cost_cents.fetch_add(cost, Ordering::Relaxed);
}
pub fn daily_report(&self) -> CostReport {
CostReport {
input_tokens: self.total_input_tokens.load(Ordering::Relaxed),
output_tokens: self.total_output_tokens.load(Ordering::Relaxed),
total_cost_usd: self.total_cost_cents.load(Ordering::Relaxed) as f64 / 100.0,
}
}
}vLLM 集成:本地模型服务
// vLLM 兼容 OpenAI API 格式,可以无缝切换
pub struct VllmClient {
client: reqwest::Client,
base_url: String,
model: String,
}
impl VllmClient {
pub async fn stream_completion(&self, req: ChatRequest)
-> impl Stream<Item = ChatChunk>
{
let response = self.client
.post(format!("{}/v1/chat/completions", self.base_url))
.json(&json!({
"model": self.model,
"messages": req.messages,
"stream": true,
"max_tokens": req.max_tokens.unwrap_or(2048),
}))
.send()
.await
.unwrap();
// 解析 SSE 流
let stream = response.bytes_stream()
.map(|chunk| parse_sse_chunk(chunk.unwrap()));
stream
}
}Prometheus 指标监控
lazy_static::lazy_static! {
static ref LLM_REQUEST_DURATION: Histogram = register_histogram!(
"llm_request_duration_seconds",
"Time to first token (TTFT)",
vec![0.1, 0.5, 1.0, 2.0, 5.0, 10.0]
).unwrap();
static ref LLM_TOKENS_TOTAL: CounterVec = register_counter_vec!(
"llm_tokens_total",
"Total tokens processed",
&["model", "type"] // type: input/output
).unwrap();
static ref LLM_COST_DOLLARS: CounterVec = register_counter_vec!(
"llm_cost_dollars_total",
"Total cost in USD",
&["model"]
).unwrap();
}监控 TTFT(Time to First Token) 比总响应时间更重要——用户感知的延迟是第一个 token 出现的时间。Prometheus Histogram 的分位数让你精确了解 P50/P95/P99 延迟分布。
Kubernetes GPU 调度
# k8s/vllm-deployment.yaml
apiVersion: apps/v1
kind: Deployment
metadata:
name: vllm-server
spec:
replicas: 2
template:
spec:
containers:
- name: vllm
image: vllm/vllm-openai:latest
args:
- --model
- Qwen/Qwen2.5-7B-Instruct
- --tensor-parallel-size
- "1"
resources:
limits:
nvidia.com/gpu: 1 # 请求 1 个 GPU
env:
- name: CUDA_VISIBLE_DEVICES
value: "0"
nodeSelector:
accelerator: nvidia-t4 # 调度到 GPU 节点AI 推理网关实战演示
构建多模型路由网关:OpenAI/Anthropic/本地 vLLM 统一接口、流式响应代理、实时 Token 成本追踪、Grafana 监控面板展示 TTFT 分布
视频即将上线
实战项目
生产级 LLM 推理网关
初级
构建 OpenAI 兼容的 AI 网关:JWT 认证、多模型智能路由(成本/延迟优化)、SSE 流式响应代理、Token 计数与成本追踪、Prometheus 指标暴露、金丝雀发布支持。
SSE 流式传输多模型路由Token 计数成本追踪vLLM 集成Prometheus