use std::sync::Arc;
use axum::extract::State;
use axum::http::StatusCode;
use axum::routing::post;
use axum::{Json, Router};
use serde::{Deserialize, Serialize};
use crate::error::{Error, Result};
use crate::frame::Frame;
use crate::onnx::InferenceModel;
#[cfg(feature = "monitor")]
use crate::monitor::DriftMonitor;
#[derive(Deserialize)]
struct PredictRequest {
rows: Vec<Vec<f64>>,
}
#[derive(Serialize)]
struct PredictResponse {
predictions: Vec<f64>,
}
struct AppState {
model: InferenceModel,
#[cfg(feature = "monitor")]
monitor: Option<Arc<DriftMonitor>>,
}
pub struct Server {
model: InferenceModel,
route: String,
#[cfg(feature = "monitor")]
monitor: Option<Arc<DriftMonitor>>,
}
impl Server {
pub fn from_onnx(path: impl AsRef<std::path::Path>) -> Result<Server> {
Ok(Server {
model: InferenceModel::load(path)?,
route: "/predict".to_string(),
#[cfg(feature = "monitor")]
monitor: None,
})
}
#[cfg(feature = "registry")]
pub fn from_registry(
registry: &crate::registry::Registry,
name: &str,
tag: &str,
) -> Result<Server> {
Server::from_onnx(registry.onnx_path(name, tag)?)
}
pub fn route(mut self, path: impl Into<String>) -> Self {
self.route = path.into();
self
}
#[cfg(feature = "monitor")]
pub fn with_monitor(mut self, monitor: DriftMonitor) -> Self {
self.monitor = Some(Arc::new(monitor));
self
}
pub fn router(self) -> Router {
let route = self.route.clone();
let state = Arc::new(AppState {
model: self.model,
#[cfg(feature = "monitor")]
monitor: self.monitor,
});
let router = Router::new().route(&route, post(predict_handler));
#[cfg(feature = "monitor")]
let router = router.route("/metrics", axum::routing::get(metrics_handler));
router.with_state(state)
}
pub async fn serve(self, addr: &str) -> Result<()> {
let router = self.router();
let listener = tokio::net::TcpListener::bind(addr)
.await
.map_err(|e| Error::Backend(format!("bind {addr}: {e}")))?;
axum::serve(listener, router)
.await
.map_err(|e| Error::Backend(format!("serve: {e}")))
}
}
async fn predict_handler(
State(state): State<Arc<AppState>>,
Json(req): Json<PredictRequest>,
) -> std::result::Result<Json<PredictResponse>, (StatusCode, String)> {
if req.rows.is_empty() {
return Err((StatusCode::BAD_REQUEST, "no rows provided".into()));
}
let ncols = req.rows[0].len();
if ncols == 0 || req.rows.iter().any(|r| r.len() != ncols) {
return Err((
StatusCode::BAD_REQUEST,
"rows must be non-empty and rectangular".into(),
));
}
let columns = (0..ncols).map(|i| format!("f{i}")).collect();
let frame = Frame::from_rows(req.rows, columns)
.map_err(|e| (StatusCode::BAD_REQUEST, e.to_string()))?;
let predictions = state
.model
.predict(&frame)
.map_err(|e| (StatusCode::UNPROCESSABLE_ENTITY, e.to_string()))?;
#[cfg(feature = "monitor")]
if let Some(monitor) = &state.monitor {
monitor.observe(&predictions);
}
Ok(Json(PredictResponse { predictions }))
}
#[cfg(feature = "monitor")]
async fn metrics_handler(
State(state): State<Arc<AppState>>,
) -> std::result::Result<Json<serde_json::Value>, (StatusCode, String)> {
let Some(monitor) = &state.monitor else {
return Err((StatusCode::NOT_FOUND, "no monitor attached".into()));
};
let status = monitor
.report()
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
Ok(Json(serde_json::json!({
"drifted": status.drifted,
"psi": status.psi,
"observed": status.observed,
})))
}
#[cfg(all(test, feature = "smartcore-backend"))]
mod tests {
use super::*;
use crate::backends::smartcore::LinearRegression;
use crate::frame::Dataset;
use crate::onnx::ExportOnnx;
use crate::traits::Estimator;
use axum::body::Body;
use axum::http::Request;
use http_body_util::BodyExt;
use tower::ServiceExt;
fn linear_onnx(tag: &str) -> std::path::PathBuf {
let rows: Vec<Vec<f64>> = (0..15).map(|i| vec![i as f64, (i % 4) as f64]).collect();
let y: Vec<f64> = rows.iter().map(|r| 2.0 * r[0] + 3.0 * r[1] + 1.0).collect();
let ds = Dataset::new(
Frame::from_rows(rows, vec!["x1".into(), "x2".into()]).unwrap(),
y,
)
.unwrap();
let mut lr = LinearRegression::new();
lr.fit(&ds).unwrap();
let path = std::env::temp_dir().join(format!("mw_serve_{}_{tag}.onnx", std::process::id()));
lr.export_onnx(&path).unwrap();
path
}
#[tokio::test]
async fn predict_endpoint_returns_predictions() {
let path = linear_onnx("predict");
let app = Server::from_onnx(&path).unwrap().router();
let body =
serde_json::to_vec(&serde_json::json!({ "rows": [[20.0, 1.0], [5.0, 2.0]] })).unwrap();
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/predict")
.header("content-type", "application/json")
.body(Body::from(body))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let bytes = response.into_body().collect().await.unwrap().to_bytes();
let parsed: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
let preds = parsed["predictions"].as_array().unwrap();
assert!((preds[0].as_f64().unwrap() - 44.0).abs() < 1e-2);
assert!((preds[1].as_f64().unwrap() - 17.0).abs() < 1e-2);
let _ = std::fs::remove_file(&path);
}
#[tokio::test]
async fn ragged_rows_are_rejected() {
let path = linear_onnx("ragged");
let app = Server::from_onnx(&path).unwrap().router();
let body = serde_json::to_vec(&serde_json::json!({ "rows": [[1.0, 2.0], [3.0]] })).unwrap();
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/predict")
.header("content-type", "application/json")
.body(Body::from(body))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let _ = std::fs::remove_file(&path);
}
}