1use 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 debug!("Failed to handle request, err: {:?}", e);
130 }
131 e
132 })
133 };
134
135 match runtime {
136 Some(runtime) => {
137 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 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
201pub(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 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
278pub(crate) struct RequestTimer {
282 start: Instant,
283 db: String,
284 request_type: String,
285 status_code: StatusCode,
286}
287
288impl RequestTimer {
289 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 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}