use std::{
future::Future,
pin::Pin,
task::{Context, Poll},
time::{Duration, Instant},
};
use axum::body::Body;
use tower::{Layer, Service};
pub use shardline_metrics::*;
pub fn record_upload(protocol: &str, bytes: u64, duration_secs: f64, ok: bool) {
let status = if ok { 200_u16 } else { 500 };
shardline_metrics::record_upload(protocol, bytes);
shardline_metrics::metrics()
.transfer
.record_upload_duration(duration_secs);
let _ = status;
}
pub fn record_download(protocol: &str, bytes: u64, duration_secs: f64, ok: bool) {
let status = if ok { 200_u16 } else { 500 };
shardline_metrics::record_download(protocol, bytes);
shardline_metrics::metrics()
.transfer
.record_download_duration(duration_secs);
let _ = status;
}
pub fn record_range_request() {
shardline_metrics::metrics().transfer.record_range_request();
}
pub fn record_webhook_event(provider: &str, event_type: &str, duration_secs: f64) {
shardline_metrics::record_provider_webhook(provider, event_type);
shardline_metrics::metrics()
.provider
.record_webhook_duration(Duration::from_secs_f64(duration_secs));
}
pub fn record_token_exchange() {
shardline_metrics::record_provider_token_exchange();
}
pub fn record_chunk_inserted(bytes: u64) {
shardline_metrics::metrics()
.storage
.record_chunk_stored(bytes);
}
pub fn record_xorb_stored(bytes: u64) {
shardline_metrics::metrics()
.storage
.record_xorb_stored(bytes);
}
pub fn record_shard_stored() {
shardline_metrics::metrics().storage.record_shard_stored();
}
pub fn record_lfs_upload() {
shardline_metrics::metrics().protocol.record_lfs_upload();
}
pub fn record_lfs_download() {
shardline_metrics::metrics().protocol.record_lfs_download();
}
pub fn record_xet_xorb_download(bytes: u64) {
shardline_metrics::record_xet_xorb_download(bytes);
}
pub fn record_dedup_saves(bytes: u64) {
shardline_metrics::metrics()
.storage
.record_dedup_saves(bytes);
}
#[cfg(test)]
mod tests {
use super::*;
use shardline_metrics::metrics;
#[test]
fn record_upload_no_panic() {
record_upload("http", 1024, 1.5, true);
record_upload("grpc", 0, 0.0, false);
}
#[test]
fn record_download_no_panic() {
record_download("http", 512, 0.5, true);
record_download("grpc", 0, 0.0, false);
}
#[test]
fn record_range_request_no_panic() {
record_range_request();
record_range_request();
}
#[test]
fn record_webhook_event_no_panic() {
record_webhook_event("github", "push", 0.25);
record_webhook_event("gitlab", "merge_request", 0.0);
}
#[test]
fn record_upload_increments_upload_counter() {
let before = metrics().transfer.upload_requests.get();
record_upload("http", 42, 0.1, true);
let after = metrics().transfer.upload_requests.get();
assert!(
after > before,
"upload_requests should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_download_increments_download_counter() {
let before = metrics().transfer.download_requests.get();
record_download("http", 99, 0.2, true);
let after = metrics().transfer.download_requests.get();
assert!(
after > before,
"download_requests should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_range_request_increments_range_counter() {
let before = metrics().transfer.range_requests.get();
record_range_request();
let after = metrics().transfer.range_requests.get();
assert!(
after > before,
"range_requests should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_token_exchange_increments_counter() {
let before = metrics().provider.token_exchanges.get();
record_token_exchange();
let after = metrics().provider.token_exchanges.get();
assert!(
after > before,
"token_exchanges should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_chunk_inserted_increments_chunk_counter() {
let before = metrics().storage.chunks_total.get();
record_chunk_inserted(64);
let after = metrics().storage.chunks_total.get();
assert!(
after > before,
"chunks_total should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_xorb_stored_increments_xorb_counter() {
let before = metrics().storage.xorbs_total.get();
record_xorb_stored(128);
let after = metrics().storage.xorbs_total.get();
assert!(
after > before,
"xorbs_total should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_shard_stored_increments_shard_counter() {
let before = metrics().storage.shards_total.get();
record_shard_stored();
let after = metrics().storage.shards_total.get();
assert!(
after > before,
"shards_total should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_dedup_saves_increments_dedup_counter() {
let before = metrics().storage.dedup_saves_bytes_total.get();
record_dedup_saves(1024);
let after = metrics().storage.dedup_saves_bytes_total.get();
assert!(
after > before,
"dedup_saves_bytes_total should increase (before: {before}, after: {after})"
);
}
#[test]
fn metrics_layer_is_cloneable() {
let layer = MetricsLayer;
let _clone = layer;
}
#[test]
fn metrics_service_construction() {
let inner = tower::util::service_fn(|_req: axum::http::Request<Body>| async {
Ok::<_, std::convert::Infallible>(axum::http::Response::new(Body::empty()))
});
let _svc = MetricsService { inner };
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn metrics_service_poll_ready_and_call_tracks_connections() {
use tower::ServiceExt;
let svc = MetricsService {
inner: tower::util::service_fn(|_req: axum::http::Request<Body>| async {
Ok::<_, std::convert::Infallible>(axum::http::Response::new(Body::empty()))
}),
};
let before = metrics().system.active_connections.get();
let _response = svc
.oneshot(axum::http::Request::builder().body(Body::empty()).unwrap())
.await;
let after = metrics().system.active_connections.get();
assert_eq!(after, before);
}
#[test]
fn record_token_exchange_no_panic() {
record_token_exchange();
}
#[test]
fn record_chunk_inserted_increments_bytes_counter() {
let before_bytes = metrics().storage.chunks_bytes_total.get();
record_chunk_inserted(128);
let after_bytes = metrics().storage.chunks_bytes_total.get();
assert!(
after_bytes >= before_bytes + 128,
"chunks_bytes_total should increase by at least 128 (before: {before_bytes}, after: {after_bytes})"
);
}
#[test]
fn record_xorb_stored_increments_bytes_counter() {
let before_bytes = metrics().storage.xorbs_bytes_total.get();
record_xorb_stored(256);
let after_bytes = metrics().storage.xorbs_bytes_total.get();
assert!(
after_bytes >= before_bytes + 256,
"xorbs_bytes_total should increase by at least 256 (before: {before_bytes}, after: {after_bytes})"
);
}
#[test]
fn record_dedup_saves_increments_bytes_counter() {
let before_bytes = metrics().storage.dedup_saves_bytes_total.get();
record_dedup_saves(512);
let after_bytes = metrics().storage.dedup_saves_bytes_total.get();
assert!(
after_bytes >= before_bytes + 512,
"dedup_saves_bytes_total should increase by at least 512 (before: {before_bytes}, after: {after_bytes})"
);
}
#[test]
fn record_webhook_event_increments_webhook_counter() {
let before = metrics().provider.webhook_events.get();
record_webhook_event("github", "push", 0.25);
let after = metrics().provider.webhook_events.get();
assert!(
after > before,
"webhook_events should increase (before: {before}, after: {after})"
);
}
#[test]
fn record_range_request_does_not_panic() {
record_range_request();
}
#[test]
fn record_upload_protocol_variants() {
record_upload("xet", 1024, 0.5, true);
record_upload("lfs", 2048, 1.0, false);
}
#[test]
fn record_download_protocol_variants() {
record_download("xet", 4096, 2.0, true);
record_download("oci", 8192, 3.0, false);
}
#[test]
fn record_token_exchange_does_not_panic() {
record_token_exchange();
}
#[test]
fn metrics_layer_creates_metrics_service() {
use axum::routing::get;
use tower::ServiceExt;
async fn handler() -> &'static str {
"ok"
}
let layer = MetricsLayer;
let svc = layer.layer(get(handler));
let response = svc.oneshot(
axum::http::Request::builder()
.uri("/")
.body(axum::body::Body::empty())
.unwrap(),
);
drop(response);
}
}
#[derive(Clone)]
pub(crate) struct MetricsLayer;
impl<S> Layer<S> for MetricsLayer {
type Service = MetricsService<S>;
fn layer(&self, inner: S) -> Self::Service {
MetricsService { inner }
}
}
#[derive(Clone, Debug)]
pub(crate) struct MetricsService<S> {
inner: S,
}
impl<S, ReqBody> Service<axum::http::Request<ReqBody>> for MetricsService<S>
where
S: Service<axum::http::Request<ReqBody>, Response = axum::http::Response<Body>>
+ Clone
+ Send
+ 'static,
S::Future: Send + 'static,
ReqBody: Send + 'static,
{
type Response = axum::http::Response<Body>;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: axum::http::Request<ReqBody>) -> Self::Future {
let start = Instant::now();
shardline_metrics::metrics().system.connection_opened();
let mut inner = self.inner.clone();
Box::pin(async move {
let result = inner.call(req).await;
shardline_metrics::metrics().system.connection_closed();
let response = result?;
let _elapsed = start.elapsed().as_secs_f64();
Ok(response)
})
}
}