Skip to main content

servers/grpc/
greptime_handler.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
15//! Handler for Greptime Database service. It's implemented by frontend.
16
17use std::collections::HashMap;
18use std::future::Future;
19use std::str::FromStr;
20use std::sync::{Arc, RwLock};
21use std::time::Instant;
22
23use api::helper::request_type;
24use api::v1::greptime_request::Request as QueryRequest;
25use api::v1::{GreptimeRequest, RequestHeader};
26use auth::UserProviderRef;
27use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
28use common_catalog::parse_catalog_and_schema_from_db_string;
29use common_error::ext::ErrorExt;
30use common_error::status_code::StatusCode;
31use common_grpc::flight::do_put::DoPutResponse;
32use common_query::Output;
33use common_runtime::Runtime;
34use common_runtime::runtime::RuntimeTrait;
35use common_session::ReadPreference;
36use common_telemetry::tracing_context::{FutureExt, TracingContext};
37use common_telemetry::{debug, error, tracing, warn};
38use common_time::timezone::parse_timezone;
39use futures_util::StreamExt;
40use session::context::{Channel, QueryContextBuilder, QueryContextRef};
41use session::hints::{INSERT_SKIP_WAL_HINT, READ_PREFERENCE_HINT, is_reserved_extension_key};
42use snafu::{OptionExt, ResultExt};
43use tokio::sync::mpsc;
44use tokio::sync::mpsc::error::TrySendError;
45use tonic::Status;
46
47use crate::error::{InvalidQuerySnafu, JoinTaskSnafu, Result, UnknownHintSnafu};
48use crate::grpc::flight::PutRecordBatchRequestStream;
49use crate::grpc::{FlightCompression, TonicResult, context_auth};
50use crate::metrics::{self, METRIC_SERVER_GRPC_DB_REQUEST_TIMER};
51use crate::query_handler::grpc::ServerGrpcQueryHandlerRef;
52
53#[derive(Clone)]
54pub struct GreptimeRequestHandler {
55    handler: ServerGrpcQueryHandlerRef,
56    pub(crate) user_provider: Option<UserProviderRef>,
57    runtime: Option<Runtime>,
58    pub(crate) flight_compression: FlightCompression,
59}
60
61impl GreptimeRequestHandler {
62    pub fn new(
63        handler: ServerGrpcQueryHandlerRef,
64        user_provider: Option<UserProviderRef>,
65        runtime: Option<Runtime>,
66        flight_compression: FlightCompression,
67    ) -> Self {
68        Self {
69            handler,
70            user_provider,
71            runtime,
72            flight_compression,
73        }
74    }
75
76    #[tracing::instrument(skip_all, fields(protocol = "grpc", request_type = get_request_type(&request)))]
77    pub(crate) async fn handle_request(
78        &self,
79        request: GreptimeRequest,
80        hints: Vec<(String, String)>,
81        channel: Channel,
82    ) -> Result<Output> {
83        let header = request.header.as_ref();
84        let query_ctx = create_query_context(channel, header, hints, HashMap::new())?;
85        let query = request.request.context(InvalidQuerySnafu {
86            reason: "Expecting non-empty GreptimeRequest.",
87        })?;
88        self.authenticate_request_with_query_ctx(request.header.as_ref(), &query_ctx)
89            .await?;
90        self.handle_request_with_query_ctx(query, query_ctx).await
91    }
92
93    pub(crate) async fn authenticate_request_with_query_ctx(
94        &self,
95        header: Option<&RequestHeader>,
96        query_ctx: &QueryContextRef,
97    ) -> Result<()> {
98        let user_info = context_auth::auth(self.user_provider.clone(), header, query_ctx).await?;
99        query_ctx.set_current_user(user_info);
100        Ok(())
101    }
102
103    pub(crate) fn handle_request_with_query_ctx(
104        &self,
105        query: QueryRequest,
106        query_ctx: QueryContextRef,
107    ) -> impl Future<Output = Result<Output>> + Send + 'static {
108        let handler = self.handler.clone();
109        let runtime = self.runtime.clone();
110        let request_type = request_type(&query).to_string();
111        let db = query_ctx.get_db_string();
112        let timer = RequestTimer::new(db.clone(), request_type);
113        let tracing_context = TracingContext::from_current_span();
114
115        async move {
116            let result_future = async move {
117                handler
118                .do_query(query, query_ctx)
119                .trace(tracing_context.attach(tracing::info_span!(
120                    "GreptimeRequestHandler::handle_request_runtime"
121                )))
122                .await
123                .map_err(|e| {
124                    if e.status_code().should_log_error() {
125                        let root_error = e.root_cause().unwrap_or(&e);
126                        error!(e; "Failed to handle request, error: {}", root_error.to_string());
127                    } else {
128                        // Currently, we still print a debug log.
129                        debug!("Failed to handle request, err: {:?}", e);
130                    }
131                    e
132                })
133            };
134
135            match runtime {
136                Some(runtime) => {
137                    // Executes requests in another runtime to
138                    // 1. prevent the execution from being cancelled unexpected by Tonic runtime;
139                    //   - Refer to our blog for the rational behind it:
140                    //     https://www.greptime.com/blogs/2023-01-12-hidden-control-flow.html
141                    //   - Obtaining a `JoinHandle` to get the panic message (if there's any).
142                    //     From its docs, `JoinHandle` is cancel safe. The task keeps running even it's handle been dropped.
143                    // 2. avoid the handler blocks the gRPC runtime incidentally.
144                    runtime
145                        .spawn(result_future)
146                        .await
147                        .context(JoinTaskSnafu)
148                        .inspect_err(|e| {
149                            timer.record(e.status_code());
150                        })?
151                }
152                None => result_future.await,
153            }
154        }
155    }
156
157    pub(crate) async fn put_record_batches(
158        &self,
159        stream: PutRecordBatchRequestStream,
160        result_sender: mpsc::Sender<TonicResult<DoPutResponse>>,
161        query_ctx: QueryContextRef,
162    ) {
163        let handler = self.handler.clone();
164        let runtime = self
165            .runtime
166            .clone()
167            .unwrap_or_else(common_runtime::global_runtime);
168        runtime.spawn(async move {
169            let mut result_stream = handler.handle_put_record_batch_stream(stream, query_ctx);
170
171            while let Some(result) = result_stream.next().await {
172                match &result {
173                    Ok(response) => {
174                        // Record the elapsed time metric from the response
175                        metrics::GRPC_BULK_INSERT_ELAPSED.observe(response.elapsed_secs());
176                    }
177                    Err(e) => {
178                        error!(e; "Failed to handle flight record batches");
179                    }
180                }
181
182                if let Err(e) = result_sender.try_send(result.map_err(Status::from))
183                    && let TrySendError::Closed(_) = e
184                {
185                    warn!(r#""DoPut" client maybe unreachable, abort handling its message"#);
186                    break;
187                }
188            }
189        });
190    }
191}
192
193pub fn get_request_type(request: &GreptimeRequest) -> &'static str {
194    request
195        .request
196        .as_ref()
197        .map(request_type)
198        .unwrap_or_default()
199}
200
201/// Creates a new `QueryContext` from the provided request header and extensions.
202/// Strongly recommend setting an appropriate channel, as this is very helpful for statistics.
203pub(crate) fn create_query_context(
204    channel: Channel,
205    header: Option<&RequestHeader>,
206    extensions: Vec<(String, String)>,
207    snapshot_seqs: HashMap<u64, u64>,
208) -> Result<QueryContextRef> {
209    let (catalog, schema) = header
210        .map(|header| {
211            // We provide dbname field in newer versions of protos/sdks
212            // parse dbname from header in priority
213            if !header.dbname.is_empty() {
214                parse_catalog_and_schema_from_db_string(&header.dbname)
215            } else {
216                (
217                    if !header.catalog.is_empty() {
218                        header.catalog.to_lowercase()
219                    } else {
220                        DEFAULT_CATALOG_NAME.to_string()
221                    },
222                    if !header.schema.is_empty() {
223                        header.schema.to_lowercase()
224                    } else {
225                        DEFAULT_SCHEMA_NAME.to_string()
226                    },
227                )
228            }
229        })
230        .unwrap_or_else(|| {
231            (
232                DEFAULT_CATALOG_NAME.to_string(),
233                DEFAULT_SCHEMA_NAME.to_string(),
234            )
235        });
236    let timezone = parse_timezone(header.map(|h| h.timezone.as_str()));
237    let mut ctx_builder = QueryContextBuilder::default()
238        .current_catalog(catalog)
239        .current_schema(schema)
240        .timezone(timezone)
241        .channel(channel)
242        .snapshot_seqs(Arc::new(RwLock::new(snapshot_seqs)));
243
244    for (key, value) in extensions {
245        match key.as_str() {
246            READ_PREFERENCE_HINT => {
247                let Ok(read_preference) = ReadPreference::from_str(&value) else {
248                    return UnknownHintSnafu {
249                        hint: format!("{key}={value}"),
250                    }
251                    .fail();
252                };
253                ctx_builder = ctx_builder.read_preference(read_preference);
254            }
255            INSERT_SKIP_WAL_HINT => {
256                let skip_wal = value.parse::<bool>().map_err(|_| {
257                    UnknownHintSnafu {
258                        hint: format!("{key}={value}"),
259                    }
260                    .build()
261                })?;
262                ctx_builder = ctx_builder.skip_wal(skip_wal);
263            }
264            _ if is_reserved_extension_key(&key) => {
265                debug!(
266                    key = key.as_str(),
267                    "Ignoring reserved external query context extension key"
268                );
269            }
270            _ => {
271                ctx_builder = ctx_builder.set_extension(key, value);
272            }
273        }
274    }
275    Ok(ctx_builder.build().into())
276}
277
278/// Histogram timer for handling gRPC request.
279///
280/// The timer records the elapsed time with [StatusCode::Success] on drop.
281pub(crate) struct RequestTimer {
282    start: Instant,
283    db: String,
284    request_type: String,
285    status_code: StatusCode,
286}
287
288impl RequestTimer {
289    /// Returns a new timer.
290    pub fn new(db: String, request_type: String) -> RequestTimer {
291        RequestTimer {
292            start: Instant::now(),
293            db,
294            request_type,
295            status_code: StatusCode::Success,
296        }
297    }
298
299    /// Consumes the timer and record the elapsed time with specific `status_code`.
300    pub fn record(mut self, status_code: StatusCode) {
301        self.status_code = status_code;
302    }
303}
304
305impl Drop for RequestTimer {
306    fn drop(&mut self) {
307        METRIC_SERVER_GRPC_DB_REQUEST_TIMER
308            .with_label_values(&[
309                self.db.as_str(),
310                self.request_type.as_str(),
311                self.status_code.as_ref(),
312            ])
313            .observe(self.start.elapsed().as_secs_f64());
314    }
315}
316
317#[cfg(test)]
318mod tests {
319    use chrono::FixedOffset;
320    use common_error::ext::BoxedError;
321    use common_error::{GREPTIME_DB_HEADER_ERROR_CODE, GREPTIME_DB_HEADER_ERROR_RETRY_HINT};
322    use common_time::Timezone;
323    use query::options::FLOW_SCHEDULED_TIME_MILLIS;
324    use session::hints::{
325        INITIAL_REMOTE_DYN_FILTER_REGISTRATIONS_EXTENSION_KEY, REMOTE_QUERY_ID_EXTENSION_KEY,
326    };
327    use snafu::ResultExt;
328    use tonic::Code;
329
330    use super::*;
331    use crate::error::{ExecuteGrpcRequestSnafu, InvalidParameterSnafu};
332
333    #[test]
334    fn test_create_query_context_typed_skip_wal() {
335        let ctx = create_query_context(Channel::Grpc, None, vec![], HashMap::new()).unwrap();
336        assert!(!ctx.skip_wal());
337        let legacy = create_query_context(
338            Channel::Grpc,
339            None,
340            vec![("skip_wal".to_string(), "true".to_string())],
341            HashMap::new(),
342        )
343        .unwrap();
344        assert!(!legacy.skip_wal());
345        assert_eq!(legacy.extension("skip_wal"), Some("true"));
346        for (value, expected) in [("true", true), ("false", false)] {
347            let ctx = create_query_context(
348                Channel::Grpc,
349                None,
350                vec![(INSERT_SKIP_WAL_HINT.to_string(), value.to_string())],
351                HashMap::new(),
352            )
353            .unwrap();
354            assert_eq!(ctx.skip_wal(), expected);
355            assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None);
356        }
357        for value in ["", "TRUE", "1", "invalid"] {
358            assert!(
359                create_query_context(
360                    Channel::Grpc,
361                    None,
362                    vec![(INSERT_SKIP_WAL_HINT.to_string(), value.to_string())],
363                    HashMap::new()
364                )
365                .is_err()
366            );
367        }
368        let ctx = create_query_context(
369            Channel::Grpc,
370            None,
371            vec![
372                (INSERT_SKIP_WAL_HINT.to_string(), "true".to_string()),
373                (INSERT_SKIP_WAL_HINT.to_string(), "false".to_string()),
374            ],
375            HashMap::new(),
376        )
377        .unwrap();
378        assert!(!ctx.skip_wal());
379        assert_eq!(ctx.extension(INSERT_SKIP_WAL_HINT), None);
380    }
381
382    #[test]
383    fn test_create_query_context_read_preference_duplicates() {
384        for (values, valid) in [
385            (["leader", "LEADER"], true),
386            (["invalid", "leader"], false),
387            (["leader", "invalid"], false),
388        ] {
389            let result = create_query_context(
390                Channel::Grpc,
391                None,
392                values
393                    .into_iter()
394                    .map(|value| (READ_PREFERENCE_HINT.to_string(), value.to_string()))
395                    .collect(),
396                HashMap::new(),
397            );
398            if valid {
399                let ctx = result.unwrap();
400                assert!(matches!(ctx.read_preference(), ReadPreference::Leader));
401                assert_eq!(ctx.extension(READ_PREFERENCE_HINT), None);
402            } else {
403                assert!(result.is_err());
404            }
405        }
406    }
407
408    #[test]
409    fn test_create_query_context() {
410        let header = RequestHeader {
411            catalog: "cat-a-log".to_string(),
412            timezone: "+01:00".to_string(),
413            ..Default::default()
414        };
415        let query_context = create_query_context(
416            Channel::Unknown,
417            Some(&header),
418            vec![
419                ("auto_create_table".to_string(), "true".to_string()),
420                ("read_preference".to_string(), "leader".to_string()),
421                (
422                    REMOTE_QUERY_ID_EXTENSION_KEY.to_string(),
423                    "spoofed".to_string(),
424                ),
425                (
426                    INITIAL_REMOTE_DYN_FILTER_REGISTRATIONS_EXTENSION_KEY.to_string(),
427                    "spoofed-regs".to_string(),
428                ),
429                (
430                    FLOW_SCHEDULED_TIME_MILLIS.to_string(),
431                    "1700000000000".to_string(),
432                ),
433            ],
434            HashMap::from([(7, 88)]),
435        )
436        .unwrap();
437        assert_eq!(query_context.get_snapshot(7), Some(88));
438        assert_eq!(query_context.current_catalog(), "cat-a-log");
439        assert_eq!(query_context.current_schema(), DEFAULT_SCHEMA_NAME);
440        assert_eq!(
441            query_context.timezone(),
442            Timezone::Offset(FixedOffset::east_opt(3600).unwrap())
443        );
444        assert!(matches!(
445            query_context.read_preference(),
446            ReadPreference::Leader
447        ));
448        assert_eq!(query_context.extension("auto_create_table"), Some("true"));
449        assert_ne!(query_context.remote_query_id(), Some("spoofed"));
450        assert!(
451            query_context
452                .extension(INITIAL_REMOTE_DYN_FILTER_REGISTRATIONS_EXTENSION_KEY)
453                .is_none()
454        );
455        assert_eq!(
456            query_context.extension(FLOW_SCHEDULED_TIME_MILLIS),
457            Some("1700000000000")
458        );
459    }
460
461    #[test]
462    fn test_create_query_context_ignores_remote_query_id_extension() {
463        let query_context = create_query_context(
464            Channel::Grpc,
465            None,
466            vec![(
467                REMOTE_QUERY_ID_EXTENSION_KEY.to_string(),
468                "spoofed-query-id".to_string(),
469            )],
470            HashMap::new(),
471        )
472        .unwrap();
473
474        assert_ne!(query_context.remote_query_id(), Some("spoofed-query-id"));
475        assert_eq!(
476            query_context.extension(REMOTE_QUERY_ID_EXTENSION_KEY),
477            query_context.remote_query_id()
478        );
479    }
480
481    #[test]
482    fn test_record_batch_error_to_status_preserves_error_details() {
483        let inner = InvalidParameterSnafu {
484            reason: "Column not found, column: new_col",
485        }
486        .build();
487        let err = Err::<(), _>(BoxedError::new(inner))
488            .context(ExecuteGrpcRequestSnafu)
489            .unwrap_err();
490
491        let status = Status::from(err);
492
493        assert_eq!(status.code(), Code::InvalidArgument);
494        assert!(
495            status
496                .message()
497                .contains("Column not found, column: new_col")
498        );
499        assert!(
500            status
501                .message()
502                .contains("Invalid request parameter: Column not found")
503        );
504        assert!(
505            status
506                .metadata()
507                .contains_key(GREPTIME_DB_HEADER_ERROR_CODE)
508        );
509        assert!(
510            status
511                .metadata()
512                .contains_key(GREPTIME_DB_HEADER_ERROR_RETRY_HINT)
513        );
514    }
515}