1use 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 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 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) = ©_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#[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 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 (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 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 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 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 return Ok(resps.remove(0));
455 } else {
456 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 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 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(¶m_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 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 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 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 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 SqlPlan::Statement(Statement::ShowFlowStatus(_), _) => Ok(vec![
628 FieldInfo::new(
629 "flow_id".to_string(),
630 None,
631 None,
632 Type::INT8, 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 SqlPlan::Shortcut(query) => {
683 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 Ok(vec![])
691 }
692 }
693 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 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 SqlPlan::Statement(Statement::DescribeTable(_), _) => {
727 schema_to_pg(&DESCRIBE_TABLE_OUTPUT_SCHEMA, format, None).map_err(convert_err)
728 }
729 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 None => Ok(vec![]),
737 }
738 }
739 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 None => Ok(vec![]),
746 }
747 }
748 _ => {
749 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}