use axum::{
extract::{Query, State},
routing::{get, post},
Router,
};
use driftwatch::{
DatasetMonitor, EqualFrequencyBinning, LiveFeature, LiveWindow, ReferenceDistribution,
};
use metrics_exporter_prometheus::{PrometheusBuilder, PrometheusHandle};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
#[derive(Clone)]
struct AppState {
window: Arc<LiveWindow>,
monitor: Arc<DatasetMonitor>,
prometheus: PrometheusHandle,
}
async fn predict(State(state): State<AppState>, Query(q): Query<HashMap<String, f64>>) -> String {
let value = q.get("v").copied().unwrap_or(0.0);
state.window.push(vec![value]);
format!("scored {value}\n")
}
async fn metrics(State(state): State<AppState>) -> String {
state.prometheus.render()
}
#[tokio::main]
async fn main() {
let prometheus = PrometheusBuilder::new()
.install_recorder()
.expect("install prometheus recorder");
let baseline: Vec<f64> = (0..1000).map(|i| (i % 100) as f64 / 100.0).collect();
let mut monitor = DatasetMonitor::new();
monitor.add_feature(
ReferenceDistribution::fit_continuous("score", &baseline, EqualFrequencyBinning::default())
.unwrap(),
);
let state = AppState {
window: Arc::new(LiveWindow::new(driftwatch::WindowMode::Count(500))),
monitor: Arc::new(monitor),
prometheus,
};
{
let checker = state.clone();
tokio::spawn(async move {
let mut ticker = tokio::time::interval(Duration::from_millis(100));
loop {
ticker.tick().await;
let live = checker.window.column(0);
if live.len() >= 2 {
let _ = checker
.monitor
.check(&[("score", LiveFeature::Continuous(&live))]);
}
}
});
}
let app = Router::new()
.route("/predict", post(predict))
.route("/metrics", get(metrics))
.with_state(state);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
for i in 0..400 {
let v = 0.5 + i as f64 / 400.0; http_post(&addr.to_string(), &format!("/predict?v={v}")).await;
}
tokio::time::sleep(Duration::from_millis(300)).await;
let scrape = http_get(&addr.to_string(), "/metrics").await;
println!("--- /metrics (driftwatch gauges) ---");
for line in scrape.lines().filter(|l| l.contains("driftwatch")) {
println!("{line}");
}
}
async fn http_post(addr: &str, path: &str) -> String {
request(addr, "POST", path).await
}
async fn http_get(addr: &str, path: &str) -> String {
request(addr, "GET", path).await
}
async fn request(addr: &str, method: &str, path: &str) -> String {
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let req =
format!("{method} {path} HTTP/1.1\r\nHost: {addr}\r\nConnection: close\r\nContent-Length: 0\r\n\r\n");
stream.write_all(req.as_bytes()).await.unwrap();
let mut buf = Vec::new();
stream.read_to_end(&mut buf).await.unwrap();
let text = String::from_utf8_lossy(&buf);
text.split_once("\r\n\r\n")
.map(|(_, body)| body.to_string())
.unwrap_or_default()
}