Skip to main content

objectstore_server/web/
app.rs

1use std::net::SocketAddr;
2
3use anyhow::Result;
4use axum::ServiceExt;
5use axum::extract::Request;
6use objectstore_log::{Level, tracing};
7use sentry::integrations::tower::{NewSentryLayer, SentryHttpLayer};
8use tokio::net::TcpListener;
9use tower::ServiceBuilder;
10use tower_http::catch_panic::CatchPanicLayer;
11use tower_http::trace::{DefaultOnFailure, TraceLayer};
12
13use crate::endpoints;
14use crate::state::ServiceState;
15use crate::web::middleware as m;
16
17/// The objectstore web server application.
18#[derive(Debug)]
19pub struct App {
20    router: axum::Router,
21    graceful_shutdown: bool,
22}
23
24impl App {
25    /// Creates a new application router for the given service state.
26    ///
27    /// The applications sets up middlewares and routes for the objectstore web API. Use
28    /// [`serve`](Self::serve) to run the server future.
29    pub fn new(state: ServiceState) -> Self {
30        // Build the router middleware into a single service which runs _after_ routing. Service
31        // builder order defines layers added first will be called first. This means:
32        //  - Requests go from top to bottom
33        //  - Responses go from bottom to top
34        let middleware = ServiceBuilder::new()
35            .layer(axum::middleware::from_fn(m::capture_request_time))
36            .layer(NewSentryLayer::new_from_top())
37            .layer(SentryHttpLayer::new().enable_transaction())
38            .layer(axum::middleware::from_fn(m::emit_request_metrics))
39            .layer(axum::middleware::from_fn(m::bind_sentry_body))
40            .layer(axum::middleware::from_fn_with_state(
41                state.request_counter.clone(),
42                m::limit_web_concurrency,
43            ))
44            .layer(state.request_counter.layer())
45            .layer(CatchPanicLayer::custom(m::handle_panic))
46            .layer(m::set_server_header())
47            .layer(
48                TraceLayer::new_for_http()
49                    .make_span_with(tracing::Span::none())
50                    .on_failure(DefaultOnFailure::new().level(Level::DEBUG)),
51            );
52
53        let router = endpoints::routes()
54            .layer(middleware)
55            .with_state(state.clone());
56
57        App {
58            router,
59            graceful_shutdown: false,
60        }
61    }
62
63    /// Enables or disables graceful shutdown for the server.
64    ///
65    /// By default, graceful shutdown is disabled.
66    pub fn graceful_shutdown(mut self, enable: bool) -> Self {
67        self.graceful_shutdown = enable;
68        self
69    }
70
71    /// Runs the web server until graceful shutdown is triggered.
72    ///
73    /// This function creates a future that runs the server. The future must be spawned or awaited for
74    /// the server to continue running.
75    pub async fn serve(self, listener: TcpListener) -> Result<()> {
76        let Self {
77            router,
78            graceful_shutdown,
79        } = self;
80
81        let service =
82            ServiceExt::<Request>::into_make_service_with_connect_info::<SocketAddr>(router);
83
84        if graceful_shutdown {
85            let guard = elegant_departure::get_shutdown_guard();
86            axum::serve(listener, service)
87                .with_graceful_shutdown(guard.wait_owned())
88                .await?;
89        } else {
90            axum::serve(listener, service).await?;
91        }
92
93        Ok(())
94    }
95}
96
97#[cfg(test)]
98mod tests {
99    use axum::body::{Body, to_bytes};
100    use axum::http::StatusCode;
101    use objectstore_log::tracing;
102    use objectstore_service::backend::local_fs::FileSystemConfig;
103    use sentry::protocol::{Context, EnvelopeItem, TraceContext, Transaction};
104    use tower::ServiceExt;
105    use tracing_subscriber::prelude::*;
106
107    use super::*;
108    use crate::config::{AuthZ, Config, StorageConfig};
109    use crate::state::Services;
110
111    fn trace<'a>(transaction: &'a Transaction<'_>) -> &'a TraceContext {
112        let Some(Context::Trace(trace)) = transaction.contexts.get("trace") else {
113            panic!("transaction is missing its trace context");
114        };
115        trace
116    }
117
118    #[test]
119    fn body_task_preserves_sentry_context() {
120        let runtime = tokio::runtime::Builder::new_current_thread()
121            .enable_all()
122            .build()
123            .unwrap();
124        let _subscriber = tracing_subscriber::registry()
125            .with(
126                sentry::integrations::tracing::layer()
127                    .span_filter(|metadata| *metadata.level() != tracing::Level::TRACE),
128            )
129            .set_default();
130        let trace_id = "11111111111111111111111111111111";
131        let caller_span = "aaaaaaaaaaaaaaaa";
132        let envelopes = sentry::test::with_captured_envelopes_options(
133            || {
134                runtime.block_on(async {
135                    let directory = tempfile::tempdir().unwrap();
136                    let state = Services::spawn(Config {
137                        storage: StorageConfig::FileSystem(FileSystemConfig {
138                            path: directory.path().into(),
139                            cogs: None,
140                        }),
141                        auth: AuthZ {
142                            enforce: false,
143                            ..Default::default()
144                        },
145                        ..Default::default()
146                    })
147                    .await
148                    .unwrap();
149                    let request = Request::builder()
150                        .method("POST")
151                        .uri("/v1/objects:batch/test/org=1/")
152                        .header("sentry-trace", format!("{trace_id}-{caller_span}-1"))
153                        .header("content-type", "multipart/form-data; boundary=boundary")
154                        .body(Body::from(concat!(
155                            "--boundary\r\n",
156                            "x-sn-batch-operation-key: missing\r\n",
157                            "x-sn-batch-operation-kind: head\r\n",
158                            "\r\n\r\n--boundary--\r\n",
159                        )))
160                        .unwrap();
161                    let response = App::new(state).router.oneshot(request).await.unwrap();
162                    assert_eq!(response.status(), StatusCode::OK);
163                    // Polling the batch body starts the task after the HTTP transaction ends.
164                    let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
165                    assert!(String::from_utf8_lossy(&body).contains("404 Not Found"));
166                })
167            },
168            sentry::ClientOptions {
169                traces_sample_rate: 1.0,
170                ..Default::default()
171            },
172        );
173        let transactions: Vec<_> = envelopes
174            .iter()
175            .flat_map(|envelope| envelope.items())
176            .filter_map(|item| match item {
177                EnvelopeItem::Transaction(transaction) => Some(transaction),
178                _ => None,
179            })
180            .collect();
181        let [http, task] = transactions.as_slice() else {
182            panic!("expected exactly one HTTP and one task transaction");
183        };
184        assert_eq!(trace(http).op.as_deref(), Some("http.server"));
185        assert_eq!(trace(task).op.as_deref(), Some("tokio.task"));
186        assert_eq!(task.name.as_deref(), Some("head"));
187        assert_eq!(trace(http).parent_span_id.unwrap().to_string(), caller_span);
188        assert_eq!(
189            trace(task).parent_span_id,
190            Some(trace(http).span_id),
191            "task parent must be the emitted HTTP span"
192        );
193        for transaction in transactions {
194            assert_eq!(trace(transaction).trace_id.to_string(), trace_id);
195            assert_eq!(
196                transaction.tags.get("usecase").map(String::as_str),
197                Some("test")
198            );
199            assert_eq!(
200                transaction.tags.get("scope.org").map(String::as_str),
201                Some("1")
202            );
203        }
204    }
205}