Skip to main content

shardline_server/
metrics.rs

1use std::{
2    future::Future,
3    pin::Pin,
4    task::{Context, Poll},
5    time::{Duration, Instant},
6};
7
8use axum::body::Body;
9use tower::{Layer, Service};
10
11// Re-export all types and convenience functions from shardline-metrics.
12pub use shardline_metrics::*;
13
14// ── Compatibility free functions ──────────────────────────────────────────
15//
16// The server originally exposed these signatures.  Wrappers convert to the
17// `shardline_metrics` API which uses `Duration` instead of raw `f64` and
18// `u16` status codes instead of string labels.
19
20pub fn record_upload(protocol: &str, bytes: u64, duration_secs: f64, ok: bool) {
21    let status = if ok { 200_u16 } else { 500 };
22    shardline_metrics::record_upload(protocol, bytes);
23    shardline_metrics::metrics()
24        .transfer
25        .record_upload_duration(duration_secs);
26    let _ = status;
27}
28
29pub fn record_download(protocol: &str, bytes: u64, duration_secs: f64, ok: bool) {
30    let status = if ok { 200_u16 } else { 500 };
31    shardline_metrics::record_download(protocol, bytes);
32    shardline_metrics::metrics()
33        .transfer
34        .record_download_duration(duration_secs);
35    let _ = status;
36}
37
38pub fn record_range_request() {
39    shardline_metrics::metrics().transfer.record_range_request();
40}
41
42pub fn record_webhook_event(provider: &str, event_type: &str, duration_secs: f64) {
43    shardline_metrics::record_provider_webhook(provider, event_type);
44    shardline_metrics::metrics()
45        .provider
46        .record_webhook_duration(Duration::from_secs_f64(duration_secs));
47}
48
49pub fn record_token_exchange() {
50    shardline_metrics::record_provider_token_exchange();
51}
52
53pub fn record_chunk_inserted(bytes: u64) {
54    shardline_metrics::metrics()
55        .storage
56        .record_chunk_stored(bytes);
57}
58
59pub fn record_xorb_stored(bytes: u64) {
60    shardline_metrics::metrics()
61        .storage
62        .record_xorb_stored(bytes);
63}
64
65pub fn record_shard_stored() {
66    shardline_metrics::metrics().storage.record_shard_stored();
67}
68
69pub fn record_lfs_upload() {
70    shardline_metrics::metrics().protocol.record_lfs_upload();
71}
72
73pub fn record_lfs_download() {
74    shardline_metrics::metrics().protocol.record_lfs_download();
75}
76
77pub fn record_xet_xorb_download(bytes: u64) {
78    shardline_metrics::record_xet_xorb_download(bytes);
79}
80
81pub fn record_dedup_saves(bytes: u64) {
82    shardline_metrics::metrics()
83        .storage
84        .record_dedup_saves(bytes);
85}
86
87#[cfg(test)]
88mod tests {
89    use super::*;
90    use shardline_metrics::metrics;
91
92    // ── Smoke tests: verify no panic ─────────────────────────────────────
93
94    #[test]
95    fn record_upload_no_panic() {
96        record_upload("http", 1024, 1.5, true);
97        record_upload("grpc", 0, 0.0, false);
98    }
99
100    #[test]
101    fn record_download_no_panic() {
102        record_download("http", 512, 0.5, true);
103        record_download("grpc", 0, 0.0, false);
104    }
105
106    #[test]
107    fn record_range_request_no_panic() {
108        record_range_request();
109        record_range_request();
110    }
111
112    #[test]
113    fn record_webhook_event_no_panic() {
114        record_webhook_event("github", "push", 0.25);
115        record_webhook_event("gitlab", "merge_request", 0.0);
116    }
117
118    // ── Counter increment tests ──────────────────────────────────────────
119
120    #[test]
121    fn record_upload_increments_upload_counter() {
122        let before = metrics().transfer.upload_requests.get();
123        record_upload("http", 42, 0.1, true);
124        let after = metrics().transfer.upload_requests.get();
125        assert!(
126            after > before,
127            "upload_requests should increase (before: {before}, after: {after})"
128        );
129    }
130
131    #[test]
132    fn record_download_increments_download_counter() {
133        let before = metrics().transfer.download_requests.get();
134        record_download("http", 99, 0.2, true);
135        let after = metrics().transfer.download_requests.get();
136        assert!(
137            after > before,
138            "download_requests should increase (before: {before}, after: {after})"
139        );
140    }
141
142    #[test]
143    fn record_range_request_increments_range_counter() {
144        let before = metrics().transfer.range_requests.get();
145        record_range_request();
146        let after = metrics().transfer.range_requests.get();
147        assert!(
148            after > before,
149            "range_requests should increase (before: {before}, after: {after})"
150        );
151    }
152
153    #[test]
154    fn record_token_exchange_increments_counter() {
155        let before = metrics().provider.token_exchanges.get();
156        record_token_exchange();
157        let after = metrics().provider.token_exchanges.get();
158        assert!(
159            after > before,
160            "token_exchanges should increase (before: {before}, after: {after})"
161        );
162    }
163
164    #[test]
165    fn record_chunk_inserted_increments_chunk_counter() {
166        let before = metrics().storage.chunks_total.get();
167        record_chunk_inserted(64);
168        let after = metrics().storage.chunks_total.get();
169        assert!(
170            after > before,
171            "chunks_total should increase (before: {before}, after: {after})"
172        );
173    }
174
175    #[test]
176    fn record_xorb_stored_increments_xorb_counter() {
177        let before = metrics().storage.xorbs_total.get();
178        record_xorb_stored(128);
179        let after = metrics().storage.xorbs_total.get();
180        assert!(
181            after > before,
182            "xorbs_total should increase (before: {before}, after: {after})"
183        );
184    }
185
186    #[test]
187    fn record_shard_stored_increments_shard_counter() {
188        let before = metrics().storage.shards_total.get();
189        record_shard_stored();
190        let after = metrics().storage.shards_total.get();
191        assert!(
192            after > before,
193            "shards_total should increase (before: {before}, after: {after})"
194        );
195    }
196
197    #[test]
198    fn record_dedup_saves_increments_dedup_counter() {
199        let before = metrics().storage.dedup_saves_bytes_total.get();
200        record_dedup_saves(1024);
201        let after = metrics().storage.dedup_saves_bytes_total.get();
202        assert!(
203            after > before,
204            "dedup_saves_bytes_total should increase (before: {before}, after: {after})"
205        );
206    }
207
208    // ── MetricsLayer & MetricsService tests ──────────────────────────────
209
210    #[test]
211    fn metrics_layer_is_cloneable() {
212        let layer = MetricsLayer;
213        let _clone = layer;
214    }
215
216    #[test]
217    fn metrics_service_construction() {
218        // MetricsService wraps an inner service; we use a simple tokio
219        // runtime to test the basic construction. The inner type doesn't
220        // need to be a real HTTP service for this test.
221        let inner = tower::util::service_fn(|_req: axum::http::Request<Body>| async {
222            Ok::<_, std::convert::Infallible>(axum::http::Response::new(Body::empty()))
223        });
224        let _svc = MetricsService { inner };
225    }
226
227    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
228    async fn metrics_service_poll_ready_and_call_tracks_connections() {
229        use tower::ServiceExt;
230        let svc = MetricsService {
231            inner: tower::util::service_fn(|_req: axum::http::Request<Body>| async {
232                Ok::<_, std::convert::Infallible>(axum::http::Response::new(Body::empty()))
233            }),
234        };
235        let before = metrics().system.active_connections.get();
236        // Use oneshot to drive poll_ready + call
237        let _response = svc
238            .oneshot(axum::http::Request::builder().body(Body::empty()).unwrap())
239            .await;
240        let after = metrics().system.active_connections.get();
241        // Connections opened and then closed, so active should be same
242        assert_eq!(after, before);
243    }
244
245    // ── No-panic for the remaining counters ───────────────────────────────
246
247    #[test]
248    fn record_token_exchange_no_panic() {
249        record_token_exchange();
250    }
251
252    #[test]
253    fn record_chunk_inserted_increments_bytes_counter() {
254        let before_bytes = metrics().storage.chunks_bytes_total.get();
255        record_chunk_inserted(128);
256        let after_bytes = metrics().storage.chunks_bytes_total.get();
257        assert!(
258            after_bytes >= before_bytes + 128,
259            "chunks_bytes_total should increase by at least 128 (before: {before_bytes}, after: {after_bytes})"
260        );
261    }
262
263    #[test]
264    fn record_xorb_stored_increments_bytes_counter() {
265        let before_bytes = metrics().storage.xorbs_bytes_total.get();
266        record_xorb_stored(256);
267        let after_bytes = metrics().storage.xorbs_bytes_total.get();
268        assert!(
269            after_bytes >= before_bytes + 256,
270            "xorbs_bytes_total should increase by at least 256 (before: {before_bytes}, after: {after_bytes})"
271        );
272    }
273
274    #[test]
275    fn record_dedup_saves_increments_bytes_counter() {
276        let before_bytes = metrics().storage.dedup_saves_bytes_total.get();
277        record_dedup_saves(512);
278        let after_bytes = metrics().storage.dedup_saves_bytes_total.get();
279        assert!(
280            after_bytes >= before_bytes + 512,
281            "dedup_saves_bytes_total should increase by at least 512 (before: {before_bytes}, after: {after_bytes})"
282        );
283    }
284
285    #[test]
286    fn record_webhook_event_increments_webhook_counter() {
287        let before = metrics().provider.webhook_events.get();
288        record_webhook_event("github", "push", 0.25);
289        let after = metrics().provider.webhook_events.get();
290        assert!(
291            after > before,
292            "webhook_events should increase (before: {before}, after: {after})"
293        );
294    }
295
296    // ── Remaining record_* functions ──────────────────────────────────────
297
298    #[test]
299    fn record_range_request_does_not_panic() {
300        record_range_request();
301    }
302
303    #[test]
304    fn record_upload_protocol_variants() {
305        record_upload("xet", 1024, 0.5, true);
306        record_upload("lfs", 2048, 1.0, false);
307    }
308
309    #[test]
310    fn record_download_protocol_variants() {
311        record_download("xet", 4096, 2.0, true);
312        record_download("oci", 8192, 3.0, false);
313    }
314
315    #[test]
316    fn record_token_exchange_does_not_panic() {
317        record_token_exchange();
318    }
319
320    // ── MetricsLayer ──────────────────────────────────────────────────────
321
322    #[test]
323    fn metrics_layer_creates_metrics_service() {
324        use axum::routing::get;
325        use tower::ServiceExt;
326        async fn handler() -> &'static str {
327            "ok"
328        }
329        let layer = MetricsLayer;
330        let svc = layer.layer(get(handler));
331        // Verify the service wraps the inner by calling it
332        let response = svc.oneshot(
333            axum::http::Request::builder()
334                .uri("/")
335                .body(axum::body::Body::empty())
336                .unwrap(),
337        );
338        drop(response);
339    }
340}
341
342// ── Axum middleware & routes ─────────────────────────────────────────────
343
344#[derive(Clone)]
345pub(crate) struct MetricsLayer;
346
347impl<S> Layer<S> for MetricsLayer {
348    type Service = MetricsService<S>;
349
350    fn layer(&self, inner: S) -> Self::Service {
351        MetricsService { inner }
352    }
353}
354
355#[derive(Clone, Debug)]
356pub(crate) struct MetricsService<S> {
357    inner: S,
358}
359
360impl<S, ReqBody> Service<axum::http::Request<ReqBody>> for MetricsService<S>
361where
362    S: Service<axum::http::Request<ReqBody>, Response = axum::http::Response<Body>>
363        + Clone
364        + Send
365        + 'static,
366    S::Future: Send + 'static,
367    ReqBody: Send + 'static,
368{
369    type Response = axum::http::Response<Body>;
370    type Error = S::Error;
371    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
372
373    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
374        self.inner.poll_ready(cx)
375    }
376
377    fn call(&mut self, req: axum::http::Request<ReqBody>) -> Self::Future {
378        let start = Instant::now();
379        shardline_metrics::metrics().system.connection_opened();
380
381        let mut inner = self.inner.clone();
382        Box::pin(async move {
383            let result = inner.call(req).await;
384            shardline_metrics::metrics().system.connection_closed();
385            let response = result?;
386            let _elapsed = start.elapsed().as_secs_f64();
387
388            Ok(response)
389        })
390    }
391}