Skip to main content

servers/http/
skip_wal.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use axum::body::Body;
16use axum::http::Request;
17use axum::middleware::Next;
18use axum::response::{IntoResponse, Response};
19use session::context::QueryContext;
20
21use crate::error::InvalidParameterSnafu;
22use crate::http::header::GREPTIME_INSERT_SKIP_WAL_HEADER_NAME;
23use crate::http::result::error_result::ErrorResponse;
24
25/// Extract the request-level WAL policy from the dedicated HTTP header.
26pub async fn extract_skip_wal(mut request: Request<Body>, next: Next) -> Response {
27    let skip_wal = match request.headers().get(&GREPTIME_INSERT_SKIP_WAL_HEADER_NAME) {
28        None => false,
29        Some(value) => match value
30            .to_str()
31            .ok()
32            .and_then(|value| value.parse::<bool>().ok())
33        {
34            Some(skip_wal) => skip_wal,
35            None => {
36                return (
37                    http::StatusCode::BAD_REQUEST,
38                    ErrorResponse::from_error(
39                        InvalidParameterSnafu {
40                            reason: format!(
41                                "{} must be true or false",
42                                GREPTIME_INSERT_SKIP_WAL_HEADER_NAME
43                            ),
44                        }
45                        .build(),
46                    ),
47                )
48                    .into_response();
49            }
50        },
51    };
52    if let Some(query_ctx) = request.extensions_mut().get_mut::<QueryContext>() {
53        query_ctx.set_skip_wal(skip_wal);
54    }
55    next.run(request).await
56}