1use 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
37pub 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 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
78pub 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
102pub 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 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
146pub 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 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 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}