Skip to main content

servers/grpc/
context_auth.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 std::sync::Arc;
16
17use api::v1::auth_header::AuthScheme;
18use api::v1::{AuthHeader, RequestHeader};
19use auth::{Identity, Password, UserInfoRef, UserProviderRef};
20use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
21use common_catalog::parse_catalog_and_schema_from_db_string;
22use common_error::ext::ErrorExt;
23use session::context::{Channel, QueryContextBuilder, QueryContextRef};
24use session::hints::INSERT_SKIP_WAL_HINT;
25use snafu::{OptionExt, ResultExt};
26use tonic::Status;
27use tonic::metadata::MetadataMap;
28
29use crate::error::Error::UnsupportedAuthScheme;
30use crate::error::{AuthSnafu, InvalidParameterSnafu, NotFoundAuthHeaderSnafu, Result};
31use crate::grpc::TonicResult;
32use crate::hint_headers;
33use crate::http::AUTHORIZATION_HEADER;
34use crate::http::header::constants::GREPTIME_DB_HEADER_NAME;
35use crate::metrics::METRIC_AUTH_FAILURE;
36
37/// Create a query context from gRPC metadata and server-owned request extensions.
38pub fn create_query_context_from_grpc_metadata(
39    headers: &MetadataMap,
40    extensions: &http::Extensions,
41) -> TonicResult<QueryContextRef> {
42    let (catalog, schema) = if let Some(db) = extract_header(headers, &[GREPTIME_DB_HEADER_NAME])? {
43        parse_catalog_and_schema_from_db_string(db)
44    } else {
45        (
46            DEFAULT_CATALOG_NAME.to_string(),
47            DEFAULT_SCHEMA_NAME.to_string(),
48        )
49    };
50
51    let ctx = QueryContextBuilder::default()
52        .current_catalog(catalog)
53        .current_schema(schema)
54        .channel(
55            extensions
56                .get::<Channel>()
57                .copied()
58                .unwrap_or(Channel::Grpc),
59        )
60        .build();
61    // OTEL Arrow uses ordinary inserts. Accept only its request-level WAL hint,
62    // leaving unrelated hints and reserved internal extensions unchanged.
63    if let Some((key, value)) = hint_headers::extract_hints(headers)
64        .into_iter()
65        .find(|(key, _)| key == INSERT_SKIP_WAL_HINT)
66    {
67        let skip_wal = value.parse::<bool>().map_err(|_| {
68            InvalidParameterSnafu {
69                reason: format!("Invalid {key} hint: expected true or false, got {value:?}"),
70            }
71            .build()
72        })?;
73        ctx.set_skip_wal(skip_wal);
74    }
75    Ok(Arc::new(ctx))
76}
77
78/// Helper function to extract a header from the metadata map.
79/// Can be multiple keys, and the first one found will be returned.
80pub fn extract_header<'a>(headers: &'a MetadataMap, keys: &[&str]) -> TonicResult<Option<&'a str>> {
81    let mut value = None;
82    for key in keys {
83        if let Some(v) = headers.get(*key) {
84            value = Some(v);
85            break;
86        }
87    }
88
89    let Some(v) = value else {
90        return Ok(None);
91    };
92    let Ok(v) = std::str::from_utf8(v.as_bytes()) else {
93        return Err(InvalidParameterSnafu {
94            reason: "expect valid UTF-8 value",
95        }
96        .build()
97        .into());
98    };
99    Ok(Some(v))
100}
101
102/// Helper function to extract the header from the metadata and authenticate the user.
103pub async fn check_auth(
104    user_provider: Option<UserProviderRef>,
105    headers: &MetadataMap,
106    query_ctx: QueryContextRef,
107) -> TonicResult<bool> {
108    if user_provider.is_none() {
109        return Ok(true);
110    }
111
112    let auth_schema = extract_header(
113        headers,
114        &[AUTHORIZATION_HEADER, http::header::AUTHORIZATION.as_str()],
115    )?
116    .map(|x| {
117        if x.len() > 5 && x[0..5].eq_ignore_ascii_case("Basic") {
118            x.try_into()
119        } else {
120            // compatible with old version
121            format!("Basic {}", x).as_str().try_into()
122        }
123    })
124    .transpose()?
125    .map(|x: crate::http::authorize::AuthScheme| x.into());
126
127    let auth_schema = auth_schema.context(NotFoundAuthHeaderSnafu)?;
128    let header = RequestHeader {
129        authorization: Some(AuthHeader {
130            auth_scheme: Some(auth_schema),
131        }),
132        catalog: query_ctx.current_catalog().to_string(),
133        schema: query_ctx.current_schema(),
134        ..Default::default()
135    };
136
137    match auth(user_provider, Some(&header), &query_ctx).await {
138        Ok(user_info) => {
139            query_ctx.set_current_user(user_info);
140            Ok(true)
141        }
142        Err(_) => Err(Status::unauthenticated("auth failed")),
143    }
144}
145
146/// Authenticate the user based on the header and query context.
147pub async fn auth(
148    user_provider: Option<UserProviderRef>,
149    header: Option<&RequestHeader>,
150    query_ctx: &QueryContextRef,
151) -> Result<UserInfoRef> {
152    let Some(user_provider) = user_provider else {
153        return Ok(auth::userinfo_by_name(None));
154    };
155
156    let auth_scheme = header
157        .and_then(|header| {
158            header
159                .authorization
160                .as_ref()
161                .and_then(|x| x.auth_scheme.clone())
162        })
163        .context(NotFoundAuthHeaderSnafu)?;
164
165    match auth_scheme {
166        AuthScheme::Basic(api::v1::Basic { username, password }) => user_provider
167            .auth(
168                Identity::UserId(&username, None),
169                Password::PlainText(password.into()),
170                query_ctx.current_catalog(),
171                &query_ctx.current_schema(),
172            )
173            .await
174            .context(AuthSnafu),
175        AuthScheme::Token(_) => Err(UnsupportedAuthScheme {
176            name: "Token AuthScheme".to_string(),
177        }),
178    }
179    .inspect_err(|e| {
180        METRIC_AUTH_FAILURE
181            .with_label_values(&[e.status_code().as_ref()])
182            .inc();
183    })
184}
185
186#[cfg(test)]
187mod tests {
188    use session::hints::{HINTS_KEY, REMOTE_QUERY_ID_EXTENSION_KEY, RESERVED_EXTENSION_KEYS};
189
190    use super::*;
191
192    #[test]
193    fn test_channel_comes_from_server_extensions() {
194        let mut headers = MetadataMap::new();
195        headers.insert(
196            "x-greptime-flow-extensions",
197            r#"[["flow.return_region_seq","true"]]"#.parse().unwrap(),
198        );
199        headers.insert(HINTS_KEY, "channel=internal".parse().unwrap());
200        let mut extensions = http::Extensions::new();
201        let ctx = create_query_context_from_grpc_metadata(&headers, &extensions).unwrap();
202        assert_eq!(ctx.channel(), Channel::Grpc);
203
204        extensions.insert(Channel::Internal);
205        let ctx = create_query_context_from_grpc_metadata(&headers, &extensions).unwrap();
206        assert_eq!(ctx.channel(), Channel::Internal);
207    }
208
209    #[test]
210    fn test_arrow_insert_hint_does_not_accept_reserved_extensions() {
211        let mut headers = MetadataMap::new();
212        assert_eq!(
213            create_query_context_from_grpc_metadata(&headers, &Default::default())
214                .unwrap()
215                .extension(INSERT_SKIP_WAL_HINT),
216            None
217        );
218        for (value, expected) in [("true", true), ("false", false)] {
219            let mut hints = format!("insert_skip_wal={value},ttl=7d");
220            for key in RESERVED_EXTENSION_KEYS {
221                hints.push_str(&format!(",{key}=external"));
222            }
223            headers.insert(HINTS_KEY, hints.parse().unwrap());
224            let ctx =
225                create_query_context_from_grpc_metadata(&headers, &Default::default()).unwrap();
226            assert_eq!(ctx.skip_wal(), expected);
227            assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None);
228            assert_eq!(ctx.extension("ttl"), None);
229            for key in RESERVED_EXTENSION_KEYS {
230                if key == REMOTE_QUERY_ID_EXTENSION_KEY {
231                    // The builder generates this ID; external hints must not replace it.
232                    assert!(ctx.remote_query_id().is_some());
233                    assert_ne!(ctx.extension(key), Some("external"));
234                } else {
235                    assert_eq!(ctx.extension(key), None);
236                }
237            }
238        }
239        // Only the first matching hint is parsed and applied.
240        for (hints, expected) in [
241            ("insert_skip_wal=true,insert_skip_wal=false", true),
242            ("insert_skip_wal=false,insert_skip_wal=true", false),
243            ("insert_skip_wal=true,insert_skip_wal=invalid", true),
244        ] {
245            headers.insert(HINTS_KEY, hints.parse().unwrap());
246            let ctx =
247                create_query_context_from_grpc_metadata(&headers, &Default::default()).unwrap();
248            assert_eq!(ctx.skip_wal(), expected);
249        }
250        for value in ["", "TRUE", "1", "invalid"] {
251            headers.insert(
252                HINTS_KEY,
253                format!("insert_skip_wal={value}").parse().unwrap(),
254            );
255            assert!(
256                create_query_context_from_grpc_metadata(&headers, &Default::default()).is_err()
257            );
258        }
259    }
260}