Skip to main content

servers/http/
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
15use 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    // (Optional) result format: [`greptimedb_v1`, `influxdb_v1`, `csv`,
69    // `arrow`],
70    // the default value is `greptimedb_v1`
71    pub format: Option<String>,
72    // Returns epoch timestamps with the specified precision.
73    // Both u and µ indicate microseconds.
74    // epoch = [ns,u,µ,ms,s],
75    //
76    // TODO(jeremy): currently, only InfluxDB result format is supported,
77    // and all columns of the `Timestamp` type will be converted to their
78    // specified time precision. Maybe greptimedb format can support this
79    // param too.
80    pub epoch: Option<String>,
81    pub limit: Option<usize>,
82    // For arrow output
83    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    // Keeping the guard in the body state makes dropping the response cancel and
131    // abort the worker instead of leaving an owned query stream detached.
132    _worker: AnalyzeStreamWorkerGuard,
133    done: bool,
134}
135
136/// Handler to execute sql
137#[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
232/// Limits each statement's response before collecting batches and converting rows.
233fn 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            // This is a response limit. Drain the input even after reaching it so
251            // execution errors and terminal query metrics are still observed.
252            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/// Handler to stream partial `EXPLAIN ANALYZE VERBOSE` metrics as SSE.
281///
282/// This endpoint is POST-only SSE, so browser `EventSource` does
283/// not apply. Each `metrics` event carries a complete snapshot (not a delta);
284/// large snapshots are throttled but never truncated. `final`, `canceled`, and
285/// `error` are terminal events. If the client disconnects it won't receive a
286/// `canceled` event, but the production frontend stream is dropped and
287/// best-effort cancels the underlying query.
288#[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                // Register before checking the slots to avoid a lost wakeup.
500                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                // Inspect the terminal slot directly because a closed watch channel
509                // can otherwise hide a value published while a worker was unwinding.
510                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    // Periodic snapshots are compact; terminal snapshots retain verbose plan details.
620    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/// Handler to parse sql
646#[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/// Handler to format sql string
672#[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    // Parse using GreptimeDB dialect then reconstruct statements via Display
687    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
709/// Create a response from query result
710pub async fn from_output(
711    outputs: Vec<crate::error::Result<Output>>,
712) -> std::result::Result<(Vec<GreptimeQueryOutput>, HashMap<String, Value>), ErrorResponse> {
713    // TODO(sunng87): this api response structure cannot represent error well.
714    //  It hides successful execution results from error response
715    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                    // TODO(sunng87): streaming response
730                    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    // (Optional) result format: [`greptimedb_v1`, `influxdb_v1`, `csv`,
789    // `arrow`],
790    // the default value is `greptimedb_v1`
791    pub format: Option<String>,
792    // For arrow output
793    pub compression: Option<String>,
794    // Returns epoch timestamps with the specified precision.
795    // Both u and µ indicate microseconds.
796    // epoch = [ns,u,µ,ms,s],
797    //
798    // For influx output only
799    //
800    // TODO(jeremy): currently, only InfluxDB result format is supported,
801    // and all columns of the `Timestamp` type will be converted to their
802    // specified time precision. Maybe greptimedb format can support this
803    // param too.
804    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            // TODO(dennis): support alias from http params or parse from query.query
818            alias: None,
819        }
820    }
821}
822
823/// Handler to execute promql
824#[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/// Handler to export metrics
882#[axum_macros::debug_handler]
883pub async fn metrics(
884    State(state): State<MetricsHandler>,
885    Query(_params): Query<HashMap<String, String>>,
886) -> String {
887    // A default ProcessCollector is registered automatically in prometheus.
888    // We do not need to explicitly collect process-related data.
889    // But ProcessCollector only support on linux.
890
891    #[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/// Handler to export healthy check
907///
908/// Currently simply return status "200 OK" (default) with an empty json payload "{}"
909#[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/// Handler to expose information info about runtime, build, etc.
924#[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/// Handler to expose configuration information info about runtime, build, etc.
940#[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                // Only retained rows should be decoded into HTTP response values.
1244                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}