Serving Models with Axum
Production ML inference over HTTP with Axum 0.8 - shared model state, validated requests, and non-blocking execution.
Busca en todas las páginas de la documentación
Production ML inference over HTTP with Axum 0.8 - shared model state, validated requests, and non-blocking execution.
use axum::{extract::State, routing::post, Json, Router};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tokio::task;
#[derive(Clone)]
struct AppState {
// session: Arc<ort::session::Session>,
}
#[derive(Deserialize)]
struct PredictRequest {
features: Vec<f32>,
}
#[derive(Serialize)]
struct PredictResponse {
score: f32,
}
async fn predict(
State(_state): State<AppState>,
Json(req): Json<PredictRequest>,
) -> Json<PredictResponse> {
let features = req.features;
let score = task::spawn_blocking(move || run_model(&features))
.await
.unwrap();
Json(PredictResponse { score })
}
fn run_model(features: &[f32]) -> f32 {
features.iter().sum::<f32>() / features.len() as f32
}When to reach for this:
use axum::{extract::State, http::StatusCode, routing::post, Json, Router};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::time::Instant;
use tokio::task;
use tracing::info;
#[derive(Clone)]
struct AppState {
model_version: Arc<String>,
}
#[derive(Deserialize)]
struct PredictRequest {
features: Vec<f32>,
}
#[derive(Serialize)]
struct PredictResponse {
score: f32,
model_version: String,
}
async fn predict(
State(state): State<AppState>,
Json(req): Json<PredictRequest>,
) -> Result<Json<PredictResponse>, StatusCode> {
if req.features.is_empty() || req.features.len() > 10_000 {
return Err(StatusCode::BAD_REQUEST);
}
let version = state.model_version.clone();
let start = Instant::now();
let features = req.features;
let score = task::spawn_blocking(move || infer(&features))
.await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
info!(elapsed_ms = start.elapsed().as_millis(), "inference");
Ok(Json(PredictResponse {
score,
model_version: (*version).clone(),
}))
}
fn infer(features: &[f32]) -> f32 {
features.iter().copied().fold(0.0f32, f32::max)
}
#[tokio::main]
async fn main() {
tracing_subscriber::fmt::init();
let state = AppState {
model_version: Arc::new("v1.0.0".into()),
};
let app = Router::new()
.route("/predict", post(predict))
.with_state(state);
let listener = tokio::net::TcpListener::bind("0.0.0.0:8080").await.unwrap();
axum::serve(listener, app).await.unwrap();
}What this demonstrates:
spawn_blocking for synchronous inferencetracing timing per requestAppState holds Arc<Session> or model weights loaded at boot.| Layer | Purpose |
|---|---|
TimeoutLayer | Kill hung inference |
RequestBodyLimitLayer | Cap huge JSON payloads |
TraceLayer | HTTP-level spans |
spawn_blocking or dedicated thread pool.Arc with watch channel on SIGHUP or file watcher./predict - open GPU burn. Fix: API keys or mTLS at gateway.| Alternative | Use When | Don't Use When |
|---|---|---|
| gRPC + tonic | High-QPS internal mesh | Browser clients need JSON |
| BentoML / Ray Serve | Python ecosystem ownership | Rust-only ops mandate |
| Serverless GPU | Spiky traffic | Cold start unacceptable |
| Batch queue worker | Throughput over latency | Interactive API SLAs |
Router::with_state + State<T> extractor - clone cheap inner Arc fields only.
#[tokio::main] async fn: load in main, pass AppState to router before serve.
Use SSE handler - see LLM Inference & Serving.
Queue in channel worker that groups requests within N ms window - advanced pattern for GPU utilization.
Liveness: process up. Readiness: model loaded and dummy inference under latency threshold.
Use validator crate or manual checks on vector length and value ranges before tensor alloc.
Minimal container with model volume mount, CPU/GPU resource limits, and HPA on queue depth or latency.
OpenTelemetry traces spanning HTTP + spawn_blocking section with model version attribute.
Separate routes or model_id field routing to HashMap<String, Arc<Session>>.
Replace infer stub with ort session call inside blocking task.
Stack versions: This page was written for Rust 1.97.0 (edition 2024), Tokio 1.x, Axum 0.8, serde 1.0, sqlx 0.8, clap 4, and Polars 0.46+.
Revisado por Chris St. John·Última actualización: 16 jul 2026