Skip to main content

servers/postgres/
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::fmt::Debug;
16use std::pin::Pin;
17use std::sync::Arc;
18
19use async_trait::async_trait;
20use common_query::{Output, OutputData};
21use common_recordbatch::RecordBatch;
22use common_recordbatch::error::Result as RecordBatchResult;
23use common_telemetry::{debug, info, tracing};
24use datafusion::sql::sqlparser::ast::{CopyOption, CopyTarget, Statement as SqlParserStatement};
25use datafusion_common::ParamValues;
26use datafusion_expr::LogicalPlan;
27use datafusion_pg_catalog::sql::PostgresCompatibilityParser;
28use datatypes::prelude::ConcreteDataType;
29use datatypes::schema::{Schema, SchemaRef};
30use futures::{Sink, SinkExt, Stream, StreamExt, future, stream};
31use operator::statement::admin_output_schema;
32use pgwire::api::portal::{Format, Portal};
33use pgwire::api::query::{ExtendedQueryHandler, SimpleQueryHandler};
34use pgwire::api::results::{
35    CopyCsvOptions, CopyEncoder, CopyResponse, CopyTextOptions, DataRowEncoder,
36    DescribePortalResponse, DescribeStatementResponse, FieldInfo, QueryResponse, Response, Tag,
37};
38use pgwire::api::stmt::{QueryParser, StoredStatement};
39use pgwire::api::{ClientInfo, ErrorHandler, Type};
40use pgwire::error::{ErrorInfo, PgWireError, PgWireResult};
41use pgwire::messages::PgWireBackendMessage;
42use pgwire::messages::copy::CopyData;
43use pgwire::messages::data::DataRow;
44use query::dist_analyze_output_schema;
45use query::planner::DfLogicalPlanner;
46use query::query_engine::DescribeResult;
47use query::sql::DESCRIBE_TABLE_OUTPUT_SCHEMA;
48use session::Session;
49use session::context::QueryContextRef;
50use snafu::ResultExt;
51use sql::dialect::PostgreSqlDialect;
52use sql::parser::{ParseOptions, ParserContext};
53use sql::statements::statement::Statement;
54
55use crate::SqlPlan;
56use crate::error::{DataFusionSnafu, InferParameterTypesSnafu, Result};
57use crate::postgres::types::*;
58use crate::postgres::utils::convert_err;
59use crate::postgres::{PostgresServerHandlerInner, fixtures};
60use crate::query_handler::sql::ServerSqlQueryHandlerRef;
61
62impl PostgresServerHandlerInner {
63    fn new_query_context(&self) -> QueryContextRef {
64        let mut ctx = self.session.new_query_context();
65        Arc::make_mut(&mut ctx).set_batching_enabled(self.batching_enabled);
66        ctx
67    }
68}
69
70#[async_trait]
71impl SimpleQueryHandler for PostgresServerHandlerInner {
72    #[tracing::instrument(skip_all, fields(protocol = "postgres"))]
73    async fn do_query<C>(&self, client: &mut C, query: &str) -> PgWireResult<Vec<Response>>
74    where
75        C: ClientInfo + Sink<PgWireBackendMessage> + Unpin + Send + Sync,
76        C::Error: Debug,
77        PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
78    {
79        let query_ctx = self.new_query_context();
80        let db = query_ctx.get_db_string();
81        let _timer = crate::metrics::METRIC_POSTGRES_QUERY_TIMER
82            .with_label_values(&[crate::metrics::METRIC_POSTGRES_SIMPLE_QUERY, db.as_str()])
83            .start_timer();
84
85        if query.is_empty() {
86            // early return if query is empty
87            return Ok(vec![Response::EmptyQuery]);
88        }
89
90        let parsed_query = self.query_parser.compatibility_parser.parse(query);
91
92        let query = if let Ok(statements) = &parsed_query {
93            // Comments, whitespace and empty statements also require EmptyQueryResponse.
94            if statements.is_empty() {
95                return Ok(vec![Response::EmptyQuery]);
96            }
97            statements
98                .iter()
99                .map(|s| s.to_string())
100                .collect::<Vec<_>>()
101                .join(";")
102        } else {
103            query.to_string()
104        };
105
106        if let Some(resps) = fixtures::process(&query, query_ctx.clone()) {
107            send_warning_opt(client, query_ctx).await?;
108            Ok(resps)
109        } else {
110            let outputs = self.query_handler.do_query(&query, query_ctx.clone()).await;
111
112            let mut results = Vec::with_capacity(outputs.len());
113
114            let statements = parsed_query.ok();
115            for (idx, output) in outputs.into_iter().enumerate() {
116                let copy_format = statements
117                    .as_ref()
118                    .and_then(|stmts| stmts.get(idx))
119                    .and_then(check_copy_to_stdout);
120                let resp = if let Some(format) = &copy_format {
121                    output_to_copy_response(query_ctx.clone(), output, format)?
122                } else {
123                    output_to_query_response(query_ctx.clone(), output, &Format::UnifiedText)?
124                };
125                results.push(resp);
126            }
127
128            send_warning_opt(client, query_ctx).await?;
129            Ok(results)
130        }
131    }
132}
133
134async fn send_warning_opt<C>(client: &mut C, query_context: QueryContextRef) -> PgWireResult<()>
135where
136    C: Sink<PgWireBackendMessage> + Unpin + Send + Sync,
137    C::Error: Debug,
138    PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
139{
140    if let Some(warning) = query_context.warning() {
141        client
142            .feed(PgWireBackendMessage::NoticeResponse(
143                ErrorInfo::new(
144                    PgErrorSeverity::Warning.to_string(),
145                    PgErrorCode::Ec01000.code(),
146                    warning.clone(),
147                )
148                .into(),
149            ))
150            .await?;
151    }
152
153    Ok(())
154}
155
156pub(crate) fn output_to_query_response(
157    query_ctx: QueryContextRef,
158    output: Result<Output>,
159    field_format: &Format,
160) -> PgWireResult<Response> {
161    match output {
162        Ok(o) => match o.data {
163            OutputData::AffectedRows(rows) => {
164                Ok(Response::Execution(Tag::new("OK").with_rows(rows)))
165            }
166            OutputData::Stream(record_stream) => {
167                let schema = record_stream.schema();
168                recordbatches_to_query_response(query_ctx, record_stream, schema, field_format)
169            }
170            OutputData::RecordBatches(recordbatches) => {
171                let schema = recordbatches.schema();
172                recordbatches_to_query_response(
173                    query_ctx,
174                    recordbatches.as_stream(),
175                    schema,
176                    field_format,
177                )
178            }
179        },
180        Err(e) => Err(convert_err(e)),
181    }
182}
183
184type RowStream<T> = Pin<Box<dyn Stream<Item = PgWireResult<T>> + Send + Unpin>>;
185
186fn recordbatches_to_query_response<S>(
187    query_ctx: QueryContextRef,
188    recordbatches_stream: S,
189    schema: SchemaRef,
190    field_format: &Format,
191) -> PgWireResult<Response>
192where
193    S: Stream<Item = RecordBatchResult<RecordBatch>> + Send + Unpin + 'static,
194{
195    let format_options = format_options_from_query_ctx(&query_ctx);
196    let pg_schema = Arc::new(
197        schema_to_pg(schema.as_ref(), field_format, Some(format_options)).map_err(convert_err)?,
198    );
199
200    let encoder = DataRowEncoder::new(pg_schema.clone());
201    let row_stream = RecordBatchRowStream::new(
202        query_ctx.clone(),
203        pg_schema.clone(),
204        schema.clone(),
205        recordbatches_stream,
206        encoder,
207    );
208
209    let data_row_stream: RowStream<DataRow> = Box::pin(
210        row_stream
211            .map(move |result| match result {
212                Ok(rows) => Box::pin(stream::iter(rows.into_iter().map(Ok))) as RowStream<DataRow>,
213                Err(e) => Box::pin(stream::once(future::ready(Err(e)))) as RowStream<DataRow>,
214            })
215            .flatten(),
216    );
217
218    Ok(Response::Query(QueryResponse::new(
219        pg_schema,
220        data_row_stream,
221    )))
222}
223
224pub(crate) fn output_to_copy_response(
225    query_ctx: QueryContextRef,
226    output: Result<Output>,
227    format: &str,
228) -> PgWireResult<Response> {
229    match output {
230        Ok(o) => match o.data {
231            OutputData::AffectedRows(_) => Err(PgWireError::UserError(Box::new(ErrorInfo::new(
232                "ERROR".to_string(),
233                "42601".to_string(),
234                "COPY cannot be used with non-query statements".to_string(),
235            )))),
236            OutputData::Stream(record_stream) => {
237                let schema = record_stream.schema();
238                recordbatches_to_copy_response(query_ctx, record_stream, schema, format)
239            }
240            OutputData::RecordBatches(recordbatches) => {
241                let schema = recordbatches.schema();
242                recordbatches_to_copy_response(query_ctx, recordbatches.as_stream(), schema, format)
243            }
244        },
245        Err(e) => Err(convert_err(e)),
246    }
247}
248
249fn recordbatches_to_copy_response<S>(
250    query_ctx: QueryContextRef,
251    recordbatches_stream: S,
252    schema: SchemaRef,
253    format: &str,
254) -> PgWireResult<Response>
255where
256    S: Stream<Item = RecordBatchResult<RecordBatch>> + Send + Unpin + 'static,
257{
258    let format_options = format_options_from_query_ctx(&query_ctx);
259    let pg_fields = schema_to_pg(schema.as_ref(), &Format::UnifiedText, Some(format_options))
260        .map_err(convert_err)?;
261
262    let copy_format = match format.to_lowercase().as_str() {
263        "binary" => 1,
264        _ => 0,
265    };
266
267    let pg_schema = Arc::new(pg_fields);
268    let num_columns = pg_schema.len();
269
270    let copy_encoder = match format.to_lowercase().as_str() {
271        "csv" => CopyEncoder::new_csv(pg_schema.clone(), CopyCsvOptions::default()),
272        "binary" => CopyEncoder::new_binary(pg_schema.clone()),
273        _ => CopyEncoder::new_text(pg_schema.clone(), CopyTextOptions::default()),
274    };
275
276    let row_stream = RecordBatchRowStream::new(
277        query_ctx.clone(),
278        pg_schema.clone(),
279        schema.clone(),
280        recordbatches_stream,
281        copy_encoder,
282    );
283
284    let copy_stream: RowStream<CopyData> = Box::pin(
285        row_stream
286            .map(move |result| match result {
287                Ok(rows) => Box::pin(stream::iter(rows.into_iter().map(Ok))) as RowStream<CopyData>,
288                Err(e) => Box::pin(stream::once(future::ready(Err(e)))) as RowStream<CopyData>,
289            })
290            .flatten(),
291    );
292
293    Ok(Response::CopyOut(CopyResponse::new(
294        copy_format,
295        num_columns,
296        copy_stream,
297    )))
298}
299
300pub struct DefaultQueryParser {
301    query_handler: ServerSqlQueryHandlerRef,
302    session: Arc<Session>,
303    compatibility_parser: PostgresCompatibilityParser,
304}
305
306impl DefaultQueryParser {
307    pub fn new(query_handler: ServerSqlQueryHandlerRef, session: Arc<Session>) -> Self {
308        DefaultQueryParser {
309            query_handler,
310            session,
311            compatibility_parser: PostgresCompatibilityParser::new(),
312        }
313    }
314}
315
316/// A container type of parse result types
317#[derive(Clone, Debug)]
318pub struct PgSqlPlan {
319    pub(crate) plan: SqlPlan,
320    pub(crate) copy_to_stdout_format: Option<String>,
321}
322
323#[async_trait]
324impl QueryParser for DefaultQueryParser {
325    type Statement = PgSqlPlan;
326
327    async fn parse_sql<C>(
328        &self,
329        _client: &C,
330        sql: &str,
331        _types: &[Option<Type>],
332    ) -> PgWireResult<Option<Self::Statement>> {
333        crate::metrics::METRIC_POSTGRES_PREPARED_COUNT.inc();
334        let query_ctx = self.session.new_query_context();
335
336        // do not parse if query is empty or matches rules
337        if sql.is_empty() {
338            return Ok(None);
339        }
340
341        if fixtures::matches(sql) {
342            return Ok(Some(PgSqlPlan {
343                plan: SqlPlan::Shortcut(sql.to_string()),
344                copy_to_stdout_format: None,
345            }));
346        }
347
348        let parsed_statements = self.compatibility_parser.parse(sql);
349        let (sql, copy_to_stdout_format) = if let Ok(mut statements) = parsed_statements {
350            if statements.is_empty() {
351                return Ok(None);
352            }
353            let first_stmt = statements.remove(0);
354            let format = check_copy_to_stdout(&first_stmt);
355            (first_stmt.to_string(), format)
356        } else {
357            // bypass the error: it can run into error because of different
358            // versions of sqlparser
359            (sql.to_string(), None)
360        };
361
362        let mut stmts = ParserContext::create_with_dialect(
363            &sql,
364            &PostgreSqlDialect {},
365            ParseOptions::default(),
366        )
367        .map_err(convert_err)?;
368        if stmts.len() != 1 {
369            Err(PgWireError::UserError(Box::new(ErrorInfo::from(
370                PgErrorCode::Ec42P14,
371            ))))
372        } else {
373            let stmt = stmts.remove(0);
374
375            if let Some(logical_plan) = self
376                .query_handler
377                .do_describe(stmt.clone(), query_ctx)
378                .await
379                .map_err(convert_err)?
380                .map(|DescribeResult { logical_plan }| logical_plan)
381            {
382                Ok(Some(PgSqlPlan {
383                    plan: SqlPlan::Plan(logical_plan, stmt),
384                    copy_to_stdout_format,
385                }))
386            } else {
387                Ok(Some(PgSqlPlan {
388                    plan: SqlPlan::Statement(stmt, sql),
389                    copy_to_stdout_format,
390                }))
391            }
392        }
393    }
394
395    fn get_parameter_types(&self, _stmt: &Self::Statement) -> PgWireResult<Vec<Type>> {
396        // we have our own implementation of describes in ExtendedQueryHandler
397        // so we don't use these methods
398        Err(PgWireError::ApiError(
399            "get_parameter_types is not expected to be called".into(),
400        ))
401    }
402
403    fn get_result_schema(
404        &self,
405        _stmt: &Self::Statement,
406        _column_format: Option<&Format>,
407    ) -> PgWireResult<Vec<FieldInfo>> {
408        // we have our own implementation of describes in ExtendedQueryHandler
409        // so we don't use these methods
410        Err(PgWireError::ApiError(
411            "get_result_schema is not expected to be called".into(),
412        ))
413    }
414}
415
416#[async_trait]
417impl ExtendedQueryHandler for PostgresServerHandlerInner {
418    type Statement = PgSqlPlan;
419    type QueryParser = DefaultQueryParser;
420
421    fn query_parser(&self) -> Arc<Self::QueryParser> {
422        self.query_parser.clone()
423    }
424
425    async fn do_query<C>(
426        &self,
427        client: &mut C,
428        portal: &Portal<Self::Statement>,
429        _max_rows: usize,
430    ) -> PgWireResult<Response>
431    where
432        C: ClientInfo + Sink<PgWireBackendMessage> + Unpin + Send + Sync,
433        C::Error: Debug,
434        PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
435    {
436        let query_ctx = self.new_query_context();
437        let db = query_ctx.get_db_string();
438        let _timer = crate::metrics::METRIC_POSTGRES_QUERY_TIMER
439            .with_label_values(&[crate::metrics::METRIC_POSTGRES_EXTENDED_QUERY, db.as_str()])
440            .start_timer();
441
442        let pg_sql_plan = &portal.statement.statement;
443        let sql_plan = &pg_sql_plan.plan;
444
445        let output = match sql_plan {
446            SqlPlan::Empty => {
447                // early return if query is empty
448                return Ok(Response::EmptyQuery);
449            }
450            SqlPlan::Shortcut(query) => {
451                if let Some(mut resps) = fixtures::process(query, query_ctx.clone()) {
452                    send_warning_opt(client, query_ctx).await?;
453                    // if the statement matches our predefined rules, return it early
454                    return Ok(resps.remove(0));
455                } else {
456                    // unreachable logic
457                    return Ok(Response::EmptyQuery);
458                }
459            }
460            SqlPlan::Plan(plan, stmt) => {
461                let values = parameters_to_scalar_values(plan, portal)?;
462                let plan = plan
463                    .clone()
464                    .replace_params_with_values(&ParamValues::List(
465                        values.into_iter().map(Into::into).collect(),
466                    ))
467                    .context(DataFusionSnafu)
468                    .map_err(convert_err)?;
469                self.query_handler
470                    .do_exec_plan(plan, Some(stmt.clone()), query_ctx.clone())
471                    .await
472            }
473            SqlPlan::Statement(_stmt, query) => {
474                // We won't replace params from statement manually any more.
475                // Newer version of datafusion can generate plan for SELECT/INSERT/UPDATE/DELETE.
476                // Only CREATE TABLE and others minor statements cannot generate sql plan,
477                // in this case, we assume these statements will not carry parameters
478                // and execute them directly.
479                self.query_handler
480                    .do_query(query, query_ctx.clone())
481                    .await
482                    .remove(0)
483            }
484        };
485
486        send_warning_opt(client, query_ctx.clone()).await?;
487
488        if let Some(format) = &pg_sql_plan.copy_to_stdout_format {
489            output_to_copy_response(query_ctx, output, format)
490        } else {
491            output_to_query_response(query_ctx, output, &portal.result_column_format)
492        }
493    }
494
495    async fn do_describe_statement<C>(
496        &self,
497        _client: &mut C,
498        stmt: &StoredStatement<Self::Statement>,
499    ) -> PgWireResult<DescribeStatementResponse>
500    where
501        C: ClientInfo + Unpin + Send + Sync,
502    {
503        let sql_plan = &stmt.statement.plan;
504        // client provided parameter types, can be empty if client doesn't try to parse statement
505        let provided_param_types = &stmt.parameter_types;
506        let server_inferenced_types = if let SqlPlan::Plan(plan, _) = &sql_plan {
507            let param_types = DfLogicalPlanner::get_inferred_parameter_types(plan)
508                .context(InferParameterTypesSnafu)
509                .map_err(convert_err)?
510                .into_iter()
511                .map(|(k, v)| (k, v.map(|v| ConcreteDataType::from_arrow_type(&v))))
512                .collect();
513
514            let types = param_types_to_pg_types(&param_types).map_err(convert_err)?;
515
516            Some(types)
517        } else {
518            None
519        };
520
521        let param_count = if provided_param_types.is_empty() {
522            server_inferenced_types
523                .as_ref()
524                .map(|types| types.len())
525                .unwrap_or(0)
526        } else {
527            provided_param_types.len()
528        };
529
530        let param_types = (0..param_count)
531            .map(|i| {
532                let client_type = provided_param_types.get(i);
533                // use server type when client provided type is None (oid: 0 or other invalid values)
534                match client_type {
535                    Some(Some(client_type)) => client_type.clone(),
536                    _ => server_inferenced_types
537                        .as_ref()
538                        .and_then(|types| types.get(i).cloned())
539                        .unwrap_or(Type::UNKNOWN),
540                }
541            })
542            .collect::<Vec<_>>();
543
544        let fields = describe_fields(sql_plan, &Format::UnifiedText, &self.session)?;
545
546        Ok(DescribeStatementResponse::new(param_types, fields))
547    }
548
549    async fn do_describe_portal<C>(
550        &self,
551        _client: &mut C,
552        portal: &Portal<Self::Statement>,
553    ) -> PgWireResult<DescribePortalResponse>
554    where
555        C: ClientInfo + Unpin + Send + Sync,
556    {
557        let sql_plan = &portal.statement.statement.plan;
558        let format = &portal.result_column_format;
559
560        let fields = describe_fields(sql_plan, format, &self.session)?;
561
562        Ok(DescribePortalResponse::new(fields))
563    }
564}
565
566fn describe_fields(
567    sql_plan: &SqlPlan,
568    format: &Format,
569    session: &Arc<Session>,
570) -> PgWireResult<Vec<FieldInfo>> {
571    match sql_plan {
572        // Execution swaps in DistAnalyzeExec (stage/node/plan), whose schema
573        // differs from the logical `Analyze` plan's (plan_type/plan).
574        SqlPlan::Plan(LogicalPlan::Analyze(_), _) => {
575            let schema: Schema =
576                Schema::try_from(dist_analyze_output_schema()).map_err(convert_err)?;
577            schema_to_pg(&schema, format, None).map_err(convert_err)
578        }
579        // query
580        SqlPlan::Plan(plan, _) if !matches!(plan, LogicalPlan::Dml(_) | LogicalPlan::Ddl(_)) => {
581            let schema: Schema = plan.schema().clone().try_into().map_err(convert_err)?;
582            schema_to_pg(&schema, format, None).map_err(convert_err)
583        }
584        // We can cover only part of show statements
585        // these show create statements will return 2 columns
586        SqlPlan::Statement(
587            Statement::ShowCreateDatabase(_)
588            | Statement::ShowCreateTable(_)
589            | Statement::ShowCreateFlow(_)
590            | Statement::ShowCreateView(_),
591            _,
592        ) => Ok(vec![
593            FieldInfo::new(
594                "name".to_string(),
595                None,
596                None,
597                Type::TEXT,
598                format.format_for(0),
599            ),
600            FieldInfo::new(
601                "create_statement".to_string(),
602                None,
603                None,
604                Type::TEXT,
605                format.format_for(1),
606            ),
607        ]),
608        #[cfg(feature = "enterprise")]
609        SqlPlan::Statement(Statement::ShowCreateTrigger(_), _) => Ok(vec![
610            FieldInfo::new(
611                "name".to_string(),
612                None,
613                None,
614                Type::TEXT,
615                format.format_for(0),
616            ),
617            FieldInfo::new(
618                "create_statement".to_string(),
619                None,
620                None,
621                Type::TEXT,
622                format.format_for(1),
623            ),
624        ]),
625        // SHOW FLOW STATUS returns six columns; return their descriptions so
626        // prepared/extended-protocol clients receive the correct row description.
627        SqlPlan::Statement(Statement::ShowFlowStatus(_), _) => Ok(vec![
628            FieldInfo::new(
629                "flow_id".to_string(),
630                None,
631                None,
632                Type::INT8, // matches type_gt_to_pg(UInt32) — do not use INT4
633                format.format_for(0),
634            ),
635            FieldInfo::new(
636                "flow_name".to_string(),
637                None,
638                None,
639                Type::TEXT,
640                format.format_for(1),
641            ),
642            FieldInfo::new(
643                "start_time".to_string(),
644                None,
645                None,
646                Type::TIMESTAMP,
647                format.format_for(2),
648            ),
649            FieldInfo::new(
650                "last_execution_time".to_string(),
651                None,
652                None,
653                Type::TIMESTAMP,
654                format.format_for(3),
655            ),
656            FieldInfo::new(
657                "uptime_seconds".to_string(),
658                None,
659                None,
660                Type::INT8,
661                format.format_for(4),
662            ),
663            FieldInfo::new(
664                "state_size".to_string(),
665                None,
666                None,
667                Type::NUMERIC,
668                format.format_for(5),
669            ),
670        ]),
671
672        #[cfg(feature = "enterprise")]
673        SqlPlan::Statement(Statement::ShowTriggers(_), _) => Ok(vec![FieldInfo::new(
674            "name".to_string(),
675            None,
676            None,
677            Type::TEXT,
678            format.format_for(0),
679        )]),
680        // we will not support other show statements for extended query protocol at least for now.
681        // because the return columns is not predictable at this stage
682        SqlPlan::Shortcut(query) => {
683            // test if query caught by fixture
684            if let Some(mut resp) = fixtures::process(query, session.new_query_context())
685                && let Response::Query(query_response) = resp.remove(0)
686            {
687                Ok((*query_response.row_schema()).clone())
688            } else {
689                // fallback to NoData
690                Ok(vec![])
691            }
692        }
693        // Single column named after the variable (see `query::sql::show_variable`).
694        SqlPlan::Statement(Statement::ShowVariables(show), _) => Ok(vec![FieldInfo::new(
695            show.variable.to_string().to_uppercase(),
696            None,
697            None,
698            Type::TEXT,
699            format.format_for(0),
700        )]),
701        // Mirrors `query::sql::show_status` (currently always empty).
702        SqlPlan::Statement(Statement::ShowStatus(_), _) => Ok(vec![
703            FieldInfo::new(
704                "Variable_name".to_string(),
705                None,
706                None,
707                Type::TEXT,
708                format.format_for(0),
709            ),
710            FieldInfo::new(
711                "Value".to_string(),
712                None,
713                None,
714                Type::TEXT,
715                format.format_for(1),
716            ),
717        ]),
718        SqlPlan::Statement(Statement::ShowSearchPath(_), _) => Ok(vec![FieldInfo::new(
719            "search_path".to_string(),
720            None,
721            None,
722            Type::TEXT,
723            format.format_for(0),
724        )]),
725        // Mirrors `query::sql::describe_table`.
726        SqlPlan::Statement(Statement::DescribeTable(_), _) => {
727            schema_to_pg(&DESCRIBE_TABLE_OUTPUT_SCHEMA, format, None).map_err(convert_err)
728        }
729        // Single column typed with the function's return type (see
730        // `operator::statement::admin_output_schema`).
731        SqlPlan::Statement(Statement::Admin(admin), _) => {
732            let query_ctx = session.new_query_context();
733            match admin_output_schema(admin, &query_ctx) {
734                Some(schema) => schema_to_pg(&schema, format, None).map_err(convert_err),
735                // Unresolvable; execution will surface the error.
736                None => Ok(vec![]),
737            }
738        }
739        // Describe from the declared cursor's schema.
740        SqlPlan::Statement(Statement::FetchCursor(fetch), _) => {
741            let cursor_name = fetch.cursor_name.to_string();
742            match session.get_cursor(&cursor_name) {
743                Some(cursor) => schema_to_pg(&cursor.schema(), format, None).map_err(convert_err),
744                // Cursor not declared yet; execution will error.
745                None => Ok(vec![]),
746            }
747        }
748        _ => {
749            // NoData
750            Ok(vec![])
751        }
752    }
753}
754
755impl ErrorHandler for PostgresServerHandlerInner {
756    fn on_error<C>(&self, _client: &C, error: &mut PgWireError)
757    where
758        C: ClientInfo,
759    {
760        match error {
761            PgWireError::IoError(e) => debug!("Postgres client disconnected: {}", e),
762            _ => info!("Postgres interface error: {}", error),
763        }
764    }
765}
766
767fn check_copy_to_stdout(statement: &SqlParserStatement) -> Option<String> {
768    if let SqlParserStatement::Copy {
769        target, options, ..
770    } = statement
771        && matches!(target, CopyTarget::Stdout)
772    {
773        for opt in options {
774            if let CopyOption::Format(format_ident) = opt {
775                return Some(format_ident.value.to_lowercase());
776            }
777        }
778        return Some("txt".to_string());
779    }
780
781    None
782}
783
784#[cfg(test)]
785mod tests {
786    use datafusion_pg_catalog::sql::PostgresCompatibilityParser;
787
788    use super::*;
789
790    fn parse_copy_statement(sql: &str) -> SqlParserStatement {
791        let parser = PostgresCompatibilityParser::new();
792        let statements = parser.parse(sql).unwrap();
793        statements.into_iter().next().unwrap()
794    }
795
796    #[test]
797    fn test_check_copy_out_with_csv_format() {
798        let statement = parse_copy_statement("COPY (SELECT 1) TO STDOUT WITH (FORMAT CSV)");
799        assert_eq!(check_copy_to_stdout(&statement), Some("csv".to_string()));
800    }
801
802    #[test]
803    fn test_check_copy_out_with_txt_format() {
804        let statement = parse_copy_statement("COPY (SELECT 1) TO STDOUT WITH (FORMAT TXT)");
805        assert_eq!(check_copy_to_stdout(&statement), Some("txt".to_string()));
806    }
807
808    #[test]
809    fn test_check_copy_out_with_binary_format() {
810        let statement = parse_copy_statement("COPY (SELECT 1) TO STDOUT WITH (FORMAT BINARY)");
811        assert_eq!(check_copy_to_stdout(&statement), Some("binary".to_string()));
812    }
813
814    #[test]
815    fn test_check_copy_out_without_format() {
816        let statement = parse_copy_statement("COPY (SELECT 1) TO STDOUT");
817        assert_eq!(check_copy_to_stdout(&statement), Some("txt".to_string()));
818    }
819
820    #[test]
821    fn test_check_copy_out_to_file() {
822        let statement =
823            parse_copy_statement("COPY (SELECT 1) TO '/path/to/file.csv' WITH (FORMAT CSV)");
824        assert_eq!(check_copy_to_stdout(&statement), None);
825    }
826
827    #[test]
828    fn test_check_copy_out_case_insensitive() {
829        let statement = parse_copy_statement("COPY (SELECT 1) TO STDOUT WITH (FORMAT csv)");
830        assert_eq!(check_copy_to_stdout(&statement), Some("csv".to_string()));
831
832        let statement = parse_copy_statement("COPY (SELECT 1) TO STDOUT WITH (FORMAT binary)");
833        assert_eq!(check_copy_to_stdout(&statement), Some("binary".to_string()));
834    }
835
836    #[test]
837    fn test_check_copy_out_with_multiple_options() {
838        let statement = parse_copy_statement(
839            "COPY (SELECT 1) TO STDOUT WITH (FORMAT csv, DELIMITER ',', HEADER)",
840        );
841        assert_eq!(check_copy_to_stdout(&statement), Some("csv".to_string()));
842
843        let statement = parse_copy_statement(
844            "COPY (SELECT 1) TO STDOUT WITH (DELIMITER ',', HEADER, FORMAT binary)",
845        );
846        assert_eq!(check_copy_to_stdout(&statement), Some("binary".to_string()));
847    }
848}