1use std::collections::HashMap;
16use std::future;
17use std::panic::AssertUnwindSafe;
18use std::sync::Arc;
19use std::sync::atomic::{AtomicU64, Ordering};
20use std::time::{Duration, Instant};
21
22use axum::extract::rejection::FormRejection;
23use axum::extract::{Json, Query, State};
24use axum::response::sse::{Event, KeepAlive, Sse};
25use axum::response::{IntoResponse, Response};
26use axum::{Extension, Form};
27use common_catalog::parse_catalog_and_schema_from_db_string;
28use common_error::ext::ErrorExt;
29use common_error::status_code::StatusCode;
30use common_plugins::GREPTIME_EXEC_WRITE_COST;
31use common_query::{Output, OutputData};
32use common_recordbatch::error::Result as RecordBatchResult;
33use common_recordbatch::{RecordBatch, RecordBatchStreamWrapper, RecordBatches, util};
34use common_telemetry::tracing;
35use datafusion::physical_plan::ExecutionPlan;
36use futures::{FutureExt, StreamExt};
37use query::parser::{DEFAULT_LOOKBACK_STRING, PromQuery};
38use serde::{Deserialize, Serialize};
39use serde_json::Value;
40use session::context::{Channel, QueryContext, QueryContextRef};
41use snafu::ResultExt;
42use sql::dialect::GreptimeDbDialect;
43use sql::parser::{ParseOptions, ParserContext};
44use sql::statements::statement::Statement;
45use tokio::sync::{Notify, watch};
46
47use crate::error::{CollectRecordbatchSnafu, FailedToParseQuerySnafu, InvalidQuerySnafu, Result};
48use crate::http::header::collect_plan_metrics;
49use crate::http::result::arrow_result::ArrowResponse;
50use crate::http::result::csv_result::CsvResponse;
51use crate::http::result::error_result::ErrorResponse;
52use crate::http::result::greptime_result_v1::GreptimedbV1Response;
53use crate::http::result::influxdb_result_v1::InfluxdbV1Response;
54use crate::http::result::json_result::JsonResponse;
55use crate::http::result::null_result::NullResponse;
56use crate::http::result::table_result::TableResponse;
57use crate::http::{
58 ApiState, Epoch, GreptimeOptionsConfigState, GreptimeQueryOutput, HttpRecordsOutput,
59 HttpResponse, ResponseFormat,
60};
61use crate::metrics_handler::MetricsHandler;
62use crate::query_handler::sql::ServerSqlQueryHandlerRef;
63
64#[derive(Debug, Default, Serialize, Deserialize)]
65pub struct SqlQuery {
66 pub db: Option<String>,
67 pub sql: Option<String>,
68 pub format: Option<String>,
72 pub epoch: Option<String>,
81 pub limit: Option<usize>,
82 pub compression: Option<String>,
84 pub snapshot_interval_ms: Option<u64>,
85}
86
87const DEFAULT_ANALYZE_SNAPSHOT_INTERVAL_MS: u64 = 5000;
88const MIN_ANALYZE_SNAPSHOT_INTERVAL_MS: u64 = 1000;
89const MAX_ANALYZE_SNAPSHOT_INTERVAL_MS: u64 = 60000;
90
91#[derive(Serialize)]
92struct AnalyzeStreamPayload {
93 seq: u64,
94 state: &'static str,
95 partial: bool,
96 elapsed_ms: u64,
97 #[serde(skip_serializing_if = "Option::is_none")]
98 metrics: Option<Value>,
99 #[serde(skip_serializing_if = "Option::is_none")]
100 output: Option<GreptimeQueryOutput>,
101 #[serde(skip_serializing_if = "Option::is_none")]
102 reason: Option<String>,
103 #[serde(skip_serializing_if = "Option::is_none")]
104 code: Option<u32>,
105}
106
107#[derive(Clone, Debug)]
108#[doc(hidden)]
109pub struct AnalyzeStreamMessage {
110 pub event_name: &'static str,
111 pub payload: String,
112}
113
114struct AnalyzeStreamWorkerGuard {
115 cancel: watch::Sender<bool>,
116 handle: common_runtime::JoinHandle<()>,
117}
118
119impl Drop for AnalyzeStreamWorkerGuard {
120 fn drop(&mut self) {
121 let _ = self.cancel.send(true);
122 self.handle.abort();
123 }
124}
125
126struct AnalyzeStreamBodyState {
127 latest_metrics: watch::Receiver<Option<String>>,
128 terminal: watch::Receiver<Option<AnalyzeStreamMessage>>,
129 notify: Arc<Notify>,
130 _worker: AnalyzeStreamWorkerGuard,
133 done: bool,
134}
135
136#[axum_macros::debug_handler]
138#[tracing::instrument(skip_all, fields(protocol = "http", request_type = "sql"))]
139pub async fn sql(
140 State(state): State<ApiState>,
141 Query(query_params): Query<SqlQuery>,
142 Extension(mut query_ctx): Extension<QueryContext>,
143 Form(form_params): Form<SqlQuery>,
144) -> HttpResponse {
145 let start = Instant::now();
146 let sql_handler = &state.sql_handler;
147 if let Some(db) = &query_params.db.or(form_params.db) {
148 let (catalog, schema) = parse_catalog_and_schema_from_db_string(db);
149 query_ctx.set_current_catalog(&catalog);
150 query_ctx.set_current_schema(&schema);
151 }
152 let db = query_ctx.get_db_string();
153
154 query_ctx.set_channel(Channel::HttpSql);
155 let query_ctx = Arc::new(query_ctx);
156
157 let _timer = crate::metrics::METRIC_HTTP_SQL_ELAPSED
158 .with_label_values(&[db.as_str()])
159 .start_timer();
160
161 let sql = query_params.sql.or(form_params.sql);
162 let format = query_params
163 .format
164 .or(form_params.format)
165 .map(|s| s.to_lowercase())
166 .map(|s| ResponseFormat::parse(s.as_str()).unwrap_or(ResponseFormat::GreptimedbV1))
167 .unwrap_or(ResponseFormat::GreptimedbV1);
168 let epoch = query_params
169 .epoch
170 .or(form_params.epoch)
171 .map(|s| s.to_lowercase())
172 .map(|s| Epoch::parse(s.as_str()).unwrap_or(Epoch::Millisecond));
173
174 let result = if let Some(sql) = &sql {
175 if let Some((status, msg)) = validate_schema(sql_handler.clone(), query_ctx.clone()).await {
176 Err((status, msg))
177 } else {
178 Ok(sql_handler.do_query(sql, query_ctx.clone()).await)
179 }
180 } else {
181 Err((
182 StatusCode::InvalidArguments,
183 "sql parameter is required.".to_string(),
184 ))
185 };
186
187 let outputs = match result {
188 Err((status, msg)) => {
189 return HttpResponse::Error(
190 ErrorResponse::from_error_message(status, msg)
191 .with_execution_time(start.elapsed().as_millis() as u64),
192 );
193 }
194 Ok(outputs) => outputs,
195 };
196
197 let outputs = match query_params.limit {
198 Some(limit)
199 if matches!(
200 format,
201 ResponseFormat::Csv(..)
202 | ResponseFormat::Table
203 | ResponseFormat::GreptimedbV1
204 | ResponseFormat::Json
205 ) =>
206 {
207 outputs
208 .into_iter()
209 .map(|output| output.and_then(|output| limit_output_rows(output, limit)))
210 .collect()
211 }
212 _ => outputs,
213 };
214
215 let resp = match format {
216 ResponseFormat::Arrow => {
217 ArrowResponse::from_output(outputs, query_params.compression).await
218 }
219 ResponseFormat::Csv(with_names, with_types) => {
220 CsvResponse::from_output(outputs, with_names, with_types).await
221 }
222 ResponseFormat::Table => TableResponse::from_output(outputs).await,
223 ResponseFormat::GreptimedbV1 => GreptimedbV1Response::from_output(outputs).await,
224 ResponseFormat::InfluxdbV1 => InfluxdbV1Response::from_output(outputs, epoch).await,
225 ResponseFormat::Json => JsonResponse::from_output(outputs).await,
226 ResponseFormat::Null => NullResponse::from_output(outputs).await,
227 };
228
229 resp.with_execution_time(start.elapsed().as_millis() as u64)
230}
231
232fn limit_output_rows(output: Output, limit: usize) -> Result<Output> {
234 let mut remaining = limit;
235 let data = match output.data {
236 OutputData::AffectedRows(rows) => OutputData::AffectedRows(rows),
237 OutputData::RecordBatches(batches) => {
238 let schema = batches.schema();
239 let batches = batches
240 .into_iter()
241 .filter_map(|batch| take_response_rows(batch, &mut remaining).transpose())
242 .collect::<RecordBatchResult<Vec<_>>>()
243 .context(CollectRecordbatchSnafu)?;
244 OutputData::RecordBatches(
245 RecordBatches::try_new(schema, batches).context(CollectRecordbatchSnafu)?,
246 )
247 }
248 OutputData::Stream(stream) => {
249 let schema = stream.schema();
250 let stream = stream.filter_map(move |batch| {
253 future::ready(
254 batch
255 .and_then(|batch| take_response_rows(batch, &mut remaining))
256 .transpose(),
257 )
258 });
259 OutputData::Stream(Box::pin(RecordBatchStreamWrapper::new(schema, stream)))
260 }
261 };
262 Ok(Output::new(data, output.meta))
263}
264
265fn take_response_rows(
266 batch: RecordBatch,
267 remaining: &mut usize,
268) -> RecordBatchResult<Option<RecordBatch>> {
269 let num_rows = batch.num_rows().min(*remaining);
270 *remaining -= num_rows;
271 if num_rows == 0 {
272 Ok(None)
273 } else if num_rows == batch.num_rows() {
274 Ok(Some(batch))
275 } else {
276 batch.slice(0, num_rows).map(Some)
277 }
278}
279
280#[axum_macros::debug_handler]
289#[tracing::instrument(
290 skip_all,
291 fields(protocol = "http", request_type = "sql_analyze_stream")
292)]
293pub async fn sql_analyze_stream(
294 State(state): State<ApiState>,
295 Query(query_params): Query<SqlQuery>,
296 Extension(mut query_ctx): Extension<QueryContext>,
297 form_params: std::result::Result<Form<SqlQuery>, FormRejection>,
298) -> Response {
299 let start = Instant::now();
300 let form_params = match form_params {
301 Ok(Form(params)) => params,
302 Err(err) => {
303 if err.status() != axum::http::StatusCode::UNSUPPORTED_MEDIA_TYPE {
304 return ErrorResponse::from_error_message(
305 StatusCode::InvalidArguments,
306 err.body_text(),
307 )
308 .with_execution_time(start.elapsed().as_millis() as u64)
309 .into_response();
310 }
311 SqlQuery::default()
312 }
313 };
314 let sql_handler = &state.sql_handler;
315 if let Some(db) = &query_params.db.or(form_params.db) {
316 let (catalog, schema) = parse_catalog_and_schema_from_db_string(db);
317 query_ctx.set_current_catalog(&catalog);
318 query_ctx.set_current_schema(&schema);
319 }
320 query_ctx.set_channel(Channel::HttpSql);
321 query_ctx.enable_live_analyze_metrics();
322 let query_ctx = Arc::new(query_ctx);
323
324 let Some(sql) = query_params.sql.or(form_params.sql) else {
325 return ErrorResponse::from_error_message(
326 StatusCode::InvalidArguments,
327 "sql parameter is required.".to_string(),
328 )
329 .with_execution_time(start.elapsed().as_millis() as u64)
330 .into_response();
331 };
332 if let Some((status, msg)) = validate_schema(sql_handler.clone(), query_ctx.clone()).await {
333 return ErrorResponse::from_error_message(status, msg)
334 .with_execution_time(start.elapsed().as_millis() as u64)
335 .into_response();
336 }
337
338 let interval_ms = query_params
339 .snapshot_interval_ms
340 .or(form_params.snapshot_interval_ms)
341 .unwrap_or(DEFAULT_ANALYZE_SNAPSHOT_INTERVAL_MS)
342 .clamp(
343 MIN_ANALYZE_SNAPSHOT_INTERVAL_MS,
344 MAX_ANALYZE_SNAPSHOT_INTERVAL_MS,
345 );
346
347 let output = match state
348 .sql_handler
349 .do_analyze_stream_query(&sql, query_ctx.clone())
350 .await
351 {
352 Ok(output) => output,
353 Err(err) => {
354 return ErrorResponse::from_error(err)
355 .with_execution_time(start.elapsed().as_millis() as u64)
356 .into_response();
357 }
358 };
359
360 let plan = output.meta.plan.clone();
361 let OutputData::Stream(stream) = output.data else {
362 return ErrorResponse::from_error_message(
363 StatusCode::InvalidArguments,
364 "analyze stream query must return a stream".to_string(),
365 )
366 .with_execution_time(start.elapsed().as_millis() as u64)
367 .into_response();
368 };
369 let schema = stream.schema();
370
371 let (metrics_tx, metrics_rx) = watch::channel::<Option<String>>(None);
372 let (terminal_tx, terminal_rx) = watch::channel::<Option<AnalyzeStreamMessage>>(None);
373 let notify = Arc::new(Notify::new());
374 let (cancel_tx, mut cancel_rx) = watch::channel(false);
375 let worker_notify = notify.clone();
376 let panic_notify = notify.clone();
377 let sequence = Arc::new(AtomicU64::new(0));
378 let worker_sequence = sequence.clone();
379 let worker_terminal_tx = terminal_tx.clone();
380 let worker = common_runtime::spawn_global(async move {
381 let worker_result = AssertUnwindSafe(async move {
382 let mut stream = stream;
383 let mut batches = Vec::new();
384 let mut current_interval_ms = interval_ms;
385 let tick = tokio::time::sleep(Duration::from_millis(current_interval_ms));
386 tokio::pin!(tick);
387
388 loop {
389 tokio::select! {
390 _ = cancel_rx.changed() => return,
391 item = stream.next() => {
392 match item {
393 Some(Ok(next_batch)) => batches.push(next_batch),
394 Some(Err(err)) => {
395 let status = err.status_code();
396 let event_name = if status == StatusCode::Cancelled { "canceled" } else { "error" };
397 let (payload, _) = make_analyze_payload(AnalyzePayloadArgs {
398 seq: worker_sequence.load(Ordering::Relaxed),
399 state: event_name,
400 partial: false,
401 start,
402 plan: plan.as_ref(),
403 output: None,
404 reason: Some(err.output_msg()),
405 code: Some(status as u32),
406 });
407 send_analyze_terminal(&terminal_tx, &worker_notify, event_name, payload);
408 return;
409 }
410 None => {
411 let output = HttpRecordsOutput::try_new(schema.clone(), batches)
412 .map(GreptimeQueryOutput::Records);
413 let (event_name, payload) = make_final_analyze_event(
414 output.map_err(|err| (err.output_msg(), err.status_code() as u32)),
415 worker_sequence.load(Ordering::Relaxed),
416 start,
417 plan.as_ref(),
418 );
419 send_analyze_terminal(&terminal_tx, &worker_notify, event_name, payload);
420 return;
421 }
422 }
423 }
424 _ = &mut tick, if plan.is_some() => {
425 let (payload, payload_bytes) = make_analyze_payload(AnalyzePayloadArgs {
426 seq: worker_sequence.load(Ordering::Relaxed),
427 state: "metrics",
428 partial: true,
429 start,
430 plan: plan.as_ref(),
431 output: None,
432 reason: None,
433 code: None,
434 });
435 current_interval_ms = adaptive_interval_ms(payload_bytes, interval_ms);
436 worker_sequence.fetch_add(1, Ordering::Relaxed);
437 if metrics_tx.send(Some(payload)).is_err() {
438 return;
439 }
440 worker_notify.notify_one();
441 tick.as_mut().reset(tokio::time::Instant::now() + Duration::from_millis(current_interval_ms));
442 }
443 }
444 }
445 })
446 .catch_unwind()
447 .await;
448
449 if worker_result.is_err() {
450 tracing::debug!("analyze stream worker panicked");
451 let (payload, _) = make_analyze_payload(AnalyzePayloadArgs {
452 seq: sequence.load(Ordering::Relaxed),
453 state: "error",
454 partial: false,
455 start,
456 plan: None,
457 output: None,
458 reason: Some("analyze stream worker panicked".to_string()),
459 code: Some(StatusCode::Internal as u32),
460 });
461 send_analyze_terminal(&worker_terminal_tx, &panic_notify, "error", payload);
462 }
463 });
464
465 let sse_stream = analyze_stream_body(metrics_rx, terminal_rx, notify, cancel_tx, worker);
466
467 Sse::new(sse_stream)
468 .keep_alive(KeepAlive::new().interval(Duration::from_secs(15)))
469 .into_response()
470}
471
472#[doc(hidden)]
473pub fn analyze_stream_body(
474 metrics: watch::Receiver<Option<String>>,
475 terminal: watch::Receiver<Option<AnalyzeStreamMessage>>,
476 notify: Arc<Notify>,
477 cancel: watch::Sender<bool>,
478 worker: common_runtime::JoinHandle<()>,
479) -> impl futures::Stream<Item = std::result::Result<Event, std::convert::Infallible>> {
480 futures::stream::unfold(
481 AnalyzeStreamBodyState {
482 latest_metrics: metrics,
483 terminal,
484 notify,
485 _worker: AnalyzeStreamWorkerGuard {
486 cancel,
487 handle: worker,
488 },
489 done: false,
490 },
491 |mut state| async move {
492 if state.done {
493 return None;
494 }
495 loop {
496 let notify = Arc::clone(&state.notify);
497 let notified = notify.notified();
498 tokio::pin!(notified);
499 notified.as_mut().enable();
501 let latest = {
502 let latest = state.latest_metrics.borrow_and_update();
503 latest.has_changed().then(|| latest.clone())
504 };
505 if let Some(Some(payload)) = latest {
506 return Some((Ok(Event::default().event("metrics").data(payload)), state));
507 }
508 if state.terminal.borrow().is_some() {
511 let terminal = { state.terminal.borrow_and_update().clone() };
512 if let Some(AnalyzeStreamMessage {
513 event_name,
514 payload,
515 }) = terminal
516 {
517 state.done = true;
518 return Some((Ok(Event::default().event(event_name).data(payload)), state));
519 }
520 }
521 notified.await;
522 }
523 },
524 )
525}
526
527#[doc(hidden)]
528pub fn send_analyze_terminal(
529 terminal_tx: &watch::Sender<Option<AnalyzeStreamMessage>>,
530 notify: &Notify,
531 event_name: &'static str,
532 payload: String,
533) {
534 if terminal_tx.send_if_modified(|terminal| {
535 if terminal.is_none() {
536 *terminal = Some(AnalyzeStreamMessage {
537 event_name,
538 payload,
539 });
540 true
541 } else {
542 false
543 }
544 }) {
545 notify.notify_one();
546 }
547}
548
549fn adaptive_interval_ms(payload_bytes: usize, requested_ms: u64) -> u64 {
550 if payload_bytes >= 10 * 1024 * 1024 {
551 requested_ms.max(30_000)
552 } else if payload_bytes >= 1024 * 1024 {
553 requested_ms.max(10_000)
554 } else {
555 requested_ms
556 }
557}
558
559fn make_final_analyze_event(
560 output: std::result::Result<GreptimeQueryOutput, (String, u32)>,
561 seq: u64,
562 start: Instant,
563 plan: Option<&Arc<dyn ExecutionPlan>>,
564) -> (&'static str, String) {
565 match output {
566 Ok(output) => (
567 "final",
568 make_analyze_payload(AnalyzePayloadArgs {
569 seq,
570 state: "final",
571 partial: false,
572 start,
573 plan,
574 output: Some(output),
575 reason: None,
576 code: None,
577 })
578 .0,
579 ),
580 Err((reason, code)) => (
581 "error",
582 make_analyze_payload(AnalyzePayloadArgs {
583 seq,
584 state: "error",
585 partial: false,
586 start,
587 plan,
588 output: None,
589 reason: Some(reason),
590 code: Some(code),
591 })
592 .0,
593 ),
594 }
595}
596
597struct AnalyzePayloadArgs<'a> {
598 seq: u64,
599 state: &'static str,
600 partial: bool,
601 start: Instant,
602 plan: Option<&'a Arc<dyn ExecutionPlan>>,
603 output: Option<GreptimeQueryOutput>,
604 reason: Option<String>,
605 code: Option<u32>,
606}
607
608fn make_analyze_payload(args: AnalyzePayloadArgs<'_>) -> (String, usize) {
609 let AnalyzePayloadArgs {
610 seq,
611 state,
612 partial,
613 start,
614 plan,
615 output,
616 reason,
617 code,
618 } = args;
619 let metrics =
621 plan.and_then(|plan| query::analyze_plan_metrics_to_json_value(plan, !partial).ok());
622 let payload = AnalyzeStreamPayload {
623 seq,
624 state,
625 partial,
626 elapsed_ms: start.elapsed().as_millis() as u64,
627 metrics,
628 output,
629 reason,
630 code,
631 };
632 let payload_string = serde_json::to_string(&payload).unwrap_or_else(|e| {
633 serde_json::json!({
634 "seq": seq,
635 "state": "error",
636 "partial": false,
637 "reason": format!("Failed to serialize SSE payload: {e}"),
638 })
639 .to_string()
640 });
641 let payload_bytes = payload_string.len();
642 (payload_string, payload_bytes)
643}
644
645#[axum_macros::debug_handler]
647#[tracing::instrument(skip_all, fields(protocol = "http", request_type = "sql"))]
648pub async fn sql_parse(
649 Query(query_params): Query<SqlQuery>,
650 Form(form_params): Form<SqlQuery>,
651) -> Result<Json<Vec<Statement>>> {
652 let Some(sql) = query_params.sql.or(form_params.sql) else {
653 return InvalidQuerySnafu {
654 reason: "sql parameter is required.",
655 }
656 .fail();
657 };
658
659 let stmts =
660 ParserContext::create_with_dialect(&sql, &GreptimeDbDialect {}, ParseOptions::default())
661 .context(FailedToParseQuerySnafu)?;
662
663 Ok(stmts.into())
664}
665
666#[derive(Debug, Serialize, Deserialize)]
667pub struct SqlFormatResponse {
668 pub formatted: String,
669}
670
671#[axum_macros::debug_handler]
673#[tracing::instrument(skip_all, fields(protocol = "http", request_type = "sql_format"))]
674pub async fn sql_format(
675 Query(query_params): Query<SqlQuery>,
676 Form(form_params): Form<SqlQuery>,
677) -> axum::response::Response {
678 let Some(sql) = query_params.sql.or(form_params.sql) else {
679 let resp = ErrorResponse::from_error_message(
680 StatusCode::InvalidArguments,
681 "sql parameter is required.".to_string(),
682 );
683 return HttpResponse::Error(resp).into_response();
684 };
685
686 let stmts = match ParserContext::create_with_dialect(
688 &sql,
689 &GreptimeDbDialect {},
690 ParseOptions::default(),
691 ) {
692 Ok(v) => v,
693 Err(e) => return HttpResponse::Error(ErrorResponse::from_error(e)).into_response(),
694 };
695
696 let mut parts: Vec<String> = Vec::with_capacity(stmts.len());
697 for stmt in stmts {
698 let mut s = format!("{stmt}");
699 if !s.trim_end().ends_with(';') {
700 s.push(';');
701 }
702 parts.push(s);
703 }
704
705 let formatted = parts.join("\n");
706 Json(SqlFormatResponse { formatted }).into_response()
707}
708
709pub async fn from_output(
711 outputs: Vec<crate::error::Result<Output>>,
712) -> std::result::Result<(Vec<GreptimeQueryOutput>, HashMap<String, Value>), ErrorResponse> {
713 let mut results = Vec::with_capacity(outputs.len());
716 let mut merge_map = HashMap::new();
717
718 for out in outputs {
719 match out {
720 Ok(o) => match o.data {
721 OutputData::AffectedRows(rows) => {
722 results.push(GreptimeQueryOutput::AffectedRows(rows));
723 if o.meta.cost > 0 {
724 merge_map.insert(GREPTIME_EXEC_WRITE_COST.to_string(), o.meta.cost as u64);
725 }
726 }
727 OutputData::Stream(stream) => {
728 let schema = stream.schema().clone();
729 let mut http_record_output = match util::collect(stream).await {
731 Ok(rows) => match HttpRecordsOutput::try_new(schema, rows) {
732 Ok(rows) => rows,
733 Err(err) => {
734 return Err(ErrorResponse::from_error(err));
735 }
736 },
737 Err(err) => {
738 return Err(ErrorResponse::from_error(err));
739 }
740 };
741 if let Some(physical_plan) = o.meta.plan {
742 let mut result_map = HashMap::new();
743
744 let mut tmp = vec![&mut merge_map, &mut result_map];
745 collect_plan_metrics(&physical_plan, &mut tmp);
746 let re = result_map
747 .into_iter()
748 .map(|(k, v)| (k, Value::from(v)))
749 .collect::<HashMap<String, Value>>();
750 http_record_output.metrics.extend(re);
751 }
752 results.push(GreptimeQueryOutput::Records(http_record_output))
753 }
754 OutputData::RecordBatches(rbs) => {
755 match HttpRecordsOutput::try_new(rbs.schema(), rbs.take()) {
756 Ok(rows) => {
757 results.push(GreptimeQueryOutput::Records(rows));
758 }
759 Err(err) => {
760 return Err(ErrorResponse::from_error(err));
761 }
762 }
763 }
764 },
765
766 Err(err) => {
767 return Err(ErrorResponse::from_error(err));
768 }
769 }
770 }
771
772 let merge_map = merge_map
773 .into_iter()
774 .map(|(k, v)| (k, Value::from(v)))
775 .collect();
776
777 Ok((results, merge_map))
778}
779
780#[derive(Debug, Default, Serialize, Deserialize)]
781pub struct PromqlQuery {
782 pub query: String,
783 pub start: String,
784 pub end: String,
785 pub step: String,
786 pub lookback: Option<String>,
787 pub db: Option<String>,
788 pub format: Option<String>,
792 pub compression: Option<String>,
794 pub epoch: Option<String>,
805}
806
807impl From<PromqlQuery> for PromQuery {
808 fn from(query: PromqlQuery) -> Self {
809 PromQuery {
810 query: query.query,
811 start: query.start,
812 end: query.end,
813 step: query.step,
814 lookback: query
815 .lookback
816 .unwrap_or_else(|| DEFAULT_LOOKBACK_STRING.to_string()),
817 alias: None,
819 }
820 }
821}
822
823#[axum_macros::debug_handler]
825#[tracing::instrument(skip_all, fields(protocol = "http", request_type = "promql"))]
826pub async fn promql(
827 State(state): State<ApiState>,
828 Query(params): Query<PromqlQuery>,
829 Extension(mut query_ctx): Extension<QueryContext>,
830) -> Response {
831 let sql_handler = &state.sql_handler;
832 let exec_start = Instant::now();
833 let db = query_ctx.get_db_string();
834
835 query_ctx.set_channel(Channel::Promql);
836 let query_ctx = Arc::new(query_ctx);
837
838 let _timer = crate::metrics::METRIC_HTTP_PROMQL_ELAPSED
839 .with_label_values(&[db.as_str()])
840 .start_timer();
841
842 let resp = if let Some((status, msg)) =
843 validate_schema(sql_handler.clone(), query_ctx.clone()).await
844 {
845 let resp = ErrorResponse::from_error_message(status, msg);
846 HttpResponse::Error(resp)
847 } else {
848 let format = params
849 .format
850 .as_ref()
851 .map(|s| s.to_lowercase())
852 .map(|s| ResponseFormat::parse(s.as_str()).unwrap_or(ResponseFormat::GreptimedbV1))
853 .unwrap_or(ResponseFormat::GreptimedbV1);
854 let epoch = params
855 .epoch
856 .as_ref()
857 .map(|s| s.to_lowercase())
858 .map(|s| Epoch::parse(s.as_str()).unwrap_or(Epoch::Millisecond));
859 let compression = params.compression.clone();
860
861 let prom_query = params.into();
862 let outputs = sql_handler.do_promql_query(&prom_query, query_ctx).await;
863
864 match format {
865 ResponseFormat::Arrow => ArrowResponse::from_output(outputs, compression).await,
866 ResponseFormat::Csv(with_names, with_types) => {
867 CsvResponse::from_output(outputs, with_names, with_types).await
868 }
869 ResponseFormat::Table => TableResponse::from_output(outputs).await,
870 ResponseFormat::GreptimedbV1 => GreptimedbV1Response::from_output(outputs).await,
871 ResponseFormat::InfluxdbV1 => InfluxdbV1Response::from_output(outputs, epoch).await,
872 ResponseFormat::Json => JsonResponse::from_output(outputs).await,
873 ResponseFormat::Null => NullResponse::from_output(outputs).await,
874 }
875 };
876
877 resp.with_execution_time(exec_start.elapsed().as_millis() as u64)
878 .into_response()
879}
880
881#[axum_macros::debug_handler]
883pub async fn metrics(
884 State(state): State<MetricsHandler>,
885 Query(_params): Query<HashMap<String, String>>,
886) -> String {
887 #[cfg(not(windows))]
892 if let Some(c) = crate::metrics::jemalloc::JEMALLOC_COLLECTOR.as_ref()
893 && let Err(e) = c.update()
894 {
895 common_telemetry::error!(e; "Failed to update jemalloc metrics");
896 }
897 state.render()
898}
899
900#[derive(Debug, Serialize, Deserialize)]
901pub struct HealthQuery {}
902
903#[derive(Debug, Serialize, Deserialize, PartialEq, Eq)]
904pub struct HealthResponse {}
905
906#[axum_macros::debug_handler]
910pub async fn health(Query(_params): Query<HealthQuery>) -> Json<HealthResponse> {
911 Json(HealthResponse {})
912}
913
914#[derive(Debug, Serialize, Deserialize, PartialEq, Eq)]
915pub struct StatusResponse<'a> {
916 pub commit: &'a str,
917 pub branch: &'a str,
918 pub rustc_version: &'a str,
919 pub hostname: String,
920 pub version: &'a str,
921}
922
923#[axum_macros::debug_handler]
925pub async fn status() -> Json<StatusResponse<'static>> {
926 let hostname = hostname::get()
927 .map(|s| s.to_string_lossy().to_string())
928 .unwrap_or_else(|_| "unknown".to_string());
929 let build_info = common_version::build_info();
930 Json(StatusResponse {
931 commit: build_info.commit,
932 branch: build_info.branch,
933 rustc_version: build_info.rustc,
934 hostname,
935 version: build_info.version,
936 })
937}
938
939#[axum_macros::debug_handler]
941pub async fn config(State(state): State<GreptimeOptionsConfigState>) -> Response {
942 (axum::http::StatusCode::OK, state.greptime_config_options).into_response()
943}
944
945async fn validate_schema(
946 sql_handler: ServerSqlQueryHandlerRef,
947 query_ctx: QueryContextRef,
948) -> Option<(StatusCode, String)> {
949 match sql_handler
950 .is_valid_schema(query_ctx.current_catalog(), &query_ctx.current_schema())
951 .await
952 {
953 Ok(true) => None,
954 Ok(false) => Some((
955 StatusCode::DatabaseNotFound,
956 format!("Database not found: {}", query_ctx.get_db_string()),
957 )),
958 Err(e) => Some((
959 StatusCode::Internal,
960 format!(
961 "Error checking database: {}, {}",
962 query_ctx.get_db_string(),
963 e.output_msg(),
964 ),
965 )),
966 }
967}
968
969pub async fn index() -> axum::response::Html<String> {
970 let name = common_version::product_name();
971 let version = common_version::version();
972 axum::response::Html(format!(
973 r#"<!DOCTYPE html>
974<html>
975<head><title>{name}</title></head>
976<body>
977<h1>{name}</h1>
978<p>Version: {version}</p>
979<ul>
980<li><a href="/dashboard">Dashboard UI</a></li>
981<li><a href="/v1/health">Health</a> (JSON)</li>
982<li><a href="/status">Status</a> (JSON)</li>
983<li><a href="/metrics">Metrics</a> (For Prometheus Scrape)</li>
984<li><a href="/config">Config</a> (TXT)</li>
985</ul>
986</body>
987</html>"#,
988 ))
989}
990
991#[cfg(test)]
992mod tests {
993 use std::sync::Mutex;
994 use std::sync::atomic::AtomicUsize;
995
996 use async_trait::async_trait;
997 use common_query::OutputMeta;
998 use datafusion_expr::LogicalPlan;
999 use datatypes::prelude::ConcreteDataType;
1000 use datatypes::schema::{ColumnSchema, Schema};
1001 use datatypes::vectors::{BinaryVector, UInt32Vector, VectorRef};
1002 use futures::stream;
1003 use query::query_engine::DescribeResult;
1004
1005 use super::*;
1006 use crate::query_handler::sql::SqlQueryHandler;
1007
1008 struct TestSqlQueryHandler {
1009 outputs: Mutex<Vec<Result<Output>>>,
1010 }
1011
1012 #[async_trait]
1013 impl SqlQueryHandler for TestSqlQueryHandler {
1014 async fn do_query(&self, _: &str, _: QueryContextRef) -> Vec<Result<Output>> {
1015 std::mem::take(&mut *self.outputs.lock().unwrap())
1016 }
1017
1018 async fn do_analyze_stream_query(&self, _: &str, _: QueryContextRef) -> Result<Output> {
1019 unimplemented!()
1020 }
1021
1022 async fn do_exec_plan(
1023 &self,
1024 _: LogicalPlan,
1025 _: Option<Statement>,
1026 _: QueryContextRef,
1027 ) -> Result<Output> {
1028 unimplemented!()
1029 }
1030
1031 async fn do_promql_query(&self, _: &PromQuery, _: QueryContextRef) -> Vec<Result<Output>> {
1032 unimplemented!()
1033 }
1034
1035 async fn do_describe(
1036 &self,
1037 _: Statement,
1038 _: QueryContextRef,
1039 ) -> Result<Option<DescribeResult>> {
1040 unimplemented!()
1041 }
1042
1043 async fn is_valid_schema(&self, _: &str, _: &str) -> Result<bool> {
1044 Ok(true)
1045 }
1046 }
1047
1048 async fn sql_response(
1049 outputs: Vec<Result<Output>>,
1050 format: &str,
1051 limit: Option<usize>,
1052 ) -> HttpResponse {
1053 sql(
1054 State(ApiState {
1055 sql_handler: Arc::new(TestSqlQueryHandler {
1056 outputs: Mutex::new(outputs),
1057 }),
1058 }),
1059 Query(SqlQuery {
1060 sql: Some("select number from numbers".to_string()),
1061 format: Some(format.to_string()),
1062 limit,
1063 ..Default::default()
1064 }),
1065 Extension(QueryContext::with_db_name(None)),
1066 Form(SqlQuery::default()),
1067 )
1068 .await
1069 }
1070
1071 fn number_batches() -> RecordBatches {
1072 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1073 "number",
1074 ConcreteDataType::uint32_datatype(),
1075 false,
1076 )]));
1077 let batches = [vec![], vec![0, 1], vec![], vec![2, 3, 4], vec![]]
1078 .into_iter()
1079 .map(|values| {
1080 let columns: Vec<VectorRef> = vec![Arc::new(UInt32Vector::from_slice(values))];
1081 RecordBatch::new(schema.clone(), columns).unwrap()
1082 })
1083 .collect();
1084 RecordBatches::try_new(schema, batches).unwrap()
1085 }
1086
1087 fn number_output(streaming: bool) -> Output {
1088 let batches = number_batches();
1089 if streaming {
1090 Output::new_with_stream(batches.as_stream())
1091 } else {
1092 Output::new_with_record_batches(batches)
1093 }
1094 }
1095
1096 #[tokio::test]
1097 async fn test_sql_response_limit_across_batches_and_formats() {
1098 for format in [
1099 "greptimedb_v1",
1100 "json",
1101 "csv",
1102 "csvWithNames",
1103 "csvWithNamesAndTypes",
1104 "table",
1105 ] {
1106 for streaming in [false, true] {
1107 for limit in [
1108 None,
1109 Some(0),
1110 Some(1),
1111 Some(2),
1112 Some(3),
1113 Some(5),
1114 Some(10),
1115 Some(usize::MAX),
1116 ] {
1117 let response =
1118 sql_response(vec![Ok(number_output(streaming))], format, limit).await;
1119 let output = match &response {
1120 HttpResponse::GreptimedbV1(response) => response.output(),
1121 HttpResponse::Json(response) => response.output(),
1122 HttpResponse::Csv(response) => response.output(),
1123 HttpResponse::Table(response) => response.output(),
1124 _ => panic!("unexpected response: {response:?}"),
1125 };
1126 let GreptimeQueryOutput::Records(records) = &output[0] else {
1127 panic!("expected records");
1128 };
1129 let expected_rows: Vec<_> = (0..limit.unwrap_or(5).min(5))
1130 .map(|number| vec![Value::from(number)])
1131 .collect();
1132 assert_eq!(records.rows(), &expected_rows);
1133 assert_eq!(records.total_rows, expected_rows.len());
1134 assert_eq!(records.schema.column_schemas[0].name, "number");
1135 }
1136 }
1137 }
1138 }
1139
1140 #[tokio::test]
1141 async fn test_sql_response_limit_is_per_statement_and_preserves_write_cost() {
1142 let affected = Output::new(
1143 OutputData::AffectedRows(7),
1144 OutputMeta {
1145 cost: 42,
1146 ..Default::default()
1147 },
1148 );
1149 let response = sql_response(
1150 vec![
1151 Ok(number_output(true)),
1152 Ok(affected),
1153 Ok(number_output(false)),
1154 ],
1155 "greptimedb_v1",
1156 Some(1),
1157 )
1158 .await;
1159 let HttpResponse::GreptimedbV1(response) = response else {
1160 panic!("expected greptimedb response");
1161 };
1162 assert_eq!(response.output().len(), 3);
1163 for index in [0, 2] {
1164 let GreptimeQueryOutput::Records(records) = &response.output()[index] else {
1165 panic!("expected records");
1166 };
1167 assert_eq!(records.rows(), &vec![vec![Value::from(0)]]);
1168 }
1169 assert!(matches!(
1170 response.output()[1],
1171 GreptimeQueryOutput::AffectedRows(7)
1172 ));
1173 assert_eq!(
1174 response.resp_metrics[GREPTIME_EXEC_WRITE_COST],
1175 Value::from(42)
1176 );
1177 }
1178
1179 #[tokio::test]
1180 async fn test_sql_response_limit_does_not_affect_other_formats() {
1181 for format in ["arrow", "influxdb_v1", "null"] {
1182 let unlimited = sql_response(vec![Ok(number_output(true))], format, None).await;
1183 let limited = sql_response(vec![Ok(number_output(true))], format, Some(0)).await;
1184 assert_eq!(
1185 serde_json::to_value(unlimited.with_execution_time(0)).unwrap(),
1186 serde_json::to_value(limited.with_execution_time(0)).unwrap(),
1187 );
1188 }
1189 }
1190
1191 #[tokio::test]
1192 async fn test_sql_response_limit_drains_input_and_propagates_late_errors() {
1193 for limit in [0, 1, 3] {
1194 for late_error in [false, true] {
1195 let batches = number_batches();
1196 let schema = batches.schema();
1197 let mut input: Vec<_> = batches.into_iter().map(Ok).collect();
1198 if late_error {
1199 input.push(
1200 common_recordbatch::error::CreateRecordBatchesSnafu {
1201 reason: "late stream error",
1202 }
1203 .fail(),
1204 );
1205 }
1206 let expected_batches = input.len();
1207 let observed = Arc::new(AtomicUsize::new(0));
1208 let counter = observed.clone();
1209 let stream = stream::iter(input).inspect(move |_| {
1210 counter.fetch_add(1, Ordering::Relaxed);
1211 });
1212 let output = Output::new_with_stream(Box::pin(RecordBatchStreamWrapper::new(
1213 schema, stream,
1214 )));
1215 let response = sql_response(vec![Ok(output)], "greptimedb_v1", Some(limit)).await;
1216 assert_eq!(observed.load(Ordering::Relaxed), expected_batches);
1217 assert_eq!(matches!(response, HttpResponse::Error(_)), late_error);
1218 }
1219 }
1220 }
1221
1222 #[tokio::test]
1223 async fn test_sql_response_limit_skips_conversion_of_discarded_rows() {
1224 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1225 "payload",
1226 ConcreteDataType::json_datatype(),
1227 false,
1228 )]));
1229 let columns: Vec<VectorRef> = vec![Arc::new(BinaryVector::from(vec![
1230 datatypes::types::parse_string_to_jsonb(r#"{"ok":1}"#).unwrap(),
1231 b"invalid jsonb".to_vec(),
1232 ]))];
1233 let batch = RecordBatch::new(schema.clone(), columns).unwrap();
1234 for streaming in [false, true] {
1235 for limit in [0, 1, 2] {
1236 let batches = RecordBatches::try_new(schema.clone(), vec![batch.clone()]).unwrap();
1237 let output = if streaming {
1238 Output::new_with_stream(batches.as_stream())
1239 } else {
1240 Output::new_with_record_batches(batches)
1241 };
1242 let response = sql_response(vec![Ok(output)], "greptimedb_v1", Some(limit)).await;
1243 assert_eq!(matches!(response, HttpResponse::Error(_)), limit == 2);
1245 }
1246 }
1247 }
1248
1249 #[test]
1250 fn test_final_analyze_event_uses_error_event_for_conversion_error() {
1251 let (event_name, payload) = make_final_analyze_event(
1252 Err((
1253 "conversion failed".to_string(),
1254 StatusCode::InvalidArguments as u32,
1255 )),
1256 7,
1257 Instant::now(),
1258 None,
1259 );
1260
1261 assert_eq!(event_name, "error");
1262 let value: Value = serde_json::from_str(&payload).unwrap();
1263 assert_eq!(value["state"], "error");
1264 assert_eq!(value["reason"], "conversion failed");
1265 }
1266}