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
11pub use shardline_metrics::*;
13
14pub 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 #[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 #[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 #[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 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 let _response = svc
238 .oneshot(axum::http::Request::builder().body(Body::empty()).unwrap())
239 .await;
240 let after = metrics().system.active_connections.get();
241 assert_eq!(after, before);
243 }
244
245 #[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 #[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 #[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 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#[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}